Compare commits
2 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 35e7670376 | |||
| 1aa47be16a |
@@ -1,162 +0,0 @@
|
||||
name: ACR Cleanup
|
||||
|
||||
on:
|
||||
schedule:
|
||||
- cron: '0 19 * * *' # UTC 19:00 = 北京时间凌晨3:00
|
||||
workflow_dispatch:
|
||||
inputs:
|
||||
pr_sha:
|
||||
description: "PR commit SHA(仅清理指定PR镜像,留空则全量清理)"
|
||||
required: false
|
||||
default: ""
|
||||
dry_run:
|
||||
description: "预览模式(dry-run),不实际删除"
|
||||
required: false
|
||||
default: "true"
|
||||
pull_request_target:
|
||||
types: [closed]
|
||||
branches: [develop, main]
|
||||
|
||||
concurrency:
|
||||
group: acr-cleanup-${{ gitea.ref }}
|
||||
cancel-in-progress: false
|
||||
|
||||
jobs:
|
||||
cleanup:
|
||||
name: ACR Image Cleanup
|
||||
runs-on: ci-l2
|
||||
timeout-minutes: 20
|
||||
permissions:
|
||||
contents: read
|
||||
env:
|
||||
ACR_REGISTRY: xiaoxia-registry.cn-hangzhou.cr.aliyuncs.com
|
||||
ACR_NAMESPACE: xiaoxiakeji
|
||||
ACR_SERVICE: registry.aliyuncs.com:cn-hangzhou:china:cri-fvec8o9q4mmxrkaa
|
||||
GITEA_URL: https://git.xiaoxiajianji.com
|
||||
GITEA_REPO: xiaoxia/xiaoxia-saas
|
||||
steps:
|
||||
- name: Checkout code
|
||||
uses: actions/checkout@v4
|
||||
|
||||
- name: Setup Python
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: '3.12'
|
||||
|
||||
# ====== Cron模式:获取staging运行中镜像作为白名单 ======
|
||||
- name: Get staging running images (whitelist)
|
||||
id: protected_images
|
||||
if: gitea.event_name != 'pull_request_target' && !gitea.event.inputs.pr_sha
|
||||
env:
|
||||
STAGING_SSH_KEY: ${{ secrets.PREVIEW_SSH_KEY }}
|
||||
run: |
|
||||
set +e
|
||||
echo "获取staging服务器运行中镜像作为白名单..."
|
||||
mkdir -p ~/.ssh
|
||||
echo "$STAGING_SSH_KEY" > ~/.ssh/id_rsa
|
||||
chmod 600 ~/.ssh/id_rsa
|
||||
|
||||
staging_host="${STAGING_SSH_HOST:-47.98.113.167}"
|
||||
staging_port="${STAGING_SSH_PORT:-22222}"
|
||||
|
||||
ssh-keyscan -p "$staging_port" -H "$staging_host" >> ~/.ssh/known_hosts 2>/dev/null
|
||||
|
||||
# 获取所有运行容器的镜像,提取tag部分
|
||||
IMAGES=$(ssh -p "$staging_port" -i ~/.ssh/id_rsa -o StrictHostKeyChecking=no \
|
||||
"root@$staging_host" "docker ps --format '{{.Image}}' 2>/dev/null" 2>/dev/null | grep -v "^$" | sort -u)
|
||||
|
||||
PROTECTED_TAGS=""
|
||||
if [ -n "$IMAGES" ]; then
|
||||
while IFS= read -r img; do
|
||||
# 从完整镜像名中提取tag(最后一个冒号后)
|
||||
tag=$(echo "$img" | rev | cut -d: -f1 | rev)
|
||||
if [ -n "$tag" ] && [ "$tag" != "latest" ] && [ ${#tag} -gt 5 ]; then
|
||||
if [ -z "$PROTECTED_TAGS" ]; then
|
||||
PROTECTED_TAGS="$tag"
|
||||
else
|
||||
PROTECTED_TAGS="$PROTECTED_TAGS,$tag"
|
||||
fi
|
||||
fi
|
||||
done <<< "$IMAGES"
|
||||
fi
|
||||
|
||||
echo "staging运行中镜像tag: ${PROTECTED_TAGS:-(无)}"
|
||||
echo "protected_tags=$PROTECTED_TAGS" >> $GITEA_OUTPUT
|
||||
|
||||
# ====== Docker登录 ======
|
||||
- name: Docker login to ACR
|
||||
env:
|
||||
ACR_USERNAME: ${{ secrets.ACR_USERNAME }}
|
||||
ACR_PASSWORD: ${{ secrets.ACR_PASSWORD }}
|
||||
run: |
|
||||
printf '%s' "$ACR_PASSWORD" | docker login "$ACR_REGISTRY" -u "$ACR_USERNAME" --password-stdin
|
||||
|
||||
# ====== 模式1:PR关闭时清理 ======
|
||||
- name: Cleanup PR images (PR closed)
|
||||
if: gitea.event_name == 'pull_request_target'
|
||||
env:
|
||||
ACR_USERNAME: ${{ secrets.ACR_USERNAME }}
|
||||
ACR_PASSWORD: ${{ secrets.ACR_PASSWORD }}
|
||||
GITEA_TOKEN: ${{ secrets.GITEA_TOKEN }}
|
||||
PR_SHA: ${{ gitea.event.pull_request.head.sha }}
|
||||
PR_NUMBER: ${{ gitea.event.pull_request.number }}
|
||||
run: |
|
||||
echo "============================================"
|
||||
echo " PR #$PR_NUMBER 已关闭,清理对应镜像"
|
||||
echo " Head SHA: ${PR_SHA::12}"
|
||||
echo "============================================"
|
||||
echo ""
|
||||
python3 scripts/ci/acr_cleanup.py \
|
||||
--pr-sha "$PR_SHA" \
|
||||
--execute
|
||||
|
||||
# ====== 模式2:Cron全量清理 ======
|
||||
- name: Full cleanup (cron / manual)
|
||||
if: gitea.event_name != 'pull_request_target' && !gitea.event.inputs.pr_sha
|
||||
env:
|
||||
ACR_USERNAME: ${{ secrets.ACR_USERNAME }}
|
||||
ACR_PASSWORD: ${{ secrets.ACR_PASSWORD }}
|
||||
GITEA_TOKEN: ${{ secrets.GITEA_TOKEN }}
|
||||
PROTECTED_TAGS: ${{ steps.protected_images.outputs.protected_tags }}
|
||||
DRY_RUN_INPUT: ${{ gitea.event.inputs.dry_run }}
|
||||
run: |
|
||||
echo "============================================"
|
||||
echo " ACR 全量清理(${{ gitea.event_name }})"
|
||||
echo "============================================"
|
||||
echo ""
|
||||
|
||||
# 决定是否dry-run
|
||||
DRY_RUN_FLAG=""
|
||||
if [ "$DRY_RUN_INPUT" = "true" ]; then
|
||||
DRY_RUN_FLAG="--dry-run"
|
||||
echo "模式: 预览模式 (dry-run)"
|
||||
else
|
||||
echo "模式: 执行模式"
|
||||
fi
|
||||
echo ""
|
||||
|
||||
python3 scripts/ci/acr_cleanup.py \
|
||||
--keep 20 \
|
||||
--protected-tags "$PROTECTED_TAGS" \
|
||||
$DRY_RUN_FLAG
|
||||
|
||||
# ====== 模式3:手动指定PR SHA清理 ======
|
||||
- name: Cleanup specific PR image (manual)
|
||||
if: gitea.event_name == 'workflow_dispatch' && gitea.event.inputs.pr_sha
|
||||
env:
|
||||
ACR_USERNAME: ${{ secrets.ACR_USERNAME }}
|
||||
ACR_PASSWORD: ${{ secrets.ACR_PASSWORD }}
|
||||
PR_SHA: ${{ gitea.event.inputs.pr_sha }}
|
||||
DRY_RUN_INPUT: ${{ gitea.event.inputs.dry_run }}
|
||||
run: |
|
||||
echo "手动清理PR镜像: ${PR_SHA::12}"
|
||||
echo ""
|
||||
|
||||
DRY_RUN_FLAG=""
|
||||
if [ "$DRY_RUN_INPUT" = "true" ]; then
|
||||
DRY_RUN_FLAG="--dry-run"
|
||||
fi
|
||||
|
||||
python3 scripts/ci/acr_cleanup.py \
|
||||
--pr-sha "$PR_SHA" \
|
||||
$DRY_RUN_FLAG
|
||||
@@ -1230,7 +1230,7 @@ jobs:
|
||||
runs-on: runtime-builder
|
||||
timeout-minutes: ${{ matrix.timeout }}
|
||||
needs:
|
||||
if: startsWith(github.ref, 'refs/tags/v') || (github.event_name == 'push' && github.ref_name == 'main')
|
||||
if: startsWith(github.ref, 'refs/tags/v')
|
||||
strategy:
|
||||
fail-fast: false
|
||||
matrix:
|
||||
@@ -1307,16 +1307,10 @@ jobs:
|
||||
run: |
|
||||
set -eu
|
||||
REGISTRY="xiaoxia-registry.cn-hangzhou.cr.aliyuncs.com/xiaoxiakeji"
|
||||
# 根据ref类型设置镜像标签:tag用版本号,分支用分支名+sha
|
||||
if [[ "$GITHUB_REF" == refs/tags/* ]]; then
|
||||
TAG_NAME="${GITHUB_REF_NAME}"
|
||||
else
|
||||
TAG_NAME="${GITHUB_REF_NAME}-${GITHUB_SHA::8}"
|
||||
fi
|
||||
IMAGE_TAG="${REGISTRY}/${{ matrix.image_name }}:${TAG_NAME}"
|
||||
IMAGE_TAG="${REGISTRY}/${{ matrix.image_name }}:${GITHUB_REF_NAME}"
|
||||
CACHE_REF="${REGISTRY}/${{ matrix.cache_name }}:main"
|
||||
|
||||
EXTRA_BUILD_ARGS="APP_VERSION=\"${TAG_NAME}\""
|
||||
EXTRA_BUILD_ARGS="APP_VERSION=\"${GITHUB_REF_NAME}\""
|
||||
if [ "${{ matrix.service }}" = "web" ]; then
|
||||
EXTRA_BUILD_ARGS="$EXTRA_BUILD_ARGS NGINX_CONF=infra/docker/nginx-production.conf"
|
||||
fi
|
||||
@@ -1614,89 +1608,4 @@ jobs:
|
||||
[ ${{ job.status }} = "success" ] || STATUS="error"
|
||||
START_TIME=""
|
||||
[ -f /tmp/ci_job_start_time ] && START_TIME=$(cat /tmp/ci_job_start_time)
|
||||
python3 scripts/ci/ci_trace_report.py --service xiaoxia-saas-ci --status $STATUS --start-time "$START_TIME" || true
|
||||
|
||||
canary-release:
|
||||
name: Canary Release to Production
|
||||
runs-on: runtime-builder
|
||||
timeout-minutes: 120
|
||||
concurrency:
|
||||
group: canary-release-production
|
||||
cancel-in-progress: false
|
||||
if: github.event_name == 'push' && github.ref_name == 'main'
|
||||
needs:
|
||||
- build-production
|
||||
- staging-api-tests
|
||||
steps:
|
||||
- name: Checkout code
|
||||
shell: sh
|
||||
env:
|
||||
GITHUB_TOKEN: ${{ github.token }}
|
||||
run: |
|
||||
curl -sH "Authorization: token $GITHUB_TOKEN" "${GITHUB_API_URL}/repos/${GITHUB_REPOSITORY}/raw/scripts/ci/step_checkout.sh?ref=${GITHUB_SHA}" | bash
|
||||
- name: Record job start time
|
||||
shell: sh
|
||||
run: bash scripts/ci/step_timer_start.sh
|
||||
- name: Notify canary release start
|
||||
continue-on-error: true
|
||||
shell: sh
|
||||
env:
|
||||
CI_NOTIFY_WEBHOOK: ${{ secrets.CI_NOTIFY_WEBHOOK }}
|
||||
run: |
|
||||
set +e
|
||||
NOTIFY_MODE=start JOB_NAME="Canary Release" python3 scripts/ci_notify.py
|
||||
- name: Install SSH client
|
||||
shell: sh
|
||||
run: |
|
||||
set -eu
|
||||
apt-get update -qq && apt-get install -y -qq openssh-client curl >/dev/null 2>&1
|
||||
echo "openssh-client installed"
|
||||
- name: Run canary release
|
||||
shell: bash
|
||||
env:
|
||||
PRODUCTION_SSH_HOST: ${{ secrets.PRODUCTION_SSH_HOST }}
|
||||
PRODUCTION_SSH_USER: ${{ secrets.PRODUCTION_SSH_USER }}
|
||||
PRODUCTION_SSH_PORT: ${{ secrets.PRODUCTION_SSH_PORT }}
|
||||
PRODUCTION_SSH_KEY: ${{ secrets.PRODUCTION_SSH_KEY }}
|
||||
ACR_USERNAME: ${{ secrets.ACR_USERNAME }}
|
||||
ACR_PASSWORD: ${{ secrets.ACR_PASSWORD }}
|
||||
CI_NOTIFY_WEBHOOK: ${{ secrets.CI_NOTIFY_WEBHOOK }}
|
||||
run: |
|
||||
set -eu
|
||||
IMAGE_TAG="main-${GITHUB_SHA::8}"
|
||||
export IMAGE_TAG
|
||||
echo "Canary release version: $IMAGE_TAG"
|
||||
bash scripts/ci/canary_release.sh
|
||||
- name: Job duration summary
|
||||
if: always()
|
||||
shell: sh
|
||||
run: bash scripts/ci/step_timer_end.sh
|
||||
- name: Notify on success
|
||||
continue-on-error: true
|
||||
if: success()
|
||||
shell: sh
|
||||
env:
|
||||
CI_NOTIFY_WEBHOOK: ${{ secrets.CI_NOTIFY_WEBHOOK }}
|
||||
run: |
|
||||
set +e
|
||||
NOTIFY_MODE=success JOB_NAME="Canary Release" python3 scripts/ci_notify.py
|
||||
- name: Notify on failure
|
||||
continue-on-error: true
|
||||
if: failure()
|
||||
shell: sh
|
||||
env:
|
||||
CI_NOTIFY_WEBHOOK: ${{ secrets.CI_NOTIFY_WEBHOOK }}
|
||||
run: |
|
||||
set +e
|
||||
NOTIFY_MODE=failure JOB_NAME="Canary Release" python3 scripts/ci_notify.py
|
||||
- name: Report CI trace
|
||||
if: always()
|
||||
shell: sh
|
||||
env:
|
||||
AGENTLOOP_LICENSE_KEY: ${{ secrets.AGENTLOOP_LICENSE_KEY }}
|
||||
run: |
|
||||
STATUS="ok"
|
||||
[ ${{ job.status }} = "success" ] || STATUS="error"
|
||||
START_TIME=""
|
||||
[ -f /tmp/ci_job_start_time ] && START_TIME=$(cat /tmp/ci_job_start_time)
|
||||
python3 scripts/ci/ci_trace_report.py --service xiaoxia-saas-ci --status $STATUS --start-time "$START_TIME" || true
|
||||
python3 scripts/ci/ci_trace_report.py --service xiaoxia-saas-ci --status $STATUS --start-time "$START_TIME" || true
|
||||
@@ -1,64 +0,0 @@
|
||||
"""#642 - 生成任务新增 bgm_config 字段
|
||||
|
||||
Revision ID: 052_generation_task_bgm_config
|
||||
Revises: 051_generation_task_resolution
|
||||
Create Date: 2026-07-25
|
||||
|
||||
Changes:
|
||||
1. generation_tasks 表新增 bgm_config 字段(JSON类型),存储用户自定义BGM配置
|
||||
2. 为空时使用默认空字典
|
||||
|
||||
背景:
|
||||
#642 一键生成支持自定义BGM 功能在 SQLAlchemy 模型中加了 bgm_config 字段,
|
||||
但遗漏了 alembic migration,导致 staging 环境数据库没有该列,
|
||||
创建生成任务时直接 500。
|
||||
"""
|
||||
|
||||
import sqlalchemy as sa
|
||||
|
||||
from alembic import context, op
|
||||
|
||||
revision = "052_generation_task_bgm_config"
|
||||
down_revision = "051_generation_task_resolution"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
conn = op.get_bind()
|
||||
if context.get_context().dialect.name == "postgresql":
|
||||
# 检查列是否已存在(幂等)
|
||||
result = conn.execute(
|
||||
sa.text(
|
||||
"SELECT column_name FROM information_schema.columns "
|
||||
"WHERE table_name = 'generation_tasks' AND column_name = 'bgm_config'"
|
||||
)
|
||||
)
|
||||
if result.scalar() is not None:
|
||||
return
|
||||
|
||||
op.add_column(
|
||||
"generation_tasks",
|
||||
sa.Column(
|
||||
"bgm_config",
|
||||
sa.JSON,
|
||||
nullable=False,
|
||||
server_default=sa.text("'{}'::json"),
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
conn = op.get_bind()
|
||||
if context.get_context().dialect.name == "postgresql":
|
||||
# 检查列是否存在(幂等)
|
||||
result = conn.execute(
|
||||
sa.text(
|
||||
"SELECT column_name FROM information_schema.columns "
|
||||
"WHERE table_name = 'generation_tasks' AND column_name = 'bgm_config'"
|
||||
)
|
||||
)
|
||||
if result.scalar() is None:
|
||||
return
|
||||
|
||||
op.drop_column("generation_tasks", "bgm_config")
|
||||
Generated
-6117
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
@@ -1,128 +0,0 @@
|
||||
import React from "react"
|
||||
import {
|
||||
PlayCircleOutlined,
|
||||
CheckOutlined,
|
||||
DeleteOutlined,
|
||||
ExperimentOutlined,
|
||||
LoadingOutlined,
|
||||
CloseCircleOutlined,
|
||||
} from "@ant-design/icons"
|
||||
import { Popconfirm } from "antd"
|
||||
import type { AssetItem } from "@/pages/assets/types"
|
||||
import { thumbGradient } from "@/pages/assets/utils/asset"
|
||||
import { kindIcon } from "@/pages/assets/utils/kindIcon"
|
||||
import { StatusPill } from "./AssetSkeleton"
|
||||
|
||||
/* ============================================================
|
||||
* AssetCard — 素材卡片(网格视图)
|
||||
* ============================================================ */
|
||||
export interface AssetCardProps {
|
||||
asset: AssetItem
|
||||
selected: boolean
|
||||
diagnosing?: boolean
|
||||
onToggle: () => void
|
||||
onDiagnose: () => void
|
||||
onPlay: () => void
|
||||
onDelete: () => void
|
||||
}
|
||||
|
||||
const AssetCard: React.FC<AssetCardProps> = ({
|
||||
asset,
|
||||
selected,
|
||||
diagnosing,
|
||||
onToggle,
|
||||
onDiagnose,
|
||||
onPlay,
|
||||
onDelete,
|
||||
}) => (
|
||||
<div className={`xx-asset-card${selected ? " xx-asset-card-selected" : ""}`} onClick={onToggle}>
|
||||
{/* 缩略图区 */}
|
||||
<div className="xx-asset-thumb" style={{ background: thumbGradient(asset.kind) }}>
|
||||
{asset.thumbUrl ? (
|
||||
<img src={asset.thumbUrl} alt={asset.name} />
|
||||
) : (
|
||||
<span className="xx-asset-thumb-placeholder">
|
||||
{asset.loading ? <LoadingOutlined /> : kindIcon(asset.kind)}
|
||||
</span>
|
||||
)}
|
||||
|
||||
{/* 处理中遮罩 */}
|
||||
{asset.loading && (
|
||||
<div className="xx-asset-thumb-overlay xx-asset-thumb-processing">
|
||||
<LoadingOutlined />
|
||||
<span>处理中</span>
|
||||
</div>
|
||||
)}
|
||||
|
||||
{/* 失败状态标识 */}
|
||||
{asset.status === "bad" && asset.statusLabel === "处理失败" && (
|
||||
<div className="xx-asset-thumb-overlay xx-asset-thumb-failed">
|
||||
<CloseCircleOutlined />
|
||||
<span>处理失败</span>
|
||||
</div>
|
||||
)}
|
||||
|
||||
{/* 视频/配音类显示播放按钮(处理中/失败不显示) */}
|
||||
{asset.kind === "video" && !asset.loading && asset.status !== "bad" && (
|
||||
<span
|
||||
className="xx-asset-play"
|
||||
onClick={(e) => {
|
||||
e.stopPropagation()
|
||||
onPlay()
|
||||
}}
|
||||
>
|
||||
<PlayCircleOutlined />
|
||||
</span>
|
||||
)}
|
||||
|
||||
{/* 删除按钮 */}
|
||||
<Popconfirm
|
||||
title="确认删除"
|
||||
description="删除后不可恢复,确定要删除这个素材吗?"
|
||||
onConfirm={(e) => {
|
||||
e?.stopPropagation()
|
||||
onDelete()
|
||||
}}
|
||||
onCancel={(e) => e?.stopPropagation()}
|
||||
okText="删除"
|
||||
cancelText="取消"
|
||||
okButtonProps={{ danger: true }}
|
||||
>
|
||||
<span className="xx-asset-delete" onClick={(e) => e.stopPropagation()}>
|
||||
<DeleteOutlined />
|
||||
</span>
|
||||
</Popconfirm>
|
||||
|
||||
{/* 选中态勾选 */}
|
||||
{selected && (
|
||||
<span className="xx-asset-check">
|
||||
<CheckOutlined />
|
||||
</span>
|
||||
)}
|
||||
</div>
|
||||
|
||||
{/* 信息区 */}
|
||||
<div className="xx-asset-info">
|
||||
<p className="xx-asset-name" title={asset.name}>
|
||||
{asset.name}
|
||||
</p>
|
||||
<div className="xx-asset-meta">
|
||||
<StatusPill status={asset.status} label={asset.statusLabel} />
|
||||
{asset.duration && <span>{asset.duration}</span>}
|
||||
</div>
|
||||
<button
|
||||
className={`xx-asset-diagnose-btn${diagnosing ? " xx-asset-diagnose-btn-loading" : ""}`}
|
||||
disabled={diagnosing || asset.loading || asset.status === "bad"}
|
||||
onClick={(e) => {
|
||||
e.stopPropagation()
|
||||
onDiagnose()
|
||||
}}
|
||||
>
|
||||
{diagnosing ? <LoadingOutlined /> : <ExperimentOutlined />}
|
||||
{diagnosing ? "诊断中..." : "诊断"}
|
||||
</button>
|
||||
</div>
|
||||
</div>
|
||||
)
|
||||
|
||||
export default AssetCard
|
||||
@@ -1,62 +0,0 @@
|
||||
import React from "react"
|
||||
import { SearchOutlined } from "@ant-design/icons"
|
||||
import { Button, Input, Select } from "@/components/ui"
|
||||
import { TYPE_FILTER_OPTIONS, TIME_FILTER_OPTIONS } from "@/pages/assets/constants"
|
||||
|
||||
/* ============================================================
|
||||
* AssetFilterBar — 筛选栏(搜索/类型/时间 + 全选/计数)
|
||||
* ============================================================ */
|
||||
export interface AssetFilterBarProps {
|
||||
searchText: string
|
||||
onSearchChange: (value: string) => void
|
||||
filterType: string
|
||||
onFilterTypeChange: (value: string) => void
|
||||
filterTime: string
|
||||
onFilterTimeChange: (value: string) => void
|
||||
onSelectAll: () => void
|
||||
totalCount: number
|
||||
}
|
||||
|
||||
const AssetFilterBar: React.FC<AssetFilterBarProps> = ({
|
||||
searchText,
|
||||
onSearchChange,
|
||||
filterType,
|
||||
onFilterTypeChange,
|
||||
filterTime,
|
||||
onFilterTimeChange,
|
||||
onSelectAll,
|
||||
totalCount,
|
||||
}) => (
|
||||
<div className="xx-assets-filters">
|
||||
<div className="xx-assets-filters-left">
|
||||
<Input
|
||||
placeholder="搜索素材名称..."
|
||||
prefix={<SearchOutlined />}
|
||||
value={searchText}
|
||||
onChange={(e) => onSearchChange(e.target.value)}
|
||||
allowClear
|
||||
style={{ width: 220 }}
|
||||
/>
|
||||
<Select
|
||||
value={filterType}
|
||||
onChange={onFilterTypeChange}
|
||||
style={{ width: 120 }}
|
||||
options={TYPE_FILTER_OPTIONS}
|
||||
/>
|
||||
<Select
|
||||
value={filterTime}
|
||||
onChange={onFilterTimeChange}
|
||||
style={{ width: 120 }}
|
||||
options={TIME_FILTER_OPTIONS}
|
||||
/>
|
||||
</div>
|
||||
<div className="xx-assets-filters-right">
|
||||
<Button buttonType="ghost" buttonSize="sm" onClick={onSelectAll}>
|
||||
全选
|
||||
</Button>
|
||||
<span className="xx-assets-filter-count">共 {totalCount} 个素材</span>
|
||||
</div>
|
||||
</div>
|
||||
)
|
||||
|
||||
export default AssetFilterBar
|
||||
@@ -1,36 +0,0 @@
|
||||
import React from "react"
|
||||
import type { StatusType } from "@/pages/assets/types"
|
||||
|
||||
/* ============================================================
|
||||
* StatusPill — 状态标签
|
||||
* ============================================================ */
|
||||
export interface StatusPillProps {
|
||||
status: StatusType
|
||||
label: string
|
||||
}
|
||||
|
||||
export const StatusPill: React.FC<StatusPillProps> = ({ status, label }) => (
|
||||
<span className={`xx-status-pill xx-status-pill-${status}`}>{label}</span>
|
||||
)
|
||||
|
||||
/* ============================================================
|
||||
* SkeletonCard — 骨架屏卡片(素材列表加载时占位)
|
||||
* ============================================================ */
|
||||
export const SkeletonCard: React.FC = () => (
|
||||
<div className="xx-asset-card xx-asset-skeleton">
|
||||
<div className="xx-asset-thumb xx-skeleton-pulse" />
|
||||
<div className="xx-asset-info">
|
||||
<div className="xx-skeleton-line xx-skeleton-pulse" style={{ width: "70%" }} />
|
||||
<div className="xx-skeleton-line xx-skeleton-pulse" style={{ width: "40%", marginTop: 8 }} />
|
||||
<div
|
||||
className="xx-skeleton-line xx-skeleton-pulse"
|
||||
style={{
|
||||
width: "100%",
|
||||
height: 28,
|
||||
marginTop: 8,
|
||||
borderRadius: "var(--radius-xs)",
|
||||
}}
|
||||
/>
|
||||
</div>
|
||||
</div>
|
||||
)
|
||||
@@ -1,49 +0,0 @@
|
||||
import React from "react"
|
||||
import { Modal as AntModal, Select as AntSelect } from "antd"
|
||||
import { CATEGORY_OPTIONS } from "@/pages/assets/constants"
|
||||
|
||||
/* ============================================================
|
||||
* BatchClassifyModal — 批量改分类弹窗
|
||||
* ============================================================ */
|
||||
export interface BatchClassifyModalProps {
|
||||
open: boolean
|
||||
selectedCount: number
|
||||
onCancel: () => void
|
||||
onOk: () => void
|
||||
category: string
|
||||
onCategoryChange: (value: string) => void
|
||||
confirmLoading?: boolean
|
||||
}
|
||||
|
||||
const BatchClassifyModal: React.FC<BatchClassifyModalProps> = ({
|
||||
open,
|
||||
selectedCount,
|
||||
onCancel,
|
||||
onOk,
|
||||
category,
|
||||
onCategoryChange,
|
||||
confirmLoading,
|
||||
}) => (
|
||||
<AntModal
|
||||
title={`批量改分类(${selectedCount} 个素材)`}
|
||||
open={open}
|
||||
onCancel={onCancel}
|
||||
onOk={onOk}
|
||||
confirmLoading={confirmLoading}
|
||||
okText="确认修改"
|
||||
cancelText="取消"
|
||||
>
|
||||
<div className="xx-batch-classify-modal">
|
||||
<p className="xx-batch-classify-hint">将选中的 {selectedCount} 个素材统一修改为以下分类:</p>
|
||||
<AntSelect
|
||||
value={category || undefined}
|
||||
onChange={(v) => onCategoryChange(v)}
|
||||
placeholder="请选择分类"
|
||||
style={{ width: "100%" }}
|
||||
options={CATEGORY_OPTIONS}
|
||||
/>
|
||||
</div>
|
||||
</AntModal>
|
||||
)
|
||||
|
||||
export default BatchClassifyModal
|
||||
@@ -1,67 +0,0 @@
|
||||
import React from "react"
|
||||
import { Modal as AntModal, Tag, Radio } from "antd"
|
||||
|
||||
/* ============================================================
|
||||
* BatchMarkModal — 批量智能标记弹窗
|
||||
* ============================================================ */
|
||||
export type SmartViewType = "recommended" | "caution" | "high_risk"
|
||||
|
||||
export interface BatchMarkModalProps {
|
||||
open: boolean
|
||||
selectedCount: number
|
||||
onCancel: () => void
|
||||
onOk: () => void
|
||||
smartView: SmartViewType
|
||||
onSmartViewChange: (value: SmartViewType) => void
|
||||
confirmLoading?: boolean
|
||||
}
|
||||
|
||||
const BatchMarkModal: React.FC<BatchMarkModalProps> = ({
|
||||
open,
|
||||
selectedCount,
|
||||
onCancel,
|
||||
onOk,
|
||||
smartView,
|
||||
onSmartViewChange,
|
||||
confirmLoading,
|
||||
}) => (
|
||||
<AntModal
|
||||
title={`批量智能标记(${selectedCount} 个素材)`}
|
||||
open={open}
|
||||
onCancel={onCancel}
|
||||
onOk={onOk}
|
||||
confirmLoading={confirmLoading}
|
||||
okText="确认标记"
|
||||
cancelText="取消"
|
||||
>
|
||||
<div className="xx-batch-mark-modal">
|
||||
<p className="xx-batch-mark-hint">将选中的 {selectedCount} 个素材标记为:</p>
|
||||
<Radio.Group
|
||||
value={smartView}
|
||||
onChange={(e) => onSmartViewChange(e.target.value)}
|
||||
className="xx-batch-mark-options"
|
||||
>
|
||||
<div className="xx-batch-mark-option">
|
||||
<Radio value="recommended">
|
||||
<Tag color="success">推荐</Tag>
|
||||
<span className="xx-batch-mark-desc">质量优良,可直接用于生产</span>
|
||||
</Radio>
|
||||
</div>
|
||||
<div className="xx-batch-mark-option">
|
||||
<Radio value="caution">
|
||||
<Tag color="warning">慎用</Tag>
|
||||
<span className="xx-batch-mark-desc">存在一定问题,需人工审核后再使用</span>
|
||||
</Radio>
|
||||
</div>
|
||||
<div className="xx-batch-mark-option">
|
||||
<Radio value="high_risk">
|
||||
<Tag color="error">高风险</Tag>
|
||||
<span className="xx-batch-mark-desc">存在严重问题,不建议使用</span>
|
||||
</Radio>
|
||||
</div>
|
||||
</Radio.Group>
|
||||
</div>
|
||||
</AntModal>
|
||||
)
|
||||
|
||||
export default BatchMarkModal
|
||||
@@ -1,58 +0,0 @@
|
||||
import React from "react"
|
||||
import {
|
||||
TagsOutlined,
|
||||
FolderOutlined,
|
||||
ThunderboltOutlined,
|
||||
DeleteOutlined,
|
||||
} from "@ant-design/icons"
|
||||
import { Popconfirm } from "antd"
|
||||
import { Button } from "@/components/ui"
|
||||
|
||||
/* ============================================================
|
||||
* BatchOperationBar — 批量操作栏
|
||||
* ============================================================ */
|
||||
export interface BatchOperationBarProps {
|
||||
selectedCount: number
|
||||
onDeselectAll: () => void
|
||||
onTagClick: () => void
|
||||
onClassifyClick: () => void
|
||||
onMarkClick: () => void
|
||||
onBatchDelete: () => void
|
||||
}
|
||||
|
||||
const BatchOperationBar: React.FC<BatchOperationBarProps> = ({
|
||||
selectedCount,
|
||||
onDeselectAll,
|
||||
onTagClick,
|
||||
onClassifyClick,
|
||||
onMarkClick,
|
||||
onBatchDelete,
|
||||
}) => (
|
||||
<div className="xx-assets-batch-bar">
|
||||
<span className="xx-assets-batch-count">已选 {selectedCount} 项</span>
|
||||
<Button buttonType="ghost" buttonSize="sm" onClick={onDeselectAll}>
|
||||
取消选择
|
||||
</Button>
|
||||
<Button buttonType="ghost" buttonSize="sm" icon={<TagsOutlined />} onClick={onTagClick}>
|
||||
打标签
|
||||
</Button>
|
||||
<Button buttonType="ghost" buttonSize="sm" icon={<FolderOutlined />} onClick={onClassifyClick}>
|
||||
改分类
|
||||
</Button>
|
||||
<Button buttonType="ghost" buttonSize="sm" icon={<ThunderboltOutlined />} onClick={onMarkClick}>
|
||||
智能标记
|
||||
</Button>
|
||||
<Popconfirm
|
||||
title={`确定删除 ${selectedCount} 个素材?`}
|
||||
onConfirm={onBatchDelete}
|
||||
okText="删除"
|
||||
cancelText="取消"
|
||||
>
|
||||
<Button buttonType="danger" buttonSize="sm" icon={<DeleteOutlined />}>
|
||||
批量删除
|
||||
</Button>
|
||||
</Popconfirm>
|
||||
</div>
|
||||
)
|
||||
|
||||
export default BatchOperationBar
|
||||
@@ -1,81 +0,0 @@
|
||||
import React from "react"
|
||||
import { Modal as AntModal, Tag, Radio, Input as AntInput } from "antd"
|
||||
import { ExclamationCircleOutlined } from "@ant-design/icons"
|
||||
|
||||
/* ============================================================
|
||||
* BatchTagModal — 批量打标签弹窗
|
||||
* ============================================================ */
|
||||
export interface BatchTagModalProps {
|
||||
open: boolean
|
||||
selectedCount: number
|
||||
onCancel: () => void
|
||||
onOk: () => void
|
||||
tags: string[]
|
||||
onTagInputChange: (value: string) => void
|
||||
onTagInputKeyDown: (e: React.KeyboardEvent) => void
|
||||
onRemoveTag: (tag: string) => void
|
||||
tagInput: string
|
||||
tagMode: "add" | "replace"
|
||||
onTagModeChange: (mode: "add" | "replace") => void
|
||||
confirmLoading?: boolean
|
||||
}
|
||||
|
||||
const BatchTagModal: React.FC<BatchTagModalProps> = ({
|
||||
open,
|
||||
selectedCount,
|
||||
onCancel,
|
||||
onOk,
|
||||
tags,
|
||||
onTagInputChange,
|
||||
onTagInputKeyDown,
|
||||
onRemoveTag,
|
||||
tagInput,
|
||||
tagMode,
|
||||
onTagModeChange,
|
||||
confirmLoading,
|
||||
}) => (
|
||||
<AntModal
|
||||
title={`批量打标签(${selectedCount} 个素材)`}
|
||||
open={open}
|
||||
onCancel={onCancel}
|
||||
onOk={onOk}
|
||||
confirmLoading={confirmLoading}
|
||||
okText="确认打标签"
|
||||
cancelText="取消"
|
||||
>
|
||||
<div className="xx-batch-tag-modal">
|
||||
<div className="xx-batch-tag-mode">
|
||||
<span className="xx-batch-tag-mode-label">模式:</span>
|
||||
<Radio.Group value={tagMode} onChange={(e) => onTagModeChange(e.target.value)}>
|
||||
<Radio value="add">追加标签</Radio>
|
||||
<Radio value="replace">替换全部标签</Radio>
|
||||
</Radio.Group>
|
||||
</div>
|
||||
<div className="xx-batch-tag-input-row">
|
||||
<AntInput
|
||||
placeholder="输入标签后按 Enter 添加"
|
||||
value={tagInput}
|
||||
onChange={(e) => onTagInputChange(e.target.value)}
|
||||
onKeyDown={onTagInputKeyDown}
|
||||
style={{ flex: 1 }}
|
||||
/>
|
||||
</div>
|
||||
{tags.length > 0 && (
|
||||
<div className="xx-batch-tag-list">
|
||||
{tags.map((tag) => (
|
||||
<Tag key={tag} closable onClose={() => onRemoveTag(tag)} color="blue">
|
||||
{tag}
|
||||
</Tag>
|
||||
))}
|
||||
</div>
|
||||
)}
|
||||
{tagMode === "replace" && tags.length > 0 && (
|
||||
<div className="xx-batch-tag-warning">
|
||||
<ExclamationCircleOutlined /> 替换模式将清除素材原有全部标签
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
</AntModal>
|
||||
)
|
||||
|
||||
export default BatchTagModal
|
||||
@@ -1,64 +0,0 @@
|
||||
import React from "react"
|
||||
import { Modal as AntModal } from "antd"
|
||||
import { Input, Select } from "@/components/ui"
|
||||
import type { AssetKind } from "@/pages/assets/types"
|
||||
import { LIBRARY_KIND_OPTIONS } from "@/pages/assets/constants"
|
||||
|
||||
/* ============================================================
|
||||
* CreateLibraryModal — 新建视频库弹窗
|
||||
* ============================================================ */
|
||||
export interface CreateLibraryModalProps {
|
||||
open: boolean
|
||||
onCancel: () => void
|
||||
onOk: () => void
|
||||
name: string
|
||||
onNameChange: (value: string) => void
|
||||
kind: AssetKind
|
||||
onKindChange: (value: AssetKind) => void
|
||||
confirmLoading?: boolean
|
||||
}
|
||||
|
||||
const CreateLibraryModal: React.FC<CreateLibraryModalProps> = ({
|
||||
open,
|
||||
onCancel,
|
||||
onOk,
|
||||
name,
|
||||
onNameChange,
|
||||
kind,
|
||||
onKindChange,
|
||||
confirmLoading,
|
||||
}) => (
|
||||
<AntModal
|
||||
title="新建视频库"
|
||||
open={open}
|
||||
onCancel={onCancel}
|
||||
onOk={onOk}
|
||||
okText="创建"
|
||||
cancelText="取消"
|
||||
destroyOnClose
|
||||
confirmLoading={confirmLoading}
|
||||
>
|
||||
<div className="xx-asset-form-body">
|
||||
<div>
|
||||
<div className="xx-asset-form-label">名称</div>
|
||||
<Input
|
||||
placeholder="请输入视频库名称"
|
||||
value={name}
|
||||
onChange={(e) => onNameChange(e.target.value)}
|
||||
maxLength={50}
|
||||
/>
|
||||
</div>
|
||||
<div>
|
||||
<div className="xx-asset-form-label">类型</div>
|
||||
<Select
|
||||
value={kind}
|
||||
onChange={(v) => onKindChange(v)}
|
||||
style={{ width: "100%" }}
|
||||
options={LIBRARY_KIND_OPTIONS}
|
||||
/>
|
||||
</div>
|
||||
</div>
|
||||
</AntModal>
|
||||
)
|
||||
|
||||
export default CreateLibraryModal
|
||||
@@ -1,70 +0,0 @@
|
||||
import React from "react"
|
||||
import { PlusOutlined, DeleteOutlined } from "@ant-design/icons"
|
||||
import { Popconfirm } from "antd"
|
||||
import type { LibraryItem } from "@/pages/assets/types"
|
||||
import { kindLabel } from "@/pages/assets/utils/asset"
|
||||
import { kindIcon } from "@/pages/assets/utils/kindIcon"
|
||||
|
||||
/* ============================================================
|
||||
* LibrarySidebar — 左侧素材库列表
|
||||
* ============================================================ */
|
||||
export interface LibrarySidebarProps {
|
||||
libraries: LibraryItem[]
|
||||
activeLibId: string
|
||||
onSelect: (id: string) => void
|
||||
onDelete: (id: string) => void
|
||||
onCreateClick: () => void
|
||||
}
|
||||
|
||||
const LibrarySidebar: React.FC<LibrarySidebarProps> = ({
|
||||
libraries,
|
||||
activeLibId,
|
||||
onSelect,
|
||||
onDelete,
|
||||
onCreateClick,
|
||||
}) => (
|
||||
<div className="xx-asset-library-list">
|
||||
{libraries.map((lib) => (
|
||||
<div
|
||||
key={lib.id}
|
||||
className={`xx-asset-library-item${lib.id === activeLibId ? " active" : ""}`}
|
||||
onClick={() => onSelect(lib.id)}
|
||||
>
|
||||
<div className="xx-asset-library-header">
|
||||
<h4>
|
||||
{kindIcon(lib.kind)} {lib.name}
|
||||
</h4>
|
||||
<Popconfirm
|
||||
title={`确定删除视频库 "${lib.name}"?`}
|
||||
onConfirm={(e) => {
|
||||
e?.stopPropagation()
|
||||
onDelete(lib.id)
|
||||
}}
|
||||
onCancel={(e) => e?.stopPropagation()}
|
||||
okText="删除"
|
||||
cancelText="取消"
|
||||
>
|
||||
<button
|
||||
className="xx-asset-library-delete"
|
||||
onClick={(e) => e.stopPropagation()}
|
||||
title="删除视频库"
|
||||
>
|
||||
<DeleteOutlined />
|
||||
</button>
|
||||
</Popconfirm>
|
||||
</div>
|
||||
<span>
|
||||
{kindLabel(lib.kind)} · {lib.count} 个素材
|
||||
</span>
|
||||
</div>
|
||||
))}
|
||||
|
||||
{/* 新建视频库 */}
|
||||
<div className="xx-asset-library-add" onClick={onCreateClick}>
|
||||
<PlusOutlined />
|
||||
新建视频库
|
||||
</div>
|
||||
</div>
|
||||
)
|
||||
|
||||
export default LibrarySidebar
|
||||
@@ -1,34 +0,0 @@
|
||||
import React from "react"
|
||||
import { Modal as AntModal } from "antd"
|
||||
import type { AssetItem } from "@/pages/assets/types"
|
||||
|
||||
/* ============================================================
|
||||
* PlayModal — 视频/音频播放弹窗
|
||||
* ============================================================ */
|
||||
export interface PlayModalProps {
|
||||
open: boolean
|
||||
asset: AssetItem | null
|
||||
onClose: () => void
|
||||
}
|
||||
|
||||
const PlayModal: React.FC<PlayModalProps> = ({ open, asset, onClose }) => (
|
||||
<AntModal
|
||||
title={asset?.name ?? "播放"}
|
||||
open={open}
|
||||
onCancel={onClose}
|
||||
footer={null}
|
||||
width={640}
|
||||
destroyOnClose
|
||||
>
|
||||
{asset?.fileUrl ? (
|
||||
<video src={asset.fileUrl} controls autoPlay className="xx-asset-video-player" />
|
||||
) : (
|
||||
<div className="xx-asset-empty-fallback">
|
||||
<p>暂无可播放的文件地址</p>
|
||||
<p className="xx-asset-empty-fallback-id">素材 ID: {asset?.id}</p>
|
||||
</div>
|
||||
)}
|
||||
</AntModal>
|
||||
)
|
||||
|
||||
export default PlayModal
|
||||
@@ -1,70 +0,0 @@
|
||||
import React from "react"
|
||||
import { Drawer } from "antd"
|
||||
import { CheckCircleOutlined, CloseCircleOutlined } from "@ant-design/icons"
|
||||
import type { BatchOperationResult } from "@/api/assets"
|
||||
|
||||
/* ============================================================
|
||||
* ResultDrawer — 操作结果抽屉
|
||||
* ============================================================ */
|
||||
export interface ResultDrawerProps {
|
||||
open: boolean
|
||||
title: string
|
||||
result: BatchOperationResult | null
|
||||
onClose: () => void
|
||||
}
|
||||
|
||||
const ResultDrawer: React.FC<ResultDrawerProps> = ({ open, title, result, onClose }) => (
|
||||
<Drawer title={`${title} — 操作结果`} open={open} onClose={onClose} width={420}>
|
||||
{result && (
|
||||
<div className="xx-batch-result">
|
||||
<div className="xx-batch-result-summary">
|
||||
<div className="xx-batch-result-stat">
|
||||
<span className="xx-batch-result-total">总计 {result.total} 个</span>
|
||||
</div>
|
||||
<div className="xx-batch-result-stat success">
|
||||
<CheckCircleOutlined />
|
||||
<span>成功 {result.success_count} 个</span>
|
||||
</div>
|
||||
{result.failure_count > 0 && (
|
||||
<div className="xx-batch-result-stat fail">
|
||||
<CloseCircleOutlined />
|
||||
<span>失败 {result.failure_count} 个</span>
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
|
||||
{(result?.succeeded?.length ?? 0) > 0 && (
|
||||
<div className="xx-batch-result-section">
|
||||
<h4 className="xx-batch-result-section-title success">
|
||||
<CheckCircleOutlined /> 成功列表
|
||||
</h4>
|
||||
<div className="xx-batch-result-ids">
|
||||
{result?.succeeded?.map((id) => (
|
||||
<div key={id} className="xx-batch-result-id">
|
||||
{id}
|
||||
</div>
|
||||
))}
|
||||
</div>
|
||||
</div>
|
||||
)}
|
||||
|
||||
{(result?.failed?.length ?? 0) > 0 && (
|
||||
<div className="xx-batch-result-section">
|
||||
<h4 className="xx-batch-result-section-title fail">
|
||||
<CloseCircleOutlined /> 失败列表
|
||||
</h4>
|
||||
<div className="xx-batch-result-ids">
|
||||
{result?.failed?.map((id) => (
|
||||
<div key={id} className="xx-batch-result-id fail">
|
||||
{id}
|
||||
</div>
|
||||
))}
|
||||
</div>
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
)}
|
||||
</Drawer>
|
||||
)
|
||||
|
||||
export default ResultDrawer
|
||||
@@ -1,56 +0,0 @@
|
||||
import React from "react"
|
||||
import { Modal as AntModal } from "antd"
|
||||
|
||||
/* ============================================================
|
||||
* UploadProgressModal — 上传进度弹窗(圆形动画 + 百分比)
|
||||
* ============================================================ */
|
||||
export interface UploadProgressModalProps {
|
||||
open: boolean
|
||||
progress: number
|
||||
}
|
||||
|
||||
const UploadProgressModal: React.FC<UploadProgressModalProps> = ({ open, progress }) => (
|
||||
<AntModal
|
||||
open={open}
|
||||
footer={null}
|
||||
closable={false}
|
||||
centered
|
||||
width={260}
|
||||
maskClosable={false}
|
||||
className="xx-upload-progress-modal"
|
||||
>
|
||||
<div className="xx-upload-progress-body">
|
||||
<svg className="xx-upload-progress-ring" viewBox="0 0 120 120" width={120} height={120}>
|
||||
{/* 背景圆环 */}
|
||||
<circle
|
||||
cx="60"
|
||||
cy="60"
|
||||
r="52"
|
||||
fill="none"
|
||||
stroke="var(--border-primary, #e5e7eb)"
|
||||
strokeWidth="8"
|
||||
/>
|
||||
{/* 进度圆弧 */}
|
||||
<circle
|
||||
cx="60"
|
||||
cy="60"
|
||||
r="52"
|
||||
fill="none"
|
||||
stroke="var(--primary-color, #6366f1)"
|
||||
strokeWidth="8"
|
||||
strokeLinecap="round"
|
||||
strokeDasharray={`${2 * Math.PI * 52}`}
|
||||
strokeDashoffset={`${2 * Math.PI * 52 * (1 - progress / 100)}`}
|
||||
transform="rotate(-90 60 60)"
|
||||
style={{ transition: "stroke-dashoffset 0.3s ease" }}
|
||||
/>
|
||||
</svg>
|
||||
<div className="xx-upload-progress-text">
|
||||
<span className="xx-upload-progress-pct">{progress}%</span>
|
||||
<span className="xx-upload-progress-label">上传中…</span>
|
||||
</div>
|
||||
</div>
|
||||
</AntModal>
|
||||
)
|
||||
|
||||
export default UploadProgressModal
|
||||
@@ -1,66 +0,0 @@
|
||||
/**
|
||||
* 素材库常量
|
||||
*/
|
||||
import type { AssetKind } from "./types"
|
||||
|
||||
/** 文件大小限制 */
|
||||
export const MAX_FILE_SIZE = 2048 * 1024 * 1024
|
||||
export const LARGE_FILE_THRESHOLD = 100 * 1024 * 1024
|
||||
|
||||
/** 素材类型标签 */
|
||||
export const KIND_LABELS: Record<AssetKind, string> = {
|
||||
video: "视频",
|
||||
image: "图片",
|
||||
voice: "配音",
|
||||
}
|
||||
|
||||
/** 类型筛选选项 */
|
||||
export const TYPE_FILTER_OPTIONS: { value: string; label: string }[] = [
|
||||
{ value: "all", label: "全部类型" },
|
||||
{ value: "video", label: "视频" },
|
||||
{ value: "image", label: "图片" },
|
||||
]
|
||||
|
||||
/** 时间筛选选项 */
|
||||
export const TIME_FILTER_OPTIONS: { value: string; label: string }[] = [
|
||||
{ value: "all", label: "全部时间" },
|
||||
{ value: "today", label: "今天" },
|
||||
{ value: "week", label: "近一周" },
|
||||
{ value: "month", label: "近一月" },
|
||||
]
|
||||
|
||||
/** 新建库类型选项(当前仅支持视频和图片) */
|
||||
export const LIBRARY_KIND_OPTIONS: { value: AssetKind; label: string }[] = [
|
||||
{ value: "video", label: "视频" },
|
||||
{ value: "image", label: "图片" },
|
||||
]
|
||||
|
||||
/** 分类选项 */
|
||||
export const CATEGORY_OPTIONS: { value: string; label: string }[] = [
|
||||
{ value: "person", label: "人物" },
|
||||
{ value: "scenic", label: "风景" },
|
||||
{ value: "product", label: "产品" },
|
||||
{ value: "food", label: "美食" },
|
||||
{ value: "animal", label: "动物" },
|
||||
{ value: "architecture", label: "建筑" },
|
||||
{ value: "other", label: "其他" },
|
||||
]
|
||||
|
||||
/** 智能标记视图 */
|
||||
export const SMART_VIEW_OPTIONS: {
|
||||
value: "recommended" | "caution" | "high_risk"
|
||||
label: string
|
||||
color: string
|
||||
}[] = [
|
||||
{ value: "recommended", label: "推荐", color: "success" },
|
||||
{ value: "caution", label: "慎用", color: "warning" },
|
||||
{ value: "high_risk", label: "高风险", color: "error" },
|
||||
]
|
||||
|
||||
/** 状态配置 */
|
||||
export const STATUS_CONFIG: Record<string, { label: string; className: string }> = {
|
||||
ok: { label: "合格", className: "xx-status-pill xx-status-pill-ok" },
|
||||
warn: { label: "待优化", className: "xx-status-pill xx-status-pill-warn" },
|
||||
bad: { label: "不合格", className: "xx-status-pill xx-status-pill-bad" },
|
||||
info: { label: "待诊断", className: "xx-status-pill xx-status-pill-info" },
|
||||
}
|
||||
@@ -1,301 +0,0 @@
|
||||
import { useState, useCallback } from "react"
|
||||
import { useQueryClient } from "@tanstack/react-query"
|
||||
import { message } from "antd"
|
||||
import {
|
||||
deleteAsset,
|
||||
getAssetDiagnosis,
|
||||
batchDeleteAssets,
|
||||
batchTagAssets,
|
||||
batchClassifyAssets,
|
||||
batchMarkAssets,
|
||||
type BatchOperationResult,
|
||||
} from "@/api/assets"
|
||||
import type { AssetItem } from "../types"
|
||||
import type { SmartViewType } from "../components/BatchMarkModal"
|
||||
|
||||
/**
|
||||
* 素材操作 Hook
|
||||
* 封装素材的诊断、删除、批量打标签、批量改分类、批量智能标记等操作,
|
||||
* 以及相关弹窗和结果展示的状态管理
|
||||
*/
|
||||
interface UseAssetOperationsProps {
|
||||
selectedIds: Set<string>
|
||||
setSelectedIds: (ids: Set<string>) => void
|
||||
}
|
||||
|
||||
export function useAssetOperations({ selectedIds, setSelectedIds }: UseAssetOperationsProps) {
|
||||
const queryClient = useQueryClient()
|
||||
|
||||
/* ── 诊断状态 ── */
|
||||
const [diagnosingId, setDiagnosingId] = useState<string | null>(null)
|
||||
|
||||
/* ── 批量操作弹窗状态 ── */
|
||||
const [tagModalOpen, setTagModalOpen] = useState(false)
|
||||
const [classifyModalOpen, setClassifyModalOpen] = useState(false)
|
||||
const [markModalOpen, setMarkModalOpen] = useState(false)
|
||||
const [resultDrawerOpen, setResultDrawerOpen] = useState(false)
|
||||
|
||||
/* ── 批量打标签表单 ── */
|
||||
const [batchTagInput, setBatchTagInput] = useState("")
|
||||
const [batchTags, setBatchTags] = useState<string[]>([])
|
||||
const [tagMode, setTagMode] = useState<"add" | "replace">("add")
|
||||
|
||||
/* ── 批量改分类表单 ── */
|
||||
const [batchCategory, setBatchCategory] = useState("")
|
||||
|
||||
/* ── 批量智能标记表单 ── */
|
||||
const [batchSmartView, setBatchSmartView] = useState<SmartViewType>("recommended")
|
||||
|
||||
/* ── 操作结果 ── */
|
||||
const [operationResult, setOperationResult] = useState<BatchOperationResult | null>(null)
|
||||
const [operationTitle, setOperationTitle] = useState("")
|
||||
|
||||
/* ── 批量操作 loading ── */
|
||||
const [batchLoading, setBatchLoading] = useState(false)
|
||||
|
||||
/* ── 刷新数据辅助函数 ── */
|
||||
const invalidateAssets = useCallback(() => {
|
||||
queryClient.invalidateQueries({ queryKey: ["assets"] })
|
||||
queryClient.invalidateQueries({ queryKey: ["asset-libraries"] })
|
||||
}, [queryClient])
|
||||
|
||||
/* ── 诊断 ── */
|
||||
const handleDiagnose = useCallback(
|
||||
async (asset: AssetItem) => {
|
||||
setDiagnosingId(asset.id)
|
||||
try {
|
||||
const result = await getAssetDiagnosis(asset.id)
|
||||
const score = result.readiness_score ?? "-"
|
||||
message.success(`"${asset.name}" 诊断完成,就绪分:${score}`)
|
||||
queryClient.invalidateQueries({ queryKey: ["assets"] })
|
||||
} catch {
|
||||
message.error(`"${asset.name}" 诊断失败`)
|
||||
} finally {
|
||||
setDiagnosingId(null)
|
||||
}
|
||||
},
|
||||
[queryClient],
|
||||
)
|
||||
|
||||
/* ── 单个素材删除 ── */
|
||||
const handleSingleDelete = useCallback(
|
||||
async (assetId: string) => {
|
||||
try {
|
||||
await deleteAsset(assetId)
|
||||
invalidateAssets()
|
||||
// 从选中集合中移除
|
||||
setSelectedIds(
|
||||
(() => {
|
||||
const next = new Set(selectedIds)
|
||||
next.delete(assetId)
|
||||
return next
|
||||
})(),
|
||||
)
|
||||
message.success("素材已删除")
|
||||
} catch {
|
||||
message.error("删除失败,请重试")
|
||||
}
|
||||
},
|
||||
[invalidateAssets, selectedIds, setSelectedIds],
|
||||
)
|
||||
|
||||
/* ── 显示操作结果 ── */
|
||||
const showOperationResult = useCallback(
|
||||
(result: BatchOperationResult, title: string, clearSelection = true) => {
|
||||
setOperationResult(result)
|
||||
setOperationTitle(title)
|
||||
setResultDrawerOpen(true)
|
||||
if (clearSelection) setSelectedIds(new Set())
|
||||
},
|
||||
[setSelectedIds],
|
||||
)
|
||||
|
||||
/* ── 批量删除 ── */
|
||||
const handleBatchDelete = useCallback(async () => {
|
||||
const ids = Array.from(selectedIds)
|
||||
setBatchLoading(true)
|
||||
try {
|
||||
const result = await batchDeleteAssets(ids)
|
||||
invalidateAssets()
|
||||
showOperationResult(result, "批量删除")
|
||||
if (result.failure_count === 0) {
|
||||
message.success(`成功删除 ${result.success_count} 个素材`)
|
||||
} else {
|
||||
message.warning(
|
||||
`删除完成:成功 ${result.success_count} 个,失败 ${result.failure_count} 个`,
|
||||
)
|
||||
}
|
||||
} catch {
|
||||
message.error("批量删除失败,请重试")
|
||||
} finally {
|
||||
setBatchLoading(false)
|
||||
}
|
||||
}, [selectedIds, invalidateAssets, showOperationResult])
|
||||
|
||||
/* ── 批量打标签 ── */
|
||||
const handleBatchTag = useCallback(async () => {
|
||||
if (batchTags.length === 0) {
|
||||
message.warning("请至少输入一个标签")
|
||||
return
|
||||
}
|
||||
const ids = Array.from(selectedIds)
|
||||
setBatchLoading(true)
|
||||
try {
|
||||
const result = await batchTagAssets({
|
||||
asset_ids: ids,
|
||||
tags: batchTags,
|
||||
mode: tagMode,
|
||||
})
|
||||
queryClient.invalidateQueries({ queryKey: ["assets"] })
|
||||
showOperationResult(result, "批量打标签")
|
||||
setTagModalOpen(false)
|
||||
setBatchTags([])
|
||||
setBatchTagInput("")
|
||||
setTagMode("add")
|
||||
if (result.failure_count === 0) {
|
||||
message.success(`成功为 ${result.success_count} 个素材打标签`)
|
||||
} else {
|
||||
message.warning(
|
||||
`打标签完成:成功 ${result.success_count} 个,失败 ${result.failure_count} 个`,
|
||||
)
|
||||
}
|
||||
} catch {
|
||||
message.error("批量打标签失败,请重试")
|
||||
} finally {
|
||||
setBatchLoading(false)
|
||||
}
|
||||
}, [batchTags, selectedIds, tagMode, queryClient, showOperationResult])
|
||||
|
||||
/* ── 标签输入处理 ── */
|
||||
const handleTagInputKeyDown = useCallback(
|
||||
(e: React.KeyboardEvent) => {
|
||||
if (e.key === "Enter" && batchTagInput.trim()) {
|
||||
e.preventDefault()
|
||||
const tag = batchTagInput.trim()
|
||||
if (!batchTags.includes(tag)) {
|
||||
setBatchTags([...batchTags, tag])
|
||||
}
|
||||
setBatchTagInput("")
|
||||
}
|
||||
},
|
||||
[batchTagInput, batchTags],
|
||||
)
|
||||
|
||||
const removeBatchTag = useCallback(
|
||||
(tag: string) => {
|
||||
setBatchTags(batchTags.filter((t) => t !== tag))
|
||||
},
|
||||
[batchTags],
|
||||
)
|
||||
|
||||
/* ── 批量改分类 ── */
|
||||
const handleBatchClassify = useCallback(async () => {
|
||||
if (!batchCategory) {
|
||||
message.warning("请选择分类")
|
||||
return
|
||||
}
|
||||
const ids = Array.from(selectedIds)
|
||||
setBatchLoading(true)
|
||||
try {
|
||||
const result = await batchClassifyAssets({
|
||||
asset_ids: ids,
|
||||
category: batchCategory,
|
||||
})
|
||||
queryClient.invalidateQueries({ queryKey: ["assets"] })
|
||||
showOperationResult(result, "批量改分类")
|
||||
setClassifyModalOpen(false)
|
||||
setBatchCategory("")
|
||||
if (result.failure_count === 0) {
|
||||
message.success(`成功将 ${result.success_count} 个素材改为「${batchCategory}」`)
|
||||
} else {
|
||||
message.warning(
|
||||
`改分类完成:成功 ${result.success_count} 个,失败 ${result.failure_count} 个`,
|
||||
)
|
||||
}
|
||||
} catch {
|
||||
message.error("批量改分类失败,请重试")
|
||||
} finally {
|
||||
setBatchLoading(false)
|
||||
}
|
||||
}, [batchCategory, selectedIds, queryClient, showOperationResult])
|
||||
|
||||
/* ── 批量智能标记 ── */
|
||||
const handleBatchMark = useCallback(async () => {
|
||||
const ids = Array.from(selectedIds)
|
||||
setBatchLoading(true)
|
||||
try {
|
||||
const result = await batchMarkAssets({
|
||||
asset_ids: ids,
|
||||
smart_view: batchSmartView,
|
||||
})
|
||||
queryClient.invalidateQueries({ queryKey: ["assets"] })
|
||||
showOperationResult(result, "批量智能标记")
|
||||
setMarkModalOpen(false)
|
||||
const labelMap: Record<SmartViewType, string> = {
|
||||
recommended: "推荐",
|
||||
caution: "慎用",
|
||||
high_risk: "高风险",
|
||||
}
|
||||
if (result.failure_count === 0) {
|
||||
message.success(
|
||||
`成功将 ${result.success_count} 个素材标记为「${labelMap[batchSmartView]}」`,
|
||||
)
|
||||
} else {
|
||||
message.warning(
|
||||
`智能标记完成:成功 ${result.success_count} 个,失败 ${result.failure_count} 个`,
|
||||
)
|
||||
}
|
||||
} catch {
|
||||
message.error("批量智能标记失败,请重试")
|
||||
} finally {
|
||||
setBatchLoading(false)
|
||||
}
|
||||
}, [batchSmartView, selectedIds, queryClient, showOperationResult])
|
||||
|
||||
/* ── 关闭结果 Drawer ── */
|
||||
const handleResultDrawerClose = useCallback(() => {
|
||||
setResultDrawerOpen(false)
|
||||
setOperationResult(null)
|
||||
}, [])
|
||||
|
||||
return {
|
||||
// 诊断
|
||||
diagnosingId,
|
||||
handleDiagnose,
|
||||
// 单个操作
|
||||
handleSingleDelete,
|
||||
// 批量操作 loading
|
||||
batchLoading,
|
||||
// 批量打标签
|
||||
tagModalOpen,
|
||||
setTagModalOpen,
|
||||
batchTagInput,
|
||||
setBatchTagInput,
|
||||
batchTags,
|
||||
setBatchTags,
|
||||
tagMode,
|
||||
setTagMode,
|
||||
handleBatchTag,
|
||||
handleTagInputKeyDown,
|
||||
removeBatchTag,
|
||||
// 批量改分类
|
||||
classifyModalOpen,
|
||||
setClassifyModalOpen,
|
||||
batchCategory,
|
||||
setBatchCategory,
|
||||
handleBatchClassify,
|
||||
// 批量智能标记
|
||||
markModalOpen,
|
||||
setMarkModalOpen,
|
||||
batchSmartView,
|
||||
setBatchSmartView,
|
||||
handleBatchMark,
|
||||
// 批量删除
|
||||
handleBatchDelete,
|
||||
// 操作结果
|
||||
resultDrawerOpen,
|
||||
operationResult,
|
||||
operationTitle,
|
||||
handleResultDrawerClose,
|
||||
}
|
||||
}
|
||||
@@ -1,40 +0,0 @@
|
||||
import { useState, useCallback } from "react"
|
||||
import type { AssetItem } from "../types"
|
||||
|
||||
/**
|
||||
* 素材选中态管理 Hook
|
||||
* 封装单选、全选、取消全选等选中逻辑
|
||||
*/
|
||||
interface UseAssetSelectionProps {
|
||||
filteredAssets: AssetItem[]
|
||||
}
|
||||
|
||||
export function useAssetSelection({ filteredAssets }: UseAssetSelectionProps) {
|
||||
const [selectedIds, setSelectedIds] = useState<Set<string>>(new Set())
|
||||
|
||||
const toggleSelect = useCallback((id: string) => {
|
||||
setSelectedIds((prev) => {
|
||||
const next = new Set(prev)
|
||||
if (next.has(id)) next.delete(id)
|
||||
else next.add(id)
|
||||
return next
|
||||
})
|
||||
}, [])
|
||||
|
||||
const selectAll = useCallback(() => {
|
||||
setSelectedIds(new Set(filteredAssets.map((a) => a.id)))
|
||||
}, [filteredAssets])
|
||||
|
||||
const deselectAll = useCallback(() => {
|
||||
setSelectedIds(new Set())
|
||||
}, [])
|
||||
|
||||
return {
|
||||
selectedIds,
|
||||
setSelectedIds,
|
||||
toggleSelect,
|
||||
selectAll,
|
||||
deselectAll,
|
||||
selectedCount: selectedIds.size,
|
||||
}
|
||||
}
|
||||
@@ -1,65 +0,0 @@
|
||||
import { useState, useCallback } from "react"
|
||||
import { useQueryClient } from "@tanstack/react-query"
|
||||
import { message } from "antd"
|
||||
import { uploadAssetDirect } from "@/api/assets"
|
||||
import { MAX_FILE_SIZE, LARGE_FILE_THRESHOLD } from "../constants"
|
||||
|
||||
/**
|
||||
* 素材上传 Hook
|
||||
* 封装上传状态、进度管理和上传逻辑
|
||||
*/
|
||||
interface UseAssetUploadProps {
|
||||
effectiveLibId: string
|
||||
}
|
||||
|
||||
export function useAssetUpload({ effectiveLibId }: UseAssetUploadProps) {
|
||||
const queryClient = useQueryClient()
|
||||
|
||||
const [uploading, setUploading] = useState(false)
|
||||
const [uploadProgress, setUploadProgress] = useState(0)
|
||||
|
||||
const handleUpload = useCallback(
|
||||
async (file: File) => {
|
||||
if (file.size > MAX_FILE_SIZE) {
|
||||
message.error(`文件 "${file.name}" 超过 2GB 限制`)
|
||||
return
|
||||
}
|
||||
if (!effectiveLibId) {
|
||||
message.warning("请先选择或创建一个视频库")
|
||||
return
|
||||
}
|
||||
|
||||
setUploading(true)
|
||||
setUploadProgress(0)
|
||||
try {
|
||||
if (file.size > LARGE_FILE_THRESHOLD) {
|
||||
message.info(`大文件 "${file.name}" 将使用直传上传`)
|
||||
}
|
||||
await uploadAssetDirect({
|
||||
file,
|
||||
library_id: effectiveLibId,
|
||||
onProgress: (pct) => setUploadProgress(pct),
|
||||
})
|
||||
message.success(`"${file.name}" 上传成功`)
|
||||
queryClient.invalidateQueries({ queryKey: ["assets"] })
|
||||
queryClient.invalidateQueries({ queryKey: ["asset-libraries"] })
|
||||
} catch (err: unknown) {
|
||||
const detail = err instanceof Error ? err.message : ""
|
||||
console.error("[handleUpload] 上传失败:", err)
|
||||
message.error(`"${file.name}" 上传失败${detail ? `:${detail}` : ""}`)
|
||||
// 错误时延迟关闭弹窗,让用户能看到错误提示
|
||||
await new Promise((r) => setTimeout(r, 1500))
|
||||
} finally {
|
||||
setUploading(false)
|
||||
setUploadProgress(0)
|
||||
}
|
||||
},
|
||||
[effectiveLibId, queryClient],
|
||||
)
|
||||
|
||||
return {
|
||||
uploading,
|
||||
uploadProgress,
|
||||
handleUpload,
|
||||
}
|
||||
}
|
||||
@@ -1,118 +0,0 @@
|
||||
import { useState, useMemo } from "react"
|
||||
import { useQuery } from "@tanstack/react-query"
|
||||
import {
|
||||
getAssetLibraries,
|
||||
getAssets,
|
||||
type AssetLibraryItem,
|
||||
type AssetItem as ApiAssetItem,
|
||||
} from "@/api/assets"
|
||||
import { mapLibrary, mapAsset, type AssetItem, type LibraryItem } from "../types"
|
||||
|
||||
/**
|
||||
* 素材库数据 Hook
|
||||
* 封装视频库列表、素材列表的数据查询,以及筛选、搜索状态管理
|
||||
*/
|
||||
export function useAssetsData() {
|
||||
/* ── 视频库列表查询 ── */
|
||||
const { data: apiLibraries = [], isLoading: libLoading } = useQuery<AssetLibraryItem[], Error>({
|
||||
queryKey: ["asset-libraries"],
|
||||
queryFn: getAssetLibraries,
|
||||
staleTime: 60_000,
|
||||
})
|
||||
|
||||
const libraries: LibraryItem[] = useMemo(
|
||||
() =>
|
||||
(Array.isArray(apiLibraries) ? apiLibraries : [])
|
||||
.map(mapLibrary)
|
||||
.filter((lib) => lib.kind === "video"),
|
||||
[apiLibraries],
|
||||
)
|
||||
|
||||
/* ── 当前选中的视频库 ── */
|
||||
const [activeLibId, setActiveLibId] = useState<string>("")
|
||||
|
||||
// 当库列表加载完成后,自动选中第一个
|
||||
const effectiveLibId = activeLibId || libraries[0]?.id || ""
|
||||
|
||||
/* ── 当前库的素材列表查询 ── */
|
||||
const {
|
||||
data: apiAssets = { items: [], total: 0 },
|
||||
isLoading: assetsLoading,
|
||||
isError: assetsError,
|
||||
error: assetsErrorObj,
|
||||
refetch: refetchAssets,
|
||||
} = useQuery<{ items: ApiAssetItem[]; total: number }, Error>({
|
||||
queryKey: ["assets", effectiveLibId],
|
||||
queryFn: () =>
|
||||
getAssets(effectiveLibId, {
|
||||
// 拉取所有非删除状态的素材,让用户上传后立刻能看到"处理中"的素材
|
||||
status: "ready,uploading,ingesting,processing,pending,error,failed",
|
||||
}),
|
||||
enabled: !!effectiveLibId,
|
||||
staleTime: 30_000,
|
||||
})
|
||||
|
||||
const assets: AssetItem[] = useMemo(
|
||||
() => (Array.isArray(apiAssets?.items) ? apiAssets.items : []).map(mapAsset),
|
||||
[apiAssets],
|
||||
)
|
||||
|
||||
/* ── 筛选状态 ── */
|
||||
const [searchText, setSearchText] = useState("")
|
||||
const [filterType, setFilterType] = useState<string>("all")
|
||||
const [filterTime, setFilterTime] = useState<string>("all")
|
||||
|
||||
/* ── 筛选后的素材列表 ── */
|
||||
const filteredAssets = useMemo(() => {
|
||||
let list = assets
|
||||
|
||||
/* 按素材类型过滤 */
|
||||
if (filterType !== "all") {
|
||||
list = list.filter((a) => a.kind === filterType)
|
||||
}
|
||||
|
||||
/* 按时间筛选 */
|
||||
if (filterTime !== "all") {
|
||||
const now = new Date()
|
||||
list = list.filter((a) => {
|
||||
const d = new Date(a.createdAt)
|
||||
const diffDays = (now.getTime() - d.getTime()) / (1000 * 60 * 60 * 24)
|
||||
if (filterTime === "today") return diffDays < 1
|
||||
if (filterTime === "week") return diffDays < 7
|
||||
if (filterTime === "month") return diffDays < 30
|
||||
return true
|
||||
})
|
||||
}
|
||||
|
||||
/* 搜索 */
|
||||
if (searchText.trim()) {
|
||||
const q = searchText.trim().toLowerCase()
|
||||
list = list.filter((a) => a.name.toLowerCase().includes(q))
|
||||
}
|
||||
|
||||
return list
|
||||
}, [assets, filterType, filterTime, searchText])
|
||||
|
||||
return {
|
||||
// 视频库
|
||||
libraries,
|
||||
libLoading,
|
||||
activeLibId,
|
||||
setActiveLibId,
|
||||
effectiveLibId,
|
||||
// 素材列表
|
||||
assets,
|
||||
assetsLoading,
|
||||
assetsError,
|
||||
assetsErrorObj,
|
||||
refetchAssets,
|
||||
// 筛选
|
||||
searchText,
|
||||
setSearchText,
|
||||
filterType,
|
||||
setFilterType,
|
||||
filterTime,
|
||||
setFilterTime,
|
||||
filteredAssets,
|
||||
}
|
||||
}
|
||||
@@ -1,106 +0,0 @@
|
||||
import { useState, useCallback } from "react"
|
||||
import { useMutation, useQueryClient } from "@tanstack/react-query"
|
||||
import { message } from "antd"
|
||||
import { createAssetLibrary, deleteAssetLibrary } from "@/api/assets"
|
||||
import type { AssetKind, LibraryItem } from "../types"
|
||||
|
||||
/**
|
||||
* 视频库管理 Hook
|
||||
* 封装视频库的创建、删除操作,以及新建弹窗的表单状态
|
||||
*/
|
||||
interface UseLibraryManagementProps {
|
||||
libraries: LibraryItem[]
|
||||
activeLibId: string
|
||||
setActiveLibId: (id: string) => void
|
||||
effectiveLibId: string
|
||||
}
|
||||
|
||||
export function useLibraryManagement({
|
||||
libraries,
|
||||
setActiveLibId,
|
||||
effectiveLibId,
|
||||
}: UseLibraryManagementProps) {
|
||||
const queryClient = useQueryClient()
|
||||
|
||||
/* ── 状态 ── */
|
||||
const [createModalOpen, setCreateModalOpen] = useState(false)
|
||||
const [newLibName, setNewLibName] = useState("")
|
||||
const [newLibKind, setNewLibKind] = useState<AssetKind>("video")
|
||||
|
||||
/* ── Mutations ── */
|
||||
const createLibMutation = useMutation({
|
||||
mutationFn: createAssetLibrary,
|
||||
onSuccess: () => {
|
||||
queryClient.invalidateQueries({ queryKey: ["asset-libraries"] })
|
||||
message.success("视频库创建成功")
|
||||
},
|
||||
onError: () => {
|
||||
message.error("创建视频库失败")
|
||||
},
|
||||
})
|
||||
|
||||
const deleteLibMutation = useMutation({
|
||||
mutationFn: deleteAssetLibrary,
|
||||
onSuccess: () => {
|
||||
queryClient.invalidateQueries({ queryKey: ["asset-libraries"] })
|
||||
message.success("视频库已删除")
|
||||
},
|
||||
onError: () => {
|
||||
message.error("删除视频库失败")
|
||||
},
|
||||
})
|
||||
|
||||
/* ── 新建视频库 ── */
|
||||
const handleCreateLibrary = useCallback(async () => {
|
||||
if (!newLibName.trim()) {
|
||||
message.warning("请输入视频库名称")
|
||||
return
|
||||
}
|
||||
try {
|
||||
const newLib = await createLibMutation.mutateAsync({
|
||||
name: newLibName.trim(),
|
||||
kind: newLibKind,
|
||||
})
|
||||
setActiveLibId(newLib.id)
|
||||
setCreateModalOpen(false)
|
||||
setNewLibName("")
|
||||
setNewLibKind("video")
|
||||
} catch {
|
||||
// error handled in mutation
|
||||
}
|
||||
}, [newLibName, newLibKind, createLibMutation, setActiveLibId])
|
||||
|
||||
/* ── 删除视频库 ── */
|
||||
const handleDeleteLibrary = useCallback(
|
||||
async (id: string) => {
|
||||
try {
|
||||
await deleteLibMutation.mutateAsync(id)
|
||||
if (effectiveLibId === id) {
|
||||
const remaining = libraries.filter((l) => l.id !== id)
|
||||
if (remaining.length > 0) setActiveLibId(remaining[0].id)
|
||||
else setActiveLibId("")
|
||||
}
|
||||
} catch {
|
||||
// error handled in mutation
|
||||
}
|
||||
},
|
||||
[deleteLibMutation, effectiveLibId, libraries, setActiveLibId],
|
||||
)
|
||||
|
||||
return {
|
||||
// 弹窗状态
|
||||
createModalOpen,
|
||||
setCreateModalOpen,
|
||||
// 表单状态
|
||||
newLibName,
|
||||
setNewLibName,
|
||||
newLibKind,
|
||||
setNewLibKind,
|
||||
// Mutations
|
||||
isCreating: createLibMutation.isPending,
|
||||
isDeleting: deleteLibMutation.isPending,
|
||||
// Handlers
|
||||
handleCreateLibrary,
|
||||
handleDeleteLibrary,
|
||||
}
|
||||
}
|
||||
@@ -1,115 +0,0 @@
|
||||
/**
|
||||
* 素材库类型定义
|
||||
*/
|
||||
import type { AssetLibraryItem, AssetItem as ApiAssetItem } from "@/api/assets"
|
||||
import { formatDuration } from "./utils/format"
|
||||
|
||||
export type AssetKind = "video" | "image" | "voice"
|
||||
export type StatusType = "ok" | "warn" | "bad" | "info"
|
||||
|
||||
export interface LibraryItem {
|
||||
id: string
|
||||
name: string
|
||||
kind: AssetKind
|
||||
count: number
|
||||
}
|
||||
|
||||
export interface AssetItem {
|
||||
id: string
|
||||
name: string
|
||||
kind: AssetKind
|
||||
thumbUrl?: string
|
||||
fileUrl?: string
|
||||
status: StatusType
|
||||
statusLabel: string
|
||||
/** 是否处于处理中状态(上传中/入库中/诊断中) */
|
||||
loading?: boolean
|
||||
duration?: string
|
||||
size: number
|
||||
createdAt: string
|
||||
}
|
||||
|
||||
/** 根据 mime_type 推断前端 AssetKind */
|
||||
export const inferKind = (mimeType: string): AssetKind => {
|
||||
if (mimeType.startsWith("video/")) return "video"
|
||||
if (mimeType.startsWith("audio/")) return "voice"
|
||||
return "image"
|
||||
}
|
||||
|
||||
/** 根据 quality_score / classification_status / asset status 推断前端状态 */
|
||||
export const inferStatus = (
|
||||
score?: number,
|
||||
classificationStatus?: string,
|
||||
assetStatus?: string,
|
||||
): { status: StatusType; label: string; loading?: boolean } => {
|
||||
// 已删除素材(正常情况列表已过滤,这里是防御性处理)
|
||||
if (assetStatus === "deleted") {
|
||||
return { status: "bad", label: "已删除" }
|
||||
}
|
||||
// 处理中状态:上传中 / 入库中 / 处理中
|
||||
if (
|
||||
assetStatus === "uploading" ||
|
||||
assetStatus === "ingesting" ||
|
||||
assetStatus === "processing" ||
|
||||
assetStatus === "pending"
|
||||
) {
|
||||
return { status: "info", label: "处理中", loading: true }
|
||||
}
|
||||
// 失败状态
|
||||
if (assetStatus === "error" || assetStatus === "failed") {
|
||||
return { status: "bad", label: "处理失败" }
|
||||
}
|
||||
// 素材已就绪(status=ready)时,不应因 classification 未执行而显示"处理中"
|
||||
if (assetStatus === "ready") {
|
||||
if (score == null) return { status: "info", label: "待诊断" }
|
||||
if (score >= 70) return { status: "ok", label: "合格" }
|
||||
if (score >= 40) return { status: "warn", label: "待优化" }
|
||||
return { status: "bad", label: "不合格" }
|
||||
}
|
||||
// 素材未就绪:classification 正在处理中
|
||||
if (classificationStatus === "processing" || classificationStatus === "pending") {
|
||||
return { status: "info", label: "处理中", loading: true }
|
||||
}
|
||||
if (score == null) return { status: "info", label: "待诊断" }
|
||||
if (score >= 70) return { status: "ok", label: "合格" }
|
||||
if (score >= 40) return { status: "warn", label: "待优化" }
|
||||
return { status: "bad", label: "不合格" }
|
||||
}
|
||||
|
||||
/** 将后端 AssetLibraryItem 映射为前端 LibraryItem */
|
||||
export const mapLibrary = (item: AssetLibraryItem): LibraryItem => ({
|
||||
id: item.id,
|
||||
name: item.name,
|
||||
kind: item.kind || inferKind("video"),
|
||||
count: item.asset_count ?? 0,
|
||||
})
|
||||
|
||||
/** 将后端 ApiAssetItem 映射为前端 AssetItem */
|
||||
export const mapAsset = (item: ApiAssetItem): AssetItem => {
|
||||
const { status, label, loading } = inferStatus(
|
||||
item.quality_score ?? undefined,
|
||||
item.classification_status ?? undefined,
|
||||
item.status ?? undefined,
|
||||
)
|
||||
const metadata = item.metadata || {}
|
||||
const kind = inferKind(item.mime_type || "")
|
||||
return {
|
||||
id: item.id,
|
||||
name: item.name,
|
||||
kind,
|
||||
// 视频类型不能用 file_url 做缩略图(是视频文件,<img> 无法渲染)
|
||||
// 处理中的素材没有缩略图,显示占位符
|
||||
thumbUrl: loading
|
||||
? undefined
|
||||
: (item.thumbnail_url as string | undefined) ||
|
||||
(metadata.thumbnail_url as string | undefined) ||
|
||||
(kind !== "video" ? (item.file_url as string | undefined) : undefined),
|
||||
fileUrl: (item.file_url as string | undefined) || (metadata.file_url as string | undefined),
|
||||
status,
|
||||
statusLabel: label,
|
||||
loading,
|
||||
duration: metadata.duration != null ? formatDuration(metadata.duration as number) : undefined,
|
||||
size: item.file_size ? +(item.file_size / (1024 * 1024)).toFixed(1) : 0,
|
||||
createdAt: item.created_at ? new Date(item.created_at).toISOString().slice(0, 10) : "—",
|
||||
}
|
||||
}
|
||||
@@ -1,19 +0,0 @@
|
||||
/**
|
||||
* 素材相关工具函数
|
||||
*/
|
||||
import type { AssetKind } from "../types"
|
||||
import { KIND_LABELS } from "../constants"
|
||||
|
||||
export const kindLabel = (kind: AssetKind): string => KIND_LABELS[kind] ?? kind
|
||||
|
||||
/** 根据素材类型返回渐变背景 */
|
||||
export const thumbGradient = (kind: AssetKind): string => {
|
||||
switch (kind) {
|
||||
case "video":
|
||||
return "linear-gradient(135deg, #312e81 0%, #4f46e5 50%, #6366f1 100%)"
|
||||
case "image":
|
||||
return "linear-gradient(135deg, #78350f 0%, #d97706 50%, #f59e0b 100%)"
|
||||
case "voice":
|
||||
return "linear-gradient(135deg, #064e3b 0%, #059669 50%, #10b981 100%)"
|
||||
}
|
||||
}
|
||||
@@ -1,22 +0,0 @@
|
||||
/**
|
||||
* 格式化工具函数
|
||||
*/
|
||||
|
||||
export const formatDuration = (seconds: number): string => {
|
||||
const m = Math.floor(seconds / 60)
|
||||
const s = Math.floor(seconds % 60)
|
||||
return `${m.toString().padStart(2, "0")}:${s.toString().padStart(2, "0")}`
|
||||
}
|
||||
|
||||
export const formatFileSize = (bytes: number): string => {
|
||||
if (bytes < 1024) return `${bytes} B`
|
||||
if (bytes < 1024 * 1024) return `${(bytes / 1024).toFixed(1)} KB`
|
||||
return `${(bytes / (1024 * 1024)).toFixed(1)} MB`
|
||||
}
|
||||
|
||||
export const formatDate = (iso: string): string =>
|
||||
new Date(iso).toLocaleDateString("zh-CN", {
|
||||
year: "numeric",
|
||||
month: "2-digit",
|
||||
day: "2-digit",
|
||||
})
|
||||
@@ -1,15 +0,0 @@
|
||||
import React from "react"
|
||||
import { VideoCameraOutlined, PictureOutlined, AudioOutlined } from "@ant-design/icons"
|
||||
import type { AssetKind } from "../types"
|
||||
|
||||
/** 根据素材类型返回对应图标 */
|
||||
export const kindIcon = (kind: AssetKind): React.ReactNode => {
|
||||
switch (kind) {
|
||||
case "video":
|
||||
return <VideoCameraOutlined />
|
||||
case "image":
|
||||
return <PictureOutlined />
|
||||
case "voice":
|
||||
return <AudioOutlined />
|
||||
}
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,186 +0,0 @@
|
||||
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
|
||||
@@ -1,163 +0,0 @@
|
||||
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
|
||||
@@ -1,245 +0,0 @@
|
||||
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
|
||||
@@ -1,186 +0,0 @@
|
||||
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
|
||||
@@ -1,142 +0,0 @@
|
||||
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,
|
||||
}
|
||||
}
|
||||
@@ -1,132 +0,0 @@
|
||||
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,
|
||||
}
|
||||
}
|
||||
@@ -1,135 +0,0 @@
|
||||
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,
|
||||
}
|
||||
}
|
||||
@@ -1,351 +0,0 @@
|
||||
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,
|
||||
}
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,17 +0,0 @@
|
||||
import React from "react"
|
||||
|
||||
/** 克隆音色卡片骨架屏 */
|
||||
const CloneCardSkeleton: React.FC = () => (
|
||||
<div className="xx-voice-card xx-skeleton-clone">
|
||||
<div className="xx-skeleton-clone-avatar" />
|
||||
<div className="xx-skeleton-clone-info">
|
||||
<div className="xx-skeleton-clone-line xx-skeleton-clone-line--name" />
|
||||
<div className="xx-skeleton-clone-line xx-skeleton-clone-line--status" />
|
||||
<div className="xx-skeleton-clone-line xx-skeleton-clone-line--desc" />
|
||||
<div className="xx-skeleton-clone-line xx-skeleton-clone-line--meta" />
|
||||
<div className="xx-skeleton-clone-line xx-skeleton-clone-line--footer" />
|
||||
</div>
|
||||
</div>
|
||||
)
|
||||
|
||||
export default CloneCardSkeleton
|
||||
@@ -1,100 +0,0 @@
|
||||
import React from "react"
|
||||
import {
|
||||
SoundOutlined,
|
||||
CloseCircleOutlined,
|
||||
ReloadOutlined,
|
||||
DeleteOutlined,
|
||||
} from "@ant-design/icons"
|
||||
import { Button } from "@/components/ui"
|
||||
import { type ClonedVoiceDisplay } from "@/pages/voices/types"
|
||||
import { CLONE_STATUS_CONFIG } from "@/pages/voices/constants"
|
||||
|
||||
export interface CloneDetailModalProps {
|
||||
voice: ClonedVoiceDisplay
|
||||
onClose: () => void
|
||||
onUse: () => void
|
||||
onDelete: () => void
|
||||
onRetry: () => void
|
||||
}
|
||||
|
||||
/** 克隆音色详情弹窗 */
|
||||
const CloneDetailModal: React.FC<CloneDetailModalProps> = ({
|
||||
voice,
|
||||
onClose,
|
||||
onUse,
|
||||
onDelete,
|
||||
onRetry,
|
||||
}) => {
|
||||
const statusCfg = CLONE_STATUS_CONFIG[voice.status]
|
||||
const genderText =
|
||||
voice.gender === "male" ? "男声" : voice.gender === "female" ? "女声" : voice.gender || "未知"
|
||||
const langText = voice.language || "未知"
|
||||
|
||||
return (
|
||||
<div className="xx-clone-detail-overlay" onClick={onClose}>
|
||||
<div className="xx-clone-detail" onClick={(e) => e.stopPropagation()}>
|
||||
<button type="button" className="xx-clone-detail-close" onClick={onClose}>
|
||||
<CloseCircleOutlined />
|
||||
</button>
|
||||
|
||||
<div className="xx-clone-detail-header">
|
||||
<div className="xx-clone-avatar">
|
||||
<SoundOutlined />
|
||||
</div>
|
||||
<div>
|
||||
<h3 className="xx-clone-detail-name">{voice.name}</h3>
|
||||
<span className={`xx-clone-status ${statusCfg.className}`}>
|
||||
<span className="xx-clone-status-dot" />
|
||||
{statusCfg.label}
|
||||
</span>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
{voice.description && <p className="xx-clone-detail-desc">{voice.description}</p>}
|
||||
|
||||
<div className="xx-clone-detail-info">
|
||||
<div className="xx-clone-detail-row">
|
||||
<span className="xx-clone-detail-label">性别</span>
|
||||
<span>{genderText}</span>
|
||||
</div>
|
||||
<div className="xx-clone-detail-row">
|
||||
<span className="xx-clone-detail-label">语言</span>
|
||||
<span>{langText}</span>
|
||||
</div>
|
||||
<div className="xx-clone-detail-row">
|
||||
<span className="xx-clone-detail-label">来源</span>
|
||||
<span>{voice.sourceName}</span>
|
||||
</div>
|
||||
<div className="xx-clone-detail-row">
|
||||
<span className="xx-clone-detail-label">创建时间</span>
|
||||
<span>{voice.createdAt}</span>
|
||||
</div>
|
||||
{voice.errorMessage && (
|
||||
<div className="xx-clone-detail-row" style={{ color: "var(--error-color, #ef4444)" }}>
|
||||
<span className="xx-clone-detail-label">错误</span>
|
||||
<span>{voice.errorMessage}</span>
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
|
||||
<div className="xx-clone-detail-actions">
|
||||
{voice.status === "failed" && (
|
||||
<Button buttonType="ghost" buttonSize="sm" icon={<ReloadOutlined />} onClick={onRetry}>
|
||||
重试
|
||||
</Button>
|
||||
)}
|
||||
<Button buttonType="ghost" buttonSize="sm" icon={<DeleteOutlined />} onClick={onDelete}>
|
||||
删除
|
||||
</Button>
|
||||
{voice.status === "ready" && (
|
||||
<Button buttonType="primary" buttonSize="sm" onClick={onUse}>
|
||||
使用此音色
|
||||
</Button>
|
||||
)}
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
)
|
||||
}
|
||||
|
||||
export default CloneDetailModal
|
||||
@@ -1,180 +0,0 @@
|
||||
import React from "react"
|
||||
import {
|
||||
SoundOutlined,
|
||||
PlayCircleOutlined,
|
||||
PauseCircleOutlined,
|
||||
DeleteOutlined,
|
||||
ReloadOutlined,
|
||||
CloseCircleOutlined,
|
||||
UserOutlined,
|
||||
} from "@ant-design/icons"
|
||||
import { Tooltip } from "antd"
|
||||
import { type ClonedVoiceDisplay } from "@/pages/voices/types"
|
||||
import { CLONE_STATUS_CONFIG } from "@/pages/voices/constants"
|
||||
|
||||
export interface CloneVoiceCardProps {
|
||||
voice: ClonedVoiceDisplay
|
||||
isPlaying: boolean
|
||||
currentTime: number
|
||||
onPlay: () => void
|
||||
onPause: () => void
|
||||
onUse: () => void
|
||||
onDelete: () => void
|
||||
onRetry: () => void
|
||||
onShowDetail: () => void
|
||||
}
|
||||
|
||||
/** 克隆音色卡片 */
|
||||
const CloneVoiceCard: React.FC<CloneVoiceCardProps> = ({
|
||||
voice,
|
||||
isPlaying,
|
||||
currentTime,
|
||||
onPlay,
|
||||
onPause,
|
||||
onUse,
|
||||
onDelete,
|
||||
onRetry,
|
||||
onShowDetail,
|
||||
}) => {
|
||||
const statusCfg = CLONE_STATUS_CONFIG[voice.status]
|
||||
const isFailed = voice.status === "failed"
|
||||
const isProcessing = voice.status === "processing"
|
||||
const genderText =
|
||||
voice.gender === "male" ? "男声" : voice.gender === "female" ? "女声" : voice.gender
|
||||
|
||||
return (
|
||||
<div
|
||||
className={`xx-clone-card${isPlaying ? " playing" : ""}${isFailed ? " failed" : ""}`}
|
||||
onClick={isFailed ? undefined : onShowDetail}
|
||||
>
|
||||
{/* 右上角操作按钮 */}
|
||||
<div className="xx-clone-card-actions">
|
||||
<Tooltip title="删除">
|
||||
<button
|
||||
type="button"
|
||||
className="xx-clone-action-btn xx-clone-action-btn--danger"
|
||||
onClick={(e) => {
|
||||
e.stopPropagation()
|
||||
onDelete()
|
||||
}}
|
||||
>
|
||||
<DeleteOutlined />
|
||||
</button>
|
||||
</Tooltip>
|
||||
{isFailed && (
|
||||
<Tooltip title="重试">
|
||||
<button
|
||||
type="button"
|
||||
className="xx-clone-action-btn"
|
||||
onClick={(e) => {
|
||||
e.stopPropagation()
|
||||
onRetry()
|
||||
}}
|
||||
>
|
||||
<ReloadOutlined />
|
||||
</button>
|
||||
</Tooltip>
|
||||
)}
|
||||
</div>
|
||||
|
||||
{/* 头部:头像 + 名称 + 状态 */}
|
||||
<div className="xx-clone-card-header">
|
||||
<div className={`xx-clone-avatar${isProcessing ? " xx-clone-avatar--processing" : ""}`}>
|
||||
<SoundOutlined />
|
||||
</div>
|
||||
<div className="xx-clone-header-info">
|
||||
<h4 className="xx-clone-name" title={voice.name}>
|
||||
{voice.name}
|
||||
</h4>
|
||||
<span className={`xx-clone-status ${statusCfg.className}`}>
|
||||
<span className="xx-clone-status-dot" />
|
||||
{statusCfg.label}
|
||||
</span>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
{/* 描述 */}
|
||||
{voice.description && <p className="xx-clone-desc">{voice.description}</p>}
|
||||
|
||||
{/* 元信息 */}
|
||||
<div className="xx-clone-meta">
|
||||
{(voice.gender || voice.language) && (
|
||||
<span className="xx-clone-meta-item">
|
||||
<UserOutlined />
|
||||
{genderText}
|
||||
{voice.language ? ` · ${voice.language}` : ""}
|
||||
</span>
|
||||
)}
|
||||
<span className="xx-clone-meta-item">{voice.createdAt}</span>
|
||||
</div>
|
||||
|
||||
{/* 错误信息 */}
|
||||
{isFailed && voice.errorMessage && (
|
||||
<div className="xx-clone-error">
|
||||
<CloseCircleOutlined />
|
||||
<span>{voice.errorMessage}</span>
|
||||
</div>
|
||||
)}
|
||||
|
||||
{/* 底部操作区 */}
|
||||
<div className="xx-clone-footer">
|
||||
{voice.status === "ready" && (
|
||||
<>
|
||||
<button
|
||||
type="button"
|
||||
className="xx-clone-play-btn"
|
||||
onClick={(e) => {
|
||||
e.stopPropagation()
|
||||
isPlaying ? onPause() : onPlay()
|
||||
}}
|
||||
title={isPlaying ? "暂停" : "试听"}
|
||||
>
|
||||
{isPlaying ? <PauseCircleOutlined /> : <PlayCircleOutlined />}
|
||||
</button>
|
||||
<div className="xx-clone-progress">
|
||||
<div
|
||||
className="xx-clone-progress-bar"
|
||||
style={{
|
||||
width: isPlaying
|
||||
? `${Math.min((currentTime / Math.max(voice.duration, 1)) * 100, 100)}%`
|
||||
: "0%",
|
||||
}}
|
||||
/>
|
||||
</div>
|
||||
<button
|
||||
type="button"
|
||||
className="xx-clone-use-btn"
|
||||
onClick={(e) => {
|
||||
e.stopPropagation()
|
||||
onUse()
|
||||
}}
|
||||
>
|
||||
使用
|
||||
</button>
|
||||
</>
|
||||
)}
|
||||
{isProcessing && (
|
||||
<div className="xx-clone-processing-hint">
|
||||
<ReloadOutlined spin />
|
||||
克隆处理中,请稍候...
|
||||
</div>
|
||||
)}
|
||||
{isFailed && (
|
||||
<button
|
||||
type="button"
|
||||
className="xx-clone-retry-btn"
|
||||
onClick={(e) => {
|
||||
e.stopPropagation()
|
||||
onRetry()
|
||||
}}
|
||||
>
|
||||
<ReloadOutlined />
|
||||
重试克隆
|
||||
</button>
|
||||
)}
|
||||
</div>
|
||||
</div>
|
||||
)
|
||||
}
|
||||
|
||||
export default CloneVoiceCard
|
||||
@@ -1,37 +0,0 @@
|
||||
import React from "react"
|
||||
import { AudioOutlined } from "@ant-design/icons"
|
||||
import { type AssetItem } from "@/api/assets"
|
||||
import { formatFileSize } from "@/pages/voices/utils/format"
|
||||
|
||||
export interface MaterialVoiceCardProps {
|
||||
asset: AssetItem
|
||||
onClick?: () => void
|
||||
}
|
||||
|
||||
/** 配音素材卡片 */
|
||||
const MaterialVoiceCard: React.FC<MaterialVoiceCardProps> = ({ asset, onClick }) => {
|
||||
const duration = (asset.metadata?.duration as number) || 0
|
||||
const minutes = Math.floor(duration / 60)
|
||||
const seconds = Math.floor(duration % 60)
|
||||
|
||||
return (
|
||||
<div className="vmat-card" onClick={onClick}>
|
||||
<div className="vmat-thumb">
|
||||
<AudioOutlined className="vmat-thumb-icon" />
|
||||
<span className="vmat-duration">
|
||||
{minutes}:{seconds.toString().padStart(2, "0")}
|
||||
</span>
|
||||
</div>
|
||||
<div className="vmat-info">
|
||||
<div className="vmat-name" title={asset.name}>
|
||||
{asset.name}
|
||||
</div>
|
||||
<div className="vmat-meta">
|
||||
<span>{asset.file_size ? formatFileSize(asset.file_size) : "--"}</span>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
)
|
||||
}
|
||||
|
||||
export default MaterialVoiceCard
|
||||
@@ -1,258 +0,0 @@
|
||||
import React from "react"
|
||||
import { RobotOutlined } from "@ant-design/icons"
|
||||
import { Modal, message } from "antd"
|
||||
import { type PresetVoiceDisplay } from "@/pages/voices/types"
|
||||
import { genderLabel } from "@/pages/voices/utils/format"
|
||||
|
||||
export type TtsStatus = "idle" | "synthesizing" | "done" | "error"
|
||||
|
||||
export interface TtsModalProps {
|
||||
open: boolean
|
||||
ttsText: string
|
||||
ttsVoiceId: string
|
||||
ttsSpeed: number
|
||||
ttsStatus: TtsStatus
|
||||
ttsAudioUrl: string | null
|
||||
ttsError: string | null
|
||||
presetVoices: PresetVoiceDisplay[]
|
||||
onClose: () => void
|
||||
onTextChange: (text: string) => void
|
||||
onVoiceChange: (voiceId: string) => void
|
||||
onSpeedChange: (speed: number) => void
|
||||
onSynthesize: () => void
|
||||
onSave: () => void
|
||||
}
|
||||
|
||||
/** AI 配音弹窗 */
|
||||
const TtsModal: React.FC<TtsModalProps> = ({
|
||||
open,
|
||||
ttsText,
|
||||
ttsVoiceId,
|
||||
ttsSpeed,
|
||||
ttsStatus,
|
||||
ttsAudioUrl,
|
||||
ttsError,
|
||||
presetVoices,
|
||||
onClose,
|
||||
onTextChange,
|
||||
onVoiceChange,
|
||||
onSpeedChange,
|
||||
onSynthesize,
|
||||
onSave,
|
||||
}) => {
|
||||
return (
|
||||
<Modal title="AI 配音" open={open} onCancel={onClose} footer={null} width={560}>
|
||||
<div
|
||||
style={{
|
||||
display: "flex",
|
||||
flexDirection: "column",
|
||||
gap: 16,
|
||||
padding: "8px 0",
|
||||
}}
|
||||
>
|
||||
{/* 文本输入 */}
|
||||
<div>
|
||||
<div
|
||||
style={{
|
||||
fontSize: 13,
|
||||
color: "var(--text-secondary)",
|
||||
marginBottom: 6,
|
||||
}}
|
||||
>
|
||||
输入文本
|
||||
</div>
|
||||
<textarea
|
||||
value={ttsText}
|
||||
onChange={(e) => onTextChange(e.target.value)}
|
||||
placeholder="输入要配音的文本内容..."
|
||||
maxLength={2000}
|
||||
rows={4}
|
||||
style={{
|
||||
width: "100%",
|
||||
padding: "10px 12px",
|
||||
border: "1px solid var(--border-color)",
|
||||
borderRadius: 8,
|
||||
fontSize: 13,
|
||||
background: "var(--bg-primary)",
|
||||
color: "var(--text-primary)",
|
||||
outline: "none",
|
||||
resize: "vertical",
|
||||
fontFamily: "inherit",
|
||||
lineHeight: 1.6,
|
||||
}}
|
||||
/>
|
||||
<div
|
||||
style={{
|
||||
fontSize: 11,
|
||||
color: "var(--text-tertiary)",
|
||||
textAlign: "right",
|
||||
marginTop: 4,
|
||||
}}
|
||||
>
|
||||
{ttsText.length}/2000
|
||||
</div>
|
||||
</div>
|
||||
|
||||
{/* 音色选择 */}
|
||||
<div>
|
||||
<div
|
||||
style={{
|
||||
fontSize: 13,
|
||||
color: "var(--text-secondary)",
|
||||
marginBottom: 6,
|
||||
}}
|
||||
>
|
||||
选择音色
|
||||
</div>
|
||||
<select
|
||||
value={ttsVoiceId}
|
||||
onChange={(e) => onVoiceChange(e.target.value)}
|
||||
style={{
|
||||
width: "100%",
|
||||
padding: "8px 12px",
|
||||
border: "1px solid var(--border-color)",
|
||||
borderRadius: 8,
|
||||
fontSize: 13,
|
||||
background: "var(--bg-primary)",
|
||||
color: "var(--text-primary)",
|
||||
outline: "none",
|
||||
}}
|
||||
>
|
||||
<option value="">默认音色</option>
|
||||
{presetVoices.map((v) => (
|
||||
<option key={v.id} value={v.id}>
|
||||
{v.name} — {genderLabel(v.gender)}
|
||||
</option>
|
||||
))}
|
||||
</select>
|
||||
</div>
|
||||
|
||||
{/* 语速 */}
|
||||
<div>
|
||||
<div
|
||||
style={{
|
||||
fontSize: 13,
|
||||
color: "var(--text-secondary)",
|
||||
marginBottom: 6,
|
||||
}}
|
||||
>
|
||||
语速:{ttsSpeed.toFixed(1)}x
|
||||
</div>
|
||||
<input
|
||||
type="range"
|
||||
min={0.5}
|
||||
max={2.0}
|
||||
step={0.1}
|
||||
value={ttsSpeed}
|
||||
onChange={(e) => onSpeedChange(parseFloat(e.target.value))}
|
||||
style={{ width: "100%", accentColor: "var(--primary-color)" }}
|
||||
/>
|
||||
<div
|
||||
style={{
|
||||
display: "flex",
|
||||
justifyContent: "space-between",
|
||||
fontSize: 11,
|
||||
color: "var(--text-tertiary)",
|
||||
}}
|
||||
>
|
||||
<span>0.5x</span>
|
||||
<span>1.0x</span>
|
||||
<span>2.0x</span>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
{/* 合成按钮 */}
|
||||
<button
|
||||
type="button"
|
||||
onClick={() => {
|
||||
if (!ttsText.trim()) {
|
||||
message.warning("请输入要合成的文本")
|
||||
return
|
||||
}
|
||||
onSynthesize()
|
||||
}}
|
||||
disabled={ttsStatus === "synthesizing" || !ttsText.trim()}
|
||||
style={{
|
||||
width: "100%",
|
||||
padding: "10px 0",
|
||||
borderRadius: 8,
|
||||
border: "none",
|
||||
background:
|
||||
ttsStatus === "synthesizing" || !ttsText.trim()
|
||||
? "var(--text-tertiary)"
|
||||
: "var(--primary-color)",
|
||||
color: "#fff",
|
||||
fontSize: 14,
|
||||
fontWeight: 600,
|
||||
cursor: ttsStatus === "synthesizing" || !ttsText.trim() ? "not-allowed" : "pointer",
|
||||
display: "flex",
|
||||
alignItems: "center",
|
||||
justifyContent: "center",
|
||||
gap: 8,
|
||||
}}
|
||||
>
|
||||
<RobotOutlined />
|
||||
{ttsStatus === "synthesizing" ? "合成中..." : "开始合成"}
|
||||
</button>
|
||||
|
||||
{/* 错误提示 */}
|
||||
{ttsError && (
|
||||
<div
|
||||
style={{
|
||||
padding: "10px 12px",
|
||||
background: "var(--error-soft, #fff2f0)",
|
||||
borderRadius: 8,
|
||||
color: "var(--error-color, #ff4d4f)",
|
||||
fontSize: 13,
|
||||
}}
|
||||
>
|
||||
{ttsError}
|
||||
</div>
|
||||
)}
|
||||
|
||||
{/* 合成结果 */}
|
||||
{ttsStatus === "done" && ttsAudioUrl && (
|
||||
<div
|
||||
style={{
|
||||
padding: 12,
|
||||
background: "var(--bg-secondary)",
|
||||
borderRadius: 8,
|
||||
display: "flex",
|
||||
flexDirection: "column",
|
||||
gap: 10,
|
||||
}}
|
||||
>
|
||||
<div
|
||||
style={{
|
||||
fontSize: 13,
|
||||
fontWeight: 500,
|
||||
color: "var(--success-color, #52c41a)",
|
||||
}}
|
||||
>
|
||||
✅ 合成完成
|
||||
</div>
|
||||
<audio controls src={ttsAudioUrl} style={{ width: "100%" }} />
|
||||
<button
|
||||
type="button"
|
||||
onClick={onSave}
|
||||
style={{
|
||||
padding: "8px 0",
|
||||
borderRadius: 8,
|
||||
border: "1px solid var(--primary-color)",
|
||||
background: "var(--primary-soft)",
|
||||
color: "var(--primary-color)",
|
||||
fontSize: 13,
|
||||
fontWeight: 600,
|
||||
cursor: "pointer",
|
||||
}}
|
||||
>
|
||||
保存到配音库
|
||||
</button>
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
</Modal>
|
||||
)
|
||||
}
|
||||
|
||||
export default TtsModal
|
||||
@@ -1,329 +0,0 @@
|
||||
import React from "react"
|
||||
import { UploadOutlined, SoundOutlined } from "@ant-design/icons"
|
||||
import { Modal, Upload, message } from "antd"
|
||||
import { type VoiceGender } from "@/pages/voices/types"
|
||||
import { formatFileSize } from "@/pages/voices/utils/format"
|
||||
|
||||
export interface UploadVoiceModalProps {
|
||||
open: boolean
|
||||
uploadFile: File | null
|
||||
uploadName: string
|
||||
uploadGender: VoiceGender
|
||||
uploadDesc: string
|
||||
uploadProgress: number | null
|
||||
onClose: () => void
|
||||
onFileSelect: (file: File) => void
|
||||
onFileRemove: () => void
|
||||
onNameChange: (name: string) => void
|
||||
onGenderChange: (gender: VoiceGender) => void
|
||||
onDescChange: (desc: string) => void
|
||||
onUpload: () => void
|
||||
}
|
||||
|
||||
/** 上传音频弹窗 */
|
||||
const UploadVoiceModal: React.FC<UploadVoiceModalProps> = ({
|
||||
open,
|
||||
uploadFile,
|
||||
uploadName,
|
||||
uploadGender,
|
||||
uploadDesc,
|
||||
uploadProgress,
|
||||
onClose,
|
||||
onFileSelect,
|
||||
onFileRemove,
|
||||
onNameChange,
|
||||
onGenderChange,
|
||||
onDescChange,
|
||||
onUpload,
|
||||
}) => {
|
||||
return (
|
||||
<Modal
|
||||
title="上传音频"
|
||||
open={open}
|
||||
onCancel={() => {
|
||||
if (uploadProgress !== null) return // 上传中不可关闭
|
||||
onClose()
|
||||
}}
|
||||
footer={null}
|
||||
width={520}
|
||||
maskClosable={uploadProgress === null}
|
||||
>
|
||||
<div
|
||||
style={{
|
||||
display: "flex",
|
||||
flexDirection: "column",
|
||||
gap: 16,
|
||||
padding: "8px 0",
|
||||
}}
|
||||
>
|
||||
{/* 拖拽上传区 */}
|
||||
<Upload.Dragger
|
||||
accept="audio/*"
|
||||
maxCount={1}
|
||||
beforeUpload={(file) => {
|
||||
onFileSelect(file)
|
||||
return false
|
||||
}}
|
||||
onRemove={() => {
|
||||
onFileRemove()
|
||||
}}
|
||||
showUploadList={false}
|
||||
disabled={uploadProgress !== null}
|
||||
>
|
||||
<p
|
||||
style={{
|
||||
fontSize: 32,
|
||||
color: "var(--primary-color)",
|
||||
marginBottom: 8,
|
||||
}}
|
||||
>
|
||||
<UploadOutlined />
|
||||
</p>
|
||||
<p style={{ fontSize: 14, fontWeight: 500, margin: "0 0 4px" }}>
|
||||
点击或拖拽音频文件到此处
|
||||
</p>
|
||||
<p
|
||||
style={{
|
||||
fontSize: 12,
|
||||
color: "var(--text-secondary)",
|
||||
margin: 0,
|
||||
}}
|
||||
>
|
||||
支持 MP3、WAV、AAC、FLAC 等格式,最大 200MB
|
||||
</p>
|
||||
</Upload.Dragger>
|
||||
|
||||
{/* 已选文件信息 */}
|
||||
{uploadFile && (
|
||||
<div
|
||||
style={{
|
||||
padding: "10px 12px",
|
||||
background: "var(--bg-secondary)",
|
||||
borderRadius: 8,
|
||||
display: "flex",
|
||||
alignItems: "center",
|
||||
gap: 10,
|
||||
}}
|
||||
>
|
||||
<SoundOutlined style={{ fontSize: 18, color: "var(--primary-color)" }} />
|
||||
<div style={{ flex: 1, minWidth: 0 }}>
|
||||
<div
|
||||
style={{
|
||||
fontSize: 13,
|
||||
fontWeight: 500,
|
||||
overflow: "hidden",
|
||||
textOverflow: "ellipsis",
|
||||
whiteSpace: "nowrap",
|
||||
}}
|
||||
>
|
||||
{uploadFile.name}
|
||||
</div>
|
||||
<div style={{ fontSize: 11, color: "var(--text-secondary)" }}>
|
||||
{formatFileSize(uploadFile.size)}
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
)}
|
||||
|
||||
{/* 上传进度 */}
|
||||
{uploadProgress !== null && (
|
||||
<div style={{ textAlign: "center", padding: "8px 0" }}>
|
||||
<div
|
||||
style={{
|
||||
fontSize: 22,
|
||||
fontWeight: 700,
|
||||
color: "var(--primary-color)",
|
||||
}}
|
||||
>
|
||||
{uploadProgress}%
|
||||
</div>
|
||||
<div style={{ fontSize: 12, color: "var(--text-secondary)" }}>
|
||||
{uploadProgress < 100 ? "上传中..." : "处理中..."}
|
||||
</div>
|
||||
<div
|
||||
style={{
|
||||
height: 4,
|
||||
background: "var(--bg-tertiary)",
|
||||
borderRadius: 2,
|
||||
marginTop: 8,
|
||||
overflow: "hidden",
|
||||
}}
|
||||
>
|
||||
<div
|
||||
style={{
|
||||
height: "100%",
|
||||
width: `${uploadProgress}%`,
|
||||
background: "var(--primary-color)",
|
||||
borderRadius: 2,
|
||||
transition: "width 0.3s ease",
|
||||
}}
|
||||
/>
|
||||
</div>
|
||||
</div>
|
||||
)}
|
||||
|
||||
{/* 名称 */}
|
||||
<div>
|
||||
<div
|
||||
style={{
|
||||
fontSize: 13,
|
||||
color: "var(--text-secondary)",
|
||||
marginBottom: 6,
|
||||
}}
|
||||
>
|
||||
素材名称
|
||||
</div>
|
||||
<input
|
||||
value={uploadName}
|
||||
onChange={(e) => onNameChange(e.target.value)}
|
||||
placeholder="输入素材名称"
|
||||
maxLength={100}
|
||||
disabled={uploadProgress !== null}
|
||||
style={{
|
||||
width: "100%",
|
||||
padding: "8px 12px",
|
||||
border: "1px solid var(--border-color)",
|
||||
borderRadius: 8,
|
||||
fontSize: 13,
|
||||
background: "var(--bg-primary)",
|
||||
color: "var(--text-primary)",
|
||||
outline: "none",
|
||||
}}
|
||||
/>
|
||||
</div>
|
||||
|
||||
{/* 性别选择 */}
|
||||
<div>
|
||||
<div
|
||||
style={{
|
||||
fontSize: 13,
|
||||
color: "var(--text-secondary)",
|
||||
marginBottom: 6,
|
||||
}}
|
||||
>
|
||||
音色性别
|
||||
</div>
|
||||
<div style={{ display: "flex", gap: 8 }}>
|
||||
{(["female", "male", "child"] as VoiceGender[]).map((g) => (
|
||||
<button
|
||||
key={g}
|
||||
type="button"
|
||||
onClick={() => onGenderChange(g)}
|
||||
disabled={uploadProgress !== null}
|
||||
style={{
|
||||
flex: 1,
|
||||
padding: "6px 0",
|
||||
borderRadius: 8,
|
||||
border: `1px solid ${uploadGender === g ? "var(--primary-color)" : "var(--border-color)"}`,
|
||||
background: uploadGender === g ? "var(--primary-soft)" : "transparent",
|
||||
color: uploadGender === g ? "var(--primary-color)" : "var(--text-secondary)",
|
||||
fontSize: 13,
|
||||
fontWeight: uploadGender === g ? 600 : 400,
|
||||
cursor: uploadProgress !== null ? "not-allowed" : "pointer",
|
||||
transition: "all 0.2s",
|
||||
}}
|
||||
>
|
||||
{g === "female" ? "女声" : g === "male" ? "男声" : "童声"}
|
||||
</button>
|
||||
))}
|
||||
</div>
|
||||
</div>
|
||||
|
||||
{/* 描述 */}
|
||||
<div>
|
||||
<div
|
||||
style={{
|
||||
fontSize: 13,
|
||||
color: "var(--text-secondary)",
|
||||
marginBottom: 6,
|
||||
}}
|
||||
>
|
||||
音色描述(可选)
|
||||
</div>
|
||||
<textarea
|
||||
value={uploadDesc}
|
||||
onChange={(e) => onDescChange(e.target.value)}
|
||||
placeholder="描述这个音色的特点..."
|
||||
maxLength={500}
|
||||
rows={2}
|
||||
disabled={uploadProgress !== null}
|
||||
style={{
|
||||
width: "100%",
|
||||
padding: "8px 12px",
|
||||
border: "1px solid var(--border-color)",
|
||||
borderRadius: 8,
|
||||
fontSize: 13,
|
||||
background: "var(--bg-primary)",
|
||||
color: "var(--text-primary)",
|
||||
outline: "none",
|
||||
resize: "vertical",
|
||||
fontFamily: "inherit",
|
||||
}}
|
||||
/>
|
||||
</div>
|
||||
|
||||
{/* 操作按钮 */}
|
||||
<div
|
||||
style={{
|
||||
display: "flex",
|
||||
justifyContent: "flex-end",
|
||||
gap: 10,
|
||||
paddingTop: 4,
|
||||
}}
|
||||
>
|
||||
<button
|
||||
type="button"
|
||||
onClick={onClose}
|
||||
disabled={uploadProgress !== null}
|
||||
style={{
|
||||
padding: "8px 20px",
|
||||
borderRadius: 8,
|
||||
border: "1px solid var(--border-color)",
|
||||
background: "transparent",
|
||||
fontSize: 13,
|
||||
cursor: uploadProgress !== null ? "not-allowed" : "pointer",
|
||||
color: "var(--text-secondary)",
|
||||
}}
|
||||
>
|
||||
取消
|
||||
</button>
|
||||
<button
|
||||
type="button"
|
||||
onClick={() => {
|
||||
if (!uploadFile) {
|
||||
message.warning("请先选择音频文件")
|
||||
return
|
||||
}
|
||||
if (!uploadName.trim()) {
|
||||
message.warning("请输入素材名称")
|
||||
return
|
||||
}
|
||||
onUpload()
|
||||
}}
|
||||
disabled={!uploadFile || !uploadName.trim() || uploadProgress !== null}
|
||||
style={{
|
||||
padding: "8px 20px",
|
||||
borderRadius: 8,
|
||||
border: "none",
|
||||
background:
|
||||
!uploadFile || !uploadName.trim() || uploadProgress !== null
|
||||
? "var(--text-tertiary)"
|
||||
: "var(--primary-color)",
|
||||
color: "#fff",
|
||||
fontSize: 13,
|
||||
fontWeight: 600,
|
||||
cursor:
|
||||
!uploadFile || !uploadName.trim() || uploadProgress !== null
|
||||
? "not-allowed"
|
||||
: "pointer",
|
||||
}}
|
||||
>
|
||||
{uploadProgress !== null ? "上传中..." : "开始上传"}
|
||||
</button>
|
||||
</div>
|
||||
</div>
|
||||
</Modal>
|
||||
)
|
||||
}
|
||||
|
||||
export default UploadVoiceModal
|
||||
@@ -1,133 +0,0 @@
|
||||
import React, { useRef } from "react"
|
||||
import {
|
||||
SoundOutlined,
|
||||
PlayCircleOutlined,
|
||||
PauseCircleOutlined,
|
||||
HeartOutlined,
|
||||
} from "@ant-design/icons"
|
||||
import { type VoiceGender } from "@/pages/voices/types"
|
||||
import { genderClass, formatTime } from "@/pages/voices/utils/format"
|
||||
|
||||
export interface VoiceCardProps {
|
||||
id: string
|
||||
name: string
|
||||
subtitle: string
|
||||
tags: string[]
|
||||
duration: number
|
||||
gender: VoiceGender
|
||||
isPlaying: boolean
|
||||
isSelected: boolean
|
||||
currentTime: number
|
||||
starred?: boolean
|
||||
status?: "ready" | "processing" | "failed"
|
||||
onPlay: () => void
|
||||
onPause: () => void
|
||||
onSeek: (time: number) => void
|
||||
onSelect?: () => void
|
||||
onToggleStar?: () => void
|
||||
}
|
||||
|
||||
/** 音色卡片组件(预置音色 + 克隆音色统一) */
|
||||
const VoiceCard: React.FC<VoiceCardProps> = ({
|
||||
id: _id,
|
||||
name,
|
||||
subtitle,
|
||||
tags,
|
||||
duration,
|
||||
gender,
|
||||
isPlaying,
|
||||
isSelected,
|
||||
currentTime,
|
||||
starred,
|
||||
status = "ready",
|
||||
onPlay,
|
||||
onPause,
|
||||
onSeek,
|
||||
onSelect,
|
||||
onToggleStar,
|
||||
}) => {
|
||||
const progressRef = useRef<HTMLDivElement>(null)
|
||||
|
||||
const handleProgressClick = (e: React.MouseEvent<HTMLDivElement>) => {
|
||||
if (!progressRef.current || status !== "ready") return
|
||||
const rect = progressRef.current.getBoundingClientRect()
|
||||
const percent = (e.clientX - rect.left) / rect.width
|
||||
onSeek(percent * duration)
|
||||
}
|
||||
|
||||
const progress = duration > 0 ? (currentTime / duration) * 100 : 0
|
||||
|
||||
return (
|
||||
<div
|
||||
className={`xx-voice-card ${genderClass(gender)}${isSelected ? " selected" : ""}${isPlaying ? " playing" : ""}`}
|
||||
onClick={onSelect}
|
||||
>
|
||||
<div className="xx-voice-avatar">
|
||||
<SoundOutlined />
|
||||
</div>
|
||||
|
||||
<div className="xx-voice-info">
|
||||
<div className="xx-voice-name-row">
|
||||
<h4 className="xx-voice-name" title={name}>
|
||||
{name}
|
||||
</h4>
|
||||
{starred !== undefined && (
|
||||
<button
|
||||
className={`xx-voice-star${starred ? " active" : ""}`}
|
||||
onClick={(e) => {
|
||||
e.stopPropagation()
|
||||
onToggleStar?.()
|
||||
}}
|
||||
title={starred ? "取消收藏" : "收藏"}
|
||||
>
|
||||
<HeartOutlined />
|
||||
</button>
|
||||
)}
|
||||
</div>
|
||||
<div className="xx-voice-subtitle">{subtitle}</div>
|
||||
<div className="xx-voice-tags">
|
||||
{tags.slice(0, 3).map((tag) => (
|
||||
<span key={tag} className="xx-voice-tag">
|
||||
{tag}
|
||||
</span>
|
||||
))}
|
||||
</div>
|
||||
</div>
|
||||
|
||||
{status === "processing" && (
|
||||
<div className="xx-voice-status xx-voice-status--processing">
|
||||
<span className="xx-voice-status-dot" />
|
||||
处理中...
|
||||
</div>
|
||||
)}
|
||||
{status === "failed" && (
|
||||
<div className="xx-voice-status xx-voice-status--failed">克隆失败</div>
|
||||
)}
|
||||
|
||||
{status === "ready" && <div className="xx-voice-wave" />}
|
||||
|
||||
{status === "ready" && (
|
||||
<div className="xx-voice-controls">
|
||||
<button
|
||||
className="xx-voice-play-btn"
|
||||
onClick={(e) => {
|
||||
e.stopPropagation()
|
||||
isPlaying ? onPause() : onPlay()
|
||||
}}
|
||||
title={isPlaying ? "暂停" : "试听"}
|
||||
>
|
||||
{isPlaying ? <PauseCircleOutlined /> : <PlayCircleOutlined />}
|
||||
</button>
|
||||
<div ref={progressRef} className="xx-voice-progress" onClick={handleProgressClick}>
|
||||
<div className="xx-voice-progress-bar" style={{ width: `${progress}%` }} />
|
||||
</div>
|
||||
<span className="xx-voice-time">
|
||||
{isPlaying ? formatTime(currentTime) : formatTime(duration)}
|
||||
</span>
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
)
|
||||
}
|
||||
|
||||
export default VoiceCard
|
||||
@@ -1,61 +0,0 @@
|
||||
import React from "react"
|
||||
import { SearchOutlined } from "@ant-design/icons"
|
||||
import { Input, Select } from "@/components/ui"
|
||||
|
||||
export interface VoiceFilterBarProps {
|
||||
searchText: string
|
||||
filterGender: string
|
||||
filterLang: string
|
||||
onSearchChange: (value: string) => void
|
||||
onGenderChange: (value: string) => void
|
||||
onLangChange: (value: string) => void
|
||||
}
|
||||
|
||||
/** 音色筛选栏(搜索 + 性别 + 语言) */
|
||||
const VoiceFilterBar: React.FC<VoiceFilterBarProps> = ({
|
||||
searchText,
|
||||
filterGender,
|
||||
filterLang,
|
||||
onSearchChange,
|
||||
onGenderChange,
|
||||
onLangChange,
|
||||
}) => {
|
||||
return (
|
||||
<div className="xx-voices-filters">
|
||||
<Input
|
||||
placeholder="搜索音色名称、标签..."
|
||||
prefix={<SearchOutlined />}
|
||||
value={searchText}
|
||||
onChange={(e) => onSearchChange(e.target.value)}
|
||||
allowClear
|
||||
style={{ width: 260 }}
|
||||
/>
|
||||
<Select
|
||||
value={filterGender}
|
||||
onChange={onGenderChange}
|
||||
style={{ width: 130 }}
|
||||
options={[
|
||||
{ value: "all", label: "全部音色" },
|
||||
{ value: "male", label: "男声" },
|
||||
{ value: "female", label: "女声" },
|
||||
{ value: "child", label: "童声" },
|
||||
{ value: "elderly", label: "老年" },
|
||||
]}
|
||||
/>
|
||||
<Select
|
||||
value={filterLang}
|
||||
onChange={onLangChange}
|
||||
style={{ width: 130 }}
|
||||
options={[
|
||||
{ value: "all", label: "全部语言" },
|
||||
{ value: "zh", label: "中文" },
|
||||
{ value: "en", label: "英文" },
|
||||
{ value: "ja", label: "日文" },
|
||||
{ value: "ko", label: "韩文" },
|
||||
]}
|
||||
/>
|
||||
</div>
|
||||
)
|
||||
}
|
||||
|
||||
export default VoiceFilterBar
|
||||
@@ -1,30 +0,0 @@
|
||||
/**
|
||||
* 配音库常量
|
||||
*/
|
||||
import type { VoiceGender, VoiceLanguage, ClonedVoiceDisplay } from "./types"
|
||||
|
||||
/** 性别选项 */
|
||||
export const GENDER_OPTIONS: { value: VoiceGender; label: string }[] = [
|
||||
{ value: "male", label: "男声" },
|
||||
{ value: "female", label: "女声" },
|
||||
{ value: "child", label: "童声" },
|
||||
{ value: "elderly", label: "老年" },
|
||||
]
|
||||
|
||||
/** 语言选项 */
|
||||
export const LANGUAGE_OPTIONS: { value: VoiceLanguage; label: string }[] = [
|
||||
{ value: "zh", label: "中文" },
|
||||
{ value: "en", label: "英文" },
|
||||
{ value: "ja", label: "日文" },
|
||||
{ value: "ko", label: "韩文" },
|
||||
]
|
||||
|
||||
/** 克隆音色状态配置 */
|
||||
export const CLONE_STATUS_CONFIG: Record<
|
||||
ClonedVoiceDisplay["status"],
|
||||
{ label: string; className: string }
|
||||
> = {
|
||||
ready: { label: "可用", className: "xx-clone-status--ready" },
|
||||
processing: { label: "处理中", className: "xx-clone-status--processing" },
|
||||
failed: { label: "失败", className: "xx-clone-status--failed" },
|
||||
}
|
||||
@@ -1,102 +0,0 @@
|
||||
import { useState, useRef, useCallback, useEffect } from "react"
|
||||
|
||||
/**
|
||||
* 音频播放控制 Hook
|
||||
* 封装当前播放状态、播放/暂停/跳转控制,使用 setInterval 模拟进度更新
|
||||
* (适用于预置音色/克隆音色卡片的播放按钮交互)
|
||||
*/
|
||||
export function useAudioPlayer() {
|
||||
const [playingId, setPlayingId] = useState<string | null>(null)
|
||||
const [currentTime, setCurrentTime] = useState(0)
|
||||
const intervalRef = useRef<number | null>(null)
|
||||
|
||||
/** 开始播放指定音色(从 startTime 开始,默认从 0 开始) */
|
||||
const handlePlay = useCallback(
|
||||
(voiceId: string, duration: number, startTime: number = 0) => {
|
||||
if (playingId === voiceId) return
|
||||
if (intervalRef.current) {
|
||||
clearInterval(intervalRef.current)
|
||||
}
|
||||
setPlayingId(voiceId)
|
||||
setCurrentTime(startTime)
|
||||
intervalRef.current = window.setInterval(() => {
|
||||
setCurrentTime((prev) => {
|
||||
if (prev >= duration) {
|
||||
if (intervalRef.current) {
|
||||
clearInterval(intervalRef.current)
|
||||
intervalRef.current = null
|
||||
}
|
||||
setPlayingId(null)
|
||||
return 0
|
||||
}
|
||||
return prev + 0.1
|
||||
})
|
||||
}, 100)
|
||||
},
|
||||
[playingId],
|
||||
)
|
||||
|
||||
/** 暂停播放 */
|
||||
const handlePause = useCallback(() => {
|
||||
if (intervalRef.current) {
|
||||
clearInterval(intervalRef.current)
|
||||
intervalRef.current = null
|
||||
}
|
||||
setPlayingId(null)
|
||||
}, [])
|
||||
|
||||
/** 跳转到指定时间 */
|
||||
const handleSeek = useCallback(
|
||||
(voiceId: string, time: number, duration: number) => {
|
||||
if (playingId !== voiceId) {
|
||||
// 不同音色:从指定时间开始播放
|
||||
handlePlay(voiceId, duration, time)
|
||||
} else {
|
||||
// 同一音色:直接跳转
|
||||
setCurrentTime(time)
|
||||
}
|
||||
},
|
||||
[playingId, handlePlay],
|
||||
)
|
||||
|
||||
/** 切换播放/暂停 */
|
||||
const handleTogglePlay = useCallback(
|
||||
(voiceId: string, duration: number) => {
|
||||
if (playingId === voiceId) {
|
||||
handlePause()
|
||||
} else {
|
||||
handlePlay(voiceId, duration)
|
||||
}
|
||||
},
|
||||
[playingId, handlePlay, handlePause],
|
||||
)
|
||||
|
||||
/** 停止所有播放(切换 Tab 时调用) */
|
||||
const stopPlayback = useCallback(() => {
|
||||
if (intervalRef.current) {
|
||||
clearInterval(intervalRef.current)
|
||||
intervalRef.current = null
|
||||
}
|
||||
setPlayingId(null)
|
||||
setCurrentTime(0)
|
||||
}, [])
|
||||
|
||||
// 组件卸载时清理
|
||||
useEffect(() => {
|
||||
return () => {
|
||||
if (intervalRef.current) {
|
||||
clearInterval(intervalRef.current)
|
||||
}
|
||||
}
|
||||
}, [])
|
||||
|
||||
return {
|
||||
playingId,
|
||||
currentTime,
|
||||
handlePlay,
|
||||
handlePause,
|
||||
handleSeek,
|
||||
handleTogglePlay,
|
||||
stopPlayback,
|
||||
}
|
||||
}
|
||||
@@ -1,105 +0,0 @@
|
||||
import { useState, useCallback } from "react"
|
||||
import { useMutation, useQueryClient } from "@tanstack/react-query"
|
||||
import { deleteVoiceClone, retryVoiceClone, type VoiceClone } from "@/api/voice-clone"
|
||||
import { type ClonedVoiceDisplay } from "../types"
|
||||
|
||||
/**
|
||||
* 克隆音色操作 Hook
|
||||
* 封装删除、重试、详情弹窗等克隆音色相关操作
|
||||
*/
|
||||
interface ToastShowFn {
|
||||
(message: string, type: "success" | "error"): void
|
||||
}
|
||||
|
||||
interface UseCloneOperationsProps {
|
||||
showToast: ToastShowFn
|
||||
}
|
||||
|
||||
export function useCloneOperations({ showToast }: UseCloneOperationsProps) {
|
||||
const queryClient = useQueryClient()
|
||||
|
||||
const [detailVoice, setDetailVoice] = useState<ClonedVoiceDisplay | null>(null)
|
||||
const [cloneModalOpen, setCloneModalOpen] = useState(false)
|
||||
|
||||
/** 删除克隆音色 */
|
||||
const deleteMutation = useMutation({
|
||||
mutationFn: deleteVoiceClone,
|
||||
onSuccess: () => {
|
||||
queryClient.invalidateQueries({ queryKey: ["voice-clones"] })
|
||||
showToast("音色已删除", "success")
|
||||
},
|
||||
onError: () => {
|
||||
showToast("删除失败", "error")
|
||||
},
|
||||
})
|
||||
|
||||
/** 重试克隆 */
|
||||
const retryMutation = useMutation({
|
||||
mutationFn: retryVoiceClone,
|
||||
onSuccess: () => {
|
||||
queryClient.invalidateQueries({ queryKey: ["voice-clones"] })
|
||||
showToast("已重新提交克隆", "success")
|
||||
},
|
||||
onError: () => {
|
||||
showToast("重试失败", "error")
|
||||
},
|
||||
})
|
||||
|
||||
const handleCloneDelete = useCallback(
|
||||
(voice: ClonedVoiceDisplay) => {
|
||||
if (window.confirm(`确定删除音色「${voice.name}」吗?`)) {
|
||||
deleteMutation.mutate(voice.id)
|
||||
if (detailVoice?.id === voice.id) setDetailVoice(null)
|
||||
}
|
||||
},
|
||||
[deleteMutation, detailVoice],
|
||||
)
|
||||
|
||||
const handleCloneRetry = useCallback(
|
||||
(voice: ClonedVoiceDisplay) => {
|
||||
retryMutation.mutate(voice.id)
|
||||
},
|
||||
[retryMutation],
|
||||
)
|
||||
|
||||
const handleCloneUse = useCallback(
|
||||
(_voice: ClonedVoiceDisplay) => {
|
||||
showToast("已选择音色", "success")
|
||||
},
|
||||
[showToast],
|
||||
)
|
||||
|
||||
const handleShowDetail = useCallback((voice: ClonedVoiceDisplay) => {
|
||||
setDetailVoice(voice)
|
||||
}, [])
|
||||
|
||||
const handleCloseDetail = useCallback(() => {
|
||||
setDetailVoice(null)
|
||||
}, [])
|
||||
|
||||
/** 克隆成功回调 — 刷新列表 */
|
||||
const handleCloneSuccess = useCallback(
|
||||
(_voice: VoiceClone) => {
|
||||
queryClient.invalidateQueries({ queryKey: ["voice-clones"] })
|
||||
showToast("克隆已提交,正在生成中", "success")
|
||||
},
|
||||
[queryClient, showToast],
|
||||
)
|
||||
|
||||
return {
|
||||
// 状态
|
||||
detailVoice,
|
||||
cloneModalOpen,
|
||||
setCloneModalOpen,
|
||||
// Mutations
|
||||
isDeleting: deleteMutation.isPending,
|
||||
isRetrying: retryMutation.isPending,
|
||||
// Handlers
|
||||
handleCloneDelete,
|
||||
handleCloneRetry,
|
||||
handleCloneUse,
|
||||
handleShowDetail,
|
||||
handleCloseDetail,
|
||||
handleCloneSuccess,
|
||||
}
|
||||
}
|
||||
@@ -1,145 +0,0 @@
|
||||
import { useState, useRef, useCallback, useEffect } from "react"
|
||||
import { useQueryClient } from "@tanstack/react-query"
|
||||
import { message } from "antd"
|
||||
import { synthesizeSpeech, getTTSJobStatus, saveTtsToLibrary } from "@/api/tts"
|
||||
import { type PresetVoiceDisplay } from "../types"
|
||||
|
||||
export type TtsStatus = "idle" | "synthesizing" | "done" | "error"
|
||||
|
||||
/**
|
||||
* TTS 合成 Hook
|
||||
* 封装 AI 配音弹窗状态、合成请求、轮询、保存到素材库等逻辑
|
||||
*/
|
||||
interface UseTtsSynthesizeProps {
|
||||
presetVoices: PresetVoiceDisplay[]
|
||||
showToast: (message: string, type: "success" | "error") => void
|
||||
}
|
||||
|
||||
export function useTtsSynthesize({ presetVoices, showToast }: UseTtsSynthesizeProps) {
|
||||
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)
|
||||
|
||||
/** 开始 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, 50) })
|
||||
queryClient.invalidateQueries({ queryKey: ["voice-materials"] })
|
||||
showToast("已保存到配音库", "success")
|
||||
setTtsOpen(false)
|
||||
} catch (err: unknown) {
|
||||
const msg = err instanceof Error ? err.message : "保存失败"
|
||||
showToast(msg, "error")
|
||||
}
|
||||
}, [ttsJobId, ttsText, queryClient, showToast])
|
||||
|
||||
/** 关闭 TTS 弹窗并清理状态 */
|
||||
const handleTtsClose = useCallback(() => {
|
||||
setTtsOpen(false)
|
||||
setTtsText("")
|
||||
setTtsVoiceId("")
|
||||
setTtsSpeed(1.0)
|
||||
setTtsStatus("idle")
|
||||
setTtsAudioUrl(null)
|
||||
setTtsError(null)
|
||||
setTtsJobId(null)
|
||||
if (ttsTimerRef.current) {
|
||||
clearInterval(ttsTimerRef.current)
|
||||
ttsTimerRef.current = null
|
||||
}
|
||||
}, [])
|
||||
|
||||
/** 打开 TTS 弹窗,可选指定音色 */
|
||||
const openTtsWithVoice = useCallback((voiceId?: string) => {
|
||||
setTtsOpen(true)
|
||||
if (voiceId) setTtsVoiceId(voiceId)
|
||||
}, [])
|
||||
|
||||
// 组件卸载时清理定时器
|
||||
useEffect(() => {
|
||||
return () => {
|
||||
if (ttsTimerRef.current) clearInterval(ttsTimerRef.current)
|
||||
}
|
||||
}, [])
|
||||
|
||||
return {
|
||||
// 状态
|
||||
ttsOpen,
|
||||
ttsText,
|
||||
ttsVoiceId,
|
||||
ttsSpeed,
|
||||
ttsJobId,
|
||||
ttsStatus,
|
||||
ttsAudioUrl,
|
||||
ttsError,
|
||||
// 可选音色列表
|
||||
ttsPresetVoices: presetVoices,
|
||||
// Setters
|
||||
setTtsText,
|
||||
setTtsVoiceId,
|
||||
setTtsSpeed,
|
||||
setTtsOpen,
|
||||
// Actions
|
||||
handleTtsSynthesize,
|
||||
handleTtsSave,
|
||||
handleTtsClose,
|
||||
openTtsWithVoice,
|
||||
}
|
||||
}
|
||||
@@ -1,131 +0,0 @@
|
||||
import { useState, useCallback } from "react"
|
||||
import { useMutation, useQueryClient } from "@tanstack/react-query"
|
||||
import { type VoiceGender } from "../types"
|
||||
import { uploadAssetDirect, getAssetLibraries, createAsset } from "@/api/assets"
|
||||
import { getAudioDuration } from "../utils/audio"
|
||||
import { buildVoiceMetadata } from "../types"
|
||||
|
||||
/**
|
||||
* 配音上传 Hook
|
||||
* 封装上传音频弹窗状态、上传进度、上传 mutation 逻辑
|
||||
*/
|
||||
interface UseVoiceUploadProps {
|
||||
showToast: (message: string, type: "success" | "error") => void
|
||||
}
|
||||
|
||||
export function useVoiceUpload({ showToast }: UseVoiceUploadProps) {
|
||||
const queryClient = useQueryClient()
|
||||
|
||||
const [uploadOpen, setUploadOpen] = useState(false)
|
||||
const [uploadFile, setUploadFile] = useState<File | null>(null)
|
||||
const [uploadName, setUploadName] = useState("")
|
||||
const [uploadGender, setUploadGender] = useState<VoiceGender>("female")
|
||||
const [uploadDesc, setUploadDesc] = useState("")
|
||||
const [uploadProgress, setUploadProgress] = useState<number | null>(null)
|
||||
|
||||
const uploadMutation = useMutation({
|
||||
mutationFn: async (data: {
|
||||
file: File
|
||||
name: string
|
||||
gender: VoiceGender
|
||||
description: string
|
||||
}) => {
|
||||
setUploadProgress(0)
|
||||
try {
|
||||
/* 获取或创建默认配音库 */
|
||||
const libs = await queryClient.fetchQuery({
|
||||
queryKey: ["asset-libraries"],
|
||||
queryFn: getAssetLibraries,
|
||||
})
|
||||
const lib = libs.find((l) => l.kind === "voice")
|
||||
if (!lib) throw new Error("配音库不存在,请先在配音库页面创建")
|
||||
|
||||
/* 直传文件 */
|
||||
const { storage_key } = await uploadAssetDirect({
|
||||
file: data.file,
|
||||
library_id: lib.id,
|
||||
onProgress: (p) => setUploadProgress(p),
|
||||
})
|
||||
|
||||
/* 获取音频时长 */
|
||||
const duration = await getAudioDuration(data.file)
|
||||
|
||||
/* 创建素材记录 */
|
||||
await createAsset({
|
||||
library_id: lib.id,
|
||||
name: data.name,
|
||||
storage_key,
|
||||
mime_type: data.file.type || "audio/mpeg",
|
||||
metadata: buildVoiceMetadata({
|
||||
gender: data.gender,
|
||||
description: data.description,
|
||||
duration,
|
||||
}),
|
||||
})
|
||||
} finally {
|
||||
setUploadProgress(null)
|
||||
}
|
||||
},
|
||||
onSuccess: () => {
|
||||
queryClient.invalidateQueries({ queryKey: ["voice-materials"] })
|
||||
showToast("上传成功", "success")
|
||||
handleUploadClose()
|
||||
},
|
||||
onError: (err: Error) => {
|
||||
showToast(err.message || "上传失败,请重试", "error")
|
||||
},
|
||||
})
|
||||
|
||||
const handleUploadClose = useCallback(() => {
|
||||
setUploadOpen(false)
|
||||
setUploadFile(null)
|
||||
setUploadName("")
|
||||
setUploadDesc("")
|
||||
setUploadProgress(null)
|
||||
}, [])
|
||||
|
||||
const handleFileSelect = useCallback(
|
||||
(file: File) => {
|
||||
setUploadFile(file)
|
||||
if (!uploadName) setUploadName(file.name.replace(/\.[^.]+$/, ""))
|
||||
},
|
||||
[uploadName],
|
||||
)
|
||||
|
||||
const handleFileRemove = useCallback(() => {
|
||||
setUploadFile(null)
|
||||
setUploadProgress(null)
|
||||
}, [])
|
||||
|
||||
const handleUpload = useCallback(() => {
|
||||
if (!uploadFile) return
|
||||
uploadMutation.mutate({
|
||||
file: uploadFile,
|
||||
name: uploadName.trim(),
|
||||
gender: uploadGender,
|
||||
description: uploadDesc.trim(),
|
||||
})
|
||||
}, [uploadFile, uploadName, uploadGender, uploadDesc, uploadMutation])
|
||||
|
||||
return {
|
||||
// 弹窗状态
|
||||
uploadOpen,
|
||||
setUploadOpen,
|
||||
// 表单状态
|
||||
uploadFile,
|
||||
uploadName,
|
||||
uploadGender,
|
||||
uploadDesc,
|
||||
uploadProgress,
|
||||
isUploading: uploadMutation.isPending,
|
||||
// Setters
|
||||
setUploadName,
|
||||
setUploadGender,
|
||||
setUploadDesc,
|
||||
// Handlers
|
||||
handleFileSelect,
|
||||
handleFileRemove,
|
||||
handleUpload,
|
||||
handleUploadClose,
|
||||
}
|
||||
}
|
||||
@@ -1,119 +0,0 @@
|
||||
import { useState, useMemo } from "react"
|
||||
import { useQuery } from "@tanstack/react-query"
|
||||
import { fetchPresetVoices, fetchVoices } from "@/api/voices"
|
||||
import { getVoiceClonesWithTotal, toVoiceClone } from "@/api/voice-clone"
|
||||
import { getAssetsByKind, type AssetItem } from "@/api/assets"
|
||||
import {
|
||||
type TabKey,
|
||||
type ClonedVoiceDisplay,
|
||||
type PresetVoiceDisplay,
|
||||
mapPresetToDisplay,
|
||||
mapCloneToDisplay,
|
||||
} from "../types"
|
||||
|
||||
/**
|
||||
* 配音库数据 Hook
|
||||
* 封装三个 Tab 的数据查询、筛选状态管理、数据映射逻辑
|
||||
*/
|
||||
export function useVoicesData() {
|
||||
// ── Tab & 筛选状态 ─────────────────────────────────────
|
||||
const [activeTab, setActiveTab] = useState<TabKey>("preset")
|
||||
const [searchText, setSearchText] = useState("")
|
||||
const [filterGender, setFilterGender] = useState<string>("all")
|
||||
const [filterLang, setFilterLang] = useState<string>("all")
|
||||
|
||||
// ── 数据查询 ───────────────────────────────────────────
|
||||
|
||||
/** 预置音色列表 */
|
||||
const { data: presetData, isLoading: presetLoading } = useQuery({
|
||||
queryKey: ["preset-voices"],
|
||||
queryFn: fetchPresetVoices,
|
||||
})
|
||||
|
||||
/** 克隆音色列表 */
|
||||
const { data: cloneData, isLoading: cloneLoading } = useQuery({
|
||||
queryKey: ["voice-clones"],
|
||||
queryFn: () => getVoiceClonesWithTotal({ limit: 50 }),
|
||||
})
|
||||
|
||||
/** 配音素材列表(用户上传音频) */
|
||||
const { data: materialData, isLoading: materialLoading } = useQuery({
|
||||
queryKey: ["voice-materials"],
|
||||
queryFn: () => getAssetsByKind("voice", { limit: 50 }),
|
||||
})
|
||||
|
||||
/** 统一统计(preset_count / clone_count) */
|
||||
const { data: unifiedStats } = useQuery({
|
||||
queryKey: ["voices-unified"],
|
||||
queryFn: () => fetchVoices({ limit: 1 }),
|
||||
})
|
||||
|
||||
// ── 数据映射 ─────────────────────────────────────────
|
||||
|
||||
const presetVoices: PresetVoiceDisplay[] = useMemo(
|
||||
() => (presetData?.items ?? []).map(mapPresetToDisplay),
|
||||
[presetData],
|
||||
)
|
||||
|
||||
const clonedVoices: ClonedVoiceDisplay[] = useMemo(
|
||||
() => (cloneData?.items ?? []).map((p) => mapCloneToDisplay(toVoiceClone(p))),
|
||||
[cloneData],
|
||||
)
|
||||
|
||||
const materials: AssetItem[] = useMemo(() => materialData ?? [], [materialData])
|
||||
|
||||
// ── 计数 ──────────────────────────────────────────────
|
||||
|
||||
const presetCount = unifiedStats?.preset_count ?? presetData?.total ?? 0
|
||||
const cloneCount = unifiedStats?.clone_count ?? cloneData?.total ?? 0
|
||||
const materialCount = materialData?.length ?? 0
|
||||
|
||||
// ── 预置音色筛选 ─────────────────────────────────────
|
||||
|
||||
const filteredPreset = useMemo(() => {
|
||||
let list = presetVoices
|
||||
if (filterGender !== "all") {
|
||||
list = list.filter((v) => v.gender === filterGender)
|
||||
}
|
||||
if (filterLang !== "all") {
|
||||
list = list.filter((v) => v.language === filterLang)
|
||||
}
|
||||
if (searchText.trim()) {
|
||||
const q = searchText.trim().toLowerCase()
|
||||
list = list.filter(
|
||||
(v) =>
|
||||
v.name.toLowerCase().includes(q) ||
|
||||
v.description.toLowerCase().includes(q) ||
|
||||
v.tags.some((tag) => tag.toLowerCase().includes(q)),
|
||||
)
|
||||
}
|
||||
return list
|
||||
}, [presetVoices, filterGender, filterLang, searchText])
|
||||
|
||||
return {
|
||||
// Tab 状态
|
||||
activeTab,
|
||||
setActiveTab,
|
||||
// 筛选状态
|
||||
searchText,
|
||||
setSearchText,
|
||||
filterGender,
|
||||
setFilterGender,
|
||||
filterLang,
|
||||
setFilterLang,
|
||||
// 加载状态
|
||||
presetLoading,
|
||||
cloneLoading,
|
||||
materialLoading,
|
||||
// 原始数据
|
||||
presetVoices,
|
||||
clonedVoices,
|
||||
materials,
|
||||
// 筛选后数据
|
||||
filteredPreset,
|
||||
// 计数
|
||||
presetCount,
|
||||
cloneCount,
|
||||
materialCount,
|
||||
}
|
||||
}
|
||||
@@ -1,92 +0,0 @@
|
||||
/**
|
||||
* 配音库类型定义
|
||||
*/
|
||||
import type { PresetVoiceItem } from "@/api/voices"
|
||||
import type { VoiceClone } from "@/api/voice-clone"
|
||||
|
||||
export type VoiceGender = "male" | "female" | "child" | "elderly"
|
||||
export type VoiceLanguage = "zh" | "en" | "ja" | "ko"
|
||||
export type TabKey = "preset" | "cloned" | "material"
|
||||
|
||||
/** 前端展示用的预置音色(从 PresetVoiceItem 映射) */
|
||||
export interface PresetVoiceDisplay {
|
||||
id: string
|
||||
name: string
|
||||
gender: VoiceGender
|
||||
language: VoiceLanguage
|
||||
duration: number
|
||||
tags: string[]
|
||||
description: string
|
||||
voiceId: string
|
||||
previewUrl: string
|
||||
starred: boolean
|
||||
}
|
||||
|
||||
/** 前端展示用的克隆音色(从 VoiceClone 映射) */
|
||||
export interface ClonedVoiceDisplay {
|
||||
id: string
|
||||
name: string
|
||||
description: string
|
||||
sourceName: string
|
||||
status: "ready" | "processing" | "failed"
|
||||
createdAt: string
|
||||
duration: number
|
||||
tags: string[]
|
||||
voiceId: string
|
||||
language: string
|
||||
gender: string
|
||||
errorMessage: string | null
|
||||
sampleUrl?: string
|
||||
}
|
||||
|
||||
/** 音色上传元数据(传递给 createAsset 的 metadata) */
|
||||
export interface VoiceUploadMetadata {
|
||||
gender?: string
|
||||
description?: string
|
||||
duration?: number
|
||||
[key: string]: unknown
|
||||
}
|
||||
|
||||
/** 后端 PresetVoiceItem → 前端 PresetVoiceDisplay */
|
||||
export const mapPresetToDisplay = (item: PresetVoiceItem): PresetVoiceDisplay => ({
|
||||
id: item.voice_id,
|
||||
name: item.name,
|
||||
gender: (item.gender as VoiceGender) || "female",
|
||||
language: (item.language as VoiceLanguage) || "zh",
|
||||
duration: 0,
|
||||
tags: item.tags,
|
||||
description: item.description,
|
||||
voiceId: item.voice_id,
|
||||
previewUrl: item.preview_url || "",
|
||||
starred: false,
|
||||
})
|
||||
|
||||
/** 后端 VoiceClone → 前端 ClonedVoiceDisplay */
|
||||
export const mapCloneToDisplay = (clone: VoiceClone): ClonedVoiceDisplay => ({
|
||||
id: clone.id,
|
||||
name: clone.name,
|
||||
description: clone.description || "",
|
||||
sourceName: clone.sample_url || "未知来源",
|
||||
status: clone.status,
|
||||
createdAt: new Date(clone.created_at).toLocaleDateString("zh-CN"),
|
||||
duration: clone.duration_seconds,
|
||||
tags: [],
|
||||
voiceId: clone.id,
|
||||
language: clone.language || "",
|
||||
gender: clone.gender || "",
|
||||
errorMessage: clone.error_message || null,
|
||||
sampleUrl: clone.sample_url || undefined,
|
||||
})
|
||||
|
||||
/** 前端表单数据 → 后端 metadata */
|
||||
export const buildVoiceMetadata = (data: {
|
||||
gender?: string
|
||||
description?: string
|
||||
duration?: number
|
||||
}): VoiceUploadMetadata => {
|
||||
const metadata: VoiceUploadMetadata = {}
|
||||
if (data.gender) metadata.gender = data.gender
|
||||
if (data.description) metadata.description = data.description
|
||||
if (data.duration) metadata.duration = Math.round(data.duration)
|
||||
return metadata
|
||||
}
|
||||
@@ -1,20 +0,0 @@
|
||||
/**
|
||||
* 音频相关工具函数
|
||||
*/
|
||||
|
||||
/** 获取音频文件时长(秒) */
|
||||
export const getAudioDuration = (file: File): Promise<number> => {
|
||||
return new Promise((resolve) => {
|
||||
const audio = new Audio()
|
||||
const url = URL.createObjectURL(file)
|
||||
audio.addEventListener("loadedmetadata", () => {
|
||||
resolve(audio.duration)
|
||||
URL.revokeObjectURL(url)
|
||||
})
|
||||
audio.addEventListener("error", () => {
|
||||
resolve(0)
|
||||
URL.revokeObjectURL(url)
|
||||
})
|
||||
audio.src = url
|
||||
})
|
||||
}
|
||||
@@ -1,24 +0,0 @@
|
||||
/**
|
||||
* 格式化工具函数
|
||||
*/
|
||||
import { GENDER_OPTIONS, LANGUAGE_OPTIONS } from "../constants"
|
||||
import type { VoiceGender, VoiceLanguage } from "../types"
|
||||
|
||||
export const genderLabel = (g: VoiceGender) => GENDER_OPTIONS.find((o) => o.value === g)?.label ?? g
|
||||
|
||||
export const languageLabel = (l: VoiceLanguage) =>
|
||||
LANGUAGE_OPTIONS.find((o) => o.value === l)?.label ?? l
|
||||
|
||||
export const genderClass = (g: VoiceGender) => `xx-voice-gender--${g}`
|
||||
|
||||
export const formatTime = (seconds: number): string => {
|
||||
const mins = Math.floor(seconds / 60)
|
||||
const secs = Math.floor(seconds % 60)
|
||||
return `${mins.toString().padStart(2, "0")}:${secs.toString().padStart(2, "0")}`
|
||||
}
|
||||
|
||||
export const formatFileSize = (bytes: number): string => {
|
||||
if (bytes < 1024) return `${bytes} B`
|
||||
if (bytes < 1024 * 1024) return `${(bytes / 1024).toFixed(1)} KB`
|
||||
return `${(bytes / (1024 * 1024)).toFixed(1)} MB`
|
||||
}
|
||||
@@ -1,45 +0,0 @@
|
||||
/**
|
||||
* AssetLibrary 模块 smoke test
|
||||
* 建立完整依赖链,确保 vitest related 模式能匹配到
|
||||
* assets 目录下所有文件的改动(包括子组件、Hook 和工具函数)
|
||||
*/
|
||||
import { describe, it, expect } from "vitest"
|
||||
|
||||
// 主组件
|
||||
import "@/pages/assets/AssetLibrary"
|
||||
|
||||
// 子组件
|
||||
import "@/pages/assets/components/AssetCard"
|
||||
import "@/pages/assets/components/AssetFilterBar"
|
||||
import "@/pages/assets/components/AssetSkeleton"
|
||||
import "@/pages/assets/components/BatchClassifyModal"
|
||||
import "@/pages/assets/components/BatchMarkModal"
|
||||
import "@/pages/assets/components/BatchOperationBar"
|
||||
import "@/pages/assets/components/BatchTagModal"
|
||||
import "@/pages/assets/components/CreateLibraryModal"
|
||||
import "@/pages/assets/components/LibrarySidebar"
|
||||
import "@/pages/assets/components/PlayModal"
|
||||
import "@/pages/assets/components/ResultDrawer"
|
||||
import "@/pages/assets/components/UploadProgressModal"
|
||||
|
||||
// 类型与常量
|
||||
import "@/pages/assets/types"
|
||||
import "@/pages/assets/constants"
|
||||
|
||||
// 工具函数
|
||||
import "@/pages/assets/utils/format"
|
||||
import "@/pages/assets/utils/asset"
|
||||
|
||||
// Hooks
|
||||
import "@/pages/assets/hooks/useAssetsData"
|
||||
import "@/pages/assets/hooks/useLibraryManagement"
|
||||
import "@/pages/assets/hooks/useAssetUpload"
|
||||
import "@/pages/assets/hooks/useAssetSelection"
|
||||
import "@/pages/assets/hooks/useAssetOperations"
|
||||
|
||||
describe("AssetLibrary module smoke test", () => {
|
||||
it("should load all asset modules", () => {
|
||||
// 纯模块加载测试,确保所有组件/工具函数能正常 import
|
||||
expect(true).toBe(true)
|
||||
})
|
||||
})
|
||||
@@ -1,71 +0,0 @@
|
||||
/**
|
||||
* 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([])
|
||||
})
|
||||
})
|
||||
@@ -1,148 +0,0 @@
|
||||
/**
|
||||
* 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")
|
||||
})
|
||||
})
|
||||
@@ -1,26 +0,0 @@
|
||||
/**
|
||||
* 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)
|
||||
})
|
||||
})
|
||||
@@ -1,232 +0,0 @@
|
||||
/**
|
||||
* useAudioPlayer hook 测试 — VoiceLibrary 版本
|
||||
*
|
||||
* 该 Hook 使用 setInterval 模拟音频播放进度,纯逻辑可测。
|
||||
* 参考 voice-materials/hooks/useAudioPlayer.test.ts 的测试结构。
|
||||
*/
|
||||
import { describe, it, expect, beforeEach, vi, afterEach } from "vitest"
|
||||
import { renderHook, act } from "@testing-library/react"
|
||||
import { useAudioPlayer } from "@/pages/voices/hooks/useAudioPlayer"
|
||||
|
||||
describe("useAudioPlayer (voices)", () => {
|
||||
beforeEach(() => {
|
||||
vi.useFakeTimers()
|
||||
})
|
||||
|
||||
afterEach(() => {
|
||||
vi.useRealTimers()
|
||||
vi.clearAllMocks()
|
||||
})
|
||||
|
||||
it("应该使用初始状态初始化", () => {
|
||||
const { result } = renderHook(() => useAudioPlayer())
|
||||
|
||||
expect(result.current.playingId).toBeNull()
|
||||
expect(result.current.currentTime).toBe(0)
|
||||
})
|
||||
|
||||
it("handlePlay 应该开始播放指定音色", () => {
|
||||
const { result } = renderHook(() => useAudioPlayer())
|
||||
|
||||
act(() => {
|
||||
result.current.handlePlay("voice-1", 10)
|
||||
})
|
||||
|
||||
expect(result.current.playingId).toBe("voice-1")
|
||||
expect(result.current.currentTime).toBe(0)
|
||||
})
|
||||
|
||||
it("handlePlay 对同一个音色不应重复启动播放", () => {
|
||||
const { result } = renderHook(() => useAudioPlayer())
|
||||
|
||||
act(() => {
|
||||
result.current.handlePlay("voice-1", 10)
|
||||
})
|
||||
|
||||
const initialTime = result.current.currentTime
|
||||
|
||||
// 推进一些时间让进度走动
|
||||
act(() => {
|
||||
vi.advanceTimersByTime(200)
|
||||
})
|
||||
|
||||
const timeAfterAdvance = result.current.currentTime
|
||||
expect(timeAfterAdvance).toBeGreaterThan(initialTime)
|
||||
|
||||
// 对同一个音色再次调用 handlePlay 不应重置
|
||||
act(() => {
|
||||
result.current.handlePlay("voice-1", 10)
|
||||
})
|
||||
|
||||
expect(result.current.playingId).toBe("voice-1")
|
||||
expect(result.current.currentTime).toBe(timeAfterAdvance)
|
||||
})
|
||||
|
||||
it("handlePlay 切换音色时应停止上一个并从头开始", () => {
|
||||
const { result } = renderHook(() => useAudioPlayer())
|
||||
|
||||
act(() => {
|
||||
result.current.handlePlay("voice-1", 10)
|
||||
})
|
||||
|
||||
act(() => {
|
||||
vi.advanceTimersByTime(500)
|
||||
})
|
||||
|
||||
expect(result.current.playingId).toBe("voice-1")
|
||||
expect(result.current.currentTime).toBeGreaterThan(0)
|
||||
|
||||
act(() => {
|
||||
result.current.handlePlay("voice-2", 15)
|
||||
})
|
||||
|
||||
expect(result.current.playingId).toBe("voice-2")
|
||||
expect(result.current.currentTime).toBe(0)
|
||||
})
|
||||
|
||||
it("播放进度应该随时间递增", () => {
|
||||
const { result } = renderHook(() => useAudioPlayer())
|
||||
|
||||
act(() => {
|
||||
result.current.handlePlay("voice-1", 10)
|
||||
})
|
||||
|
||||
// 每 100ms 增加 0.1
|
||||
act(() => {
|
||||
vi.advanceTimersByTime(300)
|
||||
})
|
||||
|
||||
expect(result.current.currentTime).toBeCloseTo(0.3, 1)
|
||||
expect(result.current.playingId).toBe("voice-1")
|
||||
})
|
||||
|
||||
it("播放到结尾应自动停止并重置", () => {
|
||||
const { result } = renderHook(() => useAudioPlayer())
|
||||
|
||||
act(() => {
|
||||
result.current.handlePlay("voice-1", 0.5) // 0.5 秒的短音频
|
||||
})
|
||||
|
||||
act(() => {
|
||||
vi.advanceTimersByTime(600) // 超过 0.5 秒
|
||||
})
|
||||
|
||||
expect(result.current.playingId).toBeNull()
|
||||
expect(result.current.currentTime).toBe(0)
|
||||
})
|
||||
|
||||
it("handlePause 应该暂停播放", () => {
|
||||
const { result } = renderHook(() => useAudioPlayer())
|
||||
|
||||
act(() => {
|
||||
result.current.handlePlay("voice-1", 10)
|
||||
})
|
||||
|
||||
act(() => {
|
||||
vi.advanceTimersByTime(200)
|
||||
})
|
||||
|
||||
const timeBeforePause = result.current.currentTime
|
||||
|
||||
act(() => {
|
||||
result.current.handlePause()
|
||||
})
|
||||
|
||||
expect(result.current.playingId).toBeNull()
|
||||
|
||||
// 暂停后时间不应再变化
|
||||
act(() => {
|
||||
vi.advanceTimersByTime(500)
|
||||
})
|
||||
|
||||
expect(result.current.currentTime).toBe(timeBeforePause)
|
||||
})
|
||||
|
||||
it("handleSeek 应该跳转到指定时间", () => {
|
||||
const { result } = renderHook(() => useAudioPlayer())
|
||||
|
||||
act(() => {
|
||||
result.current.handlePlay("voice-1", 10)
|
||||
})
|
||||
|
||||
act(() => {
|
||||
result.current.handleSeek("voice-1", 5, 10)
|
||||
})
|
||||
|
||||
expect(result.current.currentTime).toBe(5)
|
||||
})
|
||||
|
||||
it("handleSeek 对不同音色应该开始播放该音色", () => {
|
||||
const { result } = renderHook(() => useAudioPlayer())
|
||||
|
||||
act(() => {
|
||||
result.current.handlePlay("voice-1", 10)
|
||||
})
|
||||
|
||||
act(() => {
|
||||
result.current.handleSeek("voice-2", 3, 15)
|
||||
})
|
||||
|
||||
expect(result.current.playingId).toBe("voice-2")
|
||||
expect(result.current.currentTime).toBe(3)
|
||||
})
|
||||
|
||||
it("handleTogglePlay 应该在播放和暂停之间切换", () => {
|
||||
const { result } = renderHook(() => useAudioPlayer())
|
||||
|
||||
// 初始为暂停,调用应开始播放
|
||||
act(() => {
|
||||
result.current.handleTogglePlay("voice-1", 10)
|
||||
})
|
||||
|
||||
expect(result.current.playingId).toBe("voice-1")
|
||||
|
||||
// 再次调用应暂停
|
||||
act(() => {
|
||||
result.current.handleTogglePlay("voice-1", 10)
|
||||
})
|
||||
|
||||
expect(result.current.playingId).toBeNull()
|
||||
})
|
||||
|
||||
it("stopPlayback 应该重置所有播放状态", () => {
|
||||
const { result } = renderHook(() => useAudioPlayer())
|
||||
|
||||
act(() => {
|
||||
result.current.handlePlay("voice-1", 10)
|
||||
})
|
||||
|
||||
act(() => {
|
||||
vi.advanceTimersByTime(300)
|
||||
})
|
||||
|
||||
expect(result.current.playingId).toBe("voice-1")
|
||||
expect(result.current.currentTime).toBeGreaterThan(0)
|
||||
|
||||
act(() => {
|
||||
result.current.stopPlayback()
|
||||
})
|
||||
|
||||
expect(result.current.playingId).toBeNull()
|
||||
expect(result.current.currentTime).toBe(0)
|
||||
|
||||
// 停止后定时器不应再触发
|
||||
const timeAfterStop = result.current.currentTime
|
||||
act(() => {
|
||||
vi.advanceTimersByTime(500)
|
||||
})
|
||||
expect(result.current.currentTime).toBe(timeAfterStop)
|
||||
})
|
||||
|
||||
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.handleTogglePlay).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")
|
||||
})
|
||||
})
|
||||
@@ -1,41 +0,0 @@
|
||||
/**
|
||||
* VoiceLibrary 模块 smoke test
|
||||
* 建立完整依赖链,确保 vitest related 模式能匹配到
|
||||
* voices 目录下所有文件的改动(包括工具函数)
|
||||
*/
|
||||
import { describe, it, expect } from "vitest"
|
||||
|
||||
// 主组件
|
||||
import "@/pages/voices/VoiceLibrary"
|
||||
|
||||
// 子组件
|
||||
import "@/pages/voices/components/VoiceCard"
|
||||
import "@/pages/voices/components/CloneVoiceCard"
|
||||
import "@/pages/voices/components/CloneDetailModal"
|
||||
import "@/pages/voices/components/CloneCardSkeleton"
|
||||
import "@/pages/voices/components/UploadVoiceModal"
|
||||
import "@/pages/voices/components/TtsModal"
|
||||
import "@/pages/voices/components/VoiceFilterBar"
|
||||
import "@/pages/voices/components/MaterialVoiceCard"
|
||||
|
||||
// 类型与常量
|
||||
import "@/pages/voices/types"
|
||||
import "@/pages/voices/constants"
|
||||
|
||||
// 工具函数
|
||||
import "@/pages/voices/utils/format"
|
||||
import "@/pages/voices/utils/audio"
|
||||
|
||||
// Hooks
|
||||
import "@/pages/voices/hooks/useVoicesData"
|
||||
import "@/pages/voices/hooks/useAudioPlayer"
|
||||
import "@/pages/voices/hooks/useCloneOperations"
|
||||
import "@/pages/voices/hooks/useTtsSynthesize"
|
||||
import "@/pages/voices/hooks/useVoiceUpload"
|
||||
|
||||
describe("VoiceLibrary module smoke test", () => {
|
||||
it("should load all voice-library modules", () => {
|
||||
// 纯模块加载测试,确保所有组件/工具函数能正常 import
|
||||
expect(true).toBe(true)
|
||||
})
|
||||
})
|
||||
@@ -1,198 +0,0 @@
|
||||
# VoiceLibrary 页面重构方案
|
||||
|
||||
## 现状分析
|
||||
|
||||
**当前文件:** `apps/web/src/pages/voices/VoiceLibrary.tsx` — 1783 行
|
||||
|
||||
**代码质量:**
|
||||
- `any`: 0 处
|
||||
- `ts-ignore`: 0 处
|
||||
- `eslint-disable`: 0 处
|
||||
- `TODO`: 1 处
|
||||
- 整体质量良好,重构阻力小
|
||||
|
||||
## 目录结构(重构后)
|
||||
|
||||
```
|
||||
apps/web/src/pages/voices/
|
||||
├── VoiceLibrary.tsx # 主组件(目标:~550 行,-69%)
|
||||
├── types.ts # 类型定义
|
||||
├── constants.ts # 常量 + 配置
|
||||
├── utils/
|
||||
│ ├── format.ts # 格式化工具函数
|
||||
│ └── audio.ts # 音频工具函数
|
||||
├── components/
|
||||
│ ├── VoiceCard.tsx # 预设音色卡片
|
||||
│ ├── CloneVoiceCard.tsx # 克隆音色卡片
|
||||
│ ├── CloneDetailModal.tsx # 克隆详情弹窗
|
||||
│ ├── CloneCardSkeleton.tsx # 克隆卡片骨架屏
|
||||
│ ├── UploadModal.tsx # 上传音色弹窗
|
||||
│ ├── TTSModal.tsx # TTS合成弹窗
|
||||
│ ├── VoiceFilterBar.tsx # 筛选栏(搜索/性别/语言)
|
||||
│ └── VoiceTabBar.tsx # Tab切换栏
|
||||
└── hooks/
|
||||
├── useVoices.ts # 音色数据查询 + 筛选
|
||||
├── useAudioPlayer.ts # 音频播放控制
|
||||
├── useVoiceUpload.ts # 音色上传逻辑
|
||||
├── useTTS.ts # TTS合成逻辑
|
||||
└── useCloneVoice.ts # 克隆音色操作
|
||||
```
|
||||
|
||||
## 三阶段渐进式重构
|
||||
|
||||
### Phase 1:抽离类型、常量、工具函数
|
||||
|
||||
**目标:** 主文件 1783 → ~1550 行(-13%)
|
||||
|
||||
**抽出内容:**
|
||||
|
||||
1. **`types.ts`** — 类型定义(~60行)
|
||||
- `PresetVoiceDisplay` / `ClonedVoiceDisplay` / `VoiceCardProps` / `CloneVoiceCardProps`
|
||||
- `TabKey` / `Toast` / `VoiceUploadMetadata`
|
||||
- 现有 `mapPresetToDisplay` / `mapCloneToDisplay` 数据映射函数
|
||||
|
||||
2. **`constants.ts`** — 常量配置(~40行)
|
||||
- `CLONE_STATUS_CONFIG` 克隆状态配置
|
||||
- `GENDER_LABEL_MAP` / `LANGUAGE_LABEL_MAP` 性别/语言标签映射
|
||||
- Tab 配置项
|
||||
|
||||
3. **`utils/format.ts`** — 格式化工具(~25行)
|
||||
- `formatTime` 时长格式化
|
||||
- `formatFileSize` 文件大小格式化
|
||||
- `genderLabel` / `languageLabel` / `genderClass`
|
||||
|
||||
4. **`utils/audio.ts`** — 音频工具(~15行)
|
||||
- `getAudioDuration` 获取音频文件时长
|
||||
|
||||
5. **`format.test.ts`** — 工具函数单测
|
||||
- 覆盖 formatTime / formatFileSize 等纯函数
|
||||
|
||||
**Phase 1 交付:**
|
||||
- 新增文件:6 个(types.ts / constants.ts / utils/format.ts / utils/audio.ts / format.test.ts)
|
||||
- 主文件减少:~230 行
|
||||
- 纯机械抽离,无逻辑改动
|
||||
|
||||
---
|
||||
|
||||
### Phase 2:抽离子组件
|
||||
|
||||
**目标:** 主文件 1550 → ~950 行(-39%)
|
||||
|
||||
**抽出组件:**
|
||||
|
||||
1. **`components/VoiceCard.tsx`** — 预设音色卡片(~120行)
|
||||
- 卡片渲染:头像、名称、性别标签、播放按钮、进度条、收藏
|
||||
- Props:voice / playingId / currentTime / onPlay / onPause / onSeek / onToggleStar
|
||||
|
||||
2. **`components/CloneVoiceCard.tsx`** — 克隆音色卡片(~160行)
|
||||
- 三种状态:processing / success / failed
|
||||
- 操作:播放、详情、删除、重试、使用
|
||||
- Props:voice / playingId / currentTime / onPlayPause / onShowDetail / onDelete / onRetry / onUse
|
||||
|
||||
3. **`components/CloneDetailModal.tsx`** — 克隆详情弹窗(~90行)
|
||||
- 展示克隆音色详细信息
|
||||
- 状态展示、音频播放、操作按钮
|
||||
|
||||
4. **`components/CloneCardSkeleton.tsx`** — 克隆卡片骨架屏(~20行)
|
||||
- 加载状态占位
|
||||
|
||||
5. **`components/UploadModal.tsx`** — 上传音色弹窗(~180行)
|
||||
- 文件选择、名称/性别/描述填写
|
||||
- 上传进度展示
|
||||
- Props:open / onClose / onUpload
|
||||
|
||||
6. **`components/TTSModal.tsx`** — TTS合成弹窗(~170行)
|
||||
- 文本输入、音色选择、语速调节
|
||||
- 合成状态轮询(idle/synthesizing/done/error)
|
||||
- 保存到素材库
|
||||
- Props:open / onClose / onSave
|
||||
|
||||
7. **`components/VoiceFilterBar.tsx`** — 筛选栏(~80行)
|
||||
- 搜索框、性别筛选、语言筛选
|
||||
- Props:searchText / filterGender / filterLang / onChange handlers
|
||||
|
||||
**Phase 2 交付:**
|
||||
- 新增组件:7 个
|
||||
- 主文件减少:~600 行
|
||||
- 纯UI抽离,业务逻辑保留在主组件
|
||||
|
||||
---
|
||||
|
||||
### Phase 3:抽离业务逻辑 Hook
|
||||
|
||||
**目标:** 主文件 950 → ~550 行(-42%)
|
||||
|
||||
**抽出 Hook:**
|
||||
|
||||
1. **`hooks/useVoices.ts`** — 音色数据 Hook(~200行)
|
||||
- 三个 useQuery:presetVoices / cloneVoices / materialVoices / unifiedStats
|
||||
- 筛选逻辑:searchText / filterGender / filterLang → filteredPreset / filteredClone
|
||||
- 统计数据:presetCount / cloneCount / materialCount
|
||||
- 收藏切换 handleToggleStar
|
||||
- 返回:data / loading / filtered / counts / handlers
|
||||
|
||||
2. **`hooks/useAudioPlayer.ts`** — 音频播放控制 Hook(~120行)
|
||||
- playingId / currentTime / intervalRef 状态
|
||||
- handlePlay / handlePause / handleSeek
|
||||
- 自动停止(切换新音频时停止旧的)
|
||||
- 组件卸载清理
|
||||
- 注意:与 VoiceMaterialLibrary 的 useAudioPlayer 类似但有差异(一个操作DOM音频,一个操作audio元素),评估是否复用还是独立
|
||||
|
||||
3. **`hooks/useVoiceUpload.ts`** — 音色上传 Hook(~120行)
|
||||
- uploadFile / uploadName / uploadGender / uploadDesc / uploadProgress 状态
|
||||
- uploadMutation(上传文件 + 创建音色)
|
||||
- buildVoiceMetadata 元数据构建
|
||||
- getAudioDuration 时长获取
|
||||
- 成功后刷新列表 + toast 关闭弹窗
|
||||
|
||||
4. **`hooks/useTTS.ts`** — TTS合成 Hook(~110行)
|
||||
- ttsOpen / ttsText / ttsVoiceId / ttsSpeed 弹窗状态
|
||||
- ttsJobId / ttsStatus / ttsAudioUrl / ttsError 合成状态
|
||||
- handleTtsSynthesize 发起合成 + 轮询
|
||||
- handleTtsSave 保存到素材库
|
||||
- handleTtsClose 清理状态
|
||||
|
||||
5. **`hooks/useCloneVoice.ts`** — 克隆音色操作 Hook(~80行)
|
||||
- deleteMutation / retryMutation
|
||||
- handleCloneDelete / handleCloneRetry
|
||||
- detailVoice 详情状态
|
||||
- 操作后刷新列表
|
||||
|
||||
**Phase 3 交付:**
|
||||
- 新增 Hook:5 个
|
||||
- 主文件减少:~400 行
|
||||
- 主文件只剩:Tab切换逻辑 + JSX 组装 + 顶层状态编排
|
||||
|
||||
---
|
||||
|
||||
## 最终效果汇总
|
||||
|
||||
| 阶段 | 主文件行数 | 减少行数 | 减少比例 | 新增文件 |
|
||||
|------|-----------|---------|----------|----------|
|
||||
| 初始 | 1783 | — | — | — |
|
||||
| Phase 1 | ~1550 | -233 | -13.1% | 5 |
|
||||
| Phase 2 | ~950 | -600 | -38.7% | 7 |
|
||||
| Phase 3 | ~550 | -400 | -42.1% | 5 |
|
||||
| **总计** | **~550** | **-1233** | **-69.2%** | **17** |
|
||||
|
||||
## 与 GeneratePage / VoiceMaterialLibrary 的异同
|
||||
|
||||
**相同点:**
|
||||
- 三阶段渐进式重构(类型常量 → 组件 → Hook)
|
||||
- 每阶段独立PR,独立验证
|
||||
- 每阶段加 smoke test 保证 vitest related 模式可用
|
||||
- 代码质量基线高(0 any / 0 ts-ignore)
|
||||
|
||||
**不同点:**
|
||||
- VoiceLibrary 有两个数据源(预设音色 + 克隆音色),还有TTS和上传功能
|
||||
- 音频播放逻辑与 VoiceMaterialLibrary 类似但操作的音频类型不同
|
||||
- 有 CloneModal 是外部组件(已在 @/components/voice/CloneModal),不需要重写
|
||||
- 骨架屏组件比较简单
|
||||
|
||||
## 风险与注意事项
|
||||
|
||||
1. **vitest related 模式**:每阶段新增文件需建立测试覆盖,避免 "No test files found"
|
||||
2. **Preview Deploy**:环境问题导致失败,不影响代码质量
|
||||
3. **音频播放逻辑**:VoiceLibrary 和 VoiceMaterialLibrary 的播放控制类似但是两套独立实现,后续可考虑抽象为共享 Hook
|
||||
4. **克隆状态**:CLONE_STATUS_CONFIG 是核心常量,抽离时注意类型完整
|
||||
5. **上传逻辑**:uploadMutation 内部逻辑较复杂(上传文件 + 获取时长 + 创建音色),抽 Hook 时保持逻辑不变
|
||||
@@ -1,198 +0,0 @@
|
||||
# VoiceMaterialLibrary 页面拆分方案
|
||||
|
||||
## 一、现状摸底
|
||||
|
||||
**文件:** `apps/web/src/pages/voice-materials/VoiceMaterialLibrary.tsx`
|
||||
**总行数:** 1966 行
|
||||
**CSS 文件:** `apps/web/src/pages/voice-materials/voice-materials.css` (1262 行)
|
||||
|
||||
### 代码质量
|
||||
|
||||
| 指标 | 数量 | 评价 |
|
||||
|------|------|------|
|
||||
| `any` 类型 | 0 | ✅ 优秀 |
|
||||
| `@ts-ignore` | 0 | ✅ 优秀 |
|
||||
| `eslint-disable` | 0 | ✅ 优秀 |
|
||||
| TODO/FIXME | 0 | ✅ 干净 |
|
||||
|
||||
### 内联组件分布
|
||||
|
||||
| 组件 | 行数 | 职责 |
|
||||
|------|------|------|
|
||||
| `TagSelector` | 148行 (197-345) | 标签选择器(输入+建议+选择) |
|
||||
| `MaterialForm` | 187行 (356-543) | 上传/编辑表单(名称/描述/性别/标签/文件) |
|
||||
| `VoiceMaterialCard` | 220行 (543-763) | 卡片视图 + 音频播放进度条 |
|
||||
| `VoiceMaterialRow` | 153行 (763-916) | 列表视图 + 音频播放进度条 |
|
||||
| **VoiceMaterialLibrary(主组件)** | **1051行 (916-1966)** | **页面主逻辑 + 渲染** |
|
||||
|
||||
### 主组件内部结构(1051行)
|
||||
|
||||
- **状态 & hooks**:约 230行
|
||||
- 数据查询:assets / libraries / tags / presetVoices
|
||||
- 视图状态:viewMode / searchText / filterGender / filterTagId
|
||||
- 播放状态:playingId / currentTime / audioRef / volume / pausedMaterial
|
||||
- 弹窗:uploadOpen / editingMaterial / ttsOpen
|
||||
- 批量:selectedIds / uploadProgress / batchCustomTag
|
||||
- TTS:ttsText / ttsVoiceId / ttsSpeed / ttsJobId / ttsStatus / ttsAudioUrl / ttsError
|
||||
- **播放控制函数**:约 100行(stopPlayback/startPlayback/handlePlay/handlePause/handleSeek/handleVolumeChange/toggleMute)
|
||||
- **数据操作函数**:约 60行(handleUpload/handleEdit/handleDelete)
|
||||
- **筛选 & 统计**:约 40行(filtered/tagCountMap)
|
||||
- **批量操作**:约 80行(handleToggleSelect/handleSelectAll/handleBatchDelete/handleBatchTag/handleBatchCustomTag)
|
||||
- **TTS 合成**:约 70行(handleTtsSynthesize/handleTtsSave)
|
||||
- **渲染 JSX**:约 473行
|
||||
|
||||
## 二、拆分策略
|
||||
|
||||
跟 GeneratePage 同样的思路:**按职责拆分,功能不变,纯结构优化。**
|
||||
|
||||
拆分四阶段,每个阶段独立 PR,逐步降低主文件复杂度。
|
||||
|
||||
---
|
||||
|
||||
## Phase 1:抽离常量、类型、工具函数
|
||||
|
||||
**目标:** 把纯数据定义和纯函数抽出去,主文件只留组件。
|
||||
|
||||
### 抽离内容
|
||||
|
||||
| 新文件 | 内容 | 行数估计 |
|
||||
|--------|------|---------|
|
||||
| `types.ts` | `VoiceGender` / `ViewMode` / `VoiceMaterial` / `VoiceAssetMetadata` 等类型定义 | ~40行 |
|
||||
| `constants.ts` | `MAX_CARD_TAGS` / `MAX_ROW_TAGS` / `TAG_VARIANTS` / `GENDER_OPTIONS` | ~25行 |
|
||||
| `utils/format.ts` | `formatDuration` / `formatFileSize` / `formatDate` / `genderLabel` / `genderIcon` / `genderClass` | ~35行 |
|
||||
| `utils/audio.ts` | `getAudioDuration`(纯工具,跟组件无关) | ~20行 |
|
||||
| `utils/mappers.ts` | `mapAssetToMaterial` / `buildMetadata`(数据映射) | ~35行 |
|
||||
|
||||
### 预期效果
|
||||
|
||||
- 主文件减少约 **155 行**(1966 → ~1811)
|
||||
- 工具函数可复用、可单测
|
||||
|
||||
---
|
||||
|
||||
## Phase 2:抽离已有内联子组件
|
||||
|
||||
**目标:** 把 4 个已经定义好的内联组件拆成独立文件,主文件只保留 VoiceMaterialLibrary 主组件。
|
||||
|
||||
### 抽离内容
|
||||
|
||||
| 新文件 | 原位置 | 行数 | 备注 |
|
||||
|--------|--------|------|------|
|
||||
| `components/TagSelector.tsx` | 197-345行 | ~148行 | 标签选择器 |
|
||||
| `components/MaterialForm.tsx` | 356-543行 | ~187行 | 上传/编辑表单 |
|
||||
| `components/VoiceMaterialCard.tsx` | 543-763行 | ~220行 | 卡片视图,带播放控制 |
|
||||
| `components/VoiceMaterialRow.tsx` | 763-916行 | ~153行 | 列表视图,带播放控制 |
|
||||
|
||||
### 播放状态共享方案
|
||||
|
||||
Card 和 Row 都有播放进度条(共 ~60行重复代码),但播放状态在主组件里。
|
||||
|
||||
**方案:** 播放状态提升到主组件,Card/Row 只接收 prop 并回调:
|
||||
- `isPlaying` / `currentTime` / `onPlay` / `onPause` / `onSeek`
|
||||
- 主组件统一管理 audioRef 和播放状态
|
||||
|
||||
这样 Card 和 Row 是纯展示组件,播放逻辑集中在主组件。
|
||||
|
||||
### 预期效果
|
||||
|
||||
- 主文件减少约 **700 行**(1811 → ~1111)
|
||||
- 主组件专注页面逻辑,子组件专注展示
|
||||
|
||||
---
|
||||
|
||||
## Phase 3:主组件拆分 — 业务逻辑抽 Hook
|
||||
|
||||
**目标:** 把主组件里的业务逻辑按领域拆成自定义 Hook,主组件只负责组装。
|
||||
|
||||
### 抽离 Hooks
|
||||
|
||||
| Hook 文件 | 职责 | 管理的状态 |
|
||||
|-----------|------|-----------|
|
||||
| `hooks/useVoiceMaterials.ts` | 素材列表查询 + 筛选 + 增删改 | assets / searchText / filterGender / filterTagId / uploadMutation / editMutation / deleteMutation |
|
||||
| `hooks/useAudioPlayer.ts` | 音频播放控制 | playingId / currentTime / volume / audioRef / pausedMaterial |
|
||||
| `hooks/useBatchOperations.ts` | 批量操作 | selectedIds / handleToggleSelect / handleSelectAll / handleBatchDelete / handleBatchTag |
|
||||
| `hooks/useTtsSynthesize.ts` | TTS 合成逻辑 | ttsText / ttsVoiceId / ttsSpeed / ttsJobId / ttsStatus / ttsAudioUrl / ttsError |
|
||||
|
||||
### 预期效果
|
||||
|
||||
- 主组件减少约 **500 行**(1111 → ~611)
|
||||
- 每个 Hook 职责单一,可独立测试
|
||||
- 主组件变成"组装器",可读性大幅提升
|
||||
|
||||
---
|
||||
|
||||
## Phase 4:UI 组件细化 — 工具栏 & 弹窗拆分
|
||||
|
||||
**目标:** 把主组件里的大块 JSX 拆成独立 UI 组件。
|
||||
|
||||
### 抽离 UI 组件
|
||||
|
||||
| 组件文件 | 内容 | 行数估计 |
|
||||
|----------|------|---------|
|
||||
| `components/LibraryToolbar.tsx` | 顶部工具栏:搜索框 + 性别筛选 + 视图切换 + 结果计数 | ~60行 |
|
||||
| `components/TagFilterBar.tsx` | 标签筛选药丸条(全部 + 各标签计数) | ~45行 |
|
||||
| `components/BatchActionBar.tsx` | 批量操作栏:全选 + 计数 + 批量打标签 + 批量删除 | ~70行 |
|
||||
| `components/UploadModal.tsx` | 上传弹窗(包裹 MaterialForm) | ~30行 |
|
||||
| `components/EditModal.tsx` | 编辑弹窗(包裹 MaterialForm) | ~30行 |
|
||||
| `components/TtsSynthesizeModal.tsx` | AI 配音合成弹窗:文本输入 + 音色选择 + 语速 + 合成结果预览 + 保存 | ~150行 |
|
||||
| `components/EmptyState.tsx` | 空状态展示 | ~20行 |
|
||||
| `components/VoiceMaterialGrid.tsx` | 卡片视图网格容器 | ~25行 |
|
||||
| `components/VoiceMaterialList.tsx` | 列表视图表格容器 | ~25行 |
|
||||
|
||||
### 预期效果
|
||||
|
||||
- 主组件 JSX 部分减少约 **400 行**(611 → ~211)
|
||||
- 每个 UI 组件职责清晰,方便后续迭代
|
||||
|
||||
---
|
||||
|
||||
## 三、拆分后文件结构
|
||||
|
||||
```
|
||||
apps/web/src/pages/voice-materials/
|
||||
├── VoiceMaterialLibrary.tsx # 主组件(~211行,组装器)
|
||||
├── voice-materials.css # 样式(保持不变,后续再拆)
|
||||
├── types.ts # 类型定义
|
||||
├── constants.ts # 常量
|
||||
├── utils/
|
||||
│ ├── format.ts # 格式化工具
|
||||
│ ├── audio.ts # 音频工具
|
||||
│ └── mappers.ts # 数据映射
|
||||
├── hooks/
|
||||
│ ├── useVoiceMaterials.ts # 素材数据 + 筛选 + 增删改
|
||||
│ ├── useAudioPlayer.ts # 音频播放控制
|
||||
│ ├── useBatchOperations.ts # 批量操作
|
||||
│ └── useTtsSynthesize.ts # TTS合成
|
||||
└── components/
|
||||
├── TagSelector.tsx # 标签选择器
|
||||
├── MaterialForm.tsx # 上传/编辑表单
|
||||
├── VoiceMaterialCard.tsx # 卡片视图
|
||||
├── VoiceMaterialRow.tsx # 列表视图
|
||||
├── VoiceMaterialGrid.tsx # 卡片网格容器
|
||||
├── VoiceMaterialList.tsx # 列表容器
|
||||
├── LibraryToolbar.tsx # 顶部工具栏
|
||||
├── TagFilterBar.tsx # 标签筛选条
|
||||
├── BatchActionBar.tsx # 批量操作栏
|
||||
├── UploadModal.tsx # 上传弹窗
|
||||
├── EditModal.tsx # 编辑弹窗
|
||||
├── TtsSynthesizeModal.tsx # TTS合成弹窗
|
||||
└── EmptyState.tsx # 空状态
|
||||
```
|
||||
|
||||
### 拆分前后对比
|
||||
|
||||
| 指标 | 拆分前 | 拆分后 | 变化 |
|
||||
|------|--------|--------|------|
|
||||
| 主文件行数 | 1966 | ~211 | **-89%** |
|
||||
| 文件数量 | 2 | 22 | +20 |
|
||||
| 最大文件 | 1966行 | ~220行 | -89% |
|
||||
| 可测试性 | 低 | 高 | 每个Hook/组件可单测 |
|
||||
|
||||
## 四、实施顺序 & 风险
|
||||
|
||||
1. **Phase 1**:最低风险,纯抽离,零逻辑变化
|
||||
2. **Phase 2**:低风险,组件本来就是独立定义的,只是挪位置
|
||||
3. **Phase 3**:中风险,Hook 拆分需要仔细梳理状态依赖
|
||||
4. **Phase 4**:低风险,JSX 拆分,纯结构调整
|
||||
|
||||
**每个 Phase 完成后提 PR,CI 全绿再合,跟 GeneratePage 同样节奏。**
|
||||
@@ -28,28 +28,6 @@ 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,25 +31,6 @@ 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,28 +43,6 @@ 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(
|
||||
|
||||
@@ -1,58 +0,0 @@
|
||||
{
|
||||
"": "https://docs.renovatebot.com/renovate-schema.json",
|
||||
"extends": ["config:recommended"],
|
||||
|
||||
"baseBranches": ["develop"],
|
||||
"labels": ["dependencies"],
|
||||
"assignees": ["xiaoxia"],
|
||||
|
||||
"prConcurrentLimit": 3,
|
||||
"prHourlyLimit": 3,
|
||||
|
||||
"schedule": ["after 2am before 6am on monday"],
|
||||
"timezone": "Asia/Shanghai",
|
||||
|
||||
"vulnerabilityAlerts": {
|
||||
"enabled": true,
|
||||
"labels": ["dependencies", "security"],
|
||||
"schedule": ["at any time"]
|
||||
},
|
||||
|
||||
"pip_requirements": {
|
||||
"fileMatch": [
|
||||
"(^|/)requirements\.txt$",
|
||||
"(^|/)requirements-base\.txt$",
|
||||
"(^|/)requirements-dev\.txt$",
|
||||
"(^|/)requirements-worker\.txt$"
|
||||
]
|
||||
},
|
||||
|
||||
"npm": {
|
||||
"fileMatch": [
|
||||
"(^|/)apps/web/package\.json$"
|
||||
]
|
||||
},
|
||||
|
||||
"packageRules": [
|
||||
{
|
||||
"matchDepTypes": ["dependencies"],
|
||||
"matchUpdateTypes": ["patch", "minor"],
|
||||
"groupName": "production deps (minor & patch)",
|
||||
"groupSlug": "prod-deps-minor-patch"
|
||||
},
|
||||
{
|
||||
"matchDepTypes": ["devDependencies"],
|
||||
"matchUpdateTypes": ["patch", "minor"],
|
||||
"groupName": "dev deps (minor & patch)",
|
||||
"groupSlug": "dev-deps-minor-patch"
|
||||
},
|
||||
{
|
||||
"matchUpdateTypes": ["major"],
|
||||
"labels": ["dependencies", "major-update"]
|
||||
}
|
||||
],
|
||||
|
||||
"rebaseWhen": "behind-base-branch",
|
||||
"semanticCommits": "auto",
|
||||
"semanticPrefix": "chore(deps): "
|
||||
}
|
||||
Executable → Regular
+92
-319
@@ -1,32 +1,17 @@
|
||||
#!/usr/bin/env python3
|
||||
"""
|
||||
ACR 镜像清理脚本(增强版)
|
||||
|
||||
清理策略:
|
||||
ACR 镜像清理脚本
|
||||
策略:
|
||||
- 版本tag (v*): 永久保留
|
||||
- 固定tag (latest, main, develop, master): 永久保留
|
||||
- 缓存镜像 (*-cache): 永久保留
|
||||
- 受保护tag (--protected-tags): 永久保留(如当前运行中镜像)
|
||||
- PR预览tag (pr-*):
|
||||
- --pr-sha模式:删除指定PR commit的镜像(PR关闭时触发)
|
||||
- cron模式:通过Gitea API检查PR状态,已关闭/合并的删除
|
||||
- PR预览tag (pr-*): 保留 N 天(默认7天)
|
||||
- 普通commit hash tag: 保留最近 N 个(默认20),老的删除
|
||||
|
||||
使用方式:
|
||||
# 预览(不实际删除)
|
||||
python3 acr_cleanup.py --dry-run
|
||||
|
||||
# 实际执行(cron模式)
|
||||
python3 acr_cleanup.py --execute
|
||||
|
||||
# 保留最近30个commit镜像
|
||||
python3 acr_cleanup.py --keep 30 --execute
|
||||
|
||||
# PR关闭时清理指定commit的PR镜像
|
||||
python3 acr_cleanup.py --pr-sha abc123def --execute
|
||||
|
||||
# 传入受保护tag列表(运行中镜像白名单)
|
||||
python3 acr_cleanup.py --protected-tags "sha1,sha2" --execute
|
||||
python3 acr_cleanup.py --dry-run # 预览,不实际删除
|
||||
python3 acr_cleanup.py --execute # 实际执行删除
|
||||
python3 acr_cleanup.py --keep 20 --execute # 保留最近20个
|
||||
"""
|
||||
|
||||
import argparse
|
||||
@@ -38,8 +23,7 @@ import urllib.error
|
||||
import urllib.request
|
||||
from datetime import datetime, timedelta, timezone
|
||||
|
||||
# ========== 配置 ==========
|
||||
|
||||
# 配置
|
||||
REGISTRY = os.environ.get("ACR_REGISTRY", "xiaoxia-registry.cn-hangzhou.cr.aliyuncs.com")
|
||||
AUTH_URL = "https://dockerauth.cn-hangzhou.aliyuncs.com/auth"
|
||||
SERVICE = os.environ.get("ACR_SERVICE", "registry.aliyuncs.com:cn-hangzhou:china:cri-fvec8o9q4mmxrkaa")
|
||||
@@ -47,11 +31,6 @@ NAMESPACE = os.environ.get("ACR_NAMESPACE", "xiaoxiakeji")
|
||||
USERNAME = os.environ.get("ACR_USERNAME", "")
|
||||
PASSWORD = os.environ.get("ACR_PASSWORD", "")
|
||||
|
||||
# Gitea配置(用于PR状态检查)
|
||||
GITEA_URL = os.environ.get("GITEA_URL", "https://git.xiaoxiajianji.com")
|
||||
GITEA_TOKEN = os.environ.get("GITEA_TOKEN", "")
|
||||
GITEA_REPO = os.environ.get("GITEA_REPO", "xiaoxia/xiaoxia-saas")
|
||||
|
||||
REPOS = [
|
||||
"xiaoxia-saas-api",
|
||||
"xiaoxia-saas-worker",
|
||||
@@ -70,9 +49,6 @@ ACCEPT_MANIFEST_OCI = "application/vnd.oci.image.manifest.v1+json"
|
||||
ACCEPT_MANIFEST_V2 = "application/vnd.docker.distribution.manifest.v2+json"
|
||||
|
||||
|
||||
# ========== Registry API ==========
|
||||
|
||||
|
||||
def get_token(repo, action="pull"):
|
||||
"""获取仓库访问token"""
|
||||
scope = "repository:" + NAMESPACE + "/" + repo + ":" + action
|
||||
@@ -105,19 +81,22 @@ def http_get_json(url, token, accept_header):
|
||||
|
||||
def get_manifest_info(repo, tag, token):
|
||||
"""
|
||||
获取tag的manifest信息。
|
||||
返回: {digest, created, media_type, error}
|
||||
获取tag的manifest信息,支持OCI index和普通manifest两种格式。
|
||||
返回: {digest, created, media_type}
|
||||
- digest: 顶层manifest的digest(用于删除)
|
||||
- created: 镜像创建时间
|
||||
"""
|
||||
url = "https://" + REGISTRY + "/v2/" + NAMESPACE + "/" + repo + "/manifests/" + tag
|
||||
result = {"digest": "", "created": "", "media_type": "", "error": ""}
|
||||
|
||||
# 先尝试 OCI index 格式
|
||||
# 先尝试 OCI index 格式(ACR多用这种)
|
||||
try:
|
||||
data, headers = http_get_json(url, token, ACCEPT_INDEX)
|
||||
top_digest = headers.get("Docker-Content-Digest", "")
|
||||
result["digest"] = top_digest
|
||||
result["media_type"] = data.get("mediaType", ACCEPT_INDEX)
|
||||
|
||||
# OCI index:找amd64的manifest,再取config blob
|
||||
manifests = data.get("manifests", [])
|
||||
amd64_manifest = None
|
||||
for m in manifests:
|
||||
@@ -125,6 +104,7 @@ def get_manifest_info(repo, tag, token):
|
||||
if arch == "amd64":
|
||||
amd64_manifest = m
|
||||
break
|
||||
# 没有amd64就用第一个
|
||||
if not amd64_manifest and manifests:
|
||||
amd64_manifest = manifests[0]
|
||||
|
||||
@@ -134,6 +114,7 @@ def get_manifest_info(repo, tag, token):
|
||||
try:
|
||||
inner_data, _ = http_get_json(inner_url, token, ACCEPT_MANIFEST_OCI)
|
||||
except Exception:
|
||||
# 退而求其次用v2格式
|
||||
inner_data, _ = http_get_json(inner_url, token, ACCEPT_MANIFEST_V2)
|
||||
|
||||
config_digest = inner_data.get("config", {}).get("digest", "")
|
||||
@@ -200,58 +181,6 @@ def delete_manifest(repo, digest, token):
|
||||
return False, str(e.code) + " " + e.read().decode()[:200]
|
||||
|
||||
|
||||
# ========== Gitea API ==========
|
||||
|
||||
|
||||
def gitea_get_open_prs():
|
||||
"""获取所有打开的PR编号列表"""
|
||||
if not GITEA_TOKEN:
|
||||
print(" 警告: 无GITEA_TOKEN,跳过PR状态检查")
|
||||
return None
|
||||
|
||||
open_prs = set()
|
||||
page = 1
|
||||
while True:
|
||||
url = GITEA_URL + "/api/v1/repos/" + GITEA_REPO + "/pulls?state=open&page=" + str(page) + "&limit=50"
|
||||
req = urllib.request.Request(url)
|
||||
req.add_header("Authorization", "token " + GITEA_TOKEN)
|
||||
try:
|
||||
with urllib.request.urlopen(req) as resp:
|
||||
data = json.loads(resp.read())
|
||||
if not data:
|
||||
break
|
||||
for pr in data:
|
||||
open_prs.add(pr.get("number", 0))
|
||||
if len(data) < 50:
|
||||
break
|
||||
page += 1
|
||||
except Exception as e:
|
||||
print(f" 警告: 获取Gitea PR列表失败: {e}")
|
||||
return None
|
||||
|
||||
return open_prs
|
||||
|
||||
|
||||
def gitea_get_pr_commits(pr_number):
|
||||
"""获取指定PR的所有commit sha"""
|
||||
if not GITEA_TOKEN:
|
||||
return []
|
||||
|
||||
url = GITEA_URL + "/api/v1/repos/" + GITEA_REPO + "/pulls/" + str(pr_number) + "/commits?limit=100"
|
||||
req = urllib.request.Request(url)
|
||||
req.add_header("Authorization", "token " + GITEA_TOKEN)
|
||||
try:
|
||||
with urllib.request.urlopen(req) as resp:
|
||||
data = json.loads(resp.read())
|
||||
return [c.get("sha", "") for c in data]
|
||||
except Exception as e:
|
||||
print(f" 警告: 获取PR #{pr_number} commits失败: {e}")
|
||||
return []
|
||||
|
||||
|
||||
# ========== 工具函数 ==========
|
||||
|
||||
|
||||
def parse_time(created_str):
|
||||
"""解析ISO时间字符串"""
|
||||
if not created_str:
|
||||
@@ -275,49 +204,12 @@ def is_fixed_tag(tag):
|
||||
|
||||
|
||||
def is_pr_tag(tag):
|
||||
"""判断是否是PR预览tag (pr-<sha>)"""
|
||||
"""判断是否是PR预览tag"""
|
||||
return tag.startswith("pr-")
|
||||
|
||||
|
||||
def extract_sha_from_pr_tag(tag):
|
||||
"""从pr-<sha> tag中提取sha"""
|
||||
if tag.startswith("pr-"):
|
||||
return tag[3:]
|
||||
return tag
|
||||
|
||||
|
||||
def is_in_protected_list(tag, protected_set):
|
||||
"""检查tag是否在受保护列表中"""
|
||||
if not protected_set:
|
||||
return False
|
||||
# 精确匹配
|
||||
if tag in protected_set:
|
||||
return True
|
||||
# 前缀匹配(commit hash可能是完整或短的)
|
||||
for p in protected_set:
|
||||
if tag.startswith(p) or p.startswith(tag):
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
# ========== 核心清理逻辑 ==========
|
||||
|
||||
|
||||
def cleanup_repo(repo, keep_count, dry_run, protected_tags, pr_sha=None, pr_open_set=None):
|
||||
"""
|
||||
清理单个仓库
|
||||
|
||||
Args:
|
||||
repo: 仓库名
|
||||
keep_count: 保留最近N个commit tag
|
||||
dry_run: 是否预览模式
|
||||
protected_tags: 受保护tag集合(白名单)
|
||||
pr_sha: 指定PR commit sha(PR关闭模式),None表示cron模式
|
||||
pr_open_set: 打开的PR编号集合(cron模式用)
|
||||
|
||||
Returns:
|
||||
(总tag数, 删除数)
|
||||
"""
|
||||
def cleanup_repo(repo, keep_count, pr_days, dry_run):
|
||||
"""清理单个仓库"""
|
||||
print("=" * 60)
|
||||
print("仓库:", repo)
|
||||
print("=" * 60)
|
||||
@@ -333,37 +225,6 @@ def cleanup_repo(repo, keep_count, dry_run, protected_tags, pr_sha=None, pr_open
|
||||
tags = get_tags(repo, token_pull)
|
||||
print(" 总tag数:", len(tags))
|
||||
|
||||
if not tags:
|
||||
print(" 无tag,跳过")
|
||||
return 0, 0
|
||||
|
||||
# ========== PR-SHA模式:只删除指定commit的PR镜像 ==========
|
||||
if pr_sha:
|
||||
pr_tags_to_del = [
|
||||
t
|
||||
for t in tags
|
||||
if t.startswith("pr-" + pr_sha) or t == "pr-" + pr_sha or pr_sha.startswith(extract_sha_from_pr_tag(t))
|
||||
]
|
||||
if not pr_tags_to_del:
|
||||
print(f" 未找到PR镜像: pr-{pr_sha[:12]}")
|
||||
return len(tags), 0
|
||||
|
||||
print(f" 找到 {len(pr_tags_to_del)} 个PR镜像待删除:")
|
||||
for t in pr_tags_to_del:
|
||||
print(f" - {t}")
|
||||
|
||||
to_delete = []
|
||||
for tag in pr_tags_to_del:
|
||||
info = get_manifest_info(repo, tag, token_pull)
|
||||
if info["digest"]:
|
||||
to_delete.append({"tag": tag, "digest": info["digest"], "created": info["created"]})
|
||||
else:
|
||||
print(f" 警告: {tag} 无法获取digest,跳过")
|
||||
|
||||
return _execute_delete(repo, to_delete, dry_run, len(tags))
|
||||
|
||||
# ========== Cron模式:全量清理 ==========
|
||||
|
||||
# 分类
|
||||
version_tags = []
|
||||
fixed_tags = []
|
||||
@@ -382,185 +243,122 @@ def cleanup_repo(repo, keep_count, dry_run, protected_tags, pr_sha=None, pr_open
|
||||
|
||||
print(" 版本tag (v*):", len(version_tags), "-> 永久保留")
|
||||
print(" 固定tag:", len(fixed_tags), "-> 永久保留")
|
||||
print(" PR预览tag (pr-*):", len(pr_tags_list), "-> 已关闭PR的删除")
|
||||
print(" PR预览tag (pr-*):", len(pr_tags_list), "-> 保留", pr_days, "天")
|
||||
print(" Commit hash tag:", len(commit_tags), "-> 保留最近", keep_count, "个")
|
||||
print(" 白名单tag:", len(protected_tags), "个")
|
||||
|
||||
# --- PR tag清理:检查PR状态 ---
|
||||
pr_to_delete = []
|
||||
if pr_tags_list:
|
||||
print()
|
||||
print(" 检查PR镜像状态...")
|
||||
|
||||
# 策略:有Gitea token则检查PR状态,否则按时间保留7天
|
||||
if pr_open_set is not None:
|
||||
# 通过Gitea API检查每个PR镜像对应的PR是否还开着
|
||||
# 注意:pr tag是pr-<sha>,sha可能属于某个PR
|
||||
# 简化策略:收集所有打开PR的commit sha,在白名单里的保留
|
||||
print(" 模式: Gitea PR状态检查")
|
||||
open_pr_shas = set()
|
||||
# 这里做了简化:因为每个PR都查commits太慢,我们用另一种方式
|
||||
# 对于PR tag,先尝试匹配PR编号(如果tag名里有编号),否则按时间
|
||||
# 实际pr-<sha>没法直接知道PR编号,所以降级为按时间+打开PR的head sha白名单
|
||||
open_head_shas = set()
|
||||
page = 1
|
||||
while True:
|
||||
url = GITEA_URL + "/api/v1/repos/" + GITEA_REPO + "/pulls?state=open&page=" + str(page) + "&limit=50"
|
||||
req = urllib.request.Request(url)
|
||||
req.add_header("Authorization", "token " + GITEA_TOKEN)
|
||||
try:
|
||||
with urllib.request.urlopen(req) as resp:
|
||||
data = json.loads(resp.read())
|
||||
if not data:
|
||||
break
|
||||
for pr in data:
|
||||
head_sha = pr.get("head", {}).get("sha", "")
|
||||
if head_sha:
|
||||
open_head_shas.add(head_sha)
|
||||
open_head_shas.add(head_sha[:7])
|
||||
open_head_shas.add(head_sha[:12])
|
||||
if len(data) < 50:
|
||||
break
|
||||
page += 1
|
||||
except Exception:
|
||||
break
|
||||
|
||||
deleted_count = 0
|
||||
for tag in pr_tags_list:
|
||||
sha = extract_sha_from_pr_tag(tag)
|
||||
# 检查是否是打开PR的head sha
|
||||
is_open_pr = False
|
||||
for ohs in open_head_shas:
|
||||
if sha.startswith(ohs) or ohs.startswith(sha):
|
||||
is_open_pr = True
|
||||
break
|
||||
if not is_open_pr:
|
||||
info = get_manifest_info(repo, tag, token_pull)
|
||||
if info["digest"]:
|
||||
pr_to_delete.append({"tag": tag, "digest": info["digest"], "created": info["created"]})
|
||||
deleted_count += 1
|
||||
print(f" 打开PR数: {len(open_head_shas)}个head sha")
|
||||
print(f" 将删除PR镜像: {deleted_count}个")
|
||||
else:
|
||||
# 无Gitea token,降级为按7天保留
|
||||
print(" 模式: 按时间保留7天(无Gitea token降级)")
|
||||
cutoff = datetime.now(timezone.utc) - timedelta(days=7)
|
||||
for tag in pr_tags_list:
|
||||
info = get_manifest_info(repo, tag, token_pull)
|
||||
created = parse_time(info["created"])
|
||||
if created < cutoff and info["digest"]:
|
||||
pr_to_delete.append({"tag": tag, "digest": info["digest"], "created": info["created"]})
|
||||
print(f" 将删除PR镜像: {len(pr_to_delete)}个")
|
||||
|
||||
# --- Commit tag清理:保留最近N个 ---
|
||||
# 获取所有commit tag的创建时间
|
||||
print()
|
||||
print(" 获取commit tag创建时间...")
|
||||
commit_tag_infos = []
|
||||
tag_info_list = []
|
||||
errors = 0
|
||||
for i, tag in enumerate(commit_tags):
|
||||
info = get_manifest_info(repo, tag, token_pull)
|
||||
if info["error"] or not info["digest"]:
|
||||
errors += 1
|
||||
commit_tag_infos.append({"tag": tag, "digest": info["digest"], "created": info["created"]})
|
||||
# 取不到信息的tag,放到最后(最旧处理),但标记一下
|
||||
tag_info_list.append({"tag": tag, "digest": info["digest"], "created": "", "error": info.get("error", "")})
|
||||
else:
|
||||
tag_info_list.append({"tag": tag, "digest": info["digest"], "created": info["created"], "error": ""})
|
||||
if (i + 1) % 20 == 0:
|
||||
print(" 已获取", i + 1, "/", len(commit_tags), "...")
|
||||
|
||||
if errors:
|
||||
print(" 注意:", errors, "个tag获取manifest失败")
|
||||
|
||||
# 按时间倒序排序
|
||||
commit_tag_infos.sort(key=lambda x: parse_time(x["created"]), reverse=True)
|
||||
# 按时间倒序排序(空时间放最后)
|
||||
tag_info_list.sort(key=lambda x: parse_time(x["created"]), reverse=True)
|
||||
|
||||
# 确定要删除的commit tag
|
||||
commit_to_delete = []
|
||||
if len(commit_tag_infos) > keep_count:
|
||||
commit_to_delete = commit_tag_infos[keep_count:]
|
||||
print(f" 保留前{keep_count}个commit tag,删除{len(commit_to_delete)}个")
|
||||
|
||||
# 白名单过滤:受保护的tag不删除
|
||||
if protected_tags:
|
||||
before = len(commit_to_delete)
|
||||
commit_to_delete = [t for t in commit_to_delete if not is_in_protected_list(t["tag"], protected_tags)]
|
||||
removed = before - len(commit_to_delete)
|
||||
to_delete = []
|
||||
if len(tag_info_list) > keep_count:
|
||||
to_delete = tag_info_list[keep_count:]
|
||||
print(" 保留前", keep_count, "个commit tag,删除", len(to_delete), "个")
|
||||
# 打印保留范围
|
||||
kept = tag_info_list[:keep_count]
|
||||
valid_kept = [t for t in kept if t["created"]]
|
||||
if valid_kept:
|
||||
print(" 最早保留:", valid_kept[-1]["tag"][:12], "(" + valid_kept[-1]["created"][:10] + ")")
|
||||
# 保护当前构建的tag(通过PROTECTED_TAG环境变量传入,如GITHUB_SHA)
|
||||
protected_tag = os.environ.get("PROTECTED_TAG", "").strip()
|
||||
if protected_tag:
|
||||
before = len(to_delete)
|
||||
to_delete = [t for t in to_delete if not t["tag"].startswith(protected_tag)]
|
||||
removed = before - len(to_delete)
|
||||
if removed > 0:
|
||||
print(f" 白名单保护: 跳过{removed}个运行中镜像")
|
||||
print(f" 保护当前构建tag: {protected_tag[:12]} (跳过{removed}个)")
|
||||
|
||||
# 过滤无digest的
|
||||
commit_to_delete = [t for t in commit_to_delete if t["digest"]]
|
||||
print(f" 可删除(有digest): {len(commit_to_delete)}个")
|
||||
to_del_valid = [t for t in to_delete if t["digest"]]
|
||||
print(" 可删除(有digest):", len(to_del_valid), "个")
|
||||
else:
|
||||
print(f" commit tag数量不足{keep_count}个,无需清理")
|
||||
print(" commit tag数量不足", keep_count, ",无需清理")
|
||||
|
||||
# --- 合并所有待删除项 ---
|
||||
all_to_delete = commit_to_delete + pr_to_delete
|
||||
# PR tag按时间清理
|
||||
pr_to_delete = []
|
||||
if pr_tags_list:
|
||||
cutoff = datetime.now(timezone.utc) - timedelta(days=pr_days)
|
||||
print()
|
||||
print(" 检查PR预览tag(超过", pr_days, "天删除)...")
|
||||
for tag in pr_tags_list:
|
||||
info = get_manifest_info(repo, tag, token_pull)
|
||||
created = parse_time(info["created"])
|
||||
if created < cutoff:
|
||||
pr_to_delete.append({"tag": tag, "digest": info["digest"], "created": info["created"]})
|
||||
print(" PR tag将删除:", len(pr_to_delete), "个")
|
||||
|
||||
# 再次过滤白名单(PR镜像也受白名单保护)
|
||||
if protected_tags:
|
||||
before = len(all_to_delete)
|
||||
all_to_delete = [t for t in all_to_delete if not is_in_protected_list(t["tag"], protected_tags)]
|
||||
removed = before - len(all_to_delete)
|
||||
if removed > 0:
|
||||
print(f" 白名单保护(PR镜像): 跳过{removed}个")
|
||||
all_to_delete = [t for t in to_delete if t["digest"]] + [t for t in pr_to_delete if t["digest"]]
|
||||
|
||||
return _execute_delete(repo, all_to_delete, dry_run, len(tags))
|
||||
|
||||
|
||||
def _execute_delete(repo, to_delete, dry_run, total_tags):
|
||||
"""执行删除操作"""
|
||||
if not to_delete:
|
||||
if not all_to_delete:
|
||||
print()
|
||||
print(" 无需删除任何tag")
|
||||
return total_tags, 0
|
||||
|
||||
# 按digest去重
|
||||
seen_digests = set()
|
||||
unique_delete = []
|
||||
for item in to_delete:
|
||||
if item["digest"] and item["digest"] not in seen_digests:
|
||||
seen_digests.add(item["digest"])
|
||||
unique_delete.append(item)
|
||||
return len(tags), 0
|
||||
|
||||
# 执行删除
|
||||
print()
|
||||
if dry_run:
|
||||
print(f" [DRY RUN] 将删除{len(unique_delete)}个manifest(预览模式)")
|
||||
for item in unique_delete[:5]:
|
||||
print(" [DRY RUN] 将删除", len(all_to_delete), "个tag(预览模式,不实际删除)")
|
||||
# 去重digest
|
||||
unique_digests = set(t["digest"] for t in all_to_delete if t["digest"])
|
||||
print(" 去重后唯一digest数:", len(unique_digests))
|
||||
for item in all_to_delete[:5]:
|
||||
created_str = item.get("created", "")[:10] or "未知"
|
||||
print(f" - {item['tag'][:30]} ({created_str})")
|
||||
if len(unique_delete) > 5:
|
||||
print(f" ... 还有{len(unique_delete) - 5}个")
|
||||
return total_tags, len(unique_delete)
|
||||
print(" -", item["tag"][:20], "(" + created_str + ")")
|
||||
if len(all_to_delete) > 5:
|
||||
print(" ... 还有", len(all_to_delete) - 5, "个")
|
||||
return len(tags), len(unique_digests)
|
||||
|
||||
token_delete = get_token(repo, "delete")
|
||||
deleted = 0
|
||||
failed = 0
|
||||
# 按digest去重,避免重复删除同一镜像
|
||||
seen_digests = set()
|
||||
unique_delete = []
|
||||
for item in all_to_delete:
|
||||
if item["digest"] and item["digest"] not in seen_digests:
|
||||
seen_digests.add(item["digest"])
|
||||
unique_delete.append(item)
|
||||
|
||||
print(f" 开始删除{len(unique_delete)}个唯一manifest...")
|
||||
print(" 开始删除", len(unique_delete), "个唯一manifest...")
|
||||
for item in unique_delete:
|
||||
success, result = delete_manifest(repo, item["digest"], token_delete)
|
||||
if success:
|
||||
deleted += 1
|
||||
print(f" 已删除: {item['tag'][:30]}")
|
||||
print(" 已删除:", item["tag"][:20])
|
||||
else:
|
||||
failed += 1
|
||||
print(f" 删除失败: {item['tag'][:30]} - {result}")
|
||||
print(" 删除失败:", item["tag"][:20], "-", result)
|
||||
|
||||
print()
|
||||
print(f" 删除完成: 成功{deleted}个,失败{failed}个")
|
||||
return total_tags, deleted
|
||||
|
||||
|
||||
# ========== 主函数 ==========
|
||||
print(" 删除完成: 成功", deleted, "个,失败", failed, "个")
|
||||
return len(tags), deleted
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(description="ACR镜像清理工具(增强版)")
|
||||
parser = argparse.ArgumentParser(description="ACR镜像清理工具")
|
||||
parser.add_argument("--keep", type=int, default=20, help="保留最近N个commit hash tag(默认20)")
|
||||
parser.add_argument("--pr-days", type=int, default=7, help="PR预览tag保留天数(默认7天)")
|
||||
parser.add_argument("--dry-run", action="store_true", help="预览模式,不实际删除")
|
||||
parser.add_argument("--execute", action="store_true", help="实际执行删除")
|
||||
parser.add_argument("--repo", type=str, default="", help="只清理指定仓库")
|
||||
parser.add_argument("--pr-sha", type=str, default="", help="PR关闭模式:删除指定commit sha的PR镜像")
|
||||
parser.add_argument("--protected-tags", type=str, default="", help="受保护tag列表,逗号分隔(运行中镜像白名单)")
|
||||
parser.add_argument("--skip-pr-check", action="store_true", help="跳过Gitea PR状态检查(纯按时间清理PR镜像)")
|
||||
args = parser.parse_args()
|
||||
|
||||
# 必须指定 --dry-run 或 --execute
|
||||
@@ -570,12 +368,13 @@ def main():
|
||||
print("示例:")
|
||||
print(" python3 acr_cleanup.py --dry-run # 预览清理效果")
|
||||
print(" python3 acr_cleanup.py --execute # 实际执行清理")
|
||||
print(" python3 acr_cleanup.py --pr-sha abc123 --execute # PR关闭时清理")
|
||||
print(" python3 acr_cleanup.py --keep 20 --execute # 保留最近20个")
|
||||
sys.exit(1)
|
||||
|
||||
# 凭证检查
|
||||
global USERNAME, PASSWORD
|
||||
if not USERNAME or not PASSWORD:
|
||||
# 尝试从docker config读取
|
||||
try:
|
||||
docker_config_path = os.path.expanduser("~/.docker/config.json")
|
||||
with open(docker_config_path) as f:
|
||||
@@ -592,39 +391,15 @@ def main():
|
||||
print("或确保已执行 docker login", REGISTRY)
|
||||
sys.exit(1)
|
||||
|
||||
# 解析受保护tag
|
||||
protected_tags = set()
|
||||
if args.protected_tags:
|
||||
protected_tags = set(t.strip() for t in args.protected_tags.split(",") if t.strip())
|
||||
|
||||
dry_run = args.dry_run or not args.execute
|
||||
mode = "预览模式" if dry_run else "执行模式"
|
||||
|
||||
print("=" * 60)
|
||||
print("ACR 镜像清理工具(增强版)-", mode)
|
||||
print("=" * 60)
|
||||
print("ACR 镜像清理工具 -", mode)
|
||||
print("Registry:", REGISTRY)
|
||||
print("Namespace:", NAMESPACE)
|
||||
if args.pr_sha:
|
||||
print("模式: PR关闭清理")
|
||||
print("PR commit SHA:", args.pr_sha[:12])
|
||||
else:
|
||||
print("模式: Cron全量清理")
|
||||
print("保留commit tag数:", args.keep)
|
||||
print("PR状态检查:", "关闭" if args.skip_pr_check else "开启")
|
||||
if protected_tags:
|
||||
print("白名单tag数:", len(protected_tags))
|
||||
print("保留commit tag数:", args.keep)
|
||||
print("PR预览保留天数:", args.pr_days)
|
||||
print()
|
||||
|
||||
# PR模式不需要查Gitea
|
||||
pr_open_set = None
|
||||
if not args.pr_sha and not args.skip_pr_check and GITEA_TOKEN:
|
||||
print("获取打开的PR列表...")
|
||||
pr_open_set = gitea_get_open_prs()
|
||||
if pr_open_set is not None:
|
||||
print(f" 打开的PR: {len(pr_open_set)}个")
|
||||
print()
|
||||
|
||||
repos_to_clean = REPOS
|
||||
if args.repo:
|
||||
repos_to_clean = [args.repo]
|
||||
@@ -632,13 +407,11 @@ def main():
|
||||
total_deleted = 0
|
||||
total_tags = 0
|
||||
for repo in repos_to_clean:
|
||||
count, deleted = cleanup_repo(
|
||||
repo, args.keep, dry_run, protected_tags, pr_sha=args.pr_sha, pr_open_set=pr_open_set
|
||||
)
|
||||
count, deleted = cleanup_repo(repo, args.keep, args.pr_days, dry_run)
|
||||
total_tags += count
|
||||
total_deleted += deleted
|
||||
print()
|
||||
|
||||
print()
|
||||
print("=" * 60)
|
||||
print("清理完成")
|
||||
print(" 总tag数:", total_tags)
|
||||
|
||||
@@ -1,363 +0,0 @@
|
||||
#!/bin/bash
|
||||
# ===========================================
|
||||
# 金丝雀发布脚本 - 分阶段灰度到全量
|
||||
# ===========================================
|
||||
# 在 CI Runner 上执行,通过 SSH 控制生产服务器执行灰度发布。
|
||||
# 流程:5%灰度 → 20%灰度 → 50%灰度 → 100%全量
|
||||
# 每阶段自动健康检查,失败自动回滚。
|
||||
#
|
||||
# 用法:
|
||||
# IMAGE_TAG=v0.1.130 ./scripts/ci/canary_release.sh
|
||||
#
|
||||
# 环境变量:
|
||||
# IMAGE_TAG - 新版本镜像标签 (必填)
|
||||
# CANARY_STAGES - 灰度阶段配置,格式: "百分比:等待秒数" 用逗号分隔
|
||||
# 默认: "5:600,20:900,50:1200"
|
||||
# PROD_API_URL - Production API 公网地址
|
||||
# PROD_WEB_URL - Production Web 公网地址
|
||||
# PRODUCTION_SSH_HOST - 生产服务器 SSH 地址
|
||||
# PRODUCTION_SSH_USER - SSH 用户名
|
||||
# PRODUCTION_SSH_PORT - SSH 端口
|
||||
# PRODUCTION_SSH_KEY - SSH 私钥内容
|
||||
# ACR_USERNAME - 容器镜像仓库用户名
|
||||
# ACR_PASSWORD - 容器镜像仓库密码
|
||||
# CI_NOTIFY_WEBHOOK - 通知 Webhook
|
||||
# SKIP_ROLLBACK - 失败时不自动回滚 (调试用)
|
||||
|
||||
set -euo pipefail
|
||||
|
||||
SCRIPT_DIR="$(CDPATH= cd -- "$(dirname -- "$0")" && pwd)"
|
||||
REPO_ROOT="$(cd "$SCRIPT_DIR/../.." && pwd)"
|
||||
|
||||
# 配置
|
||||
IMAGE_TAG="${IMAGE_TAG:-}"
|
||||
CANARY_STAGES="${CANARY_STAGES:-5:600,20:900,50:1200}"
|
||||
PROD_API_URL="${PROD_API_URL:-https://api.xiaoxiajianji.com}"
|
||||
PROD_WEB_URL="${PROD_WEB_URL:-https://saas.xiaoxiajianji.com}"
|
||||
PRODUCTION_SSH_HOST="${PRODUCTION_SSH_HOST:-47.98.113.167}"
|
||||
PRODUCTION_SSH_USER="${PRODUCTION_SSH_USER:-root}"
|
||||
PRODUCTION_SSH_PORT="${PRODUCTION_SSH_PORT:-22222}"
|
||||
# gray_deploy.sh 的镜像命名格式是 ${REGISTRY}-component:tag
|
||||
# 需要与 ACR 镜像名 xiaoxia-registry.../xiaoxiakeji/xiaoxia-saas-api:tag 匹配
|
||||
GRAY_REGISTRY="${GRAY_REGISTRY:-xiaoxia-registry.cn-hangzhou.cr.aliyuncs.com/xiaoxiakeji/xiaoxia-saas}"
|
||||
ACR_REGISTRY_HOST="${ACR_REGISTRY_HOST:-xiaoxia-registry.cn-hangzhou.cr.aliyuncs.com}"
|
||||
ACR_USERNAME="${ACR_USERNAME:-}"
|
||||
ACR_PASSWORD="${ACR_PASSWORD:-}"
|
||||
SKIP_ROLLBACK="${SKIP_ROLLBACK:-false}"
|
||||
|
||||
if [[ -z "$IMAGE_TAG" ]]; then
|
||||
echo "ERROR: IMAGE_TAG is required"
|
||||
exit 1
|
||||
fi
|
||||
|
||||
# 颜色
|
||||
RED='\033[0;31m'
|
||||
GREEN='\033[0;32m'
|
||||
YELLOW='\033[1;33m'
|
||||
BLUE='\033[0;34m'
|
||||
NC='\033[0m'
|
||||
|
||||
log_info() { echo -e "${GREEN}[INFO]${NC} $1"; }
|
||||
log_warn() { echo -e "${YELLOW}[WARN]${NC} $1"; }
|
||||
log_error() { echo -e "${RED}[ERROR]${NC} $1"; }
|
||||
log_step() { echo -e "${BLUE}[STEP]${NC} $1"; }
|
||||
|
||||
# ===========================================
|
||||
# SSH 配置
|
||||
# ===========================================
|
||||
SSH_KEY_PATH=""
|
||||
|
||||
setup_ssh() {
|
||||
if [ -f /root/.ssh/xiaoxia_runtime_builder ]; then
|
||||
SSH_KEY_PATH="/root/.ssh/xiaoxia_runtime_builder"
|
||||
elif [ -f "$HOME/.ssh/xiaoxia_runtime_builder" ]; then
|
||||
SSH_KEY_PATH="$HOME/.ssh/xiaoxia_runtime_builder"
|
||||
elif [ -n "${PRODUCTION_SSH_KEY:-}" ]; then
|
||||
SSH_KEY_PATH="$HOME/.ssh/canary_deploy_key"
|
||||
mkdir -p "$HOME/.ssh"
|
||||
printf '%s\n' "$PRODUCTION_SSH_KEY" > "$SSH_KEY_PATH"
|
||||
chmod 600 "$SSH_KEY_PATH"
|
||||
else
|
||||
log_error "没有可用的 SSH 密钥"
|
||||
return 1
|
||||
fi
|
||||
|
||||
ssh-keyscan -p "$PRODUCTION_SSH_PORT" -H "$PRODUCTION_SSH_HOST" >> ~/.ssh/known_hosts 2>/dev/null || true
|
||||
log_info "SSH 已配置: ${PRODUCTION_SSH_USER}@${PRODUCTION_SSH_HOST}:${PRODUCTION_SSH_PORT}"
|
||||
}
|
||||
|
||||
run_ssh() {
|
||||
local cmd="$1"
|
||||
ssh -p "$PRODUCTION_SSH_PORT" -i "$SSH_KEY_PATH" -o StrictHostKeyChecking=no \
|
||||
"${PRODUCTION_SSH_USER}@${PRODUCTION_SSH_HOST}" "$cmd"
|
||||
}
|
||||
|
||||
# ===========================================
|
||||
# 上传脚本 + Docker登录
|
||||
# ===========================================
|
||||
prepare_server() {
|
||||
log_step "准备生产服务器环境"
|
||||
|
||||
# 创建临时目录
|
||||
run_ssh "mkdir -p /tmp/canary-release"
|
||||
|
||||
# 上传 gray_deploy.sh
|
||||
local gray_script="$REPO_ROOT/scripts/gray_deploy.sh"
|
||||
if [[ -f "$gray_script" ]]; then
|
||||
cat "$gray_script" | run_ssh "cat > /tmp/canary-release/gray_deploy.sh && chmod +x /tmp/canary-release/gray_deploy.sh"
|
||||
log_info " gray_deploy.sh 已上传"
|
||||
else
|
||||
log_error "找不到 gray_deploy.sh: $gray_script"
|
||||
return 1
|
||||
fi
|
||||
|
||||
# 上传 rollback_gray.sh
|
||||
local rollback_script="$REPO_ROOT/scripts/rollback_gray.sh"
|
||||
if [[ -f "$rollback_script" ]]; then
|
||||
cat "$rollback_script" | run_ssh "cat > /tmp/canary-release/rollback_gray.sh && chmod +x /tmp/canary-release/rollback_gray.sh"
|
||||
log_info " rollback_gray.sh 已上传"
|
||||
else
|
||||
log_warn "找不到 rollback_gray.sh"
|
||||
fi
|
||||
|
||||
# 上传 ci_production_deploy.sh
|
||||
local prod_deploy="$REPO_ROOT/scripts/ci_production_deploy.sh"
|
||||
if [[ -f "$prod_deploy" ]]; then
|
||||
cat "$prod_deploy" | run_ssh "cat > /tmp/canary-release/ci_production_deploy.sh && chmod +x /tmp/canary-release/ci_production_deploy.sh"
|
||||
log_info " ci_production_deploy.sh 已上传"
|
||||
else
|
||||
log_error "找不到 ci_production_deploy.sh: $prod_deploy"
|
||||
return 1
|
||||
fi
|
||||
|
||||
# Docker 登录到 ACR
|
||||
if [[ -n "$ACR_USERNAME" && -n "$ACR_PASSWORD" ]]; then
|
||||
log_info " Docker 登录到 ACR..."
|
||||
run_ssh "docker login '$ACR_REGISTRY_HOST' -u '$ACR_USERNAME' -p '$ACR_PASSWORD' 2>/dev/null" || \
|
||||
log_warn " Docker login 失败(可能已有凭证),将尝试直接 pull"
|
||||
fi
|
||||
|
||||
log_info "✅ 服务器环境准备完成"
|
||||
}
|
||||
|
||||
# ===========================================
|
||||
# 健康检查(公网访问)
|
||||
# ===========================================
|
||||
health_check() {
|
||||
local stage_name="$1"
|
||||
local timeout="${2:-120}"
|
||||
local interval=5
|
||||
local elapsed=0
|
||||
|
||||
log_step "健康检查 - $stage_name (超时 ${timeout}s)"
|
||||
|
||||
while [ $elapsed -lt $timeout ]; do
|
||||
local api_ok=false
|
||||
local web_ok=false
|
||||
|
||||
# 检查 API
|
||||
local api_code=$(curl -s -o /dev/null -w "%{http_code}" \
|
||||
--connect-timeout 5 --max-time 10 \
|
||||
"${PROD_API_URL}/health" 2>/dev/null || echo "000")
|
||||
if [[ "$api_code" == "200" ]]; then
|
||||
api_ok=true
|
||||
fi
|
||||
|
||||
# 检查 Web
|
||||
local web_code=$(curl -s -o /dev/null -w "%{http_code}" \
|
||||
--connect-timeout 5 --max-time 10 \
|
||||
"$PROD_WEB_URL" 2>/dev/null || echo "000")
|
||||
if [[ "$web_code" == "200" || "$web_code" == "301" || "$web_code" == "302" ]]; then
|
||||
web_ok=true
|
||||
fi
|
||||
|
||||
if $api_ok && $web_ok; then
|
||||
log_info "✅ 健康检查通过 (API=$api_code, Web=$web_code)"
|
||||
return 0
|
||||
fi
|
||||
|
||||
log_warn " 等待中... API=$api_code, Web=$web_code (${elapsed}s/${timeout}s)"
|
||||
sleep $interval
|
||||
elapsed=$((elapsed + interval))
|
||||
done
|
||||
|
||||
log_error "❌ 健康检查超时"
|
||||
return 1
|
||||
}
|
||||
|
||||
# ===========================================
|
||||
# 灰度发布
|
||||
# ===========================================
|
||||
gray_deploy() {
|
||||
local pct="$1"
|
||||
log_step "灰度发布 ${pct}% - $IMAGE_TAG"
|
||||
|
||||
run_ssh "cd /tmp/canary-release && \
|
||||
REGISTRY='$GRAY_REGISTRY' \
|
||||
./gray_deploy.sh '$IMAGE_TAG' '$pct'"
|
||||
}
|
||||
|
||||
# ===========================================
|
||||
# 全量部署
|
||||
# ===========================================
|
||||
full_deploy() {
|
||||
log_step "全量部署 - $IMAGE_TAG"
|
||||
|
||||
run_ssh "cd /tmp/canary-release && \
|
||||
IMAGE_TAG='$IMAGE_TAG' \
|
||||
ACR_USERNAME='$ACR_USERNAME' \
|
||||
ACR_PASSWORD='$ACR_PASSWORD' \
|
||||
sh ./ci_production_deploy.sh"
|
||||
}
|
||||
|
||||
# ===========================================
|
||||
# 灰度回滚
|
||||
# ===========================================
|
||||
rollback_gray() {
|
||||
log_error "执行灰度回滚..."
|
||||
if [[ "$SKIP_ROLLBACK" == "true" ]]; then
|
||||
log_warn "SKIP_ROLLBACK=true,跳过回滚"
|
||||
return
|
||||
fi
|
||||
|
||||
if run_ssh "test -f /tmp/canary-release/rollback_gray.sh"; then
|
||||
run_ssh "cd /tmp/canary-release && ./rollback_gray.sh" || \
|
||||
log_error "回滚脚本执行失败,请手动处理"
|
||||
else
|
||||
# 内联回滚逻辑
|
||||
log_warn "使用内联回滚逻辑"
|
||||
run_ssh '
|
||||
NGINX_CONF="/etc/nginx/sites-enabled/00-xiaoxia-saas"
|
||||
LATEST_BAK=$(ls -t "${NGINX_CONF}".bak.gray.* 2>/dev/null | head -1 || true)
|
||||
if [[ -n "$LATEST_BAK" ]]; then
|
||||
cp "$LATEST_BAK" "$NGINX_CONF"
|
||||
else
|
||||
sed -i "s|proxy_pass http://saas_api_backend|proxy_pass http://127.0.0.1:8001|g" "$NGINX_CONF"
|
||||
sed -i "s|proxy_pass http://saas_web_backend/|proxy_pass http://127.0.0.1:3002/|g" "$NGINX_CONF"
|
||||
fi
|
||||
nginx -t && nginx -s reload
|
||||
docker rm -f xiaoxia-api-canary xiaoxia-web-canary 2>/dev/null || true
|
||||
' || log_error "回滚失败,请手动处理"
|
||||
fi
|
||||
}
|
||||
|
||||
# ===========================================
|
||||
# 通知
|
||||
# ===========================================
|
||||
notify_status() {
|
||||
local status="$1"
|
||||
local message="$2"
|
||||
if [ -n "${CI_NOTIFY_WEBHOOK:-}" ]; then
|
||||
NOTIFY_MODE="$status" JOB_NAME="Canary Release - $message" \
|
||||
python3 "$REPO_ROOT/scripts/ci_notify.py" 2>/dev/null || true
|
||||
fi
|
||||
}
|
||||
|
||||
# ===========================================
|
||||
# 清理
|
||||
# ===========================================
|
||||
cleanup() {
|
||||
log_step "清理生产服务器临时文件"
|
||||
run_ssh "rm -rf /tmp/canary-release" 2>/dev/null || true
|
||||
}
|
||||
|
||||
# ===========================================
|
||||
# 主流程
|
||||
# ===========================================
|
||||
main() {
|
||||
echo "==========================================="
|
||||
echo " 🐦 金丝雀发布"
|
||||
echo " 版本: $IMAGE_TAG"
|
||||
echo " 阶段: $CANARY_STAGES"
|
||||
echo "==========================================="
|
||||
echo ""
|
||||
|
||||
setup_ssh
|
||||
prepare_server
|
||||
trap cleanup EXIT
|
||||
|
||||
# 解析灰度阶段
|
||||
IFS=',' read -ra STAGES <<< "$CANARY_STAGES"
|
||||
local total_stages=${#STAGES[@]}
|
||||
local current_stage=0
|
||||
|
||||
# 逐阶段灰度
|
||||
for stage in "${STAGES[@]}"; do
|
||||
current_stage=$((current_stage + 1))
|
||||
local pct=$(echo "$stage" | cut -d: -f1)
|
||||
local wait_time=$(echo "$stage" | cut -d: -f2)
|
||||
|
||||
echo ""
|
||||
echo "--- 阶段 $current_stage/$total_stages: ${pct}% 灰度 ---"
|
||||
|
||||
# 执行灰度发布
|
||||
if ! gray_deploy "$pct"; then
|
||||
log_error "灰度发布 ${pct}% 失败"
|
||||
rollback_gray
|
||||
notify_status "failure" "Stage ${pct}% Deploy Failed"
|
||||
exit 1
|
||||
fi
|
||||
|
||||
# 健康检查
|
||||
if ! health_check "${pct}%灰度"; then
|
||||
log_error "${pct}%灰度健康检查失败"
|
||||
rollback_gray
|
||||
notify_status "failure" "Stage ${pct}% Health Check Failed"
|
||||
exit 1
|
||||
fi
|
||||
|
||||
# 观察期
|
||||
log_info "⏳ 观察期 ${wait_time}s,监控流量稳定性..."
|
||||
local waited=0
|
||||
local check_interval=60
|
||||
while [ $waited -lt $wait_time ]; do
|
||||
sleep $check_interval
|
||||
waited=$((waited + check_interval))
|
||||
# 每隔一段时间做一次快速健康检查
|
||||
local api_code=$(curl -s -o /dev/null -w "%{http_code}" \
|
||||
--connect-timeout 5 --max-time 10 \
|
||||
"${PROD_API_URL}/health" 2>/dev/null || echo "000")
|
||||
if [[ "$api_code" != "200" ]]; then
|
||||
log_error "❌ 观察期内 API 异常 (HTTP $api_code),触发回滚"
|
||||
rollback_gray
|
||||
notify_status "failure" "Stage ${pct}% Watch Period Failed"
|
||||
exit 1
|
||||
fi
|
||||
log_info " 观察中... ${waited}s/${wait_time}s (API=$api_code)"
|
||||
done
|
||||
|
||||
log_info "✅ ${pct}%灰度阶段完成,稳定运行 ${wait_time}s"
|
||||
done
|
||||
|
||||
# 全量部署
|
||||
echo ""
|
||||
echo "--- 最终阶段: 100% 全量部署 ---"
|
||||
|
||||
if ! full_deploy; then
|
||||
log_error "全量部署失败"
|
||||
notify_status "failure" "Full Deploy Failed"
|
||||
exit 1
|
||||
fi
|
||||
|
||||
# 最终健康检查
|
||||
if ! health_check "全量部署" "180"; then
|
||||
log_error "全量部署后健康检查失败"
|
||||
notify_status "failure" "Full Deploy Health Check Failed"
|
||||
exit 1
|
||||
fi
|
||||
|
||||
# 清理 canary 容器
|
||||
log_step "清理 Canary 容器"
|
||||
run_ssh "docker rm -f xiaoxia-api-canary xiaoxia-web-canary 2>/dev/null || true" || true
|
||||
|
||||
echo ""
|
||||
echo "==========================================="
|
||||
echo " ✅ 金丝雀发布完成"
|
||||
echo " 版本: $IMAGE_TAG"
|
||||
echo " 状态: 100%全量运行"
|
||||
echo "==========================================="
|
||||
|
||||
notify_status "success" "$IMAGE_TAG Fully Deployed"
|
||||
}
|
||||
|
||||
main "$@"
|
||||
@@ -1,90 +0,0 @@
|
||||
"""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
|
||||
+339
-91
@@ -1,111 +1,359 @@
|
||||
"""BGM混音单元测试 - 配置解析等纯逻辑."""
|
||||
"""BGM 混音单元测试.
|
||||
|
||||
from __future__ import annotations
|
||||
测试:
|
||||
- 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"))
|
||||
|
||||
import pytest
|
||||
from video_processing.bgm_mixer import BGMConfig
|
||||
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 ──────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestBGMConfigDefaults:
|
||||
"""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 配置解析测试。"""
|
||||
|
||||
def test_default_values(self):
|
||||
"""默认值正确."""
|
||||
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
|
||||
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
|
||||
|
||||
|
||||
class TestBGMConfigFromConfigDict:
|
||||
"""BGMConfig.from_config_dict 解析测试."""
|
||||
# ── 预设 BGM 库测试 ─────────────────────────────────────────────────────────
|
||||
|
||||
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
|
||||
|
||||
def test_custom_volume(self):
|
||||
"""自定义音量."""
|
||||
config = BGMConfig.from_config_dict("/a.mp3", {"volume": 0.5})
|
||||
assert config.volume == 0.5
|
||||
class TestPresetBGM:
|
||||
"""预设 BGM 库查询测试。"""
|
||||
|
||||
def test_fade_in_out(self):
|
||||
"""淡入淡出."""
|
||||
config = BGMConfig.from_config_dict(
|
||||
"/a.mp3",
|
||||
{
|
||||
"fade_in": 2.0,
|
||||
"fade_out": 3.0,
|
||||
},
|
||||
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,
|
||||
)
|
||||
assert config.fade_in == 2.0
|
||||
assert config.fade_out == 3.0
|
||||
result = prepare_bgm_track(ctx, bgm, target_duration=4.0)
|
||||
|
||||
def test_loop_disabled(self):
|
||||
"""禁用循环."""
|
||||
config = BGMConfig.from_config_dict("/a.mp3", {"loop_enabled": False})
|
||||
assert config.loop_enabled is False
|
||||
assert result.exists()
|
||||
assert result.stat().st_size > 0
|
||||
|
||||
def test_sidechain_enabled(self):
|
||||
"""启用人声闪避."""
|
||||
config = BGMConfig.from_config_dict("/a.mp3", {"sidechain_enabled": True})
|
||||
assert config.sidechain_enabled is True
|
||||
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)
|
||||
|
||||
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 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,
|
||||
)
|
||||
assert config.sidechain_ratio == 0.5
|
||||
assert config.sidechain_attack == 0.05
|
||||
assert config.sidechain_release == 0.8
|
||||
assert config.sidechain_threshold == -30.0
|
||||
result = mix_bgm_with_main(ctx, Path(main_audio_path), bgm, target_duration=8.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"
|
||||
assert result.exists()
|
||||
assert result.stat().st_size > 0
|
||||
|
||||
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,
|
||||
},
|
||||
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,
|
||||
)
|
||||
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
|
||||
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(
|
||||
{
|
||||
"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,
|
||||
}
|
||||
}
|
||||
)
|
||||
|
||||
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_bgm_disabled_by_default(self):
|
||||
"""默认 BGM 是关闭的。"""
|
||||
from packages.domain.config_schemas import normalize_plan_config
|
||||
|
||||
config = normalize_plan_config({})
|
||||
assert config["bgm"]["enabled"] is False
|
||||
|
||||
+68
-127
@@ -1,164 +1,105 @@
|
||||
"""BGM工具函数单元测试。"""
|
||||
"""BGM 配置工具函数单元测试."""
|
||||
|
||||
from __future__ import annotations
|
||||
import pytest
|
||||
|
||||
from domain.bgm_utils import merge_bgm_config
|
||||
from packages.domain.bgm_utils import merge_bgm_config
|
||||
|
||||
|
||||
class TestMergeBgmConfigBothEmpty:
|
||||
"""两边都为空的情况。"""
|
||||
class TestMergeBgmConfig:
|
||||
"""merge_bgm_config 测试"""
|
||||
|
||||
def test_both_empty(self):
|
||||
result = merge_bgm_config({}, {})
|
||||
assert result == {}
|
||||
# 确保返回的是新字典,不是同一个引用
|
||||
assert result is not {}
|
||||
|
||||
def test_user_none_returns_template_copy(self):
|
||||
"""用户传 None 视为空配置,返回模板副本。"""
|
||||
result = merge_bgm_config({}, None)
|
||||
assert result == {}
|
||||
|
||||
|
||||
class TestMergeBgmConfigOnlyTemplate:
|
||||
"""只有模板配置。"""
|
||||
|
||||
def test_only_template_returns_copy(self):
|
||||
template = {"enabled": True, "volume": 0.5, "track": "default.mp3"}
|
||||
def test_user_bgm_empty_returns_template_copy(self):
|
||||
"""用户配置为空时,返回模板配置的拷贝"""
|
||||
template = {"enabled": True, "volume": 0.5, "asset_id": "tpl_123"}
|
||||
result = merge_bgm_config(template, {})
|
||||
assert result == template
|
||||
# 确保是副本,不是同一引用
|
||||
result["volume"] = 0.9
|
||||
assert template["volume"] == 0.5
|
||||
assert result is not template
|
||||
|
||||
def test_only_template_with_none_user(self):
|
||||
def test_user_bgm_none_returns_template_copy(self):
|
||||
"""用户配置为 None 时,返回模板配置的拷贝"""
|
||||
template = {"enabled": True, "volume": 0.5}
|
||||
result = merge_bgm_config(template, None)
|
||||
result = merge_bgm_config(template, None) # type: ignore
|
||||
assert result == template
|
||||
|
||||
|
||||
class TestMergeBgmConfigOnlyUser:
|
||||
"""只有用户配置。"""
|
||||
|
||||
def test_only_user_returns_copy(self):
|
||||
user = {"enabled": False, "volume": 0.8}
|
||||
def test_template_bgm_empty_returns_user_copy(self):
|
||||
"""模板配置为空时,返回用户配置的拷贝"""
|
||||
user = {"enabled": False, "volume": 0.8, "asset_id": "user_456"}
|
||||
result = merge_bgm_config({}, user)
|
||||
assert result == user
|
||||
# 确保是副本
|
||||
result["volume"] = 0.1
|
||||
assert user["volume"] == 0.8
|
||||
assert result is not user
|
||||
|
||||
def test_only_user_with_none_template(self):
|
||||
user = {"enabled": False}
|
||||
result = merge_bgm_config(None, user)
|
||||
def test_template_bgm_none_returns_user_copy(self):
|
||||
"""模板配置为 None 时,返回用户配置的拷贝"""
|
||||
user = {"enabled": False, "volume": 0.8}
|
||||
result = merge_bgm_config(None, user) # type: ignore
|
||||
assert result == user
|
||||
|
||||
|
||||
class TestMergeBgmConfigBasicOverride:
|
||||
"""用户配置覆盖模板配置。"""
|
||||
|
||||
def test_volume_override(self):
|
||||
template = {"enabled": True, "volume": 0.5}
|
||||
user = {"volume": 0.8}
|
||||
def test_user_fields_override_template(self):
|
||||
"""用户显式指定的字段覆盖模板对应字段"""
|
||||
template = {
|
||||
"enabled": True,
|
||||
"volume": 0.5,
|
||||
"asset_id": "tpl_123",
|
||||
"fade_in": 1.0,
|
||||
}
|
||||
user = {
|
||||
"volume": 0.8,
|
||||
"asset_id": "user_456",
|
||||
}
|
||||
result = merge_bgm_config(template, user)
|
||||
assert result["volume"] == 0.8
|
||||
assert result["enabled"] is True # 用户没传,保留模板
|
||||
assert result["asset_id"] == "user_456"
|
||||
assert result["fade_in"] == 1.0 # 模板值保留
|
||||
|
||||
def test_track_override(self):
|
||||
template = {"track": "default.mp3", "volume": 0.5}
|
||||
user = {"track": "custom.mp3"}
|
||||
result = merge_bgm_config(template, user)
|
||||
assert result["track"] == "custom.mp3"
|
||||
assert result["volume"] == 0.5
|
||||
|
||||
def test_multiple_fields_override(self):
|
||||
template = {"enabled": True, "volume": 0.5, "track": "a.mp3", "fade_in": 2}
|
||||
user = {"volume": 0.9, "track": "b.mp3"}
|
||||
result = merge_bgm_config(template, user)
|
||||
assert result["volume"] == 0.9
|
||||
assert result["track"] == "b.mp3"
|
||||
assert result["fade_in"] == 2
|
||||
assert result["enabled"] is True
|
||||
|
||||
|
||||
class TestMergeBgmConfigEnabledSpecialHandling:
|
||||
"""enabled 字段的特殊处理:用户没传就保留模板的。"""
|
||||
|
||||
def test_user_does_not_pass_enabled_keeps_template_true(self):
|
||||
def test_enabled_not_in_user_preserves_template_enabled(self):
|
||||
"""enabled 特殊处理:用户没传 enabled 时保留模板的 enabled 值"""
|
||||
template = {"enabled": True, "volume": 0.5}
|
||||
user = {"volume": 0.8}
|
||||
user = {"volume": 0.8} # 没传 enabled
|
||||
result = merge_bgm_config(template, user)
|
||||
assert result["enabled"] is True
|
||||
assert result["enabled"] is True # 保留模板的
|
||||
assert result["volume"] == 0.8 # 用户指定的覆盖
|
||||
|
||||
def test_user_does_not_pass_enabled_keeps_template_false(self):
|
||||
template = {"enabled": False, "volume": 0.5}
|
||||
user = {"volume": 0.8}
|
||||
def test_enabled_in_user_overrides_template(self):
|
||||
"""用户传了 enabled 时覆盖模板的 enabled"""
|
||||
template = {"enabled": True, "volume": 0.5}
|
||||
user = {"enabled": False, "volume": 0.8}
|
||||
result = merge_bgm_config(template, user)
|
||||
assert result["enabled"] is False
|
||||
assert result["volume"] == 0.8
|
||||
|
||||
def test_user_explicitly_sets_enabled_true(self):
|
||||
template = {"enabled": False, "volume": 0.5}
|
||||
user = {"enabled": True}
|
||||
def test_user_adds_new_fields(self):
|
||||
"""用户配置中的新字段会被添加到结果中"""
|
||||
template = {"enabled": True, "volume": 0.5}
|
||||
user = {"sidechain_enabled": True, "sidechain_ratio": 0.6}
|
||||
result = merge_bgm_config(template, user)
|
||||
assert result["enabled"] is True
|
||||
|
||||
def test_user_explicitly_sets_enabled_false(self):
|
||||
template = {"enabled": True, "volume": 0.5}
|
||||
user = {"enabled": False}
|
||||
result = merge_bgm_config(template, user)
|
||||
assert result["enabled"] is False
|
||||
|
||||
def test_template_no_enabled_user_no_enabled(self):
|
||||
template = {"volume": 0.5}
|
||||
user = {"track": "a.mp3"}
|
||||
result = merge_bgm_config(template, user)
|
||||
assert "enabled" not in result
|
||||
|
||||
def test_template_no_enabled_user_has_enabled(self):
|
||||
template = {"volume": 0.5}
|
||||
user = {"enabled": True}
|
||||
result = merge_bgm_config(template, user)
|
||||
assert result["enabled"] is True
|
||||
|
||||
def test_user_sets_enabled_none_explicitly(self):
|
||||
"""用户显式传 None 也视为传了,会覆盖模板。"""
|
||||
template = {"enabled": True, "volume": 0.5}
|
||||
user = {"enabled": None}
|
||||
result = merge_bgm_config(template, user)
|
||||
assert result["enabled"] is None
|
||||
|
||||
|
||||
class TestMergeBgmConfigNewFields:
|
||||
"""用户配置新增模板没有的字段。"""
|
||||
|
||||
def test_user_adds_new_field(self):
|
||||
template = {"volume": 0.5}
|
||||
user = {"fade_out": 3}
|
||||
result = merge_bgm_config(template, user)
|
||||
assert result["volume"] == 0.5
|
||||
assert result["fade_out"] == 3
|
||||
assert result["sidechain_enabled"] is True
|
||||
assert result["sidechain_ratio"] == 0.6
|
||||
|
||||
def test_user_adds_multiple_new_fields(self):
|
||||
template = {"enabled": True}
|
||||
user = {"volume": 0.7, "track": "x.mp3", "loop": True}
|
||||
def test_both_empty_returns_empty_dict(self):
|
||||
"""两者都为空时返回空字典"""
|
||||
result = merge_bgm_config({}, {})
|
||||
assert result == {}
|
||||
|
||||
def test_nested_dict_shallow_merge(self):
|
||||
"""嵌套字典是浅合并(当前设计)"""
|
||||
template = {"enabled": True, "config": {"eq": True, "compression": False}}
|
||||
user = {"config": {"compression": True, "reverb": 0.5}}
|
||||
result = merge_bgm_config(template, user)
|
||||
assert result["enabled"] is True
|
||||
assert result["volume"] == 0.7
|
||||
assert result["track"] == "x.mp3"
|
||||
assert result["loop"] is True
|
||||
# 浅合并:整个 config 被用户值覆盖
|
||||
assert result["config"] == {"compression": True, "reverb": 0.5}
|
||||
|
||||
|
||||
class TestMergeBgmConfigImmutableInput:
|
||||
"""确保输入字典不被修改。"""
|
||||
|
||||
def test_template_not_modified(self):
|
||||
def test_does_not_mutate_template(self):
|
||||
"""不修改原始模板配置"""
|
||||
template = {"enabled": True, "volume": 0.5}
|
||||
original = dict(template)
|
||||
merge_bgm_config(template, {"volume": 0.9})
|
||||
merge_bgm_config(template, {"volume": 0.8})
|
||||
assert template == original
|
||||
|
||||
def test_user_not_modified(self):
|
||||
user = {"enabled": False, "track": "x.mp3"}
|
||||
def test_does_not_mutate_user(self):
|
||||
"""不修改原始用户配置"""
|
||||
user = {"volume": 0.8}
|
||||
original = dict(user)
|
||||
merge_bgm_config({"volume": 0.5}, user)
|
||||
merge_bgm_config({"enabled": True}, user)
|
||||
assert user == original
|
||||
|
||||
@@ -1,160 +0,0 @@
|
||||
"""BGM工具函数领域层单元测试."""
|
||||
|
||||
import pytest
|
||||
|
||||
from packages.domain.bgm_utils import merge_bgm_config
|
||||
|
||||
|
||||
class TestMergeBgmConfig:
|
||||
"""merge_bgm_config 函数测试."""
|
||||
|
||||
def test_both_empty(self):
|
||||
"""两个都是空字典."""
|
||||
result = merge_bgm_config({}, {})
|
||||
assert result == {}
|
||||
|
||||
def test_user_empty_returns_template_copy(self):
|
||||
"""用户配置为空,返回模板配置的拷贝."""
|
||||
template = {"enabled": True, "volume": 0.5}
|
||||
result = merge_bgm_config(template, {})
|
||||
assert result == template
|
||||
# 确保是副本不是引用
|
||||
result["volume"] = 0.9
|
||||
assert template["volume"] == 0.5
|
||||
|
||||
def test_template_empty_returns_user_copy(self):
|
||||
"""模板配置为空,返回用户配置的拷贝."""
|
||||
user = {"enabled": False, "volume": 0.8}
|
||||
result = merge_bgm_config({}, user)
|
||||
assert result == user
|
||||
# 确保是副本
|
||||
result["volume"] = 0.1
|
||||
assert user["volume"] == 0.8
|
||||
|
||||
def test_user_overrides_template(self):
|
||||
"""用户配置覆盖模板配置."""
|
||||
template = {"enabled": True, "volume": 0.5, "track": "default"}
|
||||
user = {"volume": 0.8}
|
||||
result = merge_bgm_config(template, user)
|
||||
assert result["volume"] == 0.8
|
||||
assert result["track"] == "default"
|
||||
|
||||
def test_enabled_special_handling_user_not_set(self):
|
||||
"""enabled 特殊处理:用户没传就保留模板的."""
|
||||
template = {"enabled": True, "volume": 0.5}
|
||||
user = {"volume": 0.8} # 没传 enabled
|
||||
result = merge_bgm_config(template, user)
|
||||
# 用户没传 enabled,保留模板的 True
|
||||
assert result["enabled"] is True
|
||||
assert result["volume"] == 0.8
|
||||
|
||||
def test_enabled_user_explicit_false(self):
|
||||
"""用户显式传 enabled=False,应该覆盖模板."""
|
||||
template = {"enabled": True, "volume": 0.5}
|
||||
user = {"enabled": False}
|
||||
result = merge_bgm_config(template, user)
|
||||
assert result["enabled"] is False
|
||||
|
||||
def test_enabled_user_explicit_true(self):
|
||||
"""用户显式传 enabled=True,覆盖模板的 False."""
|
||||
template = {"enabled": False, "volume": 0.5}
|
||||
user = {"enabled": True}
|
||||
result = merge_bgm_config(template, user)
|
||||
assert result["enabled"] is True
|
||||
|
||||
def test_full_override(self):
|
||||
"""用户完全覆盖模板."""
|
||||
template = {"enabled": True, "volume": 0.3, "track": "piano"}
|
||||
user = {"enabled": False, "volume": 0.9, "track": "guitar"}
|
||||
result = merge_bgm_config(template, user)
|
||||
assert result == user
|
||||
|
||||
def test_partial_override_keep_rest(self):
|
||||
"""部分覆盖,其余保留模板值."""
|
||||
template = {
|
||||
"enabled": True,
|
||||
"volume": 0.5,
|
||||
"fade_in": 1.0,
|
||||
"fade_out": 1.0,
|
||||
"track": "default",
|
||||
}
|
||||
user = {"volume": 0.7, "fade_in": 2.0}
|
||||
result = merge_bgm_config(template, user)
|
||||
assert result["volume"] == 0.7
|
||||
assert result["fade_in"] == 2.0
|
||||
assert result["fade_out"] == 1.0
|
||||
assert result["track"] == "default"
|
||||
assert result["enabled"] is True # 用户没传,保留模板
|
||||
|
||||
def test_user_none(self):
|
||||
"""user_bgm 为 None 的情况."""
|
||||
template = {"enabled": True, "volume": 0.5}
|
||||
result = merge_bgm_config(template, None) # type: ignore
|
||||
assert result == template
|
||||
|
||||
def test_template_none(self):
|
||||
"""template_bgm 为 None 的情况."""
|
||||
user = {"enabled": False, "volume": 0.8}
|
||||
result = merge_bgm_config(None, user) # type: ignore
|
||||
assert result == user
|
||||
|
||||
def test_preserves_extra_fields(self):
|
||||
"""保留模板中的额外字段(用户没覆盖的)."""
|
||||
template = {"enabled": True, "volume": 0.5, "custom_field": "value"}
|
||||
user = {"volume": 0.6}
|
||||
result = merge_bgm_config(template, user)
|
||||
assert result["custom_field"] == "value"
|
||||
|
||||
def test_user_adds_new_fields(self):
|
||||
"""用户可以添加模板中没有的新字段."""
|
||||
template = {"enabled": True, "volume": 0.5}
|
||||
user = {"loop": True, "start_time": 5.0}
|
||||
result = merge_bgm_config(template, user)
|
||||
assert result["enabled"] is True
|
||||
assert result["volume"] == 0.5
|
||||
assert result["loop"] is True
|
||||
assert result["start_time"] == 5.0
|
||||
|
||||
def test_nested_dict_behavior(self):
|
||||
"""嵌套字典的合并行为(简单替换,不深度合并)."""
|
||||
template = {"enabled": True, "effects": {"fade": True, "reverb": False}}
|
||||
user = {"effects": {"reverb": True}}
|
||||
result = merge_bgm_config(template, user)
|
||||
# 简单合并,用户的 effects 整体覆盖模板的
|
||||
assert result["effects"] == {"reverb": True}
|
||||
|
||||
def test_enabled_in_template_only(self):
|
||||
"""只有模板有 enabled,用户没有."""
|
||||
template = {"enabled": False, "volume": 0.5}
|
||||
user = {"volume": 0.7}
|
||||
result = merge_bgm_config(template, user)
|
||||
# 用户没传 enabled,保留模板的 False
|
||||
assert result["enabled"] is False
|
||||
|
||||
def test_both_have_enabled_false(self):
|
||||
"""两边都有 enabled 且都是 False."""
|
||||
template = {"enabled": False, "volume": 0.5}
|
||||
user = {"enabled": False}
|
||||
result = merge_bgm_config(template, user)
|
||||
assert result["enabled"] is False
|
||||
|
||||
def test_return_type_is_dict(self):
|
||||
"""返回类型是 dict."""
|
||||
result = merge_bgm_config({"a": 1}, {"b": 2})
|
||||
assert isinstance(result, dict)
|
||||
|
||||
def test_does_not_mutate_template(self):
|
||||
"""不修改原始模板字典."""
|
||||
template = {"enabled": True, "volume": 0.5}
|
||||
template_copy = template.copy()
|
||||
user = {"volume": 0.9}
|
||||
merge_bgm_config(template, user)
|
||||
assert template == template_copy
|
||||
|
||||
def test_does_not_mutate_user(self):
|
||||
"""不修改原始用户字典."""
|
||||
template = {"enabled": True, "volume": 0.5}
|
||||
user = {"volume": 0.9}
|
||||
user_copy = user.copy()
|
||||
merge_bgm_config(template, user)
|
||||
assert user == user_copy
|
||||
@@ -1,210 +0,0 @@
|
||||
"""绿幕抠像引擎单元测试 - 配置解析等纯逻辑."""
|
||||
|
||||
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,6 +1,4 @@
|
||||
"""分类领域模型单元测试 - 纯逻辑部分。"""
|
||||
|
||||
from __future__ import annotations
|
||||
"""classification 模块单元测试."""
|
||||
|
||||
import pytest
|
||||
from domain.classification import (
|
||||
@@ -8,203 +6,99 @@ from domain.classification import (
|
||||
AssetLibraryKind,
|
||||
ClassificationJob,
|
||||
ClassificationJobStatus,
|
||||
ClassificationStatus,
|
||||
IngestJobStatus,
|
||||
)
|
||||
|
||||
|
||||
class TestAssetLibraryKind:
|
||||
"""素材库类型枚举。"""
|
||||
"""AssetLibraryKind 枚举测试."""
|
||||
|
||||
def test_video_value(self):
|
||||
def test_values(self):
|
||||
assert AssetLibraryKind.VIDEO == "video"
|
||||
|
||||
def test_voice_value(self):
|
||||
assert AssetLibraryKind.VOICE == "voice"
|
||||
|
||||
def test_image_value(self):
|
||||
assert AssetLibraryKind.IMAGE == "image"
|
||||
|
||||
def test_is_str_enum(self):
|
||||
assert isinstance(AssetLibraryKind.VIDEO, str)
|
||||
assert AssetLibraryKind.VIDEO + "_test" == "video_test"
|
||||
|
||||
def test_members_count(self):
|
||||
assert len(AssetLibraryKind) == 3
|
||||
|
||||
|
||||
class TestIngestJobStatus:
|
||||
"""导入任务状态枚举。"""
|
||||
"""IngestJobStatus 枚举测试."""
|
||||
|
||||
def test_pending_value(self):
|
||||
def test_values(self):
|
||||
assert IngestJobStatus.PENDING == "pending"
|
||||
|
||||
def test_processing_value(self):
|
||||
assert IngestJobStatus.PROCESSING == "processing"
|
||||
|
||||
def test_completed_value(self):
|
||||
assert IngestJobStatus.COMPLETED == "completed"
|
||||
|
||||
def test_failed_value(self):
|
||||
assert IngestJobStatus.FAILED == "failed"
|
||||
|
||||
def test_members_count(self):
|
||||
assert len(IngestJobStatus) == 4
|
||||
|
||||
class TestClassificationJobStatus:
|
||||
"""ClassificationJobStatus 枚举测试."""
|
||||
|
||||
class TestClassificationJobStatusMissing:
|
||||
"""ClassificationJobStatus._missing_ 兼容性测试。"""
|
||||
|
||||
def test_standard_values(self):
|
||||
"""标准值正常解析。"""
|
||||
assert ClassificationJobStatus("pending") == ClassificationJobStatus.PENDING
|
||||
assert ClassificationJobStatus("processing") == ClassificationJobStatus.PROCESSING
|
||||
assert ClassificationJobStatus("completed") == ClassificationJobStatus.COMPLETED
|
||||
assert ClassificationJobStatus("failed") == ClassificationJobStatus.FAILED
|
||||
|
||||
def test_done_maps_to_completed(self):
|
||||
"""历史值 done 映射到 COMPLETED。"""
|
||||
assert ClassificationJobStatus("done") == ClassificationJobStatus.COMPLETED
|
||||
|
||||
def test_success_maps_to_completed(self):
|
||||
assert ClassificationJobStatus("success") == ClassificationJobStatus.COMPLETED
|
||||
|
||||
def test_finished_maps_to_completed(self):
|
||||
assert ClassificationJobStatus("finished") == ClassificationJobStatus.COMPLETED
|
||||
|
||||
def test_complete_maps_to_completed(self):
|
||||
assert ClassificationJobStatus("complete") == ClassificationJobStatus.COMPLETED
|
||||
|
||||
def test_fail_maps_to_failed(self):
|
||||
assert ClassificationJobStatus("fail") == ClassificationJobStatus.FAILED
|
||||
|
||||
def test_error_maps_to_failed(self):
|
||||
assert ClassificationJobStatus("error") == ClassificationJobStatus.FAILED
|
||||
|
||||
def test_err_maps_to_failed(self):
|
||||
assert ClassificationJobStatus("err") == ClassificationJobStatus.FAILED
|
||||
|
||||
def test_process_maps_to_processing(self):
|
||||
assert ClassificationJobStatus("process") == ClassificationJobStatus.PROCESSING
|
||||
|
||||
def test_running_maps_to_processing(self):
|
||||
assert ClassificationJobStatus("running") == ClassificationJobStatus.PROCESSING
|
||||
|
||||
def test_run_maps_to_processing(self):
|
||||
assert ClassificationJobStatus("run") == ClassificationJobStatus.PROCESSING
|
||||
|
||||
def test_unknown_value_defaults_to_pending(self):
|
||||
"""未知值兜底为 PENDING。"""
|
||||
assert ClassificationJobStatus("unknown") == ClassificationJobStatus.PENDING
|
||||
assert ClassificationJobStatus("whatever") == ClassificationJobStatus.PENDING
|
||||
assert ClassificationJobStatus("") == ClassificationJobStatus.PENDING
|
||||
|
||||
def test_case_insensitive(self):
|
||||
"""大小写不敏感。"""
|
||||
assert ClassificationJobStatus("DONE") == ClassificationJobStatus.COMPLETED
|
||||
assert ClassificationJobStatus("Done") == ClassificationJobStatus.COMPLETED
|
||||
assert ClassificationJobStatus("FAIL") == ClassificationJobStatus.FAILED
|
||||
assert ClassificationJobStatus("Error") == ClassificationJobStatus.FAILED
|
||||
|
||||
def test_stripped(self):
|
||||
"""前后空白字符被忽略。"""
|
||||
assert ClassificationJobStatus(" done ") == ClassificationJobStatus.COMPLETED
|
||||
assert ClassificationJobStatus("\tfail\n") == ClassificationJobStatus.FAILED
|
||||
|
||||
def test_none_returns_pending(self):
|
||||
"""None 值也返回 PENDING(不报错)。"""
|
||||
assert ClassificationJobStatus(None) == ClassificationJobStatus.PENDING
|
||||
|
||||
def test_integer_returns_pending(self):
|
||||
"""非字符串值返回 PENDING。"""
|
||||
assert ClassificationJobStatus(123) == ClassificationJobStatus.PENDING
|
||||
|
||||
|
||||
class TestClassificationStatusAlias:
|
||||
"""向后兼容别名。"""
|
||||
|
||||
def test_alias_same_class(self):
|
||||
assert ClassificationStatus is ClassificationJobStatus
|
||||
|
||||
def test_alias_values_same(self):
|
||||
assert ClassificationStatus.PENDING == ClassificationJobStatus.PENDING
|
||||
assert ClassificationStatus.COMPLETED == ClassificationJobStatus.COMPLETED
|
||||
def test_values(self):
|
||||
assert ClassificationJobStatus.PENDING == "pending"
|
||||
assert ClassificationJobStatus.PROCESSING == "processing"
|
||||
assert ClassificationJobStatus.COMPLETED == "completed"
|
||||
assert ClassificationJobStatus.FAILED == "failed"
|
||||
|
||||
|
||||
class TestAssetClassification:
|
||||
"""素材分类枚举。"""
|
||||
"""AssetClassification 枚举测试."""
|
||||
|
||||
def test_scenic(self):
|
||||
def test_values(self):
|
||||
assert AssetClassification.SCENIC == "scenic"
|
||||
|
||||
def test_product(self):
|
||||
assert AssetClassification.PRODUCT == "product"
|
||||
|
||||
def test_person(self):
|
||||
assert AssetClassification.PERSON == "person"
|
||||
|
||||
def test_animal(self):
|
||||
assert AssetClassification.ANIMAL == "animal"
|
||||
|
||||
def test_food(self):
|
||||
assert AssetClassification.FOOD == "food"
|
||||
|
||||
def test_tech(self):
|
||||
assert AssetClassification.TECH == "tech"
|
||||
|
||||
def test_sport(self):
|
||||
assert AssetClassification.SPORT == "sport"
|
||||
|
||||
def test_music(self):
|
||||
assert AssetClassification.MUSIC == "music"
|
||||
|
||||
def test_other(self):
|
||||
assert AssetClassification.OTHER == "other"
|
||||
|
||||
def test_members_count(self):
|
||||
assert len(AssetClassification) == 9
|
||||
|
||||
|
||||
class TestClassificationJobCreate:
|
||||
"""ClassificationJob.create 工厂方法。"""
|
||||
"""ClassificationJob.create 工厂方法测试."""
|
||||
|
||||
def test_create_basic(self):
|
||||
job = ClassificationJob.create(project_id="proj-1", asset_id="asset-1")
|
||||
assert job.project_id == "proj-1"
|
||||
assert job.asset_id == "asset-1"
|
||||
def test_create_with_valid_params(self):
|
||||
job = ClassificationJob.create(project_id="proj_001", asset_id="asset_001")
|
||||
assert job.id
|
||||
assert len(job.id) == 32
|
||||
assert job.project_id == "proj_001"
|
||||
assert job.asset_id == "asset_001"
|
||||
assert job.status == ClassificationJobStatus.PENDING
|
||||
assert job.classification == ""
|
||||
assert job.confidence == 0.0
|
||||
assert job.error_message == ""
|
||||
assert job.id # 自动生成的 ID 非空
|
||||
assert job.created_at is not None
|
||||
assert job.updated_at is not None
|
||||
|
||||
def test_create_strips_whitespace(self):
|
||||
job = ClassificationJob.create(project_id=" proj-1 ", asset_id="\tasset-1\n")
|
||||
assert job.project_id == "proj-1"
|
||||
assert job.asset_id == "asset-1"
|
||||
def test_create_strips_strings(self):
|
||||
job = ClassificationJob.create(
|
||||
project_id=" proj_002 ",
|
||||
asset_id=" asset_002 ",
|
||||
)
|
||||
assert job.project_id == "proj_002"
|
||||
assert job.asset_id == "asset_002"
|
||||
|
||||
def test_create_empty_project_id_raises(self):
|
||||
with pytest.raises(ValueError, match="project_id 不能为空"):
|
||||
ClassificationJob.create(project_id="", asset_id="asset-1")
|
||||
with pytest.raises(ValueError, match="project_id"):
|
||||
ClassificationJob.create(project_id="", asset_id="a")
|
||||
|
||||
def test_create_whitespace_project_id_raises(self):
|
||||
with pytest.raises(ValueError, match="project_id 不能为空"):
|
||||
ClassificationJob.create(project_id=" ", asset_id="asset-1")
|
||||
with pytest.raises(ValueError, match="project_id"):
|
||||
ClassificationJob.create(project_id=" ", asset_id="a")
|
||||
|
||||
def test_create_empty_asset_id_raises(self):
|
||||
with pytest.raises(ValueError, match="asset_id 不能为空"):
|
||||
ClassificationJob.create(project_id="proj-1", asset_id="")
|
||||
with pytest.raises(ValueError, match="asset_id"):
|
||||
ClassificationJob.create(project_id="p", asset_id="")
|
||||
|
||||
def test_create_whitespace_asset_id_raises(self):
|
||||
with pytest.raises(ValueError, match="asset_id 不能为空"):
|
||||
ClassificationJob.create(project_id="proj-1", asset_id=" \t ")
|
||||
with pytest.raises(ValueError, match="asset_id"):
|
||||
ClassificationJob.create(project_id="p", asset_id=" ")
|
||||
|
||||
def test_create_generates_unique_ids(self):
|
||||
job1 = ClassificationJob.create(project_id="p", asset_id="a")
|
||||
job2 = ClassificationJob.create(project_id="p", asset_id="a")
|
||||
assert job1.id != job2.id
|
||||
def test_create_ids_are_unique(self):
|
||||
j1 = ClassificationJob.create(project_id="p", asset_id="a")
|
||||
j2 = ClassificationJob.create(project_id="p", asset_id="b")
|
||||
assert j1.id != j2.id
|
||||
|
||||
def test_create_id_is_hex(self):
|
||||
def test_create_timestamps_are_utc(self):
|
||||
job = ClassificationJob.create(project_id="p", asset_id="a")
|
||||
assert job.created_at.tzinfo is not None
|
||||
assert job.updated_at.tzinfo is not None
|
||||
@@ -243,130 +137,3 @@ class TestClassificationJobState:
|
||||
job = ClassificationJob.create(project_id="proj-1", asset_id="asset-1")
|
||||
job.confidence = 1.0
|
||||
assert job.confidence == 1.0
|
||||
|
||||
|
||||
class TestClassificationJobStatusMissing:
|
||||
"""ClassificationJobStatus._missing_ 兼容行为测试"""
|
||||
|
||||
def test_done_maps_to_completed(self):
|
||||
assert ClassificationJobStatus("done") == ClassificationJobStatus.COMPLETED
|
||||
|
||||
def test_success_maps_to_completed(self):
|
||||
assert ClassificationJobStatus("success") == ClassificationJobStatus.COMPLETED
|
||||
|
||||
def test_finished_maps_to_completed(self):
|
||||
assert ClassificationJobStatus("finished") == ClassificationJobStatus.COMPLETED
|
||||
|
||||
def test_complete_maps_to_completed(self):
|
||||
assert ClassificationJobStatus("complete") == ClassificationJobStatus.COMPLETED
|
||||
|
||||
def test_fail_maps_to_failed(self):
|
||||
assert ClassificationJobStatus("fail") == ClassificationJobStatus.FAILED
|
||||
|
||||
def test_error_maps_to_failed(self):
|
||||
assert ClassificationJobStatus("error") == ClassificationJobStatus.FAILED
|
||||
|
||||
def test_err_maps_to_failed(self):
|
||||
assert ClassificationJobStatus("err") == ClassificationJobStatus.FAILED
|
||||
|
||||
def test_unknown_maps_to_pending(self):
|
||||
assert ClassificationJobStatus("unknown_status") == ClassificationJobStatus.PENDING
|
||||
|
||||
def test_case_insensitive_mapping(self):
|
||||
assert ClassificationJobStatus("DONE") == ClassificationJobStatus.COMPLETED
|
||||
assert ClassificationJobStatus("Done") == ClassificationJobStatus.COMPLETED
|
||||
|
||||
def test_whitespace_stripped(self):
|
||||
assert ClassificationJobStatus(" done ") == ClassificationJobStatus.COMPLETED
|
||||
|
||||
|
||||
class TestClassificationJobExtended:
|
||||
"""ClassificationJob 深度补充测试"""
|
||||
|
||||
def test_id_is_hex(self):
|
||||
job = ClassificationJob.create(project_id="p", asset_id="a")
|
||||
int(job.id, 16)
|
||||
|
||||
def test_empty_classification(self):
|
||||
job = ClassificationJob.create(project_id="p", asset_id="a")
|
||||
assert job.classification == ""
|
||||
|
||||
def test_zero_confidence(self):
|
||||
job = ClassificationJob.create(project_id="p", asset_id="a")
|
||||
assert job.confidence == 0.0
|
||||
|
||||
def test_high_confidence(self):
|
||||
job = ClassificationJob.create(project_id="p", asset_id="a")
|
||||
job.confidence = 0.99
|
||||
assert job.confidence == pytest.approx(0.99)
|
||||
|
||||
def test_negative_confidence(self):
|
||||
job = ClassificationJob.create(project_id="p", asset_id="a")
|
||||
job.confidence = -0.1
|
||||
assert job.confidence == pytest.approx(-0.1)
|
||||
|
||||
def test_confidence_over_one(self):
|
||||
job = ClassificationJob.create(project_id="p", asset_id="a")
|
||||
job.confidence = 1.5
|
||||
assert job.confidence == pytest.approx(1.5)
|
||||
|
||||
def test_empty_error_message(self):
|
||||
job = ClassificationJob.create(project_id="p", asset_id="a")
|
||||
assert job.error_message == ""
|
||||
|
||||
def test_long_error_message(self):
|
||||
job = ClassificationJob.create(project_id="p", asset_id="a")
|
||||
long_msg = "error" * 100
|
||||
job.error_message = long_msg
|
||||
assert job.error_message == long_msg
|
||||
assert len(job.error_message) == 500
|
||||
|
||||
def test_status_with_string_assignment(self):
|
||||
job = ClassificationJob.create(project_id="p", asset_id="a")
|
||||
job.status = "processing"
|
||||
assert job.status == ClassificationJobStatus.PROCESSING
|
||||
|
||||
|
||||
class TestAssetLibraryKindExtended:
|
||||
"""AssetLibraryKind 深度补充测试"""
|
||||
|
||||
def test_image_value(self):
|
||||
assert AssetLibraryKind.IMAGE == "image"
|
||||
|
||||
def test_all_three_kinds(self):
|
||||
assert len(AssetLibraryKind) == 3
|
||||
|
||||
def test_is_string_enum(self):
|
||||
assert isinstance(AssetLibraryKind.VIDEO, str)
|
||||
|
||||
def test_from_string(self):
|
||||
assert AssetLibraryKind("video") == AssetLibraryKind.VIDEO
|
||||
|
||||
|
||||
class TestIngestJobStatusExtended:
|
||||
"""IngestJobStatus 深度补充测试"""
|
||||
|
||||
def test_is_string_enum(self):
|
||||
assert isinstance(IngestJobStatus.PENDING, str)
|
||||
|
||||
def test_total_count(self):
|
||||
assert len(IngestJobStatus) == 4
|
||||
|
||||
def test_from_string(self):
|
||||
assert IngestJobStatus("pending") == IngestJobStatus.PENDING
|
||||
|
||||
|
||||
class TestAssetClassificationExtended:
|
||||
"""AssetClassification 深度补充测试"""
|
||||
|
||||
def test_total_count(self):
|
||||
assert len(AssetClassification) == 9
|
||||
|
||||
def test_is_string_enum(self):
|
||||
assert isinstance(AssetClassification.SCENIC, str)
|
||||
|
||||
def test_from_string(self):
|
||||
assert AssetClassification("scenic") == AssetClassification.SCENIC
|
||||
|
||||
def test_other_category(self):
|
||||
assert AssetClassification.OTHER == "other"
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
"""调色引擎单元测试 - 配置解析等纯逻辑."""
|
||||
"""滤镜调色引擎单元测试."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
@@ -6,19 +6,92 @@ 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,
|
||||
VALID_PRESETS,
|
||||
PRESET_VINTAGE,
|
||||
PRESET_WARM,
|
||||
ColorGradeConfig,
|
||||
ColorGradeEngine,
|
||||
get_preset_names,
|
||||
get_preset_params,
|
||||
)
|
||||
|
||||
# ── 预设常量测试 ──────────────────────────────────────────────────────────────
|
||||
|
||||
class TestColorGradeConfigDefaults:
|
||||
"""默认配置测试."""
|
||||
|
||||
def test_default_values(self):
|
||||
"""默认值正确."""
|
||||
config = ColorGradeConfig()
|
||||
assert config.enabled is False
|
||||
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
|
||||
assert config.preset == ""
|
||||
assert config.brightness is None
|
||||
assert config.contrast is None
|
||||
@@ -26,83 +99,50 @@ class TestColorGradeConfigDefaults:
|
||||
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": "fresh",
|
||||
}
|
||||
)
|
||||
assert config.enabled is True
|
||||
assert config.preset == "fresh"
|
||||
config = ColorGradeConfig.from_dict({"enabled": True, "preset": PRESET_FRESH})
|
||||
assert config.enabled
|
||||
assert config.preset == PRESET_FRESH
|
||||
|
||||
def test_invalid_preset_ignored(self):
|
||||
"""无效预设被忽略."""
|
||||
config = ColorGradeConfig.from_dict(
|
||||
{
|
||||
"enabled": True,
|
||||
"preset": "unknown_preset",
|
||||
}
|
||||
)
|
||||
assert config.preset == ""
|
||||
"""无效预设名应该被忽略."""
|
||||
config = ColorGradeConfig.from_dict({"enabled": True, "preset": "invalid_preset"})
|
||||
assert config.preset == "" # 被清空
|
||||
|
||||
def test_custom_brightness(self):
|
||||
"""自定义亮度."""
|
||||
def test_with_custom_params(self):
|
||||
"""自定义参数覆盖."""
|
||||
config = ColorGradeConfig.from_dict(
|
||||
{
|
||||
"enabled": True,
|
||||
"brightness": 20,
|
||||
"contrast": -10,
|
||||
"saturation": 150,
|
||||
"temperature": 25,
|
||||
"hue": 30,
|
||||
}
|
||||
)
|
||||
assert config.brightness == 20.0
|
||||
assert config.enabled
|
||||
assert config.brightness == 20
|
||||
assert config.contrast == -10
|
||||
assert config.saturation == 150
|
||||
assert config.temperature == 25
|
||||
assert config.hue == 30
|
||||
|
||||
def test_custom_all_params(self):
|
||||
"""所有参数自定义."""
|
||||
def test_string_numeric_values(self):
|
||||
"""字符串形式的数字应该能解析."""
|
||||
config = ColorGradeConfig.from_dict(
|
||||
{
|
||||
"enabled": True,
|
||||
"brightness": 10,
|
||||
"contrast": 15,
|
||||
"saturation": 120,
|
||||
"temperature": -5,
|
||||
"hue": 10,
|
||||
"brightness": "20.5",
|
||||
"saturation": "150",
|
||||
}
|
||||
)
|
||||
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
|
||||
assert config.brightness == 20.5
|
||||
assert config.saturation == 150.0
|
||||
|
||||
def test_invalid_param_value_returns_none(self):
|
||||
"""无效参数值返回None(不覆盖)."""
|
||||
def test_invalid_value_returns_none(self):
|
||||
"""无效值应该返回None(不覆盖)."""
|
||||
config = ColorGradeConfig.from_dict(
|
||||
{
|
||||
"enabled": True,
|
||||
@@ -111,151 +151,422 @@ 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
|
||||
|
||||
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
|
||||
# ── ColorGradeConfig.resolve_params 测试 ──────────────────────────────────────
|
||||
|
||||
|
||||
class TestResolveParams:
|
||||
"""resolve_params 参数解析测试."""
|
||||
"""参数解析与边界钳制测试."""
|
||||
|
||||
def test_disabled_returns_defaults(self):
|
||||
"""禁用配置也返回默认参数."""
|
||||
config = ColorGradeConfig(enabled=False)
|
||||
def test_default_params_when_empty(self):
|
||||
"""无预设无自定义时返回默认值."""
|
||||
config = ColorGradeConfig(enabled=True)
|
||||
params = config.resolve_params()
|
||||
for key, val in DEFAULT_PARAMS.items():
|
||||
assert params[key] == val
|
||||
|
||||
def test_no_preset_no_custom_returns_defaults(self):
|
||||
"""无预设无自定义返回默认值."""
|
||||
config = ColorGradeConfig(enabled=True)
|
||||
def test_preset_params_applied(self):
|
||||
"""预设参数应该被应用."""
|
||||
config = ColorGradeConfig(enabled=True, preset=PRESET_FRESH)
|
||||
params = config.resolve_params()
|
||||
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
|
||||
preset = PRESET_PARAMS[PRESET_FRESH]
|
||||
for key, val in preset.items():
|
||||
assert params[key] == val
|
||||
|
||||
def test_custom_overrides_preset(self):
|
||||
"""自定义参数覆盖预设."""
|
||||
"""自定义参数应该覆盖预设值."""
|
||||
config = ColorGradeConfig(
|
||||
enabled=True,
|
||||
preset="fresh",
|
||||
preset=PRESET_FRESH,
|
||||
brightness=50, # 覆盖预设的8
|
||||
)
|
||||
params = config.resolve_params()
|
||||
assert params["brightness"] == 50
|
||||
# 其他参数仍用预设值
|
||||
assert params["saturation"] == 120
|
||||
# 其他参数还是预设值
|
||||
assert params["contrast"] == PRESET_PARAMS[PRESET_FRESH]["contrast"]
|
||||
|
||||
def test_brightness_clamped(self):
|
||||
"""亮度边界钳制."""
|
||||
def test_clamp_brightness_high(self):
|
||||
"""亮度超过上限应该被钳制."""
|
||||
config = ColorGradeConfig(enabled=True, brightness=200)
|
||||
params = config.resolve_params()
|
||||
assert params["brightness"] == 100.0
|
||||
assert params["brightness"] == 100
|
||||
|
||||
def test_saturation_clamped_low(self):
|
||||
"""饱和度下限钳制."""
|
||||
config = ColorGradeConfig(enabled=True, saturation=-10)
|
||||
def test_clamp_brightness_low(self):
|
||||
"""亮度低于下限应该被钳制."""
|
||||
config = ColorGradeConfig(enabled=True, brightness=-200)
|
||||
params = config.resolve_params()
|
||||
assert params["saturation"] == 0.0
|
||||
assert params["brightness"] == -100
|
||||
|
||||
def test_saturation_clamped_high(self):
|
||||
"""饱和度上限钳制."""
|
||||
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应该被钳制."""
|
||||
config = ColorGradeConfig(enabled=True, saturation=300)
|
||||
params = config.resolve_params()
|
||||
assert params["saturation"] == 200.0
|
||||
assert params["saturation"] == 200
|
||||
|
||||
def test_hue_clamped(self):
|
||||
"""色调边界钳制."""
|
||||
config = ColorGradeConfig(enabled=True, hue=200)
|
||||
def test_clamp_hue_high(self):
|
||||
"""色调超过180应该被钳制."""
|
||||
config = ColorGradeConfig(enabled=True, hue=270)
|
||||
params = config.resolve_params()
|
||||
assert params["hue"] == 180.0
|
||||
assert params["hue"] == 180
|
||||
|
||||
def test_hue_negative_clamped(self):
|
||||
"""负色调边界钳制."""
|
||||
config = ColorGradeConfig(enabled=True, hue=-200)
|
||||
def test_clamp_hue_low(self):
|
||||
"""色调低于-180应该被钳制."""
|
||||
config = ColorGradeConfig(enabled=True, hue=-270)
|
||||
params = config.resolve_params()
|
||||
assert params["hue"] == -180.0
|
||||
assert params["hue"] == -180
|
||||
|
||||
def test_returns_all_five_params(self):
|
||||
"""返回所有5个参数."""
|
||||
config = ColorGradeConfig(enabled=True)
|
||||
def test_clamp_contrast(self):
|
||||
"""对比度越界应该被钳制."""
|
||||
config = ColorGradeConfig(enabled=True, contrast=150)
|
||||
params = config.resolve_params()
|
||||
assert set(params.keys()) == {"brightness", "contrast", "saturation", "temperature", "hue"}
|
||||
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 测试 ──────────────────────────────────────────
|
||||
|
||||
|
||||
class TestHasEffect:
|
||||
"""has_effect 方法测试."""
|
||||
"""是否有实际效果判断测试."""
|
||||
|
||||
def test_default_no_effect(self):
|
||||
"""默认配置无效果."""
|
||||
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."""
|
||||
config = ColorGradeConfig(enabled=True)
|
||||
assert config.has_effect() is False
|
||||
assert not config.has_effect()
|
||||
|
||||
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):
|
||||
"""自定义亮度有效果."""
|
||||
def test_brightness_change_has_effect(self):
|
||||
"""亮度变化应该有效果."""
|
||||
config = ColorGradeConfig(enabled=True, brightness=10)
|
||||
assert config.has_effect() is True
|
||||
assert 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_100_no_effect(self):
|
||||
"""饱和度100是默认值,无效果."""
|
||||
config = ColorGradeConfig(enabled=True, saturation=100)
|
||||
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
|
||||
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()
|
||||
|
||||
|
||||
class TestPresets:
|
||||
"""预设常量测试."""
|
||||
# ── ColorGradeEngine 参数映射测试 ─────────────────────────────────────────────
|
||||
|
||||
def test_eight_valid_presets(self):
|
||||
"""8个有效预设."""
|
||||
assert len(VALID_PRESETS) == 8
|
||||
|
||||
def test_preset_params_match_valid(self):
|
||||
"""所有预设都在有效列表中."""
|
||||
for name in PRESET_PARAMS:
|
||||
assert name in VALID_PRESETS
|
||||
class TestParameterMapping:
|
||||
"""FFmpeg参数映射测试."""
|
||||
|
||||
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_zero(self):
|
||||
"""亮度0 → 0.0."""
|
||||
assert ColorGradeEngine._map_brightness(0) == 0.0
|
||||
|
||||
def test_param_ranges_defined(self):
|
||||
"""参数范围定义完整."""
|
||||
assert set(PARAM_RANGES.keys()) == {"brightness", "contrast", "saturation", "temperature", "hue"}
|
||||
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
|
||||
|
||||
+163
-203
@@ -1,4 +1,9 @@
|
||||
"""拼接引擎单元测试 - 配置解析等纯逻辑."""
|
||||
"""
|
||||
视频拼接引擎配置与纯逻辑测试.
|
||||
|
||||
覆盖 ConcatSegment.from_dict / ConcatConfig.from_config_dict / has_effect / total_segments 等纯逻辑.
|
||||
引擎核心 render 方法依赖 FFmpeg,由集成测试覆盖.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
@@ -6,306 +11,261 @@ import pytest
|
||||
from video_processing.concat_engine import ConcatConfig, ConcatSegment
|
||||
|
||||
|
||||
class TestConcatSegmentDefaults:
|
||||
"""ConcatSegment 默认值测试."""
|
||||
class TestConcatSegmentFromDict:
|
||||
"""ConcatSegment.from_dict 构造逻辑."""
|
||||
|
||||
def test_default_values(self):
|
||||
"""默认值正确."""
|
||||
seg = ConcatSegment(video_path="/a.mp4")
|
||||
assert seg.video_path == "/a.mp4"
|
||||
def test_basic(self):
|
||||
seg = ConcatSegment.from_dict({"video_path": "/tmp/a.mp4"})
|
||||
assert seg.video_path == "/tmp/a.mp4"
|
||||
assert seg.start_time == 0.0
|
||||
assert seg.duration == 0.0
|
||||
assert seg.has_audio is True
|
||||
|
||||
|
||||
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):
|
||||
"""自定义开始时间."""
|
||||
def test_full_fields(self):
|
||||
seg = ConcatSegment.from_dict(
|
||||
{
|
||||
"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",
|
||||
"video_path": "/tmp/b.mp4",
|
||||
"start_time": 5.5,
|
||||
"duration": 10.0,
|
||||
"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_start_time_negative_clamped(self):
|
||||
"""负开始时间钳制到0."""
|
||||
def test_negative_start_time_clamped(self):
|
||||
seg = ConcatSegment.from_dict(
|
||||
{
|
||||
"video_path": "/a.mp4",
|
||||
"start_time": -5.0,
|
||||
"video_path": "/tmp/a.mp4",
|
||||
"start_time": -1.0,
|
||||
}
|
||||
)
|
||||
assert seg.start_time == 0.0
|
||||
|
||||
def test_duration_negative_clamped(self):
|
||||
"""负时长钳制到0."""
|
||||
def test_negative_duration_clamped(self):
|
||||
seg = ConcatSegment.from_dict(
|
||||
{
|
||||
"video_path": "/a.mp4",
|
||||
"duration": -3.0,
|
||||
"video_path": "/tmp/a.mp4",
|
||||
"duration": -5.0,
|
||||
}
|
||||
)
|
||||
assert seg.duration == 0.0
|
||||
|
||||
def test_invalid_start_time_falls_back(self):
|
||||
"""无效start_time回退到0."""
|
||||
def test_invalid_start_time_type_falls_back(self):
|
||||
seg = ConcatSegment.from_dict(
|
||||
{
|
||||
"video_path": "/a.mp4",
|
||||
"start_time": "invalid",
|
||||
"video_path": "/tmp/a.mp4",
|
||||
"start_time": "not_a_number",
|
||||
}
|
||||
)
|
||||
assert seg.start_time == 0.0
|
||||
|
||||
def test_invalid_duration_falls_back(self):
|
||||
"""无效duration回退到0."""
|
||||
def test_invalid_duration_type_falls_back(self):
|
||||
seg = ConcatSegment.from_dict(
|
||||
{
|
||||
"video_path": "/a.mp4",
|
||||
"duration": "not_a_number",
|
||||
"video_path": "/tmp/a.mp4",
|
||||
"duration": "abc",
|
||||
}
|
||||
)
|
||||
assert seg.duration == 0.0
|
||||
|
||||
def test_no_audio(self):
|
||||
"""无音频."""
|
||||
def test_start_time_none_falls_back(self):
|
||||
seg = ConcatSegment.from_dict(
|
||||
{
|
||||
"video_path": "/a.mp4",
|
||||
"has_audio": False,
|
||||
"video_path": "/tmp/a.mp4",
|
||||
"start_time": None,
|
||||
}
|
||||
)
|
||||
assert seg.has_audio is False
|
||||
assert seg.start_time == 0.0
|
||||
|
||||
def test_full_config(self):
|
||||
"""完整配置."""
|
||||
seg = ConcatSegment.from_dict(
|
||||
{
|
||||
"video_path": "/video.mp4",
|
||||
"start_time": 2.5,
|
||||
"duration": 15.0,
|
||||
"has_audio": False,
|
||||
}
|
||||
)
|
||||
assert seg.video_path == "/video.mp4"
|
||||
assert seg.start_time == 2.5
|
||||
assert seg.duration == 15.0
|
||||
assert seg.has_audio is False
|
||||
|
||||
|
||||
class TestConcatConfigDefaults:
|
||||
"""ConcatConfig 默认值测试."""
|
||||
|
||||
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
|
||||
def test_empty_video_path_stored(self):
|
||||
seg = ConcatSegment.from_dict({"video_path": ""})
|
||||
assert seg.video_path == ""
|
||||
|
||||
|
||||
class TestConcatConfigFromConfigDict:
|
||||
"""ConcatConfig.from_config_dict 测试."""
|
||||
"""ConcatConfig.from_config_dict 构造逻辑."""
|
||||
|
||||
def test_none_returns_default(self):
|
||||
"""None返回默认配置."""
|
||||
config = ConcatConfig.from_config_dict(None)
|
||||
assert config.segments == []
|
||||
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
|
||||
|
||||
def test_empty_dict_returns_default(self):
|
||||
"""空dict返回默认."""
|
||||
config = ConcatConfig.from_config_dict({})
|
||||
assert config.segments == []
|
||||
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 == []
|
||||
|
||||
def test_single_segment(self):
|
||||
"""单片段."""
|
||||
config = ConcatConfig.from_config_dict(
|
||||
cfg = ConcatConfig.from_config_dict(
|
||||
{
|
||||
"segments": [{"video_path": "/a.mp4"}],
|
||||
"segments": [
|
||||
{"video_path": "/tmp/a.mp4", "duration": 5.0},
|
||||
],
|
||||
}
|
||||
)
|
||||
assert len(config.segments) == 1
|
||||
assert config.segments[0].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
|
||||
|
||||
def test_multiple_segments(self):
|
||||
"""多片段."""
|
||||
config = ConcatConfig.from_config_dict(
|
||||
cfg = ConcatConfig.from_config_dict(
|
||||
{
|
||||
"segments": [
|
||||
{"video_path": "/a.mp4", "start_time": 1.0},
|
||||
{"video_path": "/b.mp4", "duration": 5.0},
|
||||
{"video_path": "/c.mp4"},
|
||||
{"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},
|
||||
],
|
||||
}
|
||||
)
|
||||
assert len(config.segments) == 3
|
||||
assert config.segments[0].start_time == 1.0
|
||||
assert config.segments[1].duration == 5.0
|
||||
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
|
||||
|
||||
def test_skips_no_path(self):
|
||||
"""跳过无video_path的片段."""
|
||||
config = ConcatConfig.from_config_dict(
|
||||
def test_invalid_segments_filtered(self):
|
||||
cfg = ConcatConfig.from_config_dict(
|
||||
{
|
||||
"segments": [
|
||||
{"video_path": "/a.mp4"},
|
||||
{"other": "value"},
|
||||
{"video_path": ""},
|
||||
{"video_path": "/tmp/valid.mp4"},
|
||||
{"video_path": ""}, # 空路径被过滤
|
||||
{"not_video_path": "xxx"}, # 没有video_path被过滤
|
||||
"not_a_dict", # 不是dict被过滤
|
||||
None, # None被过滤
|
||||
],
|
||||
}
|
||||
)
|
||||
assert len(config.segments) == 1
|
||||
assert len(cfg.segments) == 1
|
||||
assert cfg.segments[0].video_path == "/tmp/valid.mp4"
|
||||
|
||||
def test_segments_not_list_ignored(self):
|
||||
"""segments不是列表忽略."""
|
||||
config = ConcatConfig.from_config_dict(
|
||||
def test_segments_not_a_list(self):
|
||||
cfg = ConcatConfig.from_config_dict(
|
||||
{
|
||||
"segments": "not_a_list",
|
||||
}
|
||||
)
|
||||
assert config.segments == []
|
||||
assert cfg.segments == []
|
||||
|
||||
def test_output_size(self):
|
||||
"""输出尺寸."""
|
||||
config = ConcatConfig.from_config_dict(
|
||||
def test_output_params(self):
|
||||
cfg = ConcatConfig.from_config_dict(
|
||||
{
|
||||
"segments": [{"video_path": "/a.mp4"}],
|
||||
"segments": [],
|
||||
"output_width": 1920,
|
||||
"output_height": 1080,
|
||||
}
|
||||
)
|
||||
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"}],
|
||||
"output_fps": 30.0,
|
||||
"force_reencode": True,
|
||||
}
|
||||
)
|
||||
assert config.force_reencode is 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
|
||||
|
||||
def test_transition_config(self):
|
||||
"""转场配置."""
|
||||
config = ConcatConfig.from_config_dict(
|
||||
cfg = ConcatConfig.from_config_dict(
|
||||
{
|
||||
"segments": [{"video_path": "/a.mp4"}, {"video_path": "/b.mp4"}],
|
||||
"segments": [],
|
||||
"transition": "crossfade",
|
||||
"transition_duration": 1.0,
|
||||
}
|
||||
)
|
||||
assert config.transition == "crossfade"
|
||||
assert config.transition_duration == 1.0
|
||||
assert cfg.transition == "crossfade"
|
||||
assert cfg.transition_duration == 1.0
|
||||
|
||||
def test_non_dict_config_returns_default(self):
|
||||
"""非dict配置返回默认."""
|
||||
config = ConcatConfig.from_config_dict("not_a_dict")
|
||||
assert config.segments == []
|
||||
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
|
||||
|
||||
|
||||
class TestHasEffect:
|
||||
"""has_effect 属性测试."""
|
||||
class TestConcatConfigProperties:
|
||||
"""has_effect / total_segments 属性."""
|
||||
|
||||
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(
|
||||
def test_has_effect_two_or_more_valid(self):
|
||||
cfg = ConcatConfig(
|
||||
segments=[
|
||||
ConcatSegment(video_path="/a.mp4"),
|
||||
ConcatSegment(video_path="/tmp/a.mp4"),
|
||||
ConcatSegment(video_path="/tmp/b.mp4"),
|
||||
]
|
||||
)
|
||||
assert config.has_effect is False
|
||||
assert cfg.has_effect is True
|
||||
|
||||
def test_two_segments_has_effect(self):
|
||||
"""两段及以上有效果."""
|
||||
config = ConcatConfig(
|
||||
def test_no_effect_one_segment(self):
|
||||
cfg = ConcatConfig(
|
||||
segments=[
|
||||
ConcatSegment(video_path="/a.mp4"),
|
||||
ConcatSegment(video_path="/b.mp4"),
|
||||
ConcatSegment(video_path="/tmp/a.mp4"),
|
||||
]
|
||||
)
|
||||
assert config.has_effect is True
|
||||
assert cfg.has_effect is False
|
||||
|
||||
def test_no_effect_zero_segments(self):
|
||||
cfg = ConcatConfig(segments=[])
|
||||
assert cfg.has_effect is False
|
||||
|
||||
class TestTotalSegments:
|
||||
"""total_segments 属性测试."""
|
||||
|
||||
def test_no_segments(self):
|
||||
"""零片段."""
|
||||
config = ConcatConfig()
|
||||
assert config.total_segments == 0
|
||||
|
||||
def test_three_segments(self):
|
||||
"""三个片段."""
|
||||
config = ConcatConfig(
|
||||
def test_no_effect_empty_paths(self):
|
||||
cfg = ConcatConfig(
|
||||
segments=[
|
||||
ConcatSegment(video_path="/a.mp4"),
|
||||
ConcatSegment(video_path="/b.mp4"),
|
||||
ConcatSegment(video_path="/c.mp4"),
|
||||
ConcatSegment(video_path=""),
|
||||
ConcatSegment(video_path=""),
|
||||
]
|
||||
)
|
||||
assert config.total_segments == 3
|
||||
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
|
||||
|
||||
@@ -1,539 +0,0 @@
|
||||
"""CoverGenerator 纯逻辑单测 — 时间钳制 + 智能选帧算法.
|
||||
|
||||
通过 mock run_ffmpeg 和 probe_video_info 验证纯逻辑部分,
|
||||
不实际执行 FFmpeg,确保测试轻量快速。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from pathlib import Path
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from video_processing.cover_generator import (
|
||||
CoverGenerator,
|
||||
DEFAULT_COVER_HEIGHT,
|
||||
DEFAULT_COVER_QUALITY,
|
||||
DEFAULT_COVER_TIME,
|
||||
DEFAULT_COVER_WIDTH,
|
||||
SMART_COVER_FRAME_COUNT,
|
||||
)
|
||||
|
||||
|
||||
class TestCoverGeneratorConstants:
|
||||
"""常量默认值测试."""
|
||||
|
||||
def test_default_cover_time(self):
|
||||
"""默认抽帧时间为 1.0 秒."""
|
||||
assert DEFAULT_COVER_TIME == 1.0
|
||||
|
||||
def test_default_dimensions(self):
|
||||
"""默认封面尺寸 1080x1920 (竖屏)."""
|
||||
assert DEFAULT_COVER_WIDTH == 1080
|
||||
assert DEFAULT_COVER_HEIGHT == 1920
|
||||
|
||||
def test_default_quality(self):
|
||||
"""默认质量为 5 (JPEG q:v, 越小越好)."""
|
||||
assert DEFAULT_COVER_QUALITY == 5
|
||||
|
||||
def test_smart_cover_frame_count(self):
|
||||
"""智能封面默认抽 3 帧."""
|
||||
assert SMART_COVER_FRAME_COUNT == 3
|
||||
|
||||
|
||||
class TestExtractFrameCommand:
|
||||
"""extract_frame 命令构建测试."""
|
||||
|
||||
def _probe_video_info_mock(self, duration=10.0):
|
||||
"""创建 probe_video_info 的 mock."""
|
||||
return {"duration": duration, "width": 1920, "height": 1080, "fps": 25.0}
|
||||
|
||||
def test_default_params_command(self, tmp_path):
|
||||
"""默认参数下 FFmpeg 命令正确."""
|
||||
video_file = tmp_path / "test.mp4"
|
||||
video_file.write_bytes(b"fake")
|
||||
output_file = tmp_path / "cover.jpg"
|
||||
|
||||
with (
|
||||
patch(
|
||||
"video_processing.cover_generator.probe_video_info",
|
||||
return_value=self._probe_video_info_mock(),
|
||||
),
|
||||
patch("video_processing.cover_generator.run_ffmpeg") as mock_run,
|
||||
):
|
||||
# 让 output_path 在 run_ffmpeg 后存在
|
||||
def fake_run(cmd):
|
||||
output_file.write_bytes(b"fake jpg")
|
||||
|
||||
mock_run.side_effect = fake_run
|
||||
result = CoverGenerator.extract_frame(str(video_file), str(output_file))
|
||||
|
||||
assert result == Path(output_file)
|
||||
mock_run.assert_called_once()
|
||||
cmd = mock_run.call_args[0][0]
|
||||
|
||||
# 基本结构
|
||||
assert cmd[0].endswith("ffmpeg") or "ffmpeg" in cmd[0]
|
||||
assert "-y" in cmd
|
||||
assert "-vframes" in cmd
|
||||
assert cmd[cmd.index("-vframes") + 1] == "1"
|
||||
assert "-f" in cmd
|
||||
assert "mjpeg" in cmd[cmd.index("-f") + 1]
|
||||
|
||||
# 时间点
|
||||
ss_idx = cmd.index("-ss")
|
||||
assert float(cmd[ss_idx + 1]) == pytest.approx(DEFAULT_COVER_TIME, abs=0.001)
|
||||
|
||||
# 输入文件
|
||||
i_idx = cmd.index("-i")
|
||||
assert cmd[i_idx + 1] == str(video_file)
|
||||
|
||||
# 输出文件
|
||||
assert cmd[-1] == str(output_file)
|
||||
|
||||
# scale + crop 滤镜
|
||||
vf_idx = cmd.index("-vf")
|
||||
vf_value = cmd[vf_idx + 1]
|
||||
assert "scale=" in vf_value
|
||||
assert "crop=" in vf_value
|
||||
assert "force_original_aspect_ratio=increase" in vf_value
|
||||
|
||||
def test_custom_time(self, tmp_path):
|
||||
"""自定义抽帧时间点."""
|
||||
video_file = tmp_path / "test.mp4"
|
||||
video_file.write_bytes(b"fake")
|
||||
output_file = tmp_path / "cover.jpg"
|
||||
|
||||
with (
|
||||
patch(
|
||||
"video_processing.cover_generator.probe_video_info",
|
||||
return_value=self._probe_video_info_mock(duration=30.0),
|
||||
),
|
||||
patch("video_processing.cover_generator.run_ffmpeg") as mock_run,
|
||||
):
|
||||
|
||||
def fake_run(cmd):
|
||||
output_file.write_bytes(b"fake jpg")
|
||||
|
||||
mock_run.side_effect = fake_run
|
||||
CoverGenerator.extract_frame(str(video_file), str(output_file), time_sec=5.5)
|
||||
|
||||
cmd = mock_run.call_args[0][0]
|
||||
ss_idx = cmd.index("-ss")
|
||||
assert float(cmd[ss_idx + 1]) == pytest.approx(5.5, abs=0.001)
|
||||
|
||||
def test_custom_dimensions(self, tmp_path):
|
||||
"""自定义输出尺寸."""
|
||||
video_file = tmp_path / "test.mp4"
|
||||
video_file.write_bytes(b"fake")
|
||||
output_file = tmp_path / "cover.jpg"
|
||||
|
||||
with (
|
||||
patch(
|
||||
"video_processing.cover_generator.probe_video_info",
|
||||
return_value=self._probe_video_info_mock(),
|
||||
),
|
||||
patch("video_processing.cover_generator.run_ffmpeg") as mock_run,
|
||||
):
|
||||
|
||||
def fake_run(cmd):
|
||||
output_file.write_bytes(b"fake jpg")
|
||||
|
||||
mock_run.side_effect = fake_run
|
||||
CoverGenerator.extract_frame(str(video_file), str(output_file), width=1920, height=1080)
|
||||
|
||||
cmd = mock_run.call_args[0][0]
|
||||
vf_idx = cmd.index("-vf")
|
||||
vf_value = cmd[vf_idx + 1]
|
||||
assert "scale=1920:1080:" in vf_value
|
||||
assert "crop=1920:1080" in vf_value
|
||||
|
||||
def test_custom_quality(self, tmp_path):
|
||||
"""自定义 JPEG 质量."""
|
||||
video_file = tmp_path / "test.mp4"
|
||||
video_file.write_bytes(b"fake")
|
||||
output_file = tmp_path / "cover.jpg"
|
||||
|
||||
with (
|
||||
patch(
|
||||
"video_processing.cover_generator.probe_video_info",
|
||||
return_value=self._probe_video_info_mock(),
|
||||
),
|
||||
patch("video_processing.cover_generator.run_ffmpeg") as mock_run,
|
||||
):
|
||||
|
||||
def fake_run(cmd):
|
||||
output_file.write_bytes(b"fake jpg")
|
||||
|
||||
mock_run.side_effect = fake_run
|
||||
CoverGenerator.extract_frame(str(video_file), str(output_file), quality=2)
|
||||
|
||||
cmd = mock_run.call_args[0][0]
|
||||
q_idx = cmd.index("-q:v")
|
||||
assert cmd[q_idx + 1] == "2"
|
||||
|
||||
def test_time_exceeds_duration_clamps_to_midpoint(self, tmp_path):
|
||||
"""抽帧时间超过视频时长时,钳制到中间帧."""
|
||||
video_file = tmp_path / "test.mp4"
|
||||
video_file.write_bytes(b"fake")
|
||||
output_file = tmp_path / "cover.jpg"
|
||||
|
||||
with (
|
||||
patch(
|
||||
"video_processing.cover_generator.probe_video_info",
|
||||
return_value=self._probe_video_info_mock(duration=5.0),
|
||||
),
|
||||
patch("video_processing.cover_generator.run_ffmpeg") as mock_run,
|
||||
):
|
||||
|
||||
def fake_run(cmd):
|
||||
output_file.write_bytes(b"fake jpg")
|
||||
|
||||
mock_run.side_effect = fake_run
|
||||
CoverGenerator.extract_frame(str(video_file), str(output_file), time_sec=10.0)
|
||||
|
||||
cmd = mock_run.call_args[0][0]
|
||||
ss_idx = cmd.index("-ss")
|
||||
# 钳制到 duration/2 = 2.5
|
||||
assert float(cmd[ss_idx + 1]) == pytest.approx(2.5, abs=0.001)
|
||||
|
||||
def test_negative_time_clamps_to_zero(self, tmp_path):
|
||||
"""负时间钳制到 0."""
|
||||
video_file = tmp_path / "test.mp4"
|
||||
video_file.write_bytes(b"fake")
|
||||
output_file = tmp_path / "cover.jpg"
|
||||
|
||||
with (
|
||||
patch(
|
||||
"video_processing.cover_generator.probe_video_info",
|
||||
return_value=self._probe_video_info_mock(duration=10.0),
|
||||
),
|
||||
patch("video_processing.cover_generator.run_ffmpeg") as mock_run,
|
||||
):
|
||||
|
||||
def fake_run(cmd):
|
||||
output_file.write_bytes(b"fake jpg")
|
||||
|
||||
mock_run.side_effect = fake_run
|
||||
CoverGenerator.extract_frame(str(video_file), str(output_file), time_sec=-2.0)
|
||||
|
||||
cmd = mock_run.call_args[0][0]
|
||||
ss_idx = cmd.index("-ss")
|
||||
assert float(cmd[ss_idx + 1]) == pytest.approx(0.0, abs=0.001)
|
||||
|
||||
def test_time_equals_duration_clamps_to_midpoint(self, tmp_path):
|
||||
"""时间点等于时长时钳制到中间帧."""
|
||||
video_file = tmp_path / "test.mp4"
|
||||
video_file.write_bytes(b"fake")
|
||||
output_file = tmp_path / "cover.jpg"
|
||||
|
||||
with (
|
||||
patch(
|
||||
"video_processing.cover_generator.probe_video_info",
|
||||
return_value=self._probe_video_info_mock(duration=10.0),
|
||||
),
|
||||
patch("video_processing.cover_generator.run_ffmpeg") as mock_run,
|
||||
):
|
||||
|
||||
def fake_run(cmd):
|
||||
output_file.write_bytes(b"fake jpg")
|
||||
|
||||
mock_run.side_effect = fake_run
|
||||
CoverGenerator.extract_frame(str(video_file), str(output_file), time_sec=10.0)
|
||||
|
||||
cmd = mock_run.call_args[0][0]
|
||||
ss_idx = cmd.index("-ss")
|
||||
assert float(cmd[ss_idx + 1]) == pytest.approx(5.0, abs=0.001)
|
||||
|
||||
def test_zero_duration_video(self, tmp_path):
|
||||
"""视频时长为 0 时的行为(不钳制,用原始时间)."""
|
||||
video_file = tmp_path / "test.mp4"
|
||||
video_file.write_bytes(b"fake")
|
||||
output_file = tmp_path / "cover.jpg"
|
||||
|
||||
with (
|
||||
patch(
|
||||
"video_processing.cover_generator.probe_video_info",
|
||||
return_value=self._probe_video_info_mock(duration=0.0),
|
||||
),
|
||||
patch("video_processing.cover_generator.run_ffmpeg") as mock_run,
|
||||
):
|
||||
|
||||
def fake_run(cmd):
|
||||
output_file.write_bytes(b"fake jpg")
|
||||
|
||||
mock_run.side_effect = fake_run
|
||||
CoverGenerator.extract_frame(str(video_file), str(output_file), time_sec=0.5)
|
||||
|
||||
cmd = mock_run.call_args[0][0]
|
||||
ss_idx = cmd.index("-ss")
|
||||
assert float(cmd[ss_idx + 1]) == pytest.approx(0.5, abs=0.001)
|
||||
|
||||
def test_video_not_found_raises(self, tmp_path):
|
||||
"""视频文件不存在时抛出 FileNotFoundError."""
|
||||
output_file = tmp_path / "cover.jpg"
|
||||
|
||||
with pytest.raises(FileNotFoundError):
|
||||
CoverGenerator.extract_frame(str(tmp_path / "nonexistent.mp4"), str(output_file))
|
||||
|
||||
def test_output_creates_parent_dir(self, tmp_path):
|
||||
"""输出目录不存在时自动创建."""
|
||||
video_file = tmp_path / "test.mp4"
|
||||
video_file.write_bytes(b"fake")
|
||||
out_dir = tmp_path / "deep" / "nested"
|
||||
output_file = out_dir / "cover.jpg"
|
||||
|
||||
with (
|
||||
patch(
|
||||
"video_processing.cover_generator.probe_video_info",
|
||||
return_value=self._probe_video_info_mock(),
|
||||
),
|
||||
patch("video_processing.cover_generator.run_ffmpeg") as mock_run,
|
||||
):
|
||||
|
||||
def fake_run(cmd):
|
||||
output_file.write_bytes(b"fake jpg")
|
||||
|
||||
mock_run.side_effect = fake_run
|
||||
CoverGenerator.extract_frame(str(video_file), str(output_file))
|
||||
|
||||
assert out_dir.exists()
|
||||
assert out_dir.is_dir()
|
||||
|
||||
def test_ffmpeg_failure_propagates(self, tmp_path):
|
||||
"""FFmpeg 失败时异常向上传递."""
|
||||
video_file = tmp_path / "test.mp4"
|
||||
video_file.write_bytes(b"fake")
|
||||
output_file = tmp_path / "cover.jpg"
|
||||
|
||||
with (
|
||||
patch(
|
||||
"video_processing.cover_generator.probe_video_info",
|
||||
return_value=self._probe_video_info_mock(),
|
||||
),
|
||||
patch(
|
||||
"video_processing.cover_generator.run_ffmpeg",
|
||||
side_effect=RuntimeError("FFmpeg error"),
|
||||
),
|
||||
):
|
||||
with pytest.raises(RuntimeError, match="FFmpeg error"):
|
||||
CoverGenerator.extract_frame(str(video_file), str(output_file))
|
||||
|
||||
|
||||
class TestSmartCoverTimePoints:
|
||||
"""智能封面时间点计算测试."""
|
||||
|
||||
def test_single_frame_falls_back_to_default(self, tmp_path):
|
||||
"""只有 1 帧时退化为普通抽帧(取 DEFAULT_COVER_TIME 和 midpoint 中较小值)."""
|
||||
video_file = tmp_path / "test.mp4"
|
||||
video_file.write_bytes(b"fake")
|
||||
output_file = tmp_path / "cover.jpg"
|
||||
|
||||
with (
|
||||
patch(
|
||||
"video_processing.cover_generator.probe_video_info",
|
||||
return_value={"duration": 20.0},
|
||||
),
|
||||
patch("video_processing.cover_generator.run_ffmpeg") as mock_run,
|
||||
):
|
||||
|
||||
def fake_run(cmd):
|
||||
output_file.write_bytes(b"fake jpg")
|
||||
|
||||
mock_run.side_effect = fake_run
|
||||
# frame_count=1 时退化为普通抽帧
|
||||
CoverGenerator.extract_smart_cover(str(video_file), str(output_file), frame_count=1)
|
||||
|
||||
# 只调用一次(退化路径)
|
||||
assert mock_run.call_count == 1
|
||||
cmd = mock_run.call_args[0][0]
|
||||
ss_idx = cmd.index("-ss")
|
||||
# min(DEFAULT_COVER_TIME=1.0, duration/2=10.0) = 1.0
|
||||
assert float(cmd[ss_idx + 1]) == pytest.approx(1.0, abs=0.001)
|
||||
|
||||
def test_zero_duration_falls_back(self, tmp_path):
|
||||
"""视频时长为 0 时退化为普通抽帧."""
|
||||
video_file = tmp_path / "test.mp4"
|
||||
video_file.write_bytes(b"fake")
|
||||
output_file = tmp_path / "cover.jpg"
|
||||
|
||||
with (
|
||||
patch(
|
||||
"video_processing.cover_generator.probe_video_info",
|
||||
return_value={"duration": 0.0},
|
||||
),
|
||||
patch("video_processing.cover_generator.run_ffmpeg") as mock_run,
|
||||
):
|
||||
|
||||
def fake_run(cmd):
|
||||
output_file.write_bytes(b"fake jpg")
|
||||
|
||||
mock_run.side_effect = fake_run
|
||||
CoverGenerator.extract_smart_cover(str(video_file), str(output_file))
|
||||
|
||||
# 只调用一次(退化路径)
|
||||
assert mock_run.call_count == 1
|
||||
|
||||
def test_three_frames_uniform_distribution(self, tmp_path):
|
||||
"""3 帧均匀分布在 5%~95% 区间."""
|
||||
video_file = tmp_path / "test.mp4"
|
||||
video_file.write_bytes(b"fake")
|
||||
output_file = tmp_path / "cover.jpg"
|
||||
|
||||
call_times = []
|
||||
|
||||
with (
|
||||
patch(
|
||||
"video_processing.cover_generator.probe_video_info",
|
||||
return_value={"duration": 100.0},
|
||||
),
|
||||
patch("video_processing.cover_generator.run_ffmpeg") as mock_run,
|
||||
):
|
||||
|
||||
def fake_run(cmd):
|
||||
# 记录抽帧时间
|
||||
ss_idx = cmd.index("-ss")
|
||||
call_times.append(float(cmd[ss_idx + 1]))
|
||||
# 在输出路径写文件
|
||||
output_arg = cmd[-1]
|
||||
Path(output_arg).parent.mkdir(parents=True, exist_ok=True)
|
||||
# 不同文件大小,让第三帧"最清晰"
|
||||
idx = len(call_times) - 1
|
||||
size = 1000 * (idx + 1) # 递增的文件大小
|
||||
Path(output_arg).write_bytes(b"x" * size)
|
||||
|
||||
mock_run.side_effect = fake_run
|
||||
CoverGenerator.extract_smart_cover(str(video_file), str(output_file))
|
||||
|
||||
# 3 帧:5%、50%、95%
|
||||
assert len(call_times) == 3
|
||||
assert call_times[0] == pytest.approx(5.0, abs=0.1) # 5%
|
||||
assert call_times[1] == pytest.approx(50.0, abs=0.1) # 50%
|
||||
assert call_times[2] == pytest.approx(95.0, abs=0.1) # 95%
|
||||
|
||||
def test_five_frames_distribution(self, tmp_path):
|
||||
"""5 帧均匀分布."""
|
||||
video_file = tmp_path / "test.mp4"
|
||||
video_file.write_bytes(b"fake")
|
||||
output_file = tmp_path / "cover.jpg"
|
||||
|
||||
call_times = []
|
||||
|
||||
with (
|
||||
patch(
|
||||
"video_processing.cover_generator.probe_video_info",
|
||||
return_value={"duration": 100.0},
|
||||
),
|
||||
patch("video_processing.cover_generator.run_ffmpeg") as mock_run,
|
||||
):
|
||||
|
||||
def fake_run(cmd):
|
||||
ss_idx = cmd.index("-ss")
|
||||
call_times.append(float(cmd[ss_idx + 1]))
|
||||
output_arg = cmd[-1]
|
||||
Path(output_arg).parent.mkdir(parents=True, exist_ok=True)
|
||||
idx = len(call_times) - 1
|
||||
Path(output_arg).write_bytes(b"x" * (1000 * (idx + 1)))
|
||||
|
||||
mock_run.side_effect = fake_run
|
||||
CoverGenerator.extract_smart_cover(str(video_file), str(output_file), frame_count=5)
|
||||
|
||||
assert len(call_times) == 5
|
||||
# step = (95-5) / (5-1) = 22.5
|
||||
# times: 5, 27.5, 50, 72.5, 95
|
||||
assert call_times[0] == pytest.approx(5.0, abs=0.1)
|
||||
assert call_times[1] == pytest.approx(27.5, abs=0.1)
|
||||
assert call_times[2] == pytest.approx(50.0, abs=0.1)
|
||||
assert call_times[3] == pytest.approx(72.5, abs=0.1)
|
||||
assert call_times[4] == pytest.approx(95.0, abs=0.1)
|
||||
|
||||
def test_selects_largest_file_as_best(self, tmp_path):
|
||||
"""选择文件最大的帧作为最佳封面(清晰度近似)."""
|
||||
video_file = tmp_path / "test.mp4"
|
||||
video_file.write_bytes(b"fake")
|
||||
output_file = tmp_path / "cover.jpg"
|
||||
|
||||
sizes = [5000, 15000, 8000] # 第二帧最大
|
||||
|
||||
with (
|
||||
patch(
|
||||
"video_processing.cover_generator.probe_video_info",
|
||||
return_value={"duration": 100.0},
|
||||
),
|
||||
patch("video_processing.cover_generator.run_ffmpeg") as mock_run,
|
||||
):
|
||||
call_idx = [0]
|
||||
|
||||
def fake_run(cmd):
|
||||
output_arg = cmd[-1]
|
||||
Path(output_arg).parent.mkdir(parents=True, exist_ok=True)
|
||||
idx = call_idx[0]
|
||||
Path(output_arg).write_bytes(b"x" * sizes[idx])
|
||||
call_idx[0] += 1
|
||||
|
||||
mock_run.side_effect = fake_run
|
||||
result = CoverGenerator.extract_smart_cover(str(video_file), str(output_file))
|
||||
|
||||
# 第二帧(索引1)应该是最佳
|
||||
assert result == output_file
|
||||
# 输出文件大小应等于第二帧大小
|
||||
assert output_file.stat().st_size == 15000
|
||||
|
||||
|
||||
class TestProcessCustomCover:
|
||||
"""自定义封面处理测试."""
|
||||
|
||||
def test_custom_cover_resize_command(self, tmp_path):
|
||||
"""自定义封面调整尺寸命令正确."""
|
||||
input_file = tmp_path / "upload.jpg"
|
||||
input_file.write_bytes(b"fake")
|
||||
output_file = tmp_path / "cover.jpg"
|
||||
|
||||
with patch("video_processing.cover_generator.run_ffmpeg") as mock_run:
|
||||
|
||||
def fake_run(cmd):
|
||||
output_file.write_bytes(b"fake jpg")
|
||||
|
||||
mock_run.side_effect = fake_run
|
||||
CoverGenerator.process_custom_cover(str(input_file), str(output_file))
|
||||
|
||||
mock_run.assert_called_once()
|
||||
cmd = mock_run.call_args[0][0]
|
||||
|
||||
assert "-i" in cmd
|
||||
assert cmd[cmd.index("-i") + 1] == str(input_file)
|
||||
assert cmd[-1] == str(output_file)
|
||||
|
||||
# scale + crop
|
||||
vf_idx = cmd.index("-vf")
|
||||
vf_value = cmd[vf_idx + 1]
|
||||
assert "scale=" in vf_value
|
||||
assert "crop=" in vf_value
|
||||
|
||||
def test_custom_cover_not_found_raises(self, tmp_path):
|
||||
"""自定义封面文件不存在时抛出 FileNotFoundError."""
|
||||
output_file = tmp_path / "cover.jpg"
|
||||
|
||||
with pytest.raises(FileNotFoundError):
|
||||
CoverGenerator.process_custom_cover(str(tmp_path / "nonexistent.jpg"), str(output_file))
|
||||
|
||||
def test_custom_cover_custom_dimensions(self, tmp_path):
|
||||
"""自定义封面自定义输出尺寸."""
|
||||
input_file = tmp_path / "upload.jpg"
|
||||
input_file.write_bytes(b"fake")
|
||||
output_file = tmp_path / "cover.jpg"
|
||||
|
||||
with patch("video_processing.cover_generator.run_ffmpeg") as mock_run:
|
||||
|
||||
def fake_run(cmd):
|
||||
output_file.write_bytes(b"fake jpg")
|
||||
|
||||
mock_run.side_effect = fake_run
|
||||
CoverGenerator.process_custom_cover(str(input_file), str(output_file), width=800, height=600)
|
||||
|
||||
cmd = mock_run.call_args[0][0]
|
||||
vf_idx = cmd.index("-vf")
|
||||
vf_value = cmd[vf_idx + 1]
|
||||
assert "scale=800:600:" in vf_value
|
||||
assert "crop=800:600" in vf_value
|
||||
@@ -1,223 +0,0 @@
|
||||
"""去重纯算法测试 — hamming_distance + histogram_similarity + VideoFingerprint."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import sys
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
|
||||
# 模块级mock有副作用的依赖(纯算法测试不需要db/celery/cv2)
|
||||
# 注意:必须在 import dedup 前全部 mock 完,避免链式导入触发db连接
|
||||
|
||||
# cv2(视频处理依赖,纯算法测试不需要)
|
||||
sys.modules["cv2"] = MagicMock()
|
||||
# 模块级mock worker_app.db(dedup模块import时会触发数据库初始化,纯算法测试不需要)
|
||||
sys.modules["worker_app.db"] = MagicMock()
|
||||
sys.modules["worker_app.db"].SessionLocal = MagicMock()
|
||||
|
||||
# celery 及其子模块
|
||||
_mock_celery = MagicMock()
|
||||
_mock_celery.Task = MagicMock
|
||||
_mock_celery.Celery = MagicMock
|
||||
sys.modules["celery"] = _mock_celery
|
||||
|
||||
# sqlalchemy 作为包结构 mock
|
||||
_mock_sqla = MagicMock()
|
||||
_mock_sqla.__path__ = []
|
||||
_mock_sqla.__package__ = "sqlalchemy"
|
||||
_mock_sqla_orm = MagicMock()
|
||||
_mock_sqla_orm.__path__ = []
|
||||
_mock_sqla_orm.Session = MagicMock
|
||||
_mock_sqla_engine = MagicMock()
|
||||
sys.modules["sqlalchemy"] = _mock_sqla
|
||||
sys.modules["sqlalchemy.orm"] = _mock_sqla_orm
|
||||
sys.modules["sqlalchemy.engine"] = _mock_sqla_engine
|
||||
sys.modules["sqlalchemy.ext"] = MagicMock()
|
||||
sys.modules["sqlalchemy.ext.declarative"] = MagicMock()
|
||||
|
||||
# worker_app 及其子模块(避免导入时触发数据库连接)
|
||||
_mock_worker_app = MagicMock()
|
||||
_mock_worker_app.__path__ = []
|
||||
_mock_worker_db = MagicMock()
|
||||
_mock_worker_db.SessionLocal = MagicMock()
|
||||
_mock_worker_celery = MagicMock()
|
||||
_mock_worker_celery.celery_app = MagicMock()
|
||||
_mock_worker_core = MagicMock()
|
||||
_mock_worker_core.__path__ = []
|
||||
_mock_worker_config = MagicMock()
|
||||
_mock_worker_config.get_settings = MagicMock(return_value=MagicMock())
|
||||
sys.modules["worker_app"] = _mock_worker_app
|
||||
sys.modules["worker_app.db"] = _mock_worker_db
|
||||
sys.modules["worker_app.celery_app"] = _mock_worker_celery
|
||||
sys.modules["worker_app.core"] = _mock_worker_core
|
||||
sys.modules["worker_app.core.config"] = _mock_worker_config
|
||||
|
||||
# packages.adapters.sqlalchemy_impl(整个包mock掉)
|
||||
_mock_sqla_impl = MagicMock()
|
||||
_mock_sqla_impl.__path__ = []
|
||||
sys.modules["packages.adapters.sqlalchemy_impl"] = _mock_sqla_impl
|
||||
sys.modules["packages.adapters.sqlalchemy_impl.generated_video_repository"] = MagicMock()
|
||||
sys.modules["packages.adapters.sqlalchemy_impl.schema_guard"] = MagicMock()
|
||||
|
||||
# packages.shared
|
||||
_mock_packages_shared = MagicMock()
|
||||
_mock_packages_shared.__path__ = []
|
||||
_mock_shared_config = MagicMock()
|
||||
_mock_shared_config.get_shared_settings = MagicMock(return_value=MagicMock())
|
||||
sys.modules["packages.shared"] = _mock_packages_shared
|
||||
sys.modules["packages.shared.config"] = _mock_shared_config
|
||||
sys.modules["packages.shared.storage"] = MagicMock()
|
||||
|
||||
from video_processing.dedup import ( # noqa: E402
|
||||
VideoDeduplicator,
|
||||
VideoFingerprint,
|
||||
hamming_distance,
|
||||
)
|
||||
|
||||
|
||||
class TestHammingDistance:
|
||||
"""hamming_distance 汉明距离计算测试."""
|
||||
|
||||
def test_identical_hashes_zero(self):
|
||||
"""相同哈希距离为0."""
|
||||
assert hamming_distance("ff", "ff") == 0
|
||||
assert hamming_distance("00", "00") == 0
|
||||
|
||||
def test_all_different(self):
|
||||
"""全不同的8bit哈希距离为8."""
|
||||
assert hamming_distance("00", "ff") == 8
|
||||
|
||||
def test_single_bit_diff(self):
|
||||
"""1个bit不同."""
|
||||
# 0x01 = 00000001, 0x00 = 00000000 → 1 bit不同
|
||||
assert hamming_distance("01", "00") == 1
|
||||
|
||||
def test_four_bits_diff(self):
|
||||
"""4个bit不同."""
|
||||
# 0x0F = 00001111, 0xF0 = 11110000 → 8 bits都不同
|
||||
assert hamming_distance("0f", "f0") == 8
|
||||
|
||||
def test_longer_hashes(self):
|
||||
"""更长的哈希(如64-bit pHash)."""
|
||||
# 两个完全不同的64-bit哈希
|
||||
assert hamming_distance("0000000000000000", "ffffffffffffffff") == 64
|
||||
|
||||
def test_partial_difference(self):
|
||||
"""部分bit不同."""
|
||||
# a = 1010, 5 = 0101 → 4 bits不同(每个hex digit)
|
||||
assert hamming_distance("aa", "55") == 8
|
||||
|
||||
def test_case_insensitive(self):
|
||||
"""十六进制不区分大小写."""
|
||||
assert hamming_distance("FF", "ff") == 0
|
||||
assert hamming_distance("AbC123", "aBc123") == 0
|
||||
|
||||
def test_different_length_hashes(self):
|
||||
"""不同长度的哈希(短的前补零)."""
|
||||
# "ff" = 0xff = 255, "0ff" = 0x0ff = 255
|
||||
# int("ff", 16) = 255, int("0ff", 16) = 255
|
||||
assert hamming_distance("ff", "0ff") == 0
|
||||
|
||||
|
||||
class TestVideoFingerprint:
|
||||
"""VideoFingerprint 数据结构测试."""
|
||||
|
||||
def test_to_dict_contains_all_fields(self):
|
||||
"""to_dict返回完整字典."""
|
||||
fp = VideoFingerprint(
|
||||
md5="abc123",
|
||||
keyframe_phashes=["hash1", "hash2"],
|
||||
color_histograms=[[0.1, 0.2], [0.3, 0.4]],
|
||||
duration=30.5,
|
||||
resolution=(1920, 1080),
|
||||
)
|
||||
d = fp.to_dict()
|
||||
assert d["md5"] == "abc123"
|
||||
assert d["keyframe_phashes"] == ["hash1", "hash2"]
|
||||
assert d["duration"] == 30.5
|
||||
assert d["resolution"] == [1920, 1080]
|
||||
assert "color_histograms" in d
|
||||
|
||||
def test_empty_phashes(self):
|
||||
"""空关键帧列表."""
|
||||
fp = VideoFingerprint(
|
||||
md5="test",
|
||||
keyframe_phashes=[],
|
||||
color_histograms=[],
|
||||
duration=0.0,
|
||||
resolution=(0, 0),
|
||||
)
|
||||
d = fp.to_dict()
|
||||
assert d["keyframe_phashes"] == []
|
||||
assert d["color_histograms"] == []
|
||||
|
||||
|
||||
class TestAverageHistogramSimilarity:
|
||||
"""_average_histogram_similarity 直方图相似度测试."""
|
||||
|
||||
def test_identical_histograms(self):
|
||||
"""完全相同的直方图相似度为1.0."""
|
||||
hist = [[0.5, 0.5, 0.0], [0.3, 0.4, 0.3]]
|
||||
sim = VideoDeduplicator._average_histogram_similarity(hist, hist)
|
||||
assert sim == pytest.approx(1.0)
|
||||
|
||||
def test_empty_first_list(self):
|
||||
"""第一组为空返回0."""
|
||||
sim = VideoDeduplicator._average_histogram_similarity([], [[0.5, 0.5]])
|
||||
assert sim == 0.0
|
||||
|
||||
def test_empty_second_list(self):
|
||||
"""第二组为空返回0."""
|
||||
sim = VideoDeduplicator._average_histogram_similarity([[0.5, 0.5]], [])
|
||||
assert sim == 0.0
|
||||
|
||||
def test_both_empty(self):
|
||||
"""两组都为空返回0."""
|
||||
sim = VideoDeduplicator._average_histogram_similarity([], [])
|
||||
assert sim == 0.0
|
||||
|
||||
def test_orthogonal_histograms(self):
|
||||
"""正交直方图相似度为0."""
|
||||
# [1, 0] 和 [0, 1] 正交
|
||||
sim = VideoDeduplicator._average_histogram_similarity([[1.0, 0.0]], [[0.0, 1.0]])
|
||||
assert sim == pytest.approx(0.0)
|
||||
|
||||
def test_partial_similarity(self):
|
||||
"""部分相似."""
|
||||
# [1, 1] 和 [1, 0] 的余弦相似度 = 1/√2 ≈ 0.707
|
||||
sim = VideoDeduplicator._average_histogram_similarity([[1.0, 1.0]], [[1.0, 0.0]])
|
||||
assert sim == pytest.approx(1.0 / (2**0.5), rel=0.01)
|
||||
|
||||
def test_multiple_frames_best_match(self):
|
||||
"""多帧时取最佳匹配."""
|
||||
# 第一帧完全不同,第二帧完全相同 → 平均 best = (0 + 1) / 2 = 0.5
|
||||
sim = VideoDeduplicator._average_histogram_similarity(
|
||||
[[1.0, 0.0], [0.0, 1.0]],
|
||||
[[0.0, 1.0]], # 只有一帧,和第一帧0相似,和第二帧1相似
|
||||
)
|
||||
# 第一帧最佳匹配=0,第二帧最佳匹配=1,平均=0.5
|
||||
assert sim == pytest.approx(0.5)
|
||||
|
||||
def test_zero_norm_histogram_skipped(self):
|
||||
"""零范数直方图被跳过."""
|
||||
sim = VideoDeduplicator._average_histogram_similarity([[0.0, 0.0]], [[1.0, 1.0]])
|
||||
# 第一组的零范数被跳过,similarities为空,返回0
|
||||
assert sim == 0.0
|
||||
|
||||
def test_different_length_histograms(self):
|
||||
"""不同长度的直方图取最小长度对齐."""
|
||||
sim = VideoDeduplicator._average_histogram_similarity(
|
||||
[[1.0, 1.0, 0.0, 0.0]], # 4维
|
||||
[[1.0, 1.0]], # 2维
|
||||
)
|
||||
# 对齐到前2维,都是[1,1],相似度1.0
|
||||
assert sim == pytest.approx(1.0)
|
||||
|
||||
def test_similarity_in_zero_one_range(self):
|
||||
"""相似度在[0, 1]范围内."""
|
||||
hist_a = [np.random.rand(96).tolist() for _ in range(5)]
|
||||
hist_b = [np.random.rand(96).tolist() for _ in range(5)]
|
||||
sim = VideoDeduplicator._average_histogram_similarity(hist_a, hist_b)
|
||||
assert 0.0 <= sim <= 1.0
|
||||
@@ -1,523 +0,0 @@
|
||||
"""领域实体测试 — Project + Asset + AssetStatus + IngestJob."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
|
||||
from packages.domain.entities import (
|
||||
Asset,
|
||||
AssetLibrary,
|
||||
AssetStatus,
|
||||
ClassificationStatus,
|
||||
IngestJob,
|
||||
Project,
|
||||
User,
|
||||
)
|
||||
|
||||
|
||||
class TestUser:
|
||||
"""User 实体测试."""
|
||||
|
||||
def test_create_user(self):
|
||||
"""创建用户."""
|
||||
user = User(
|
||||
id="u1",
|
||||
email="test@example.com",
|
||||
display_name="测试用户",
|
||||
username="testuser",
|
||||
)
|
||||
assert user.id == "u1"
|
||||
assert user.email == "test@example.com"
|
||||
assert user.display_name == "测试用户"
|
||||
assert user.username == "testuser"
|
||||
|
||||
def test_default_subscription_free(self):
|
||||
"""默认订阅免费版."""
|
||||
user = User(id="u1", email="a@b.com", display_name="A")
|
||||
assert user.subscription_plan == "free"
|
||||
assert user.subscription_status == "active"
|
||||
assert user.max_projects == 3
|
||||
assert user.max_storage_gb == 10
|
||||
|
||||
def test_default_not_admin(self):
|
||||
"""默认不是管理员."""
|
||||
user = User(id="u1", email="a@b.com", display_name="A")
|
||||
assert user.is_admin is False
|
||||
|
||||
|
||||
class TestProject:
|
||||
"""Project 实体测试."""
|
||||
|
||||
def test_create_project_success(self):
|
||||
"""创建项目成功."""
|
||||
project = Project.create(
|
||||
owner_user_id="user_1",
|
||||
name="我的项目",
|
||||
description="测试项目描述",
|
||||
)
|
||||
assert project.id is not None
|
||||
assert len(project.id) == 32 # uuid hex
|
||||
assert project.owner_user_id == "user_1"
|
||||
assert project.name == "我的项目"
|
||||
assert project.description == "测试项目描述"
|
||||
assert project.shared_users == []
|
||||
|
||||
def test_create_project_name_stripped(self):
|
||||
"""项目名称首尾空白被去除."""
|
||||
project = Project.create(owner_user_id="u1", name=" 测试项目 ")
|
||||
assert project.name == "测试项目"
|
||||
|
||||
def test_create_project_empty_name_raises(self):
|
||||
"""空项目名抛异常."""
|
||||
with pytest.raises(ValueError, match="项目名称不能为空"):
|
||||
Project.create(owner_user_id="u1", name="")
|
||||
|
||||
def test_create_project_whitespace_name_raises(self):
|
||||
"""全空白项目名抛异常."""
|
||||
with pytest.raises(ValueError, match="项目名称不能为空"):
|
||||
Project.create(owner_user_id="u1", name=" ")
|
||||
|
||||
def test_is_owner_true(self):
|
||||
"""是项目所有者."""
|
||||
project = Project.create(owner_user_id="owner_1", name="项目")
|
||||
assert project.is_owner("owner_1") is True
|
||||
|
||||
def test_is_owner_false(self):
|
||||
"""不是项目所有者."""
|
||||
project = Project.create(owner_user_id="owner_1", name="项目")
|
||||
assert project.is_owner("other_user") is False
|
||||
|
||||
def test_is_shared_with_true(self):
|
||||
"""项目已共享给用户."""
|
||||
project = Project.create(owner_user_id="u1", name="项目")
|
||||
project.shared_users.append("user_2")
|
||||
assert project.is_shared_with("user_2") is True
|
||||
|
||||
def test_is_shared_with_false(self):
|
||||
"""项目未共享给用户."""
|
||||
project = Project.create(owner_user_id="u1", name="项目")
|
||||
assert project.is_shared_with("nobody") is False
|
||||
|
||||
def test_can_access_owner(self):
|
||||
"""所有者可以访问."""
|
||||
project = Project.create(owner_user_id="u1", name="项目")
|
||||
assert project.can_access("u1") is True
|
||||
|
||||
def test_can_access_shared_user(self):
|
||||
"""共享用户可以访问."""
|
||||
project = Project.create(owner_user_id="u1", name="项目")
|
||||
project.shared_users.append("u2")
|
||||
assert project.can_access("u2") is True
|
||||
|
||||
def test_can_access_outsider(self):
|
||||
"""无关用户不能访问."""
|
||||
project = Project.create(owner_user_id="u1", name="项目")
|
||||
assert project.can_access("stranger") is False
|
||||
|
||||
|
||||
class TestAssetLibrary:
|
||||
"""AssetLibrary 实体测试."""
|
||||
|
||||
def test_create_library(self):
|
||||
"""创建素材库."""
|
||||
from packages.domain.classification import AssetLibraryKind
|
||||
|
||||
lib = AssetLibrary.create(
|
||||
project_id="p1",
|
||||
name="默认库",
|
||||
kind=AssetLibraryKind.VIDEO,
|
||||
)
|
||||
assert lib.id is not None
|
||||
assert lib.project_id == "p1"
|
||||
assert lib.name == "默认库"
|
||||
assert lib.kind == AssetLibraryKind.VIDEO
|
||||
assert lib.asset_count == 0
|
||||
|
||||
def test_create_default_counts(self):
|
||||
"""默认素材数量和大小为0."""
|
||||
from packages.domain.classification import AssetLibraryKind
|
||||
|
||||
lib = AssetLibrary.create(
|
||||
project_id="p1",
|
||||
name="我的素材",
|
||||
kind=AssetLibraryKind.VOICE,
|
||||
)
|
||||
assert lib.asset_count == 0
|
||||
assert lib.total_size == 0
|
||||
|
||||
def test_empty_name_raises(self):
|
||||
"""空名称抛异常."""
|
||||
from packages.domain.classification import AssetLibraryKind
|
||||
|
||||
with pytest.raises(ValueError, match="素材库名称不能为空"):
|
||||
AssetLibrary.create(project_id="p1", name="", kind=AssetLibraryKind.VIDEO)
|
||||
|
||||
|
||||
class TestAssetStatus:
|
||||
"""AssetStatus 枚举兼容测试."""
|
||||
|
||||
def test_direct_values(self):
|
||||
"""直接枚举值."""
|
||||
assert AssetStatus.UPLOADING.value == "uploading"
|
||||
assert AssetStatus.READY.value == "ready"
|
||||
assert AssetStatus.PROCESSING.value == "processing"
|
||||
assert AssetStatus.ERROR.value == "error"
|
||||
assert AssetStatus.DELETED.value == "deleted"
|
||||
|
||||
def test_missing_uploaded_maps_to_ready(self):
|
||||
"""历史值 uploaded → READY."""
|
||||
assert AssetStatus("uploaded") == AssetStatus.READY
|
||||
|
||||
def test_missing_success_maps_to_ready(self):
|
||||
"""success → READY."""
|
||||
assert AssetStatus("success") == AssetStatus.READY
|
||||
|
||||
def test_missing_ok_maps_to_ready(self):
|
||||
"""ok → READY."""
|
||||
assert AssetStatus("ok") == AssetStatus.READY
|
||||
|
||||
def test_missing_done_maps_to_ready(self):
|
||||
"""done → READY."""
|
||||
assert AssetStatus("done") == AssetStatus.READY
|
||||
|
||||
def test_missing_upload_maps_to_uploading(self):
|
||||
"""upload → UPLOADING."""
|
||||
assert AssetStatus("upload") == AssetStatus.UPLOADING
|
||||
|
||||
def test_missing_uploading_start_maps_to_uploading(self):
|
||||
"""uploading_start → UPLOADING."""
|
||||
assert AssetStatus("uploading_start") == AssetStatus.UPLOADING
|
||||
|
||||
def test_missing_failed_maps_to_error(self):
|
||||
"""failed → ERROR."""
|
||||
assert AssetStatus("failed") == AssetStatus.ERROR
|
||||
|
||||
def test_missing_fail_maps_to_error(self):
|
||||
"""fail → ERROR."""
|
||||
assert AssetStatus("fail") == AssetStatus.ERROR
|
||||
|
||||
def test_missing_process_maps_to_processing(self):
|
||||
"""process → PROCESSING."""
|
||||
assert AssetStatus("process") == AssetStatus.PROCESSING
|
||||
|
||||
def test_missing_running_maps_to_processing(self):
|
||||
"""running → PROCESSING."""
|
||||
assert AssetStatus("running") == AssetStatus.PROCESSING
|
||||
|
||||
def test_unknown_value_fallback_to_ready(self):
|
||||
"""完全未知值 → READY兜底."""
|
||||
assert AssetStatus("completely_unknown") == AssetStatus.READY
|
||||
|
||||
def test_case_insensitive(self):
|
||||
"""大小写不敏感."""
|
||||
assert AssetStatus("UPLOADED") == AssetStatus.READY
|
||||
assert AssetStatus("Failed") == AssetStatus.ERROR
|
||||
|
||||
def test_whitespace_stripped(self):
|
||||
"""首尾空白被去除."""
|
||||
assert AssetStatus(" ready ") == AssetStatus.READY
|
||||
|
||||
def test_none_value_fallback(self):
|
||||
"""None值 → READY兜底(不抛异常)."""
|
||||
assert AssetStatus(None) == AssetStatus.READY
|
||||
|
||||
|
||||
class TestAsset:
|
||||
"""Asset 实体测试."""
|
||||
|
||||
def test_create_asset_success(self):
|
||||
"""创建素材成功."""
|
||||
asset = Asset.create(
|
||||
project_id="p1",
|
||||
library_id="lib1",
|
||||
name="测试视频.mp4",
|
||||
storage_key="projects/p1/assets/test.mp4",
|
||||
mime_type="video/mp4",
|
||||
file_size=1024000,
|
||||
duration=30.5,
|
||||
width=1920,
|
||||
height=1080,
|
||||
)
|
||||
assert asset.id is not None
|
||||
assert len(asset.id) == 32
|
||||
assert asset.project_id == "p1"
|
||||
assert asset.name == "测试视频.mp4"
|
||||
assert asset.storage_key == "projects/p1/assets/test.mp4"
|
||||
assert asset.mime_type == "video/mp4"
|
||||
assert asset.file_size == 1024000
|
||||
assert asset.duration == pytest.approx(30.5)
|
||||
assert asset.width == 1920
|
||||
assert asset.height == 1080
|
||||
assert asset.status == AssetStatus.UPLOADING
|
||||
assert asset.tag_ids == []
|
||||
|
||||
def test_file_type_video(self):
|
||||
"""video类型从mime_type推导."""
|
||||
asset = Asset.create(
|
||||
project_id="p1",
|
||||
library_id="lib1",
|
||||
name="v.mp4",
|
||||
storage_key="k",
|
||||
mime_type="video/mp4",
|
||||
)
|
||||
assert asset.file_type == "video"
|
||||
|
||||
def test_file_type_image(self):
|
||||
"""image类型."""
|
||||
asset = Asset.create(
|
||||
project_id="p1",
|
||||
library_id="lib1",
|
||||
name="i.jpg",
|
||||
storage_key="k",
|
||||
mime_type="image/jpeg",
|
||||
)
|
||||
assert asset.file_type == "image"
|
||||
|
||||
def test_file_type_audio(self):
|
||||
"""audio类型."""
|
||||
asset = Asset.create(
|
||||
project_id="p1",
|
||||
library_id="lib1",
|
||||
name="a.mp3",
|
||||
storage_key="k",
|
||||
mime_type="audio/mpeg",
|
||||
)
|
||||
assert asset.file_type == "audio"
|
||||
|
||||
def test_file_type_no_slash(self):
|
||||
"""mime_type没有斜杠时返回原值."""
|
||||
asset = Asset.create(
|
||||
project_id="p1",
|
||||
library_id="lib1",
|
||||
name="f.txt",
|
||||
storage_key="k",
|
||||
mime_type="text",
|
||||
)
|
||||
assert asset.file_type == "text"
|
||||
|
||||
def test_empty_name_raises(self):
|
||||
"""空名称抛异常."""
|
||||
with pytest.raises(ValueError, match="素材名称不能为空"):
|
||||
Asset.create(
|
||||
project_id="p1",
|
||||
library_id="lib1",
|
||||
name="",
|
||||
storage_key="k",
|
||||
mime_type="video/mp4",
|
||||
)
|
||||
|
||||
def test_empty_storage_key_raises(self):
|
||||
"""空storage_key抛异常."""
|
||||
with pytest.raises(ValueError, match="storage_key"):
|
||||
Asset.create(
|
||||
project_id="p1",
|
||||
library_id="lib1",
|
||||
name="v.mp4",
|
||||
storage_key=" ",
|
||||
mime_type="video/mp4",
|
||||
)
|
||||
|
||||
def test_empty_mime_type_raises(self):
|
||||
"""空mime_type抛异常."""
|
||||
with pytest.raises(ValueError, match="mime_type"):
|
||||
Asset.create(
|
||||
project_id="p1",
|
||||
library_id="lib1",
|
||||
name="v.mp4",
|
||||
storage_key="k",
|
||||
mime_type="",
|
||||
)
|
||||
|
||||
def test_name_stripped(self):
|
||||
"""名称首尾空白被去除."""
|
||||
asset = Asset.create(
|
||||
project_id="p1",
|
||||
library_id="lib1",
|
||||
name=" 视频.mp4 ",
|
||||
storage_key="k",
|
||||
mime_type="video/mp4",
|
||||
)
|
||||
assert asset.name == "视频.mp4"
|
||||
|
||||
def test_add_tag_success(self):
|
||||
"""添加标签成功."""
|
||||
asset = Asset.create(
|
||||
project_id="p1",
|
||||
library_id="lib1",
|
||||
name="v.mp4",
|
||||
storage_key="k",
|
||||
mime_type="video/mp4",
|
||||
)
|
||||
asset.add_tag("tag_1")
|
||||
assert "tag_1" in asset.tag_ids
|
||||
assert len(asset.tag_ids) == 1
|
||||
|
||||
def test_add_duplicate_tag_deduped(self):
|
||||
"""重复添加标签自动去重."""
|
||||
asset = Asset.create(
|
||||
project_id="p1",
|
||||
library_id="lib1",
|
||||
name="v.mp4",
|
||||
storage_key="k",
|
||||
mime_type="video/mp4",
|
||||
)
|
||||
asset.add_tag("tag_1")
|
||||
asset.add_tag("tag_1")
|
||||
assert asset.tag_ids.count("tag_1") == 1
|
||||
|
||||
def test_add_empty_tag_raises(self):
|
||||
"""空标签ID抛异常."""
|
||||
asset = Asset.create(
|
||||
project_id="p1",
|
||||
library_id="lib1",
|
||||
name="v.mp4",
|
||||
storage_key="k",
|
||||
mime_type="video/mp4",
|
||||
)
|
||||
with pytest.raises(ValueError, match="标签 ID 不能为空"):
|
||||
asset.add_tag(" ")
|
||||
|
||||
def test_remove_tag_success(self):
|
||||
"""删除标签成功."""
|
||||
asset = Asset.create(
|
||||
project_id="p1",
|
||||
library_id="lib1",
|
||||
name="v.mp4",
|
||||
storage_key="k",
|
||||
mime_type="video/mp4",
|
||||
)
|
||||
asset.add_tag("tag_1")
|
||||
asset.remove_tag("tag_1")
|
||||
assert "tag_1" not in asset.tag_ids
|
||||
|
||||
def test_remove_nonexistent_tag_no_error(self):
|
||||
"""删除不存在的标签不报错(幂等)."""
|
||||
asset = Asset.create(
|
||||
project_id="p1",
|
||||
library_id="lib1",
|
||||
name="v.mp4",
|
||||
storage_key="k",
|
||||
mime_type="video/mp4",
|
||||
)
|
||||
# 不抛异常
|
||||
asset.remove_tag("nonexistent_tag")
|
||||
|
||||
def test_default_status_uploading(self):
|
||||
"""默认状态UPLOADING."""
|
||||
asset = Asset.create(
|
||||
project_id="p1",
|
||||
library_id="lib1",
|
||||
name="v.mp4",
|
||||
storage_key="k",
|
||||
mime_type="video/mp4",
|
||||
)
|
||||
assert asset.status == AssetStatus.UPLOADING
|
||||
|
||||
def test_custom_status(self):
|
||||
"""自定义状态."""
|
||||
asset = Asset.create(
|
||||
project_id="p1",
|
||||
library_id="lib1",
|
||||
name="v.mp4",
|
||||
storage_key="k",
|
||||
mime_type="video/mp4",
|
||||
status=AssetStatus.READY,
|
||||
)
|
||||
assert asset.status == AssetStatus.READY
|
||||
|
||||
def test_metadata_default_empty_dict(self):
|
||||
"""metadata默认为空dict."""
|
||||
asset = Asset.create(
|
||||
project_id="p1",
|
||||
library_id="lib1",
|
||||
name="v.mp4",
|
||||
storage_key="k",
|
||||
mime_type="video/mp4",
|
||||
)
|
||||
assert asset.metadata == {}
|
||||
|
||||
def test_metadata_none_becomes_empty(self):
|
||||
"""metadata=None → {}."""
|
||||
asset = Asset.create(
|
||||
project_id="p1",
|
||||
library_id="lib1",
|
||||
name="v.mp4",
|
||||
storage_key="k",
|
||||
mime_type="video/mp4",
|
||||
metadata=None,
|
||||
)
|
||||
assert asset.metadata == {}
|
||||
|
||||
|
||||
class TestIngestJob:
|
||||
"""IngestJob 实体测试."""
|
||||
|
||||
def test_create_ingest_job(self):
|
||||
"""创建入库任务."""
|
||||
job = IngestJob.create(
|
||||
project_id="p1",
|
||||
library_id="lib1",
|
||||
storage_key="projects/p1/uploads/temp.mp4",
|
||||
file_hash="abc123",
|
||||
)
|
||||
assert job.id is not None
|
||||
assert job.project_id == "p1"
|
||||
assert job.library_id == "lib1"
|
||||
assert job.storage_key == "projects/p1/uploads/temp.mp4"
|
||||
assert job.file_hash == "abc123"
|
||||
|
||||
def test_default_status_pending(self):
|
||||
"""默认状态PENDING."""
|
||||
from packages.domain.classification import IngestJobStatus
|
||||
|
||||
job = IngestJob.create(
|
||||
project_id="p1",
|
||||
library_id="lib1",
|
||||
storage_key="projects/p1/v.mp4",
|
||||
)
|
||||
assert job.status == IngestJobStatus.PENDING
|
||||
|
||||
def test_empty_project_id_raises(self):
|
||||
"""空project_id抛异常."""
|
||||
with pytest.raises(ValueError, match="project_id"):
|
||||
IngestJob.create(
|
||||
project_id=" ",
|
||||
library_id="lib1",
|
||||
storage_key="k",
|
||||
)
|
||||
|
||||
def test_empty_library_id_raises(self):
|
||||
"""空library_id抛异常."""
|
||||
with pytest.raises(ValueError, match="library_id"):
|
||||
IngestJob.create(
|
||||
project_id="p1",
|
||||
library_id="",
|
||||
storage_key="k",
|
||||
)
|
||||
|
||||
def test_empty_storage_key_raises(self):
|
||||
"""空storage_key抛异常."""
|
||||
with pytest.raises(ValueError, match="storage_key"):
|
||||
IngestJob.create(
|
||||
project_id="p1",
|
||||
library_id="lib1",
|
||||
storage_key=" ",
|
||||
)
|
||||
|
||||
def test_default_result_asset_id_empty(self):
|
||||
"""默认result_asset_id为空."""
|
||||
job = IngestJob.create(
|
||||
project_id="p1",
|
||||
library_id="lib1",
|
||||
storage_key="k",
|
||||
)
|
||||
assert job.result_asset_id == ""
|
||||
|
||||
def test_error_message_empty_by_default(self):
|
||||
"""默认error_message为空."""
|
||||
job = IngestJob.create(
|
||||
project_id="p1",
|
||||
library_id="lib1",
|
||||
storage_key="k",
|
||||
)
|
||||
assert job.error_message == ""
|
||||
@@ -1,6 +1,4 @@
|
||||
"""
|
||||
领域层异常类单元测试
|
||||
"""
|
||||
"""领域层通用异常单元测试."""
|
||||
|
||||
import pytest
|
||||
|
||||
@@ -16,99 +14,118 @@ class TestDomainError:
|
||||
"""DomainError 基类测试"""
|
||||
|
||||
def test_is_exception(self):
|
||||
"""DomainError 是 Exception 的子类"""
|
||||
assert issubclass(DomainError, Exception)
|
||||
|
||||
def test_raise_and_catch(self):
|
||||
def test_can_raise_and_catch(self):
|
||||
"""可以抛出和捕获"""
|
||||
with pytest.raises(DomainError):
|
||||
raise DomainError("test error")
|
||||
raise DomainError("something went wrong")
|
||||
|
||||
def test_error_message(self):
|
||||
err = DomainError("something went wrong")
|
||||
assert str(err) == "something went wrong"
|
||||
|
||||
def test_empty_message(self):
|
||||
err = DomainError("")
|
||||
assert str(err) == ""
|
||||
def test_message(self):
|
||||
"""异常消息正确"""
|
||||
err = DomainError("test message")
|
||||
assert str(err) == "test message"
|
||||
|
||||
|
||||
class TestNotFoundError:
|
||||
"""NotFoundError 测试"""
|
||||
|
||||
def test_is_domain_error(self):
|
||||
"""NotFoundError 继承自 DomainError"""
|
||||
assert issubclass(NotFoundError, DomainError)
|
||||
|
||||
def test_raise_and_catch_as_domain(self):
|
||||
def test_can_raise_as_domain_error(self):
|
||||
"""可以作为 DomainError 捕获"""
|
||||
with pytest.raises(DomainError):
|
||||
raise NotFoundError("resource not found")
|
||||
|
||||
def test_raise_and_catch_specific(self):
|
||||
with pytest.raises(NotFoundError):
|
||||
raise NotFoundError("not found")
|
||||
def test_default_message(self):
|
||||
"""无参构造"""
|
||||
err = NotFoundError()
|
||||
assert isinstance(err, NotFoundError)
|
||||
|
||||
def test_error_message(self):
|
||||
def test_custom_message(self):
|
||||
"""自定义消息"""
|
||||
err = NotFoundError("user 123 not found")
|
||||
assert "user 123 not found" in str(err)
|
||||
assert str(err) == "user 123 not found"
|
||||
|
||||
|
||||
class TestValidationError:
|
||||
"""ValidationError 测试"""
|
||||
|
||||
def test_is_domain_error(self):
|
||||
"""ValidationError 继承自 DomainError"""
|
||||
assert issubclass(ValidationError, DomainError)
|
||||
|
||||
def test_raise_and_catch_as_domain(self):
|
||||
def test_can_raise_as_domain_error(self):
|
||||
"""可以作为 DomainError 捕获"""
|
||||
with pytest.raises(DomainError):
|
||||
raise ValidationError("invalid input")
|
||||
|
||||
def test_raise_and_catch_specific(self):
|
||||
with pytest.raises(ValidationError):
|
||||
raise ValidationError("validation failed")
|
||||
|
||||
def test_error_message(self):
|
||||
err = ValidationError("name cannot be empty")
|
||||
assert "name cannot be empty" in str(err)
|
||||
def test_custom_message(self):
|
||||
"""自定义消息"""
|
||||
err = ValidationError("duration must be positive")
|
||||
assert str(err) == "duration must be positive"
|
||||
|
||||
|
||||
class TestQuotaExceededError:
|
||||
"""QuotaExceededError 测试"""
|
||||
|
||||
def test_is_domain_error(self):
|
||||
"""QuotaExceededError 继承自 DomainError"""
|
||||
assert issubclass(QuotaExceededError, DomainError)
|
||||
|
||||
def test_constructor_sets_attributes(self):
|
||||
err = QuotaExceededError(dimension="storage", limit=1024.0, used=2048.0)
|
||||
assert err.dimension == "storage"
|
||||
def test_can_raise_as_domain_error(self):
|
||||
"""可以作为 DomainError 捕获"""
|
||||
with pytest.raises(DomainError):
|
||||
raise QuotaExceededError("storage", 1024.0, 2048.0)
|
||||
|
||||
def test_stores_dimension_limit_used(self):
|
||||
"""保存 dimension、limit、used 属性"""
|
||||
err = QuotaExceededError("storage_mb", 1024.0, 1500.0)
|
||||
assert err.dimension == "storage_mb"
|
||||
assert err.limit == 1024.0
|
||||
assert err.used == 2048.0
|
||||
assert err.used == 1500.0
|
||||
|
||||
def test_error_message_format(self):
|
||||
err = QuotaExceededError(dimension="storage", limit=1024.0, used=2048.0)
|
||||
"""异常消息格式正确"""
|
||||
err = QuotaExceededError("storage_mb", 1024.0, 1500.0)
|
||||
msg = str(err)
|
||||
assert "storage" in msg
|
||||
assert "2048.0" in msg
|
||||
assert "storage_mb" in msg
|
||||
assert "1500.0" in msg
|
||||
assert "1024.0" in msg
|
||||
assert "Quota exceeded" in msg
|
||||
|
||||
def test_raise_and_catch_as_domain(self):
|
||||
with pytest.raises(DomainError):
|
||||
raise QuotaExceededError("projects", 10, 15)
|
||||
|
||||
def test_raise_and_catch_specific(self):
|
||||
with pytest.raises(QuotaExceededError):
|
||||
raise QuotaExceededError("render", 5, 10)
|
||||
def test_integer_values(self):
|
||||
"""整数值也能正常工作"""
|
||||
err = QuotaExceededError("projects", 10, 15)
|
||||
assert err.dimension == "projects"
|
||||
assert err.limit == 10
|
||||
assert err.used == 15
|
||||
|
||||
def test_zero_limit(self):
|
||||
err = QuotaExceededError(dimension="test", limit=0.0, used=1.0)
|
||||
assert err.limit == 0.0
|
||||
assert err.used == 1.0
|
||||
"""限制为 0 时也能正常工作"""
|
||||
err = QuotaExceededError("custom_templates", 0, 1)
|
||||
assert err.limit == 0
|
||||
assert err.used == 1
|
||||
|
||||
def test_negative_values(self):
|
||||
"""负数也能存(领域层不做额外校验)"""
|
||||
err = QuotaExceededError(dimension="test", limit=-5.0, used=-3.0)
|
||||
assert err.limit == -5.0
|
||||
assert err.used == -3.0
|
||||
|
||||
def test_large_values(self):
|
||||
err = QuotaExceededError(dimension="storage", limit=1e9, used=1.5e9)
|
||||
assert err.limit == 1e9
|
||||
assert err.used == 1.5e9
|
||||
class TestExceptionHierarchy:
|
||||
"""异常继承关系测试"""
|
||||
|
||||
def test_all_are_domain_errors(self):
|
||||
"""所有异常都可以作为 DomainError 捕获"""
|
||||
errors = [
|
||||
NotFoundError(),
|
||||
ValidationError("bad"),
|
||||
QuotaExceededError("x", 10.0, 20.0),
|
||||
]
|
||||
for err in errors:
|
||||
assert isinstance(err, DomainError)
|
||||
|
||||
def test_distinct_types(self):
|
||||
"""不同异常类型可以区分"""
|
||||
assert not issubclass(NotFoundError, ValidationError)
|
||||
assert not issubclass(ValidationError, QuotaExceededError)
|
||||
assert not issubclass(NotFoundError, QuotaExceededError)
|
||||
|
||||
@@ -1,5 +1,3 @@
|
||||
from dataclasses import FrozenInstanceError
|
||||
|
||||
"""domain 层剩余小模块单测 - bgm_utils/exceptions/editing_mode/recipe/template/template_version/template_clip_config/preset_bgm/preset_voices/title_library/voice_library"""
|
||||
|
||||
import pytest
|
||||
@@ -399,7 +397,7 @@ class TestPresetBGM:
|
||||
|
||||
def test_preset_bgm_frozen(self):
|
||||
bgm = PresetBGM(id="t1", name="t", style="x", duration=10.0)
|
||||
with pytest.raises(FrozenInstanceError):
|
||||
with pytest.raises(Exception):
|
||||
bgm.name = "改名"
|
||||
|
||||
def test_library_not_empty(self):
|
||||
@@ -489,7 +487,7 @@ class TestPresetVoices:
|
||||
|
||||
def test_preset_voice_frozen(self):
|
||||
v = PresetVoice(voice_id="v1", name="t", description="d", gender="female")
|
||||
with pytest.raises(FrozenInstanceError):
|
||||
with pytest.raises(Exception):
|
||||
v.name = "改名"
|
||||
|
||||
def test_preset_voices_list_not_empty(self):
|
||||
|
||||
@@ -152,103 +152,3 @@ class TestEditTemplateBumpVersion:
|
||||
old_updated = template.updated_at
|
||||
template.bump_version()
|
||||
assert template.updated_at > old_updated or template.updated_at == old_updated
|
||||
|
||||
|
||||
class TestEditTemplateExtended:
|
||||
"""EditTemplate 深度补充测试"""
|
||||
|
||||
def test_id_is_hex(self):
|
||||
t = EditTemplate.create(name="test")
|
||||
int(t.id, 16)
|
||||
|
||||
def test_ids_are_unique(self):
|
||||
t1 = EditTemplate.create(name="test1")
|
||||
t2 = EditTemplate.create(name="test2")
|
||||
assert t1.id != t2.id
|
||||
|
||||
def test_empty_description(self):
|
||||
t = EditTemplate.create(name="test", description="")
|
||||
assert t.description == ""
|
||||
|
||||
def test_long_description(self):
|
||||
desc = "描述" * 200
|
||||
t = EditTemplate.create(name="test", description=desc)
|
||||
assert t.description == desc
|
||||
assert len(t.description) == 400
|
||||
|
||||
def test_unicode_name(self):
|
||||
t = EditTemplate.create(name="🎬 口播 Vlog 模板")
|
||||
assert "🎬" in t.name
|
||||
assert "口播" in t.name
|
||||
|
||||
def test_special_characters_name(self):
|
||||
special = "模!@#$%板"
|
||||
t = EditTemplate.create(name=special)
|
||||
assert t.name == special
|
||||
|
||||
def test_long_name(self):
|
||||
long_name = "模板名称" * 50
|
||||
t = EditTemplate.create(name=long_name)
|
||||
assert t.name == long_name
|
||||
assert len(t.name) == 200
|
||||
|
||||
def test_config_independence(self):
|
||||
t1 = EditTemplate.create(name="test1")
|
||||
t2 = EditTemplate.create(name="test2")
|
||||
t1.config["key"] = "val"
|
||||
assert "key" not in t2.config
|
||||
|
||||
def test_sort_weight_negative(self):
|
||||
t = EditTemplate.create(name="test", sort_weight=-100)
|
||||
assert t.sort_weight == -100
|
||||
|
||||
def test_sort_weight_large(self):
|
||||
t = EditTemplate.create(name="test", sort_weight=99999)
|
||||
assert t.sort_weight == 99999
|
||||
|
||||
def test_sort_weight_zero(self):
|
||||
t = EditTemplate.create(name="test", sort_weight=0)
|
||||
assert t.sort_weight == 0
|
||||
|
||||
def test_preview_url_empty(self):
|
||||
t = EditTemplate.create(name="test", preview_url="")
|
||||
assert t.preview_url == ""
|
||||
|
||||
def test_version_zero(self):
|
||||
t = EditTemplate.create(name="test", version=0)
|
||||
assert t.version == 0
|
||||
|
||||
def test_version_large(self):
|
||||
t = EditTemplate.create(name="test", version=999)
|
||||
assert t.version == 999
|
||||
|
||||
def test_status_is_active_property(self):
|
||||
t = EditTemplate.create(name="test", status=EditTemplateStatus.ACTIVE)
|
||||
assert t.is_active is True
|
||||
t.deactivate()
|
||||
assert t.is_active is False
|
||||
t.activate()
|
||||
assert t.is_active is True
|
||||
|
||||
def test_bump_version_from_zero(self):
|
||||
t = EditTemplate.create(name="test", version=0)
|
||||
t.bump_version()
|
||||
assert t.version == 1
|
||||
|
||||
def test_template_type_custom(self):
|
||||
t = EditTemplate.create(name="test", template_type="custom_type")
|
||||
assert t.template_type == "custom_type"
|
||||
|
||||
def test_template_type_strips_and_default(self):
|
||||
"""空格的 template_type 回退到 default"""
|
||||
t = EditTemplate.create(name="test", template_type=" ")
|
||||
assert t.template_type == "default"
|
||||
|
||||
def test_create_with_empty_editing_mode_defaults(self):
|
||||
"""空字符串 editing_mode 回退到 one_take"""
|
||||
t = EditTemplate.create(name="test", editing_mode="")
|
||||
assert t.editing_mode == "one_take"
|
||||
|
||||
def test_preview_url_strips_whitespace(self):
|
||||
t = EditTemplate.create(name="test", preview_url=" https://example.com/v.mp4 ")
|
||||
assert t.preview_url == "https://example.com/v.mp4"
|
||||
|
||||
@@ -38,57 +38,3 @@ class TestEditingMode:
|
||||
modes = list(EditingMode)
|
||||
assert len(modes) == 4
|
||||
assert EditingMode.ONE_TAKE in modes
|
||||
|
||||
|
||||
class TestEditingModeExtended:
|
||||
"""EditingMode 深度补充测试"""
|
||||
|
||||
def test_from_string_value(self):
|
||||
"""可以从字符串值构造枚举"""
|
||||
assert EditingMode("one_take") == EditingMode.ONE_TAKE
|
||||
assert EditingMode("pip") == EditingMode.PIP
|
||||
assert EditingMode("voice_over") == EditingMode.VOICE_OVER
|
||||
assert EditingMode("voice_pip") == EditingMode.VOICE_PIP
|
||||
|
||||
def test_invalid_string_raises(self):
|
||||
import pytest
|
||||
|
||||
with pytest.raises(ValueError):
|
||||
EditingMode("invalid_mode")
|
||||
|
||||
def test_string_concatenation(self):
|
||||
"""StrEnum 支持字符串拼接"""
|
||||
result = "mode_" + EditingMode.ONE_TAKE
|
||||
assert result == "mode_one_take"
|
||||
|
||||
def test_dict_key_usage(self):
|
||||
"""可以作为字典 key 使用"""
|
||||
mapping = {
|
||||
EditingMode.ONE_TAKE: "顺序拼接",
|
||||
EditingMode.PIP: "画中画",
|
||||
}
|
||||
assert mapping[EditingMode.ONE_TAKE] == "顺序拼接"
|
||||
assert mapping[EditingMode.PIP] == "画中画"
|
||||
assert len(mapping) == 2
|
||||
|
||||
def test_value_lowercase(self):
|
||||
"""所有枚举值都是小写字母+下划线"""
|
||||
for mode in EditingMode:
|
||||
assert mode.value == mode.value.lower()
|
||||
assert " " not in mode.value
|
||||
|
||||
def test_unique_values(self):
|
||||
"""所有枚举值唯一"""
|
||||
values = [m.value for m in EditingMode]
|
||||
assert len(values) == len(set(values))
|
||||
|
||||
def test_membership_test(self):
|
||||
assert EditingMode.ONE_TAKE in EditingMode
|
||||
assert "one_take" in [m.value for m in EditingMode]
|
||||
|
||||
def test_comparison_with_string(self):
|
||||
"""和字符串直接比较"""
|
||||
mode = EditingMode.VOICE_OVER
|
||||
assert mode == "voice_over"
|
||||
assert mode != "pip"
|
||||
assert "voice_over" == mode
|
||||
|
||||
@@ -267,4 +267,4 @@ class TestGetEmailService:
|
||||
svc1 = get_email_service()
|
||||
svc2 = get_email_service()
|
||||
# 两个都可能是 Noop 或 EmailService,取决于环境
|
||||
assert type(svc1) is type(svc2)
|
||||
assert type(svc1) == type(svc2)
|
||||
|
||||
@@ -1,166 +0,0 @@
|
||||
"""领域层异常类单元测试."""
|
||||
|
||||
import pytest
|
||||
|
||||
from packages.domain.exceptions import (
|
||||
DomainError,
|
||||
NotFoundError,
|
||||
ValidationError,
|
||||
QuotaExceededError,
|
||||
)
|
||||
|
||||
|
||||
class TestDomainError:
|
||||
"""领域异常基类测试."""
|
||||
|
||||
def test_is_exception(self):
|
||||
"""DomainError 继承自 Exception."""
|
||||
err = DomainError("test")
|
||||
assert isinstance(err, Exception)
|
||||
|
||||
def test_message(self):
|
||||
"""可以设置错误消息."""
|
||||
err = DomainError("something wrong")
|
||||
assert str(err) == "something wrong"
|
||||
|
||||
def test_empty_message(self):
|
||||
"""支持空消息."""
|
||||
err = DomainError()
|
||||
assert str(err) == ""
|
||||
|
||||
def test_can_be_raised(self):
|
||||
"""可以被 raise 和 catch."""
|
||||
with pytest.raises(DomainError) as exc_info:
|
||||
raise DomainError("oops")
|
||||
assert str(exc_info.value) == "oops"
|
||||
|
||||
|
||||
class TestNotFoundError:
|
||||
"""资源不存在异常测试."""
|
||||
|
||||
def test_inherits_domain_error(self):
|
||||
"""NotFoundError 继承自 DomainError."""
|
||||
err = NotFoundError("user not found")
|
||||
assert isinstance(err, DomainError)
|
||||
assert isinstance(err, Exception)
|
||||
|
||||
def test_message(self):
|
||||
"""错误消息正确."""
|
||||
err = NotFoundError("project 123 not found")
|
||||
assert str(err) == "project 123 not found"
|
||||
assert "123" in str(err)
|
||||
|
||||
def test_can_catch_as_domain_error(self):
|
||||
"""可以用 DomainError 捕获."""
|
||||
with pytest.raises(DomainError):
|
||||
raise NotFoundError("not found")
|
||||
|
||||
|
||||
class TestValidationError:
|
||||
"""校验失败异常测试."""
|
||||
|
||||
def test_inherits_domain_error(self):
|
||||
"""ValidationError 继承自 DomainError."""
|
||||
err = ValidationError("invalid input")
|
||||
assert isinstance(err, DomainError)
|
||||
|
||||
def test_message(self):
|
||||
"""错误消息正确."""
|
||||
msg = "name must not be empty"
|
||||
err = ValidationError(msg)
|
||||
assert str(err) == msg
|
||||
|
||||
def test_not_not_found(self):
|
||||
"""ValidationError 不是 NotFoundError."""
|
||||
err = ValidationError("bad")
|
||||
assert not isinstance(err, NotFoundError)
|
||||
|
||||
|
||||
class TestQuotaExceededError:
|
||||
"""配额超限异常测试."""
|
||||
|
||||
def test_inherits_domain_error(self):
|
||||
"""QuotaExceededError 继承自 DomainError."""
|
||||
err = QuotaExceededError("storage", 100.0, 150.0)
|
||||
assert isinstance(err, DomainError)
|
||||
|
||||
def test_dimension_attribute(self):
|
||||
"""保存 dimension 属性."""
|
||||
err = QuotaExceededError("storage", 100.0, 150.0)
|
||||
assert err.dimension == "storage"
|
||||
|
||||
def test_limit_attribute(self):
|
||||
"""保存 limit 属性."""
|
||||
err = QuotaExceededError("storage", 100.0, 150.0)
|
||||
assert err.limit == 100.0
|
||||
|
||||
def test_used_attribute(self):
|
||||
"""保存 used 属性."""
|
||||
err = QuotaExceededError("storage", 100.0, 150.0)
|
||||
assert err.used == 150.0
|
||||
|
||||
def test_message_format(self):
|
||||
"""错误消息格式正确."""
|
||||
err = QuotaExceededError("credits", 50.0, 75.0)
|
||||
msg = str(err)
|
||||
assert "credits" in msg
|
||||
assert "50" in msg
|
||||
assert "75" in msg
|
||||
assert "Quota exceeded" in msg
|
||||
|
||||
def test_zero_limit(self):
|
||||
"""limit 为 0 的情况."""
|
||||
err = QuotaExceededError("test", 0.0, 1.0)
|
||||
assert err.limit == 0.0
|
||||
assert err.used == 1.0
|
||||
assert "0" in str(err)
|
||||
|
||||
def test_equal_limit_and_used(self):
|
||||
"""used 刚好等于 limit(边界情况)."""
|
||||
err = QuotaExceededError("test", 100.0, 100.0)
|
||||
assert err.used == 100.0
|
||||
assert err.limit == 100.0
|
||||
|
||||
def test_integer_values(self):
|
||||
"""整数值也能正常工作."""
|
||||
err = QuotaExceededError("count", 10, 20)
|
||||
assert err.dimension == "count"
|
||||
assert err.limit == 10
|
||||
assert err.used == 20
|
||||
|
||||
def test_can_catch_as_domain_error(self):
|
||||
"""可以用 DomainError 捕获."""
|
||||
with pytest.raises(DomainError):
|
||||
raise QuotaExceededError("x", 1.0, 2.0)
|
||||
|
||||
|
||||
class TestExceptionHierarchy:
|
||||
"""异常继承关系验证."""
|
||||
|
||||
def test_all_are_domain_errors(self):
|
||||
"""所有领域异常都是 DomainError."""
|
||||
errors = [
|
||||
NotFoundError("test"),
|
||||
ValidationError("test"),
|
||||
QuotaExceededError("test", 1, 2),
|
||||
]
|
||||
for err in errors:
|
||||
assert isinstance(err, DomainError)
|
||||
|
||||
def test_all_are_exceptions(self):
|
||||
"""所有领域异常都是 Exception."""
|
||||
errors = [
|
||||
DomainError("test"),
|
||||
NotFoundError("test"),
|
||||
ValidationError("test"),
|
||||
QuotaExceededError("test", 1, 2),
|
||||
]
|
||||
for err in errors:
|
||||
assert isinstance(err, Exception)
|
||||
|
||||
def test_not_found_is_not_validation(self):
|
||||
"""不同异常类型不能互相混淆."""
|
||||
assert not isinstance(NotFoundError("x"), ValidationError)
|
||||
assert not isinstance(ValidationError("x"), NotFoundError)
|
||||
assert not isinstance(QuotaExceededError("x", 1, 2), NotFoundError)
|
||||
assert not isinstance(QuotaExceededError("x", 1, 2), ValidationError)
|
||||
@@ -1,206 +0,0 @@
|
||||
"""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]")
|
||||
@@ -1,5 +1,3 @@
|
||||
from dataclasses import FrozenInstanceError
|
||||
|
||||
"""filter_presets 领域层单元测试 - 滤镜预设库"""
|
||||
|
||||
import pytest
|
||||
@@ -60,7 +58,7 @@ class TestFilterPreset:
|
||||
def test_frozen_immutable(self):
|
||||
"""frozen dataclass 不可修改"""
|
||||
preset = FilterPreset(id="test", name="测试", category="basic")
|
||||
with pytest.raises(FrozenInstanceError):
|
||||
with pytest.raises(Exception): # FrozenInstanceError
|
||||
preset.name = "改名"
|
||||
|
||||
def test_tags_default_empty_list(self):
|
||||
|
||||
@@ -194,117 +194,3 @@ class TestGeneratedVideoProperties:
|
||||
fingerprint = {"phash": "abc123", "md5": "def456"}
|
||||
gv.video_fingerprint = fingerprint
|
||||
assert gv.video_fingerprint == fingerprint
|
||||
|
||||
|
||||
class TestGeneratedVideoExtended:
|
||||
"""GeneratedVideo 深度补充测试"""
|
||||
|
||||
def test_id_is_hex(self):
|
||||
v = GeneratedVideo.create(project_id="p", generation_task_id="t", name="n", file_url="u")
|
||||
int(v.id, 16)
|
||||
|
||||
def test_zero_file_size(self):
|
||||
v = GeneratedVideo.create(project_id="p", generation_task_id="t", name="n", file_url="u", file_size=0)
|
||||
assert v.file_size == 0
|
||||
|
||||
def test_large_file_size(self):
|
||||
large = 1024 * 1024 * 1024 # 1GB
|
||||
v = GeneratedVideo.create(project_id="p", generation_task_id="t", name="n", file_url="u", file_size=large)
|
||||
assert v.file_size == large
|
||||
|
||||
def test_zero_duration(self):
|
||||
v = GeneratedVideo.create(project_id="p", generation_task_id="t", name="n", file_url="u", duration=0.0)
|
||||
assert v.duration == 0.0
|
||||
|
||||
def test_large_duration(self):
|
||||
v = GeneratedVideo.create(project_id="p", generation_task_id="t", name="n", file_url="u", duration=9999.99)
|
||||
assert v.duration == 9999.99
|
||||
|
||||
def test_zero_dimensions(self):
|
||||
v = GeneratedVideo.create(project_id="p", generation_task_id="t", name="n", file_url="u", width=0, height=0)
|
||||
assert v.width == 0
|
||||
assert v.height == 0
|
||||
|
||||
def test_4k_dimensions(self):
|
||||
v = GeneratedVideo.create(
|
||||
project_id="p", generation_task_id="t", name="n", file_url="u", width=3840, height=2160
|
||||
)
|
||||
assert v.width == 3840
|
||||
assert v.height == 2160
|
||||
|
||||
def test_zero_fps(self):
|
||||
v = GeneratedVideo.create(project_id="p", generation_task_id="t", name="n", file_url="u", fps=0.0)
|
||||
assert v.fps == 0.0
|
||||
|
||||
def test_high_fps(self):
|
||||
v = GeneratedVideo.create(project_id="p", generation_task_id="t", name="n", file_url="u", fps=120.0)
|
||||
assert v.fps == 120.0
|
||||
|
||||
def test_empty_user_id(self):
|
||||
v = GeneratedVideo.create(project_id="p", generation_task_id="t", name="n", file_url="u", user_id="")
|
||||
assert v.user_id == ""
|
||||
|
||||
def test_long_name(self):
|
||||
long_name = "视频" * 100
|
||||
v = GeneratedVideo.create(project_id="p", generation_task_id="t", name=long_name, file_url="u")
|
||||
assert v.name == long_name
|
||||
assert len(v.name) == 200
|
||||
|
||||
def test_unicode_name(self):
|
||||
v = GeneratedVideo.create(
|
||||
project_id="p", generation_task_id="t", name="🎬 我的精彩视频 · 旅行vlog", file_url="u"
|
||||
)
|
||||
assert "🎬" in v.name
|
||||
assert "旅行vlog" in v.name
|
||||
|
||||
def test_special_characters_name(self):
|
||||
special = "视!@#$%^&*()频"
|
||||
v = GeneratedVideo.create(project_id="p", generation_task_id="t", name=special, file_url="u")
|
||||
assert v.name == special
|
||||
|
||||
def test_thumbnail_url_none(self):
|
||||
v = GeneratedVideo.create(project_id="p", generation_task_id="t", name="n", file_url="u", thumbnail_url=None)
|
||||
assert v.thumbnail_url is None
|
||||
|
||||
def test_generation_params_complex(self):
|
||||
params = {
|
||||
"mode": "voice_over",
|
||||
"quality": "high",
|
||||
"resolution": {"width": 1920, "height": 1080},
|
||||
"effects": ["filter", "transition"],
|
||||
}
|
||||
v = GeneratedVideo.create(
|
||||
project_id="p", generation_task_id="t", name="n", file_url="u", generation_params=params
|
||||
)
|
||||
assert v.generation_params["mode"] == "voice_over"
|
||||
assert v.generation_params["resolution"]["width"] == 1920
|
||||
assert len(v.generation_params["effects"]) == 2
|
||||
|
||||
def test_generation_params_independence(self):
|
||||
v1 = GeneratedVideo.create(project_id="p", generation_task_id="t1", name="n1", file_url="u1")
|
||||
v2 = GeneratedVideo.create(project_id="p", generation_task_id="t2", name="n2", file_url="u2")
|
||||
v1.generation_params["key"] = "val"
|
||||
assert "key" not in v2.generation_params
|
||||
|
||||
def test_status_failed(self):
|
||||
v = GeneratedVideo.create(project_id="p", generation_task_id="t", name="n", file_url="u")
|
||||
v.status = "failed"
|
||||
assert v.status == "failed"
|
||||
|
||||
def test_review_status_approved(self):
|
||||
v = GeneratedVideo.create(project_id="p", generation_task_id="t", name="n", file_url="u")
|
||||
v.review_status = "approved"
|
||||
assert v.review_status == "approved"
|
||||
|
||||
def test_is_duplicate_default_false(self):
|
||||
v = GeneratedVideo.create(project_id="p", generation_task_id="t", name="n", file_url="u")
|
||||
assert v.is_duplicate is False
|
||||
assert v.duplicate_of is None
|
||||
|
||||
def test_fingerprint_none_default(self):
|
||||
v = GeneratedVideo.create(project_id="p", generation_task_id="t", name="n", file_url="u")
|
||||
assert v.video_fingerprint is None
|
||||
|
||||
def test_empty_file_url_raises(self):
|
||||
with pytest.raises(ValueError):
|
||||
GeneratedVideo.create(project_id="p", generation_task_id="t", name="n", file_url=" ")
|
||||
|
||||
@@ -1,323 +0,0 @@
|
||||
"""GeneratedVideo + VerificationCode 领域模型测试."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import timedelta
|
||||
|
||||
import pytest
|
||||
|
||||
from packages.domain.generated_video import GeneratedVideo
|
||||
from packages.domain.verification_code import VerificationCode
|
||||
|
||||
|
||||
class TestGeneratedVideo:
|
||||
"""GeneratedVideo 生成视频实体测试."""
|
||||
|
||||
def test_create_success(self):
|
||||
"""创建成功."""
|
||||
video = GeneratedVideo.create(
|
||||
project_id="p1",
|
||||
generation_task_id="task_1",
|
||||
name="我的视频.mp4",
|
||||
file_url="https://example.com/out.mp4",
|
||||
user_id="u1",
|
||||
file_size=1024000,
|
||||
duration=30.5,
|
||||
width=1080,
|
||||
height=1920,
|
||||
fps=25.0,
|
||||
)
|
||||
assert video.id is not None
|
||||
assert len(video.id) == 32
|
||||
assert video.project_id == "p1"
|
||||
assert video.generation_task_id == "task_1"
|
||||
assert video.name == "我的视频.mp4"
|
||||
assert video.file_url == "https://example.com/out.mp4"
|
||||
assert video.user_id == "u1"
|
||||
assert video.file_size == 1024000
|
||||
assert video.duration == pytest.approx(30.5)
|
||||
assert video.width == 1080
|
||||
assert video.height == 1920
|
||||
assert video.fps == pytest.approx(25.0)
|
||||
assert video.status == "completed"
|
||||
assert video.review_status == "pending_review"
|
||||
assert video.is_duplicate is False
|
||||
assert video.duplicate_of is None
|
||||
assert video.generation_params == {}
|
||||
|
||||
def test_create_empty_project_id_raises(self):
|
||||
"""空project_id抛异常."""
|
||||
with pytest.raises(ValueError, match="project_id"):
|
||||
GeneratedVideo.create(
|
||||
project_id=" ",
|
||||
generation_task_id="t1",
|
||||
name="v.mp4",
|
||||
file_url="https://x.com/v.mp4",
|
||||
)
|
||||
|
||||
def test_create_empty_task_id_raises(self):
|
||||
"""空generation_task_id抛异常."""
|
||||
with pytest.raises(ValueError, match="generation_task_id"):
|
||||
GeneratedVideo.create(
|
||||
project_id="p1",
|
||||
generation_task_id="",
|
||||
name="v.mp4",
|
||||
file_url="https://x.com/v.mp4",
|
||||
)
|
||||
|
||||
def test_create_empty_name_raises(self):
|
||||
"""空name抛异常."""
|
||||
with pytest.raises(ValueError, match="name"):
|
||||
GeneratedVideo.create(
|
||||
project_id="p1",
|
||||
generation_task_id="t1",
|
||||
name=" ",
|
||||
file_url="https://x.com/v.mp4",
|
||||
)
|
||||
|
||||
def test_create_empty_file_url_raises(self):
|
||||
"""空file_url抛异常."""
|
||||
with pytest.raises(ValueError, match="file_url"):
|
||||
GeneratedVideo.create(
|
||||
project_id="p1",
|
||||
generation_task_id="t1",
|
||||
name="v.mp4",
|
||||
file_url="",
|
||||
)
|
||||
|
||||
def test_create_strips_whitespace(self):
|
||||
"""首尾空白被去除."""
|
||||
video = GeneratedVideo.create(
|
||||
project_id=" p1 ",
|
||||
generation_task_id=" t1 ",
|
||||
name=" 视频.mp4 ",
|
||||
file_url=" https://x.com/v.mp4 ",
|
||||
)
|
||||
assert video.project_id == "p1"
|
||||
assert video.generation_task_id == "t1"
|
||||
assert video.name == "视频.mp4"
|
||||
assert video.file_url == "https://x.com/v.mp4"
|
||||
|
||||
def test_default_values(self):
|
||||
"""默认值正确."""
|
||||
video = GeneratedVideo.create(
|
||||
project_id="p1",
|
||||
generation_task_id="t1",
|
||||
name="v.mp4",
|
||||
file_url="https://x.com/v.mp4",
|
||||
)
|
||||
assert video.user_id == ""
|
||||
assert video.file_size == 0
|
||||
assert video.duration == 0.0
|
||||
assert video.width == 0
|
||||
assert video.height == 0
|
||||
assert video.fps == 0.0
|
||||
assert video.thumbnail_url is None
|
||||
assert video.generation_params == {}
|
||||
|
||||
def test_generation_params_none_becomes_empty(self):
|
||||
"""generation_params=None → {}."""
|
||||
video = GeneratedVideo.create(
|
||||
project_id="p1",
|
||||
generation_task_id="t1",
|
||||
name="v.mp4",
|
||||
file_url="https://x.com/v.mp4",
|
||||
generation_params=None,
|
||||
)
|
||||
assert video.generation_params == {}
|
||||
|
||||
def test_custom_generation_params(self):
|
||||
"""自定义生成参数."""
|
||||
params = {"mode": "smart", "resolution": "1080x1920"}
|
||||
video = GeneratedVideo.create(
|
||||
project_id="p1",
|
||||
generation_task_id="t1",
|
||||
name="v.mp4",
|
||||
file_url="https://x.com/v.mp4",
|
||||
generation_params=params,
|
||||
)
|
||||
assert video.generation_params == params
|
||||
|
||||
def test_duplicate_flag(self):
|
||||
"""重复标记可以设置."""
|
||||
video = GeneratedVideo.create(
|
||||
project_id="p1",
|
||||
generation_task_id="t1",
|
||||
name="v.mp4",
|
||||
file_url="https://x.com/v.mp4",
|
||||
)
|
||||
video.is_duplicate = True
|
||||
video.duplicate_of = "other_video_id"
|
||||
assert video.is_duplicate is True
|
||||
assert video.duplicate_of == "other_video_id"
|
||||
|
||||
def test_custom_status(self):
|
||||
"""自定义状态."""
|
||||
video = GeneratedVideo.create(
|
||||
project_id="p1",
|
||||
generation_task_id="t1",
|
||||
name="v.mp4",
|
||||
file_url="https://x.com/v.mp4",
|
||||
)
|
||||
video.status = "failed"
|
||||
assert video.status == "failed"
|
||||
|
||||
def test_thumbnail_url(self):
|
||||
"""缩略图URL."""
|
||||
video = GeneratedVideo.create(
|
||||
project_id="p1",
|
||||
generation_task_id="t1",
|
||||
name="v.mp4",
|
||||
file_url="https://x.com/v.mp4",
|
||||
thumbnail_url="https://x.com/thumb.jpg",
|
||||
)
|
||||
assert video.thumbnail_url == "https://x.com/thumb.jpg"
|
||||
|
||||
|
||||
class TestVerificationCodeCreate:
|
||||
"""VerificationCode 创建测试."""
|
||||
|
||||
def test_create_success(self):
|
||||
"""创建验证码成功."""
|
||||
vc = VerificationCode.create(
|
||||
recipient="test@example.com",
|
||||
code_type="email_bind",
|
||||
ttl_seconds=300,
|
||||
)
|
||||
assert vc.id is not None
|
||||
assert len(vc.id) == 32
|
||||
assert vc.recipient == "test@example.com"
|
||||
assert vc.code_type == "email_bind"
|
||||
assert len(vc.code) == 6 # 默认6位数字
|
||||
assert vc.code.isdigit() # 纯数字
|
||||
assert vc.used_at is None
|
||||
assert vc.attempts == 0
|
||||
|
||||
def test_create_with_custom_code(self):
|
||||
"""自定义验证码."""
|
||||
vc = VerificationCode.create(
|
||||
recipient="u@test.com",
|
||||
code_type="email_login",
|
||||
custom_code="123456",
|
||||
)
|
||||
assert vc.code == "123456"
|
||||
|
||||
def test_create_recipient_stripped(self):
|
||||
"""收件人空白被去除."""
|
||||
vc = VerificationCode.create(
|
||||
recipient=" test@example.com ",
|
||||
code_type="email_bind",
|
||||
)
|
||||
assert vc.recipient == "test@example.com"
|
||||
|
||||
def test_expiry_time_correct(self):
|
||||
"""过期时间正确(5分钟后)."""
|
||||
from datetime import datetime, timezone
|
||||
|
||||
before = datetime.now(timezone.utc) + timedelta(seconds=299)
|
||||
vc = VerificationCode.create(
|
||||
recipient="u@test.com",
|
||||
code_type="reset_password",
|
||||
ttl_seconds=300,
|
||||
)
|
||||
after = datetime.now(timezone.utc) + timedelta(seconds=301)
|
||||
assert before <= vc.expires_at <= after
|
||||
|
||||
def test_custom_ttl(self):
|
||||
"""自定义过期时间."""
|
||||
vc = VerificationCode.create(
|
||||
recipient="u@test.com",
|
||||
code_type="phone_login",
|
||||
ttl_seconds=60,
|
||||
)
|
||||
from datetime import datetime, timezone
|
||||
|
||||
now = datetime.now(timezone.utc)
|
||||
# 应该在1分钟左右过期
|
||||
diff = (vc.expires_at - now).total_seconds()
|
||||
assert 0 < diff < 70
|
||||
|
||||
|
||||
class TestVerificationCodeProperties:
|
||||
"""VerificationCode 属性方法测试."""
|
||||
|
||||
def test_is_expired_false_for_new(self):
|
||||
"""新创建的验证码未过期."""
|
||||
vc = VerificationCode.create(
|
||||
recipient="u@test.com",
|
||||
code_type="email_bind",
|
||||
ttl_seconds=300,
|
||||
)
|
||||
assert vc.is_expired is False
|
||||
|
||||
def test_is_expired_true_when_past(self):
|
||||
"""已过期的验证码is_expired=True."""
|
||||
vc = VerificationCode.create(
|
||||
recipient="u@test.com",
|
||||
code_type="email_bind",
|
||||
ttl_seconds=1,
|
||||
)
|
||||
# 手动改过期时间到过去
|
||||
from datetime import datetime, timedelta, timezone
|
||||
|
||||
vc.expires_at = datetime.now(timezone.utc) - timedelta(seconds=1)
|
||||
assert vc.is_expired is True
|
||||
|
||||
def test_is_used_false_by_default(self):
|
||||
"""默认未使用."""
|
||||
vc = VerificationCode.create(recipient="u@test.com", code_type="email_bind")
|
||||
assert vc.is_used is False
|
||||
|
||||
def test_is_valid_fresh_code(self):
|
||||
"""新验证码有效."""
|
||||
vc = VerificationCode.create(recipient="u@test.com", code_type="email_bind")
|
||||
assert vc.is_valid is True
|
||||
|
||||
def test_is_valid_expired(self):
|
||||
"""过期的验证码无效."""
|
||||
vc = VerificationCode.create(recipient="u@test.com", code_type="email_bind", ttl_seconds=1)
|
||||
from datetime import datetime, timedelta, timezone
|
||||
|
||||
vc.expires_at = datetime.now(timezone.utc) - timedelta(seconds=1)
|
||||
assert vc.is_valid is False
|
||||
|
||||
def test_is_valid_used(self):
|
||||
"""已使用的验证码无效."""
|
||||
vc = VerificationCode.create(recipient="u@test.com", code_type="email_bind")
|
||||
vc.mark_used()
|
||||
assert vc.is_valid is False
|
||||
|
||||
|
||||
class TestVerificationCodeActions:
|
||||
"""VerificationCode 操作方法测试."""
|
||||
|
||||
def test_mark_used(self):
|
||||
"""标记使用."""
|
||||
vc = VerificationCode.create(recipient="u@test.com", code_type="email_bind")
|
||||
vc.mark_used()
|
||||
assert vc.is_used is True
|
||||
assert vc.used_at is not None
|
||||
|
||||
def test_mark_used_twice(self):
|
||||
"""标记两次也没问题."""
|
||||
vc = VerificationCode.create(recipient="u@test.com", code_type="email_bind")
|
||||
vc.mark_used()
|
||||
first_time = vc.used_at
|
||||
vc.mark_used()
|
||||
# 第二次会覆盖时间
|
||||
assert vc.used_at >= first_time
|
||||
|
||||
def test_increment_attempts(self):
|
||||
"""增加尝试次数."""
|
||||
vc = VerificationCode.create(recipient="u@test.com", code_type="email_bind")
|
||||
assert vc.attempts == 0
|
||||
vc.increment_attempts()
|
||||
assert vc.attempts == 1
|
||||
vc.increment_attempts()
|
||||
assert vc.attempts == 2
|
||||
|
||||
def test_all_code_types_supported(self):
|
||||
"""支持所有code_type."""
|
||||
for code_type in ["email_bind", "phone_bind", "email_login", "phone_login", "reset_password"]:
|
||||
vc = VerificationCode.create(recipient="u@test.com", code_type=code_type)
|
||||
assert vc.code_type == code_type
|
||||
@@ -1,81 +0,0 @@
|
||||
"""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,308 +1,301 @@
|
||||
"""片头片尾引擎单元测试 - 配置解析等纯逻辑."""
|
||||
"""
|
||||
片头片尾引擎配置与纯逻辑测试.
|
||||
|
||||
from __future__ import annotations
|
||||
覆盖 IntroOutroConfig.from_dict / validate / has_intro / has_outro 等纯逻辑.
|
||||
引擎核心 render 方法依赖 FFmpeg,由集成测试覆盖.
|
||||
"""
|
||||
|
||||
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(self):
|
||||
"""None 返回默认配置."""
|
||||
config = IntroOutroConfig.from_dict(None)
|
||||
assert config.enabled is False
|
||||
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_empty_dict_returns_default(self):
|
||||
"""空 dict 返回默认."""
|
||||
config = IntroOutroConfig.from_dict({})
|
||||
assert config.enabled is False
|
||||
def test_empty_dict_returns_default_disabled(self):
|
||||
cfg = IntroOutroConfig.from_dict({})
|
||||
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_false_returns_default_disabled(self):
|
||||
cfg = IntroOutroConfig.from_dict({"enabled": False})
|
||||
assert cfg.enabled is False
|
||||
|
||||
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_text_intro(self):
|
||||
"""文字片头配置."""
|
||||
config = IntroOutroConfig.from_dict({
|
||||
"enabled": True,
|
||||
"intro": {
|
||||
"type": "text",
|
||||
"title": "我的片头",
|
||||
"subtitle": "欢迎收看",
|
||||
},
|
||||
})
|
||||
assert config.intro_type == "text"
|
||||
assert config.intro_title == "我的片头"
|
||||
assert config.intro_subtitle == "欢迎收看"
|
||||
|
||||
def test_video_intro(self):
|
||||
"""视频片头配置."""
|
||||
config = IntroOutroConfig.from_dict({
|
||||
def test_enabled_with_video_intro(self):
|
||||
cfg = IntroOutroConfig.from_dict({
|
||||
"enabled": True,
|
||||
"intro": {
|
||||
"type": "video",
|
||||
"video_path": "/videos/intro.mp4",
|
||||
"video_path": "/tmp/intro.mp4",
|
||||
"duration": 5.0,
|
||||
},
|
||||
"outro": {"type": "none"},
|
||||
})
|
||||
assert config.intro_type == "video"
|
||||
assert config.intro_video_path == "/videos/intro.mp4"
|
||||
assert config.intro_duration == 5.0
|
||||
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_video_intro_video_alias(self):
|
||||
"""video 字段作为 video_path 别名."""
|
||||
config = IntroOutroConfig.from_dict({
|
||||
def test_enabled_with_text_intro(self):
|
||||
cfg = 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,
|
||||
},
|
||||
})
|
||||
assert cfg.enabled is True
|
||||
assert cfg.outro_type == "video"
|
||||
assert cfg.outro_video_path == "/tmp/outro.mp4"
|
||||
assert cfg.outro_duration == 4.0
|
||||
|
||||
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({
|
||||
"enabled": True,
|
||||
"intro": {
|
||||
"type": "video",
|
||||
"video": "/videos/intro.mp4",
|
||||
"video": "/tmp/fallback.mp4",
|
||||
},
|
||||
"outro": {"type": "none"},
|
||||
})
|
||||
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
|
||||
assert cfg.intro_video_path == "/tmp/fallback.mp4"
|
||||
|
||||
def test_transition_config(self):
|
||||
"""转场配置."""
|
||||
config = IntroOutroConfig.from_dict({
|
||||
cfg = IntroOutroConfig.from_dict({
|
||||
"enabled": True,
|
||||
"transition": "dissolve",
|
||||
"intro": {"type": "none"},
|
||||
"outro": {"type": "none"},
|
||||
"transition": "fade",
|
||||
"transition_duration": 1.0,
|
||||
})
|
||||
assert config.transition_effect == "dissolve"
|
||||
assert config.transition_duration == 1.0
|
||||
assert cfg.transition_effect == "fade"
|
||||
assert cfg.transition_duration == 1.0
|
||||
|
||||
def test_empty_intro_dict(self):
|
||||
"""空 intro dict."""
|
||||
config = IntroOutroConfig.from_dict({
|
||||
def test_default_transition(self):
|
||||
cfg = IntroOutroConfig.from_dict({
|
||||
"enabled": True,
|
||||
"intro": {},
|
||||
"intro": {"type": "none"},
|
||||
"outro": {"type": "none"},
|
||||
})
|
||||
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"
|
||||
assert cfg.transition_effect == "fade"
|
||||
assert cfg.transition_duration == 0.5
|
||||
|
||||
|
||||
class TestHasIntroOutro:
|
||||
"""has_intro / has_outro 属性测试."""
|
||||
class TestIntroOutroConfigProperties:
|
||||
"""has_intro / has_outro 属性."""
|
||||
|
||||
def test_no_intro_when_disabled(self):
|
||||
"""禁用时无片头."""
|
||||
config = IntroOutroConfig()
|
||||
assert config.has_intro is False
|
||||
assert config.has_outro is False
|
||||
|
||||
def test_video_intro_has_intro(self):
|
||||
"""视频片头有has_intro."""
|
||||
config = IntroOutroConfig(
|
||||
def test_has_intro_video_type(self):
|
||||
cfg = IntroOutroConfig(
|
||||
enabled=True,
|
||||
intro_type="video",
|
||||
intro_video_path="/a.mp4",
|
||||
intro_video_path="/tmp/a.mp4",
|
||||
)
|
||||
assert config.has_intro is True
|
||||
assert cfg.has_intro is True
|
||||
|
||||
def test_text_intro_has_intro(self):
|
||||
"""文字片头有has_intro."""
|
||||
config = IntroOutroConfig(
|
||||
def test_has_intro_text_type(self):
|
||||
cfg = IntroOutroConfig(
|
||||
enabled=True,
|
||||
intro_type="text",
|
||||
intro_title="test",
|
||||
intro_title="Hi",
|
||||
)
|
||||
assert config.has_intro is True
|
||||
assert cfg.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_no_intro_when_disabled(self):
|
||||
cfg = IntroOutroConfig(
|
||||
enabled=False,
|
||||
intro_type="video",
|
||||
intro_video_path="/tmp/a.mp4",
|
||||
)
|
||||
assert cfg.has_intro is False
|
||||
|
||||
def test_video_outro_has_outro(self):
|
||||
"""视频片尾有has_outro."""
|
||||
config = IntroOutroConfig(
|
||||
def test_no_intro_when_none_type(self):
|
||||
cfg = IntroOutroConfig(
|
||||
enabled=True,
|
||||
intro_type="none",
|
||||
)
|
||||
assert cfg.has_intro is False
|
||||
|
||||
def test_has_outro_video_type(self):
|
||||
cfg = IntroOutroConfig(
|
||||
enabled=True,
|
||||
outro_type="video",
|
||||
outro_video_path="/a.mp4",
|
||||
outro_video_path="/tmp/a.mp4",
|
||||
)
|
||||
assert config.has_outro is True
|
||||
assert cfg.has_outro is True
|
||||
|
||||
def test_text_outro_has_outro(self):
|
||||
"""文字片尾有has_outro."""
|
||||
config = IntroOutroConfig(
|
||||
def test_has_outro_text_type(self):
|
||||
cfg = IntroOutroConfig(
|
||||
enabled=True,
|
||||
outro_type="text",
|
||||
outro_title="test",
|
||||
outro_title="Bye",
|
||||
)
|
||||
assert config.has_outro is True
|
||||
assert cfg.has_outro is True
|
||||
|
||||
def test_follow_outro_has_outro(self):
|
||||
"""follow类型片尾有has_outro."""
|
||||
config = IntroOutroConfig(
|
||||
def test_has_outro_follow_type(self):
|
||||
cfg = IntroOutroConfig(
|
||||
enabled=True,
|
||||
outro_type="follow",
|
||||
outro_title="test",
|
||||
outro_title="Follow me",
|
||||
)
|
||||
assert config.has_outro is True
|
||||
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
|
||||
|
||||
|
||||
class TestValidate:
|
||||
"""validate 配置校验测试."""
|
||||
class TestIntroOutroConfigValidate:
|
||||
"""validate 校验逻辑."""
|
||||
|
||||
def test_disabled_valid(self):
|
||||
"""禁用配置合法."""
|
||||
config = IntroOutroConfig()
|
||||
ok, msg = config.validate()
|
||||
def test_disabled_is_valid(self):
|
||||
cfg = IntroOutroConfig(enabled=False)
|
||||
ok, msg = cfg.validate()
|
||||
assert ok is True
|
||||
assert msg == ""
|
||||
|
||||
def test_video_intro_missing_path(self):
|
||||
"""视频片头缺少路径."""
|
||||
config = IntroOutroConfig(
|
||||
cfg = IntroOutroConfig(
|
||||
enabled=True,
|
||||
intro_type="video",
|
||||
intro_video_path="",
|
||||
outro_type="none",
|
||||
)
|
||||
ok, msg = config.validate()
|
||||
ok, msg = cfg.validate()
|
||||
assert ok is False
|
||||
assert "video_path" in msg
|
||||
|
||||
def test_text_intro_missing_title(self):
|
||||
"""文字片头缺少标题."""
|
||||
config = IntroOutroConfig(
|
||||
cfg = IntroOutroConfig(
|
||||
enabled=True,
|
||||
intro_type="text",
|
||||
intro_title="",
|
||||
outro_type="none",
|
||||
)
|
||||
ok, msg = config.validate()
|
||||
ok, msg = cfg.validate()
|
||||
assert ok is False
|
||||
assert "title" in msg
|
||||
|
||||
def test_video_outro_missing_path(self):
|
||||
"""视频片尾缺少路径."""
|
||||
config = IntroOutroConfig(
|
||||
cfg = IntroOutroConfig(
|
||||
enabled=True,
|
||||
intro_type="none",
|
||||
outro_type="video",
|
||||
outro_video_path="",
|
||||
)
|
||||
ok, msg = config.validate()
|
||||
ok, msg = cfg.validate()
|
||||
assert ok is False
|
||||
assert "video_path" in msg
|
||||
|
||||
def test_text_outro_missing_title(self):
|
||||
"""文字片尾缺少标题."""
|
||||
config = IntroOutroConfig(
|
||||
cfg = IntroOutroConfig(
|
||||
enabled=True,
|
||||
intro_type="none",
|
||||
outro_type="text",
|
||||
outro_title="",
|
||||
)
|
||||
ok, msg = config.validate()
|
||||
ok, msg = cfg.validate()
|
||||
assert ok is False
|
||||
assert "title" in msg
|
||||
|
||||
def test_zero_intro_duration_invalid(self):
|
||||
"""片头时长为0无效."""
|
||||
config = IntroOutroConfig(
|
||||
def test_intro_duration_zero(self):
|
||||
cfg = IntroOutroConfig(
|
||||
enabled=True,
|
||||
intro_type="text",
|
||||
intro_title="test",
|
||||
intro_title="Hi",
|
||||
intro_duration=0,
|
||||
outro_type="none",
|
||||
)
|
||||
ok, msg = config.validate()
|
||||
ok, msg = cfg.validate()
|
||||
assert ok is False
|
||||
assert "时长" in msg
|
||||
assert "片头时长" in msg
|
||||
|
||||
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(
|
||||
def test_intro_duration_negative(self):
|
||||
cfg = IntroOutroConfig(
|
||||
enabled=True,
|
||||
intro_type="text",
|
||||
intro_title="片头",
|
||||
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_duration=3.0,
|
||||
outro_type="text",
|
||||
outro_title="片尾",
|
||||
outro_duration=3.0,
|
||||
outro_title="Thanks",
|
||||
outro_duration=2.0,
|
||||
)
|
||||
ok, msg = config.validate()
|
||||
ok, msg = cfg.validate()
|
||||
assert ok is True
|
||||
assert msg == ""
|
||||
|
||||
@@ -5,7 +5,6 @@ from __future__ import annotations
|
||||
import time
|
||||
|
||||
import pytest
|
||||
from jwt.exceptions import ExpiredSignatureError, InvalidTokenError
|
||||
|
||||
from packages.application.auth.jwt_handler import (
|
||||
JWTHandler,
|
||||
@@ -87,17 +86,17 @@ class TestJWTHandler:
|
||||
# 等待一小段时间确保过期
|
||||
time.sleep(0.1)
|
||||
|
||||
with pytest.raises(ExpiredSignatureError):
|
||||
with pytest.raises(Exception):
|
||||
handler.verify_access_token(token)
|
||||
|
||||
def test_invalid_token_raises_error(self, jwt_handler):
|
||||
"""无效 token 验证失败"""
|
||||
with pytest.raises(InvalidTokenError):
|
||||
with pytest.raises(Exception):
|
||||
jwt_handler.verify_access_token("invalid.token.here")
|
||||
|
||||
def test_empty_token_raises_error(self, jwt_handler):
|
||||
"""空字符串 token 验证失败"""
|
||||
with pytest.raises(InvalidTokenError):
|
||||
with pytest.raises(Exception):
|
||||
jwt_handler.verify_access_token("")
|
||||
|
||||
def test_different_secret_fails_verification(self):
|
||||
@@ -107,7 +106,7 @@ class TestJWTHandler:
|
||||
|
||||
token = handler1.create_access_token(user_id="user_001")
|
||||
|
||||
with pytest.raises(InvalidTokenError):
|
||||
with pytest.raises(Exception):
|
||||
handler2.verify_access_token(token)
|
||||
|
||||
def test_custom_algorithm(self):
|
||||
|
||||
+453
-268
@@ -1,6 +1,12 @@
|
||||
"""Module Registry 单元测试."""
|
||||
"""
|
||||
Module Registry 模块注册中心单元测试
|
||||
|
||||
from __future__ import annotations
|
||||
覆盖:
|
||||
- ModuleStatus 枚举
|
||||
- QuotaRule / ModuleCapability / Module 数据类
|
||||
- Module.activate / disable 状态转换
|
||||
- ModuleRegistry 注册/注销/查询/能力发现/依赖检查
|
||||
"""
|
||||
|
||||
import pytest
|
||||
|
||||
@@ -13,372 +19,551 @@ from packages.infrastructure.module_registry import (
|
||||
module_registry,
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def clean_registry():
|
||||
"""每个测试前后清空全局单例,避免测试间干扰."""
|
||||
module_registry.clear()
|
||||
yield
|
||||
module_registry.clear()
|
||||
# ============================================================
|
||||
# ModuleStatus
|
||||
# ============================================================
|
||||
|
||||
|
||||
# ── Module 数据类测试 ──────────────────────────────────────────────────
|
||||
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
|
||||
|
||||
|
||||
class TestModuleDataclass:
|
||||
"""Module 数据类基本行为测试."""
|
||||
# ============================================================
|
||||
# QuotaRule
|
||||
# ============================================================
|
||||
|
||||
def test_create_module_defaults(self):
|
||||
"""创建模块,默认值正确."""
|
||||
mod = Module(name="test_module")
|
||||
assert mod.name == "test_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")
|
||||
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_create_module_full(self):
|
||||
"""创建模块,完整参数."""
|
||||
def test_full_module(self):
|
||||
cap = ModuleCapability(name="do_something")
|
||||
mod = Module(
|
||||
name="ai_voice",
|
||||
name="full_module",
|
||||
version="2.0.0",
|
||||
description="AI配音模块",
|
||||
capabilities=[ModuleCapability(name="gen_voice")],
|
||||
dependencies=["core"],
|
||||
description="完整模块",
|
||||
capabilities=[cap],
|
||||
dependencies=["dep1", "dep2"],
|
||||
status=ModuleStatus.ACTIVE,
|
||||
config={"key": "value"},
|
||||
)
|
||||
assert mod.name == "ai_voice"
|
||||
assert mod.version == "2.0.0"
|
||||
assert mod.description == "AI配音模块"
|
||||
assert mod.description == "完整模块"
|
||||
assert len(mod.capabilities) == 1
|
||||
assert mod.dependencies == ["core"]
|
||||
assert mod.dependencies == ["dep1", "dep2"]
|
||||
assert mod.status == ModuleStatus.ACTIVE
|
||||
assert mod.config == {"key": "value"}
|
||||
assert mod.config["key"] == "value"
|
||||
|
||||
def test_module_activate(self):
|
||||
"""激活模块."""
|
||||
mod = Module(name="m1")
|
||||
assert mod.status == ModuleStatus.REGISTERED
|
||||
|
||||
class TestModuleActivate:
|
||||
"""Module.activate 状态转换"""
|
||||
|
||||
def test_activate_from_registered(self):
|
||||
mod = Module(name="test")
|
||||
mod.activate()
|
||||
assert mod.status == ModuleStatus.ACTIVE
|
||||
|
||||
def test_module_activate_error_state_ignored(self):
|
||||
"""error状态的模块不能激活."""
|
||||
mod = Module(name="m1", status=ModuleStatus.ERROR)
|
||||
def test_activate_from_disabled(self):
|
||||
mod = Module(name="test", status=ModuleStatus.DISABLED)
|
||||
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_module_disable(self):
|
||||
"""禁用模块."""
|
||||
mod = Module(name="m1", status=ModuleStatus.ACTIVE)
|
||||
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)
|
||||
mod.disable()
|
||||
assert mod.status == ModuleStatus.DISABLED
|
||||
|
||||
|
||||
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 == ""
|
||||
# ============================================================
|
||||
# ModuleRegistry - 基础操作
|
||||
# ============================================================
|
||||
|
||||
|
||||
class TestModuleCapability:
|
||||
"""ModuleCapability 测试."""
|
||||
class TestModuleRegistryBasic:
|
||||
"""ModuleRegistry 基础操作"""
|
||||
|
||||
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_empty_registry(self):
|
||||
registry = ModuleRegistry()
|
||||
assert registry.list_modules() == []
|
||||
assert registry.get_active_capabilities() == {}
|
||||
|
||||
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="m1"))
|
||||
registry.register(Module(name="test_mod"))
|
||||
with pytest.raises(ValueError, match="already registered"):
|
||||
registry.register(Module(name="m1"))
|
||||
|
||||
def test_register_auto_activate_no_deps(self):
|
||||
"""无依赖的模块注册后自动激活."""
|
||||
registry = ModuleRegistry()
|
||||
registry.register(Module(name="m1"))
|
||||
assert registry.get("m1").status == ModuleStatus.ACTIVE
|
||||
|
||||
def test_register_with_missing_dependency(self):
|
||||
"""有未满足依赖的模块保持REGISTERED."""
|
||||
registry = ModuleRegistry()
|
||||
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("nonexistent")
|
||||
|
||||
def test_unregister_with_dependents_raises(self):
|
||||
"""被其他模块依赖时不能注销."""
|
||||
registry = ModuleRegistry()
|
||||
registry.register(Module(name="core"))
|
||||
registry.register(Module(name="plugin", dependencies=["core"]))
|
||||
with pytest.raises(ValueError, match="depended on by"):
|
||||
registry.unregister("core")
|
||||
|
||||
|
||||
class TestModuleRegistryQuery:
|
||||
"""模块查询测试."""
|
||||
registry.register(Module(name="test_mod"))
|
||||
|
||||
def test_get_nonexistent_returns_none(self):
|
||||
"""获取不存在的模块返回None."""
|
||||
registry = ModuleRegistry()
|
||||
assert registry.get("nonexistent") is None
|
||||
assert registry.get("no_such_module") is None
|
||||
|
||||
def test_list_modules_all(self):
|
||||
"""列出所有模块."""
|
||||
def test_unregister_success(self):
|
||||
registry = ModuleRegistry()
|
||||
registry.register(Module(name="m1"))
|
||||
registry.register(Module(name="m2"))
|
||||
assert len(registry.list_modules()) == 2
|
||||
registry.register(Module(name="test_mod"))
|
||||
registry.unregister("test_mod")
|
||||
assert registry.get("test_mod") is None
|
||||
|
||||
def test_list_modules_by_status(self):
|
||||
"""按状态过滤模块."""
|
||||
def test_unregister_nonexistent_raises(self):
|
||||
registry = ModuleRegistry()
|
||||
registry.register(Module(name="m1")) # ACTIVE
|
||||
m2 = Module(name="m2", status=ModuleStatus.DISABLED)
|
||||
registry.register(m2)
|
||||
m2.disable()
|
||||
with pytest.raises(KeyError, match="not found"):
|
||||
registry.unregister("no_such_module")
|
||||
|
||||
def test_unregister_with_dependents_raises(self):
|
||||
registry = ModuleRegistry()
|
||||
registry.register(Module(name="base_module"))
|
||||
registry.register(Module(name="dependent_module", dependencies=["base_module"]))
|
||||
with pytest.raises(ValueError, match="depended on by"):
|
||||
registry.unregister("base_module")
|
||||
|
||||
def test_clear(self):
|
||||
registry = ModuleRegistry()
|
||||
registry.register(Module(name="mod1"))
|
||||
registry.register(Module(name="mod2"))
|
||||
registry.clear()
|
||||
assert registry.list_modules() == []
|
||||
|
||||
|
||||
# ============================================================
|
||||
# ModuleRegistry - 自动激活 & 依赖
|
||||
# ============================================================
|
||||
|
||||
|
||||
class TestModuleRegistryAutoActivate:
|
||||
"""注册时自动激活逻辑"""
|
||||
|
||||
def test_no_deps_auto_activates(self):
|
||||
registry = ModuleRegistry()
|
||||
mod = Module(name="standalone")
|
||||
registry.register(mod)
|
||||
assert mod.status == ModuleStatus.ACTIVE
|
||||
|
||||
def test_with_deps_all_satisfied_auto_activates(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()
|
||||
|
||||
active = registry.list_modules(status=ModuleStatus.ACTIVE)
|
||||
assert len(active) == 1
|
||||
assert active[0].name == "m1"
|
||||
assert active[0].name == "active_mod"
|
||||
|
||||
def test_list_modules_disabled(self):
|
||||
"""列出已禁用模块."""
|
||||
def test_filter_by_disabled(self):
|
||||
registry = 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"
|
||||
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 - 能力发现
|
||||
# ============================================================
|
||||
|
||||
|
||||
class TestModuleRegistryCapabilities:
|
||||
"""能力查询测试."""
|
||||
"""能力发现:has_capability / get_capability / get_quota_rules"""
|
||||
|
||||
def test_has_capability_true(self):
|
||||
"""检查已存在的能力."""
|
||||
registry = ModuleRegistry()
|
||||
registry.register(
|
||||
Module(
|
||||
name="ai_mod",
|
||||
name="voice_module",
|
||||
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="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")],
|
||||
registry.register(
|
||||
Module(
|
||||
name="voice_module",
|
||||
capabilities=[ModuleCapability(name="generate_voice")],
|
||||
)
|
||||
)
|
||||
registry._modules["ai_mod"] = m
|
||||
assert registry.has_capability("generate_voice") is False
|
||||
assert registry.has_capability("generate_video") is False
|
||||
|
||||
def test_get_capability_returns_definition(self):
|
||||
"""获取能力定义."""
|
||||
def test_has_capability_inactive_module_not_counted(self):
|
||||
registry = ModuleRegistry()
|
||||
cap = ModuleCapability(name="gen_voice", description="配音")
|
||||
registry.register(Module(name="ai_mod", capabilities=[cap]))
|
||||
result = registry.get_capability("gen_voice")
|
||||
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")
|
||||
assert result is not None
|
||||
assert result.name == "gen_voice"
|
||||
assert result.description == "配音"
|
||||
assert result.name == "export"
|
||||
# 返回第一个匹配的(mod1)
|
||||
assert result.description == "导出1"
|
||||
|
||||
def test_get_capability_nonexistent(self):
|
||||
"""获取不存在的能力返回None."""
|
||||
def test_get_capability_nonexistent_returns_none(self):
|
||||
registry = ModuleRegistry()
|
||||
assert registry.get_capability("nonexistent") is None
|
||||
assert registry.get_capability("no_such_cap") is None
|
||||
|
||||
def test_get_quota_rules_empty(self):
|
||||
"""没有配额规则时返回空列表."""
|
||||
def test_get_quota_rules(self):
|
||||
rules = [
|
||||
QuotaRule(dimension="credits", per_operation=1.0),
|
||||
QuotaRule(dimension="storage", per_operation=0.5),
|
||||
]
|
||||
registry = ModuleRegistry()
|
||||
registry.register(
|
||||
Module(
|
||||
name="m1",
|
||||
capabilities=[ModuleCapability(name="do_something")],
|
||||
name="voice_mod",
|
||||
capabilities=[ModuleCapability(name="gen", quota_rules=rules)],
|
||||
)
|
||||
)
|
||||
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
|
||||
result = registry.get_quota_rules("gen")
|
||||
assert len(result) == 2
|
||||
assert result[0].dimension == "credits"
|
||||
assert result[0].per_operation == 2.0
|
||||
assert result[1].dimension == "storage"
|
||||
|
||||
def test_get_active_capabilities(self):
|
||||
"""获取所有已激活模块的能力."""
|
||||
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):
|
||||
registry = ModuleRegistry()
|
||||
registry.register(
|
||||
Module(
|
||||
name="mod_a",
|
||||
name="voice_mod",
|
||||
capabilities=[
|
||||
ModuleCapability(name="cap_a1"),
|
||||
ModuleCapability(name="cap_a2"),
|
||||
ModuleCapability(name="generate_voice"),
|
||||
ModuleCapability(name="clone_voice"),
|
||||
],
|
||||
)
|
||||
)
|
||||
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="mod_b",
|
||||
capabilities=[ModuleCapability(name="cap_b1")],
|
||||
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")],
|
||||
)
|
||||
)
|
||||
result = registry.get_active_capabilities()
|
||||
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"]
|
||||
assert len(result) == 2
|
||||
assert result["mod1"] == ["cap_a"]
|
||||
assert set(result["mod2"]) == {"cap_b", "cap_c"}
|
||||
|
||||
|
||||
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 TestModuleRegistryClear:
|
||||
"""清空注册测试."""
|
||||
class TestGlobalSingleton:
|
||||
"""全局 module_registry 单例"""
|
||||
|
||||
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_exists(self):
|
||||
assert module_registry is not None
|
||||
assert isinstance(module_registry, ModuleRegistry)
|
||||
|
||||
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
|
||||
def test_singleton_is_same_instance(self):
|
||||
from packages.infrastructure.module_registry import module_registry as 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"
|
||||
assert module_registry is mr2
|
||||
|
||||
@@ -1,4 +1,9 @@
|
||||
"""多轨道混音单元测试 - 配置解析等纯逻辑."""
|
||||
"""
|
||||
多轨道混音引擎配置与纯逻辑测试.
|
||||
|
||||
覆盖 AudioTrack.from_dict / MultiTrackMixConfig.from_config_dict / has_effect 等纯逻辑.
|
||||
引擎核心混音方法依赖 FFmpeg,由集成测试覆盖.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
@@ -16,367 +21,383 @@ from video_processing.multi_track_mixer import (
|
||||
)
|
||||
|
||||
|
||||
class TestConstants:
|
||||
"""常量测试."""
|
||||
class TestTrackConstants:
|
||||
"""轨道类型常量与默认值."""
|
||||
|
||||
def test_track_types(self):
|
||||
"""5种轨道类型."""
|
||||
def test_track_types_exist(self):
|
||||
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_parsing(self):
|
||||
"""基本解析."""
|
||||
def test_basic(self):
|
||||
track = AudioTrack.from_dict(
|
||||
{
|
||||
"track_id": "t1",
|
||||
"track_type": "bgm",
|
||||
"audio_path": "/bgm.mp3",
|
||||
"audio_path": "/tmp/bgm.mp3",
|
||||
}
|
||||
)
|
||||
assert track.track_id == "t1"
|
||||
assert track.track_type == "bgm"
|
||||
assert track.audio_path == "/bgm.mp3"
|
||||
assert track.audio_path == "/tmp/bgm.mp3"
|
||||
assert track.volume == 0.3 # bgm 默认音量
|
||||
|
||||
def test_default_volume_by_type_bgm(self):
|
||||
"""bgm默认音量0.3."""
|
||||
def test_custom_volume(self):
|
||||
track = AudioTrack.from_dict(
|
||||
{
|
||||
"track_id": "t1",
|
||||
"track_type": "bgm",
|
||||
"audio_path": "/a.mp3",
|
||||
"track_type": "main",
|
||||
"audio_path": "/tmp/main.wav",
|
||||
"volume": 0.8,
|
||||
}
|
||||
)
|
||||
assert track.volume == 0.3
|
||||
assert track.volume == 0.8
|
||||
|
||||
def test_default_volume_by_type_sfx(self):
|
||||
"""sfx默认音量0.7."""
|
||||
def test_volume_clamped_to_zero(self):
|
||||
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": "bgm",
|
||||
"audio_path": "/a.mp3",
|
||||
"volume": 0.5,
|
||||
}
|
||||
)
|
||||
assert track.volume == 0.5
|
||||
|
||||
def test_volume_clamped_high(self):
|
||||
"""音量上限钳制."""
|
||||
track = AudioTrack.from_dict(
|
||||
{
|
||||
"track_id": "t1",
|
||||
"track_type": "bgm",
|
||||
"audio_path": "/a.mp3",
|
||||
"volume": 3.0,
|
||||
}
|
||||
)
|
||||
assert track.volume == 2.0
|
||||
|
||||
def test_volume_clamped_low(self):
|
||||
"""音量下限钳制."""
|
||||
track = AudioTrack.from_dict(
|
||||
{
|
||||
"track_id": "t1",
|
||||
"audio_path": "/a.mp3",
|
||||
"audio_path": "/tmp/sfx.wav",
|
||||
"volume": -1.0,
|
||||
}
|
||||
)
|
||||
assert track.volume == 0.0
|
||||
|
||||
def test_volume_invalid_falls_back(self):
|
||||
"""无效音量回退到类型默认值."""
|
||||
def test_volume_clamped_to_max(self):
|
||||
track = AudioTrack.from_dict(
|
||||
{
|
||||
"track_id": "t1",
|
||||
"track_type": "sfx",
|
||||
"audio_path": "/tmp/sfx.wav",
|
||||
"volume": 3.0,
|
||||
}
|
||||
)
|
||||
assert track.volume == 2.0
|
||||
|
||||
def test_invalid_volume_falls_back_to_default(self):
|
||||
track = AudioTrack.from_dict(
|
||||
{
|
||||
"track_id": "t1",
|
||||
"track_type": "bgm",
|
||||
"audio_path": "/a.mp3",
|
||||
"audio_path": "/tmp/bgm.mp3",
|
||||
"volume": "not_a_number",
|
||||
}
|
||||
)
|
||||
assert track.volume == 0.3
|
||||
assert track.volume == 0.3 # bgm 默认
|
||||
|
||||
def test_fade_in(self):
|
||||
"""淡入时长."""
|
||||
def test_none_volume_falls_back(self):
|
||||
track = AudioTrack.from_dict(
|
||||
{
|
||||
"track_id": "t1",
|
||||
"audio_path": "/a.mp3",
|
||||
"fade_in": 2.5,
|
||||
"track_type": "voiceover",
|
||||
"audio_path": "/tmp/vo.wav",
|
||||
"volume": None,
|
||||
}
|
||||
)
|
||||
assert track.fade_in == 2.5
|
||||
assert track.volume == 1.0 # voiceover 默认
|
||||
|
||||
def test_fade_negative_clamped(self):
|
||||
"""负淡入钳制到0."""
|
||||
def test_unknown_track_type_default_volume(self):
|
||||
track = AudioTrack.from_dict(
|
||||
{
|
||||
"track_id": "t1",
|
||||
"audio_path": "/a.mp3",
|
||||
"fade_in": -1.0,
|
||||
"fade_out": -2.0,
|
||||
"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,
|
||||
}
|
||||
)
|
||||
assert track.fade_in == 0.0
|
||||
assert track.fade_out == 0.0
|
||||
|
||||
def test_start_time(self):
|
||||
"""开始时间."""
|
||||
def test_invalid_fade_falls_back(self):
|
||||
track = AudioTrack.from_dict(
|
||||
{
|
||||
"track_id": "t1",
|
||||
"audio_path": "/a.mp3",
|
||||
"start_time": 5.5,
|
||||
"track_type": "bgm",
|
||||
"audio_path": "/tmp/bgm.mp3",
|
||||
"fade_in": "abc",
|
||||
"fade_out": None,
|
||||
}
|
||||
)
|
||||
assert track.start_time == 5.5
|
||||
assert track.fade_in == 0.0
|
||||
assert track.fade_out == 0.0
|
||||
|
||||
def test_start_time_negative_clamped(self):
|
||||
"""负开始时间钳制到0."""
|
||||
def test_start_time_and_duration(self):
|
||||
track = AudioTrack.from_dict(
|
||||
{
|
||||
"track_id": "t1",
|
||||
"audio_path": "/a.mp3",
|
||||
"start_time": -3.0,
|
||||
"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,
|
||||
}
|
||||
)
|
||||
assert track.start_time == 0.0
|
||||
assert track.duration == 0.0
|
||||
|
||||
def test_disabled_track(self):
|
||||
"""禁用轨道."""
|
||||
def test_invalid_time_values_fall_back(self):
|
||||
track = AudioTrack.from_dict(
|
||||
{
|
||||
"track_id": "t1",
|
||||
"audio_path": "/a.mp3",
|
||||
"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",
|
||||
"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 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 测试."""
|
||||
class TestMultiTrackMixConfigFromDict:
|
||||
"""MultiTrackMixConfig.from_config_dict 构造逻辑."""
|
||||
|
||||
def test_none_returns_default(self):
|
||||
"""None返回默认配置."""
|
||||
config = MultiTrackMixConfig.from_config_dict(None)
|
||||
assert config.tracks == []
|
||||
assert config.master_volume == 1.0
|
||||
cfg = MultiTrackMixConfig.from_config_dict(None)
|
||||
assert cfg.tracks == []
|
||||
assert cfg.master_volume == 1.0
|
||||
assert cfg.normalize is True
|
||||
|
||||
def test_empty_dict_returns_default(self):
|
||||
"""空dict返回默认."""
|
||||
config = MultiTrackMixConfig.from_config_dict({})
|
||||
assert config.tracks == []
|
||||
cfg = MultiTrackMixConfig.from_config_dict({})
|
||||
assert cfg.tracks == []
|
||||
|
||||
def test_non_dict_returns_default(self):
|
||||
cfg = MultiTrackMixConfig.from_config_dict([])
|
||||
assert cfg.tracks == []
|
||||
|
||||
def test_single_track(self):
|
||||
"""单轨道."""
|
||||
config = MultiTrackMixConfig.from_config_dict(
|
||||
cfg = MultiTrackMixConfig.from_config_dict(
|
||||
{
|
||||
"tracks": [
|
||||
{
|
||||
"track_id": "bgm1",
|
||||
"track_type": "bgm",
|
||||
"audio_path": "/bgm.mp3",
|
||||
"audio_path": "/tmp/bgm.mp3",
|
||||
"volume": 0.5,
|
||||
},
|
||||
],
|
||||
}
|
||||
)
|
||||
assert len(config.tracks) == 1
|
||||
assert config.tracks[0].track_id == "bgm1"
|
||||
assert len(cfg.tracks) == 1
|
||||
assert cfg.tracks[0].track_id == "bgm1"
|
||||
assert cfg.tracks[0].volume == 0.5
|
||||
|
||||
def test_multiple_tracks(self):
|
||||
"""多轨道."""
|
||||
config = MultiTrackMixConfig.from_config_dict(
|
||||
cfg = MultiTrackMixConfig.from_config_dict(
|
||||
{
|
||||
"tracks": [
|
||||
{"track_id": "t1", "track_type": "bgm", "audio_path": "/a.mp3"},
|
||||
{"track_id": "t2", "track_type": "sfx", "audio_path": "/b.mp3"},
|
||||
{"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"},
|
||||
],
|
||||
}
|
||||
)
|
||||
assert len(config.tracks) == 2
|
||||
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"
|
||||
|
||||
def test_skips_disabled_tracks(self):
|
||||
"""跳过禁用轨道."""
|
||||
config = MultiTrackMixConfig.from_config_dict(
|
||||
def test_disabled_tracks_filtered(self):
|
||||
cfg = MultiTrackMixConfig.from_config_dict(
|
||||
{
|
||||
"tracks": [
|
||||
{"track_id": "t1", "audio_path": "/a.mp3", "enabled": True},
|
||||
{"track_id": "t2", "audio_path": "/b.mp3", "enabled": False},
|
||||
{"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"},
|
||||
],
|
||||
}
|
||||
)
|
||||
assert len(config.tracks) == 1
|
||||
assert config.tracks[0].track_id == "t1"
|
||||
assert len(cfg.tracks) == 2
|
||||
assert all(t.track_id != "b" for t in cfg.tracks)
|
||||
|
||||
def test_skips_no_audio_path(self):
|
||||
"""跳过无audio_path的轨道."""
|
||||
config = MultiTrackMixConfig.from_config_dict(
|
||||
def test_empty_audio_path_filtered(self):
|
||||
cfg = MultiTrackMixConfig.from_config_dict(
|
||||
{
|
||||
"tracks": [
|
||||
{"track_id": "t1", "audio_path": "/a.mp3"},
|
||||
{"track_id": "t2", "audio_path": ""},
|
||||
{"track_id": "t3"},
|
||||
{"track_id": "valid", "track_type": "sfx", "audio_path": "/tmp/a.wav"},
|
||||
{"track_id": "empty", "track_type": "sfx", "audio_path": ""},
|
||||
],
|
||||
}
|
||||
)
|
||||
assert len(config.tracks) == 1
|
||||
assert len(cfg.tracks) == 1
|
||||
assert cfg.tracks[0].track_id == "valid"
|
||||
|
||||
def test_master_volume(self):
|
||||
"""主音量."""
|
||||
config = MultiTrackMixConfig.from_config_dict(
|
||||
def test_invalid_tracks_skipped(self):
|
||||
cfg = MultiTrackMixConfig.from_config_dict(
|
||||
{
|
||||
"master_volume": 0.8,
|
||||
"tracks": [{"track_id": "t1", "audio_path": "/a.mp3"}],
|
||||
"tracks": [
|
||||
{"track_id": "ok", "track_type": "sfx", "audio_path": "/tmp/a.wav"},
|
||||
"not_a_dict",
|
||||
None,
|
||||
{"no_audio_path": "xxx"},
|
||||
],
|
||||
}
|
||||
)
|
||||
assert config.master_volume == 0.8
|
||||
assert len(cfg.tracks) == 1
|
||||
|
||||
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(
|
||||
def test_tracks_not_a_list(self):
|
||||
cfg = MultiTrackMixConfig.from_config_dict(
|
||||
{
|
||||
"tracks": "not_a_list",
|
||||
}
|
||||
)
|
||||
assert config.tracks == []
|
||||
assert cfg.tracks == []
|
||||
|
||||
def test_non_dict_track_skipped(self):
|
||||
"""非dict轨道跳过."""
|
||||
config = MultiTrackMixConfig.from_config_dict(
|
||||
def test_master_volume(self):
|
||||
cfg = MultiTrackMixConfig.from_config_dict(
|
||||
{
|
||||
"tracks": [
|
||||
{"track_id": "t1", "audio_path": "/a.mp3"},
|
||||
"not_a_dict",
|
||||
],
|
||||
"tracks": [],
|
||||
"master_volume": 0.8,
|
||||
}
|
||||
)
|
||||
assert len(config.tracks) == 1
|
||||
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
|
||||
|
||||
|
||||
class TestHasEffect:
|
||||
"""has_effect 属性测试."""
|
||||
class TestMultiTrackMixConfigProperties:
|
||||
"""has_effect 属性."""
|
||||
|
||||
def test_no_tracks_no_effect(self):
|
||||
"""无轨道无效果."""
|
||||
config = MultiTrackMixConfig()
|
||||
assert config.has_effect is False
|
||||
|
||||
def test_with_tracks_has_effect(self):
|
||||
"""有轨道有效果."""
|
||||
config = MultiTrackMixConfig(
|
||||
def test_has_effect_with_tracks(self):
|
||||
cfg = MultiTrackMixConfig(
|
||||
tracks=[
|
||||
AudioTrack(track_id="t1", track_type="bgm", audio_path="/a.mp3"),
|
||||
AudioTrack(track_id="t1", track_type="bgm", audio_path="/tmp/a.mp3"),
|
||||
]
|
||||
)
|
||||
assert config.has_effect is True
|
||||
assert cfg.has_effect is True
|
||||
|
||||
def test_disabled_tracks_no_effect(self):
|
||||
"""所有轨道都禁用无效果."""
|
||||
config = MultiTrackMixConfig(
|
||||
def test_no_effect_empty(self):
|
||||
cfg = MultiTrackMixConfig(tracks=[])
|
||||
assert cfg.has_effect is False
|
||||
|
||||
def test_no_effect_all_disabled(self):
|
||||
cfg = MultiTrackMixConfig(
|
||||
tracks=[
|
||||
AudioTrack(track_id="t1", track_type="bgm", audio_path="/a.mp3", enabled=False),
|
||||
AudioTrack(track_id="t1", track_type="bgm", audio_path="/tmp/a.mp3", enabled=False),
|
||||
]
|
||||
)
|
||||
assert config.has_effect is 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
|
||||
|
||||
@@ -1,232 +0,0 @@
|
||||
"""降噪引擎单元测试 - 配置解析等纯逻辑."""
|
||||
|
||||
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
|
||||
@@ -1,146 +0,0 @@
|
||||
"""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
|
||||
@@ -242,147 +242,3 @@ class TestPaginateFunction:
|
||||
|
||||
assert len(result.data) == 2
|
||||
assert result.data[0]["id"] == 1
|
||||
|
||||
|
||||
# ── PaginationParams 补充边界 ───────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestPaginationParamsEdgeCases:
|
||||
"""PaginationParams 补充边界场景."""
|
||||
|
||||
def test_page_size_1_minimum(self):
|
||||
"""page_size=1 是允许的最小值."""
|
||||
params = PaginationParams(page_size=1)
|
||||
assert params.page_size == 1
|
||||
assert params.limit == 1
|
||||
|
||||
def test_page_size_100_maximum(self):
|
||||
"""page_size=100 是允许的最大值."""
|
||||
params = PaginationParams(page_size=100)
|
||||
assert params.page_size == 100
|
||||
|
||||
def test_offset_page_1_size_100(self):
|
||||
"""第1页每页100条 offset=0."""
|
||||
params = PaginationParams(page=1, page_size=100)
|
||||
assert params.offset == 0
|
||||
|
||||
def test_offset_page_100_size_100(self):
|
||||
"""第100页每页100条 offset=9900."""
|
||||
params = PaginationParams(page=100, page_size=100)
|
||||
assert params.offset == 9900
|
||||
|
||||
def test_large_page_number_accepted(self):
|
||||
"""极大页码(超过实际页数)允许."""
|
||||
params = PaginationParams(page=999999, page_size=20)
|
||||
assert params.page == 999999
|
||||
assert params.offset == (999999 - 1) * 20
|
||||
|
||||
|
||||
# ── PaginationMeta 补充边界 ─────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestPaginationMetaEdgeCases:
|
||||
"""PaginationMeta 补充边界场景."""
|
||||
|
||||
def test_total_0_page_1(self):
|
||||
"""total=0, page=1 时 total_pages=0, 无上下页."""
|
||||
params = PaginationParams(page=1, page_size=20)
|
||||
meta = PaginationMeta.from_params(params, total=0)
|
||||
assert meta.total_pages == 0
|
||||
assert meta.has_next is False
|
||||
assert meta.has_prev is False
|
||||
|
||||
def test_total_0_page_beyond(self):
|
||||
"""total=0, page>1 时 has_prev=True(因为page>1)."""
|
||||
params = PaginationParams(page=3, page_size=20)
|
||||
meta = PaginationMeta.from_params(params, total=0)
|
||||
assert meta.total_pages == 0
|
||||
assert meta.has_next is False
|
||||
assert meta.has_prev is True
|
||||
|
||||
def test_exact_last_page(self):
|
||||
"""刚好是最后一页时 has_next=False."""
|
||||
params = PaginationParams(page=5, page_size=10)
|
||||
meta = PaginationMeta.from_params(params, total=50)
|
||||
assert meta.total_pages == 5
|
||||
assert meta.has_next is False
|
||||
assert meta.has_prev is True
|
||||
|
||||
def test_one_more_than_exact(self):
|
||||
"""比整数页多1条时总页数+1."""
|
||||
params = PaginationParams(page=1, page_size=10)
|
||||
meta = PaginationMeta.from_params(params, total=51)
|
||||
assert meta.total_pages == 6
|
||||
|
||||
def test_page_exactly_total_pages(self):
|
||||
"""page == total_pages 时 has_next=False."""
|
||||
params = PaginationParams(page=3, page_size=10)
|
||||
meta = PaginationMeta.from_params(params, total=30)
|
||||
assert meta.has_next is False
|
||||
|
||||
def test_total_1_page_1_size_1(self):
|
||||
"""1条数据1页."""
|
||||
params = PaginationParams(page=1, page_size=1)
|
||||
meta = PaginationMeta.from_params(params, total=1)
|
||||
assert meta.total_pages == 1
|
||||
assert meta.has_next is False
|
||||
assert meta.has_prev is False
|
||||
|
||||
|
||||
# ── paginate 补充边界 ──────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestPaginateEdgeCases:
|
||||
"""paginate 补充边界场景."""
|
||||
|
||||
def test_single_item_list(self):
|
||||
"""单元素列表."""
|
||||
result = paginate([42], PaginationParams(page=1, page_size=10))
|
||||
assert result.data == [42]
|
||||
assert result.pagination.total == 1
|
||||
assert result.pagination.total_pages == 1
|
||||
|
||||
def test_page_exactly_last(self):
|
||||
"""刚好在最后一页."""
|
||||
items = list(range(25))
|
||||
result = paginate(items, PaginationParams(page=3, page_size=10))
|
||||
assert result.data == list(range(20, 25))
|
||||
assert result.pagination.has_next is False
|
||||
|
||||
def test_page_past_end_returns_empty(self):
|
||||
"""页码超过总数返回空."""
|
||||
items = list(range(5))
|
||||
result = paginate(items, PaginationParams(page=10, page_size=10))
|
||||
assert result.data == []
|
||||
assert result.pagination.total == 5
|
||||
|
||||
def test_empty_list_page_1(self):
|
||||
"""空列表第1页."""
|
||||
result = paginate([], PaginationParams(page=1, page_size=10))
|
||||
assert result.data == []
|
||||
assert result.pagination.total == 0
|
||||
assert result.pagination.total_pages == 0
|
||||
|
||||
def test_page_size_1_iterates_all(self):
|
||||
"""page_size=1 时每页1条."""
|
||||
items = ["a", "b", "c"]
|
||||
r1 = paginate(items, PaginationParams(page=1, page_size=1))
|
||||
r2 = paginate(items, PaginationParams(page=2, page_size=1))
|
||||
r3 = paginate(items, PaginationParams(page=3, page_size=1))
|
||||
assert r1.data == ["a"]
|
||||
assert r2.data == ["b"]
|
||||
assert r3.data == ["c"]
|
||||
|
||||
def test_does_not_mutate_input(self):
|
||||
"""不修改输入列表."""
|
||||
items = [1, 2, 3, 4, 5]
|
||||
original = items[:]
|
||||
paginate(items, PaginationParams(page=1, page_size=2))
|
||||
assert items == original
|
||||
|
||||
def test_page_size_greater_than_total(self):
|
||||
"""每页条数大于总数."""
|
||||
items = list(range(5))
|
||||
result = paginate(items, PaginationParams(page=1, page_size=100))
|
||||
assert result.data == items
|
||||
assert result.pagination.total_pages == 1
|
||||
|
||||
+552
-226
@@ -1,271 +1,597 @@
|
||||
"""画中画引擎单元测试 - 配置解析+校验等纯逻辑."""
|
||||
"""画中画(PiP)引擎单元测试."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from pathlib import Path
|
||||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
from video_processing.pip_engine import PiPConfig, PiPLayerConfig
|
||||
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,
|
||||
)
|
||||
|
||||
|
||||
class TestPiPLayerConfigDefaults:
|
||||
"""PiPLayerConfig 默认配置测试."""
|
||||
|
||||
def test_default_values(self):
|
||||
"""默认值正确."""
|
||||
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
|
||||
# ── PiPLayerConfig.validate 测试 ──────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestPiPLayerConfigValidate:
|
||||
"""PiPLayerConfig.validate 校验测试."""
|
||||
"""PiP图层配置校验测试."""
|
||||
|
||||
def test_valid_config(self):
|
||||
"""合法配置."""
|
||||
layer = PiPLayerConfig(source="asset_123")
|
||||
ok, msg = layer.validate()
|
||||
assert ok is True
|
||||
assert msg == ""
|
||||
"""正常配置应该通过校验."""
|
||||
layer = PiPLayerConfig(source="asset_001")
|
||||
ok, err = layer.validate()
|
||||
assert ok
|
||||
assert err == ""
|
||||
|
||||
def test_empty_source_invalid(self):
|
||||
"""空source非法."""
|
||||
def test_empty_source(self):
|
||||
"""空source应该失败."""
|
||||
layer = PiPLayerConfig(source="")
|
||||
ok, msg = layer.validate()
|
||||
assert ok is False
|
||||
assert "source" in msg
|
||||
ok, err = layer.validate()
|
||||
assert not ok
|
||||
assert "source" in err
|
||||
|
||||
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
|
||||
"""无效位置应该失败."""
|
||||
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_123",
|
||||
position="custom",
|
||||
x=100,
|
||||
y=100,
|
||||
)
|
||||
ok, _ = layer.validate()
|
||||
assert ok is True
|
||||
"""custom位置应该通过."""
|
||||
layer = PiPLayerConfig(source="asset_001", position="custom", x=100, y=50)
|
||||
ok, err = layer.validate()
|
||||
assert ok
|
||||
|
||||
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
|
||||
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_negative(self):
|
||||
"""透明度为负."""
|
||||
layer = PiPLayerConfig(source="a", opacity=-0.1)
|
||||
ok, msg = layer.validate()
|
||||
assert ok is False
|
||||
assert "opacity" in msg
|
||||
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_zero(self):
|
||||
"""透明度边界值0."""
|
||||
layer = PiPLayerConfig(source="a", opacity=0.0)
|
||||
ok, _ = layer.validate()
|
||||
assert ok is True
|
||||
|
||||
def test_opacity_boundary_one(self):
|
||||
"""透明度边界值1."""
|
||||
layer = PiPLayerConfig(source="a", opacity=1.0)
|
||||
ok, _ = layer.validate()
|
||||
assert ok is True
|
||||
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="a", corner_radius=-5)
|
||||
ok, msg = layer.validate()
|
||||
assert ok is False
|
||||
assert "corner_radius" in msg
|
||||
"""负圆角应该失败."""
|
||||
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="a", start_time=-1.0)
|
||||
ok, msg = layer.validate()
|
||||
assert ok is False
|
||||
assert "start_time" in msg
|
||||
"""负开始时间应该失败."""
|
||||
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="a", duration=-2.0)
|
||||
ok, msg = layer.validate()
|
||||
assert ok is False
|
||||
assert "duration" in msg
|
||||
|
||||
def test_zero_duration_valid(self):
|
||||
"""零时长(全程显示)合法."""
|
||||
layer = PiPLayerConfig(source="a", duration=0.0)
|
||||
ok, _ = layer.validate()
|
||||
assert ok is True
|
||||
"""负持续时间应该失败."""
|
||||
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="a", animation_in="invalid_anim")
|
||||
ok, msg = layer.validate()
|
||||
assert ok is False
|
||||
assert "入场动画" in msg
|
||||
"""无效入场动画应该失败."""
|
||||
layer = PiPLayerConfig(source="asset_001", animation_in="spin")
|
||||
ok, err = layer.validate()
|
||||
assert not ok
|
||||
assert "入场动画" in err
|
||||
|
||||
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
|
||||
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_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_zero_duration_valid(self):
|
||||
"""duration=0(全程显示)应该通过."""
|
||||
layer = PiPLayerConfig(source="asset_001", duration=0.0)
|
||||
ok, _ = layer.validate()
|
||||
assert ok
|
||||
|
||||
|
||||
class TestPiPConfigDefaults:
|
||||
"""PiPConfig 默认配置测试."""
|
||||
|
||||
def test_default_values(self):
|
||||
"""默认值正确."""
|
||||
config = PiPConfig()
|
||||
assert config.enabled is False
|
||||
assert config.layers == []
|
||||
# ── PiPConfig.from_dict 测试 ──────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestPiPConfigFromDict:
|
||||
"""PiPConfig.from_dict 解析测试."""
|
||||
"""PiP配置字典解析测试."""
|
||||
|
||||
def test_none_returns_disabled(self):
|
||||
"""None 返回禁用配置."""
|
||||
def test_none_config(self):
|
||||
"""None配置应该返回disabled."""
|
||||
config = PiPConfig.from_dict(None)
|
||||
assert config.enabled is False
|
||||
assert config.layers == []
|
||||
assert not config.enabled
|
||||
assert len(config.layers) == 0
|
||||
|
||||
def test_empty_dict_returns_disabled(self):
|
||||
"""空 dict 返回禁用."""
|
||||
def test_empty_config(self):
|
||||
"""空字典应该返回disabled."""
|
||||
config = PiPConfig.from_dict({})
|
||||
assert config.enabled is False
|
||||
assert not config.enabled
|
||||
|
||||
def test_disabled_returns_disabled(self):
|
||||
"""enabled=False 返回禁用."""
|
||||
config = PiPConfig.from_dict({"enabled": False})
|
||||
assert config.enabled is False
|
||||
|
||||
def test_enabled_no_layers(self):
|
||||
"""启用但无图层,disabled."""
|
||||
config = PiPConfig.from_dict({"enabled": True, "layers": []})
|
||||
assert config.enabled is False
|
||||
assert config.layers == []
|
||||
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):
|
||||
"""单个图层."""
|
||||
config = PiPConfig.from_dict(
|
||||
{
|
||||
"enabled": True,
|
||||
"layers": [
|
||||
{"source": "asset_001", "position": "top_left"},
|
||||
],
|
||||
}
|
||||
)
|
||||
assert config.enabled is True
|
||||
assert len(config.layers) == 1
|
||||
assert config.layers[0].source == "asset_001"
|
||||
assert config.layers[0].position == "top_left"
|
||||
|
||||
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},
|
||||
],
|
||||
}
|
||||
)
|
||||
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
|
||||
|
||||
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_all_invalid_layers_disabled(self):
|
||||
"""全部无效则disabled."""
|
||||
config = PiPConfig.from_dict(
|
||||
{
|
||||
"enabled": True,
|
||||
"layers": [
|
||||
{"source": ""},
|
||||
{"source": ""},
|
||||
],
|
||||
}
|
||||
)
|
||||
assert config.enabled is False
|
||||
assert config.layers == []
|
||||
|
||||
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,
|
||||
},
|
||||
],
|
||||
}
|
||||
)
|
||||
"""单图层解析."""
|
||||
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 == "https://example.com/video.mp4"
|
||||
assert layer.source_type == "url"
|
||||
assert layer.source == "asset_001"
|
||||
assert layer.position == POSITION_TOP_RIGHT
|
||||
assert layer.width == "30%"
|
||||
assert layer.opacity == 0.8
|
||||
assert layer.opacity == 0.9
|
||||
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
|
||||
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
|
||||
|
||||
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
|
||||
assert layer.width == "25%"
|
||||
assert layer.opacity == 1.0
|
||||
assert layer.corner_radius == 0
|
||||
assert layer.start_time == 0.0
|
||||
assert layer.duration == 0.0
|
||||
assert layer.z_index == 1
|
||||
|
||||
|
||||
# ── PiPEngine 位置计算测试 ────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestPiPEnginePosition:
|
||||
"""PiP引擎位置计算测试."""
|
||||
|
||||
@pytest.fixture
|
||||
def engine(self):
|
||||
return PiPEngine(output_width=1920, output_height=1080, output_fps=30)
|
||||
|
||||
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):
|
||||
"""单图层基础滤镜构建."""
|
||||
layer = PiPLayerConfig(
|
||||
source="asset_001",
|
||||
position=POSITION_TOP_RIGHT,
|
||||
width="25%",
|
||||
)
|
||||
sources = [("pip_src_0", layer, fake_video)]
|
||||
|
||||
filters, inputs, final_label = engine.build_pip_filters("base_video", sources, base_input_idx=3)
|
||||
|
||||
# 应该有2个滤镜: 预处理 + overlay
|
||||
assert len(filters) == 2
|
||||
# 输入参数应该有2个(-i + path)
|
||||
assert len(inputs) == 2
|
||||
assert inputs[0] == "-i"
|
||||
assert inputs[1] == str(fake_video)
|
||||
|
||||
# 预处理滤镜应该使用正确的输入索引
|
||||
assert "3:v" in filters[0]
|
||||
# 应该包含scale
|
||||
assert "scale=" in filters[0]
|
||||
# 应该有pip_pre_0标签
|
||||
assert "[pip_pre_0]" in filters[0]
|
||||
|
||||
# overlay滤镜
|
||||
assert "overlay=" in filters[1]
|
||||
assert "[base_video][pip_pre_0]" in filters[1]
|
||||
|
||||
def test_single_layer_final_label(self, engine, fake_video):
|
||||
"""最终输出标签应该正确."""
|
||||
layer = PiPLayerConfig(source="a", position=POSITION_CENTER)
|
||||
sources = [("s0", layer, fake_video)]
|
||||
|
||||
_, _, final_label = engine.build_pip_filters("main_v", sources)
|
||||
assert final_label == "pip_combined_0"
|
||||
|
||||
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),
|
||||
]
|
||||
|
||||
filters, inputs, final_label = engine.build_pip_filters("base", sources, base_input_idx=0)
|
||||
|
||||
# 2层 × 2个滤镜(预处理+overlay)= 4个滤镜
|
||||
assert len(filters) == 4
|
||||
# 2个输入文件
|
||||
assert len(inputs) == 4 # 2 × (-i + path)
|
||||
|
||||
# 输入索引应该连续
|
||||
assert "0:v" in filters[0]
|
||||
assert "1:v" in filters[2]
|
||||
|
||||
# 最终标签应该是第二个overlay的输出
|
||||
assert final_label == "pip_combined_1"
|
||||
|
||||
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
|
||||
|
||||
def test_with_corner_radius(self, engine, fake_video):
|
||||
"""圆角裁剪应该在滤镜中体现."""
|
||||
layer = PiPLayerConfig(source="a", position=POSITION_CENTER, corner_radius=20)
|
||||
sources = [("s0", layer, fake_video)]
|
||||
|
||||
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)]
|
||||
|
||||
filters, _, _ = engine.build_pip_filters("base", sources)
|
||||
pre_filter = filters[0]
|
||||
assert "pad=" in pre_filter
|
||||
assert "red" in pre_filter
|
||||
|
||||
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)]
|
||||
|
||||
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_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)]
|
||||
|
||||
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,
|
||||
)
|
||||
sources = [("s0", layer, fake_video)]
|
||||
|
||||
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,
|
||||
)
|
||||
sources = [("s0", layer, fake_video)]
|
||||
|
||||
filters, _, _ = engine.build_pip_filters("base", sources)
|
||||
overlay_filter = filters[1]
|
||||
# x表达式应该包含动态变化
|
||||
assert "overlay=" in overlay_filter
|
||||
|
||||
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)]
|
||||
|
||||
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
|
||||
|
||||
@@ -82,102 +82,3 @@ class TestRecipe:
|
||||
for itype in ["asset", "title", "voice"]:
|
||||
item = RecipeItem(id="i1", recipe_id="r1", item_type=itype, item_id="x", position=0)
|
||||
assert item.item_type == itype
|
||||
|
||||
def test_item_metadata_independence(self):
|
||||
"""不同 RecipeItem 的 metadata_ 互不影响"""
|
||||
item1 = RecipeItem(id="i1", recipe_id="r1", item_type="asset", item_id="a1")
|
||||
item2 = RecipeItem(id="i2", recipe_id="r1", item_type="asset", item_id="a2")
|
||||
item1.metadata_["key"] = "val"
|
||||
assert "key" not in item2.metadata_
|
||||
|
||||
def test_item_position_negative(self):
|
||||
"""负数 position 也能存"""
|
||||
item = RecipeItem(id="i1", recipe_id="r1", item_type="asset", item_id="a1", position=-5)
|
||||
assert item.position == -5
|
||||
|
||||
def test_item_position_large(self):
|
||||
item = RecipeItem(id="i1", recipe_id="r1", item_type="asset", item_id="a1", position=9999)
|
||||
assert item.position == 9999
|
||||
|
||||
def test_item_empty_item_id(self):
|
||||
item = RecipeItem(id="i1", recipe_id="r1", item_type="title", item_id="")
|
||||
assert item.item_id == ""
|
||||
|
||||
def test_item_type_voice(self):
|
||||
item = RecipeItem(id="i1", recipe_id="r1", item_type="voice", item_id="v1")
|
||||
assert item.item_type == "voice"
|
||||
|
||||
def test_item_type_title(self):
|
||||
item = RecipeItem(id="i1", recipe_id="r1", item_type="title", item_id="t1")
|
||||
assert item.item_type == "title"
|
||||
|
||||
|
||||
class TestRecipeExtended:
|
||||
"""Recipe 深度补充测试"""
|
||||
|
||||
def test_items_order_preserved(self):
|
||||
items = [
|
||||
RecipeItem(id="i3", recipe_id="r1", item_type="asset", item_id="a1", position=2),
|
||||
RecipeItem(id="i1", recipe_id="r1", item_type="title", item_id="t1", position=0),
|
||||
RecipeItem(id="i2", recipe_id="r1", item_type="voice", item_id="v1", position=1),
|
||||
]
|
||||
r = Recipe(id="r1", user_id="u1", name="n", items=items)
|
||||
assert len(r.items) == 3
|
||||
assert r.items[0].position == 2
|
||||
assert r.items[1].position == 0
|
||||
assert r.items[2].position == 1
|
||||
|
||||
def test_empty_items_list(self):
|
||||
r = Recipe(id="r1", user_id="u1", name="n", items=[])
|
||||
assert r.items == []
|
||||
|
||||
def test_many_items(self):
|
||||
items = [
|
||||
RecipeItem(id=f"i{i}", recipe_id="r1", item_type="asset", item_id=f"a{i}", position=i) for i in range(50)
|
||||
]
|
||||
r = Recipe(id="r1", user_id="u1", name="n", items=items)
|
||||
assert len(r.items) == 50
|
||||
assert r.items[0].position == 0
|
||||
assert r.items[49].position == 49
|
||||
|
||||
def test_generation_params_independence(self):
|
||||
params = {"mode": "one_take", "duration": 30}
|
||||
r1 = Recipe(id="r1", user_id="u1", name="n1", generation_params=params)
|
||||
r2 = Recipe(id="r2", user_id="u1", name="n2")
|
||||
r1.generation_params["new_key"] = "new_val"
|
||||
# 传入同一个 dict 会共享,但默认生成的互不影响
|
||||
assert r2.generation_params == {}
|
||||
|
||||
def test_metadata_independence_default(self):
|
||||
r1 = Recipe(id="r1", user_id="u1", name="n1")
|
||||
r2 = Recipe(id="r2", user_id="u1", name="n2")
|
||||
r1.metadata_["key"] = "val"
|
||||
assert "key" not in r2.metadata_
|
||||
|
||||
def test_with_description(self):
|
||||
r = Recipe(id="r1", user_id="u1", name="n", description="这是一个测试配方")
|
||||
assert r.description == "这是一个测试配方"
|
||||
|
||||
def test_with_template_id(self):
|
||||
r = Recipe(id="r1", user_id="u1", name="n", template_id="tpl-123")
|
||||
assert r.template_id == "tpl-123"
|
||||
|
||||
def test_empty_name(self):
|
||||
r = Recipe(id="r1", user_id="u1", name="")
|
||||
assert r.name == ""
|
||||
|
||||
def test_long_name(self):
|
||||
long_name = "配方" * 200
|
||||
r = Recipe(id="r1", user_id="u1", name=long_name)
|
||||
assert r.name == long_name
|
||||
assert len(r.name) == 400
|
||||
|
||||
def test_special_characters_in_name(self):
|
||||
special = "配!@#$%方"
|
||||
r = Recipe(id="r1", user_id="u1", name=special)
|
||||
assert r.name == special
|
||||
|
||||
def test_unicode_name(self):
|
||||
r = Recipe(id="r1", user_id="u1", name="🎬 一键生成配方 · 美食探店")
|
||||
assert "🎬" in r.name
|
||||
assert "美食探店" in r.name
|
||||
|
||||
@@ -1,153 +0,0 @@
|
||||
"""渲染适配器纯逻辑测试 — _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
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user