Compare commits

..

3 Commits

Author SHA1 Message Date
CI Bot 276342520f cleanup: Backend Phase 1 — 5项代码清理任务
CI/CD Pipeline / Validate Code Quality And Tests (pull_request) Failing after 8s
CI/CD Pipeline / Production Browser E2E (pull_request) Failing after 1562h1m1s
CI/CD Pipeline / Build Production Runtime Images (pull_request) Failing after 1562h1m5s
CI/CD Pipeline / Deploy Production (pull_request) Failing after 1562h1m3s
CI/CD Pipeline / Build & Push Staging (Watchtower auto-deploy) (pull_request) Failing after 1562h1m5s
CI/CD Pipeline / Staging API Integration Tests (pull_request) Failing after 1562h1m3s
CI/CD Pipeline / Frontend Lint (pull_request) Has been skipped
CI/CD Pipeline / Staging E2E Tests (pull_request) Failing after 1562h32m38s
1. 修复密码重置接口路径不一致 Bug (auth.py)
   - forgot_password 硬编码 localhost → 使用 APP_BASE_URL 配置

2. 删除 8 处死代码(未使用 import/变量)
   - 清理多个文件中的未使用导入和变量

3. 删除 8 个空文件/空模块
   - 删除无内容的 __init__.py 文件

4. 合并 3 对 100% 完全重复的函数
   - 提取 check_project_access/get_user_plan/require_project_and_library
   - 新建 apps/api/app/api/routes/_helpers.py 作为共享模块
   - 6 个路由文件改为从 _helpers 导入

5. 对齐 6 个废弃/异常环境变量
   - 修复 DATABASE_POOL_RECYLE 拼写错误 → DATABASE_POOL_RECYCLE
   - 添加 JWT_ALGORITHM/JWT_ACCESS_TOKEN_EXPIRE_MINUTES/JWT_REFRESH_TOKEN_EXPIRE_DAYS 到 Settings
   - 修复 jwt_service.py hasattr 字段名匹配
   - .env.example: CORS_ORIGINS → CORS_ORIGINS_RAW(逗号分隔格式)
   - .env.example: 启用 APP_ENV
   - 修复 OSS_ENDPOINT 默认值拼写错误 (aliiyuncs.com → aliyuncs.com)
   - 添加 COSYVOICE_* 变量来源注释

修改文件: 52 个(新增 1,删除 8,修改 43)
2026-07-13 13:23:30 +08:00
用户CI Test 6feb541127 fix: P0-2 sign_url 返回 HTTPS URL(endpoint 加 https:// 前缀)
CI/CD Pipeline / Validate Code Quality And Tests (pull_request) Failing after 10s
CI/CD Pipeline / Production Browser E2E (pull_request) Failing after 1630h1m2s
CI/CD Pipeline / Staging API Integration Tests (pull_request) Failing after 1630h1m4s
CI/CD Pipeline / Staging E2E Tests (pull_request) Failing after 1630h1m4s
CI/CD Pipeline / Build & Push Staging (Watchtower auto-deploy) (pull_request) Failing after 1630h1m6s
CI/CD Pipeline / Build Production Runtime Images (pull_request) Failing after 1630h1m6s
CI/CD Pipeline / Frontend Lint (pull_request) Has been skipped
CI/CD Pipeline / Deploy Production (pull_request) Failing after 1630h32m36s
根因:oss2.Bucket 的 endpoint 参数不带 scheme 时,sign_url() 默认生成
HTTP URL(如 http://bucket.oss-cn-hangzhou.aliyuncs.com/...),
前端/浏览器视为不安全请求拒绝加载。

修复:初始化 oss2.Bucket 前检查 endpoint 是否带 http(s):// 前缀,
不带则自动补 https://,确保 sign_url 输出 HTTPS URL。

新增 2 个测试验证 endpoint scheme 处理逻辑。
2026-07-10 17:27:57 +08:00
用户CI Test 6efac8de4b fix: P0-2 OSS 凭证启动验证 + 诊断日志
CI/CD Pipeline / Validate Code Quality And Tests (pull_request) Failing after 16s
CI/CD Pipeline / Production Browser E2E (pull_request) Failing after 1630h4m17s
CI/CD Pipeline / Deploy Production (pull_request) Failing after 1630h4m19s
CI/CD Pipeline / Staging API Integration Tests (pull_request) Failing after 1630h4m19s
CI/CD Pipeline / Staging E2E Tests (pull_request) Failing after 1630h4m19s
CI/CD Pipeline / Build Production Runtime Images (pull_request) Failing after 1630h4m25s
CI/CD Pipeline / Build & Push Staging (Watchtower auto-deploy) (pull_request) Failing after 1630h4m25s
CI/CD Pipeline / Frontend Lint (pull_request) Has been skipped
- config.py: 非开发环境 OSS_ACCESS_KEY_ID/SECRET 为空时启动失败(fail-fast)
- storage.py: 新增 diagnose() 方法,启动时输出 OSS 配置状态
- 新增 7 个单元测试覆盖凭证验证和诊断逻辑

注意:staging 实际已有 OSS 凭证配置,预签名 URL 可正常生成。
真正问题是 sign_url 返回 HTTP 而非 HTTPS,后续修复。
2026-07-10 17:23:58 +08:00
126 changed files with 8923 additions and 11693 deletions
+3 -8
View File
@@ -46,15 +46,10 @@ OSS_ACCESS_KEY_SECRET=your-access-key-secret
OSS_BUCKET_NAME=xiaoxia-autocut
# ==================== CosyVoice 语音合成配置 ====================
# 注意:base_url 只需写到 /api/v1,具体路径由代码拼接
# 模型: cosyvoice-v3-flash (推荐,支持系统音色,性价比高)
# cosyvoice-v3-plus (高质量,系统音色少)
# cosyvoice-v3.5-flash / cosyvoice-v3.5-plus (仅支持克隆/设计音色,无系统音色)
# 音色: v3系列系统音色带 _v3 后缀,如 longxiaochun_v3, longxiaoxia_v3, longanyang (无后缀)
# 注意:COSYVOICE_* 变量由 packages/shared/config.py 的 SharedSettings 读取
COSYVOICE_API_KEY=your-cosyvoice-api-key
COSYVOICE_BASE_URL=https://dashscope.aliyuncs.com/api/v1
COSYVOICE_MODEL=cosyvoice-v3-flash
COSYVOICE_VOICE=longxiaochun_v3
COSYVOICE_BASE_URL=https://dashscope.aliyuncs.com/api/v1/services/aigc/text2audio
COSYVOICE_MODEL=cosyvoice-v1
COSYVOICE_VOICE=longxiaochun
COSYVOICE_SAMPLE_RATE=22050
COSYVOICE_FORMAT=mp3
Executable → Regular
+3 -8
View File
@@ -42,15 +42,10 @@ OSS_DIRECT_UPLOAD_MAX_MB=2000
OSS_DIRECT_UPLOAD_EXPIRE_SECONDS=900
# ==================== CosyVoice 语音合成(必须配置)====================
# 注意:base_url 只需写到 /api/v1,具体路径由代码拼接
# 模型: cosyvoice-v3-flash (推荐,支持系统音色,性价比高)
# cosyvoice-v3-plus (高质量,系统音色少)
# cosyvoice-v3.5-flash / cosyvoice-v3.5-plus (仅支持克隆/设计音色,无系统音色)
# 音色: v3系列系统音色带 _v3 后缀,如 longxiaochun_v3, longxiaoxia_v3, longanyang (无后缀)
COSYVOICE_API_KEY=CHANGE_ME_COSYVOICE_API_KEY
COSYVOICE_BASE_URL=https://dashscope.aliyuncs.com/api/v1
COSYVOICE_MODEL=cosyvoice-v3-flash
COSYVOICE_VOICE=longxiaochun_v3
COSYVOICE_BASE_URL=https://dashscope.aliyuncs.com/api/v1/services/aigc/text2audio
COSYVOICE_MODEL=cosyvoice-v1
COSYVOICE_VOICE=longxiaochun
COSYVOICE_SAMPLE_RATE=22050
COSYVOICE_FORMAT=mp3
-1
View File
@@ -2,7 +2,6 @@
max-line-length = 120
exclude =
.git,
.cache,
__pycache__,
.venv,
venv,
+65
View File
@@ -0,0 +1,65 @@
name: Auto Merge PRs
on:
schedule:
- cron: '0 */6 * * *'
workflow_dispatch:
jobs:
auto-merge:
runs-on: saas
timeout-minutes: 10
steps:
- name: Checkout code
shell: sh
env:
GITHUB_TOKEN: ${{ github.token }}
run: |
set -eu
python3 - <<'PY'
import io, os, tarfile, time, urllib.request, urllib.error
url = f"{os.environ['GITHUB_API_URL']}/repos/{os.environ['GITHUB_REPOSITORY']}/{os.environ['GITHUB_SHA']}.tar.gz"
request = urllib.request.Request(url, headers={"Authorization": f"token {os.environ['GITHUB_TOKEN']}"})
last_err = None
for attempt in range(5):
try:
with urllib.request.urlopen(request, timeout=120) as response:
archive = response.read()
break
except urllib.error.HTTPError as e:
last_err = e
if e.code >= 500 and attempt < 4:
wait = 2 ** attempt
print(f"Checkout HTTP {e.code}, retrying in {wait}s (attempt {attempt+1}/5)...")
time.sleep(wait)
continue
raise
except Exception as e:
last_err = e
if attempt < 4:
wait = 2 ** attempt
print(f"Checkout error: {e}, retrying in {wait}s (attempt {attempt+1}/5)...")
time.sleep(wait)
continue
raise
else:
raise last_err
with tarfile.open(fileobj=io.BytesIO(archive), mode='r:gz') as tar:
top_prefix = tar.getmembers()[0].name.split('/', 1)[0] + '/'
for member in tar.getmembers():
name = member.name
if name == top_prefix[:-1]:
continue
if name.startswith(top_prefix):
member.name = name[len(top_prefix):]
if member.name:
tar.extract(member, '.')
PY
- name: Auto merge develop PRs
run: |
bash scripts/auto_merge_prs.sh develop
- name: Auto merge main PRs (release only)
run: |
bash scripts/auto_merge_prs.sh main
+17 -302
View File
@@ -22,7 +22,7 @@ permissions:
jobs:
validate:
name: Validate Code Quality And Tests
runs-on: host
runs-on: ubuntu-22.04
timeout-minutes: 10
env:
@@ -80,7 +80,7 @@ jobs:
shell: sh
run: |
set -eu
python3 --version
python --version
python3 -m pip --version
echo "CI environment is ready"
@@ -158,195 +158,14 @@ jobs:
python3 scripts/check_migration_safety.py --allow-medium-risk
fi
unit-tests:
name: Unit Tests
runs-on: host
timeout-minutes: 8
env:
USE_IN_MEMORY_DB: "true"
steps:
- name: Checkout code
- name: Run unit tests
shell: sh
env:
GITHUB_TOKEN: ${{ github.token }}
USE_IN_MEMORY_DB: "true"
run: |
set -eu
python3 - <<'PY'
import io, os, tarfile, time, urllib.request, urllib.error
url = f"{os.environ['GITHUB_API_URL']}/repos/{os.environ['GITHUB_REPOSITORY']}/archive/{os.environ['GITHUB_SHA']}.tar.gz"
request = urllib.request.Request(url, headers={"Authorization": f"token {os.environ['GITHUB_TOKEN']}"})
last_err = None
for attempt in range(5):
try:
with urllib.request.urlopen(request, timeout=120) as response:
archive = response.read()
break
except urllib.error.HTTPError as e:
last_err = e
if e.code >= 500 and attempt < 4:
wait = 2 ** attempt
print(f"Checkout HTTP {e.code}, retrying in {wait}s (attempt {attempt+1}/5)...")
time.sleep(wait)
continue
raise
except Exception as e:
last_err = e
if attempt < 4:
wait = 2 ** attempt
print(f"Checkout error: {e}, retrying in {wait}s (attempt {attempt+1}/5)...")
time.sleep(wait)
continue
raise
else:
raise last_err
with tarfile.open(fileobj=io.BytesIO(archive), mode='r:gz') as tar:
root_prefix = tar.getmembers()[0].name.split('/', 1)[0] + '/'
for member in tar.getmembers():
name = member.name
if name == root_prefix[:-1]:
continue
if name.startswith(root_prefix):
member.name = name[len(root_prefix):]
if member.name:
tar.extract(member, '.')
PY
- name: Install dependencies
shell: sh
run: |
set -eu
python3 -m pip install -q -r requirements-base.txt
python3 -m pip install -q -r requirements.txt
python3 -m pip install -q -r requirements-dev.txt
pytest --version
- name: Run unit tests with coverage
shell: sh
run: |
set -eu
PYTHONPATH="$PWD/apps/api:$PWD" python3 -m coverage run \
--source=apps/api/app,packages \
--omit="*/migrations/*,*/tests/*,*/test_*.py,*/site-packages/*" \
--branch \
-m pytest tests/unit -q
python3 -m coverage report --show-missing
python3 -m coverage xml -o coverage.xml
python3 -m coverage report --fail-under=60 > /dev/null
- name: CI failure notification
if: failure()
shell: sh
env:
GITEA_TOKEN: ${{ secrets.GITEA_TOKEN }}
CI_WEBHOOK_URL: ${{ secrets.CI_WEBHOOK_URL }}
run: |
set +e
FAILED_JOB="Unit Tests" python3 scripts/ci_notify_failure.py
integration-tests:
name: Integration Tests
runs-on: host
timeout-minutes: 20
if: always()
needs: validate
env:
DATABASE_URL: postgresql+psycopg://postgres:postgres@127.0.0.1:5432/xiaoxia_saas
USE_IN_MEMORY_DB: "false"
steps:
- name: Checkout code
shell: sh
env:
GITHUB_TOKEN: ${{ github.token }}
run: |
set -eu
python3 - <<'PY'
import io, os, tarfile, time, urllib.request, urllib.error
url = f"{os.environ['GITHUB_API_URL']}/repos/{os.environ['GITHUB_REPOSITORY']}/archive/{os.environ['GITHUB_SHA']}.tar.gz"
request = urllib.request.Request(url, headers={"Authorization": f"token {os.environ['GITHUB_TOKEN']}"})
last_err = None
for attempt in range(5):
try:
with urllib.request.urlopen(request, timeout=120) as response:
archive = response.read()
break
except urllib.error.HTTPError as e:
last_err = e
if e.code >= 500 and attempt < 4:
wait = 2 ** attempt
print(f"Checkout HTTP {e.code}, retrying in {wait}s (attempt {attempt+1}/5)...")
time.sleep(wait)
continue
raise
except Exception as e:
last_err = e
if attempt < 4:
wait = 2 ** attempt
print(f"Checkout error: {e}, retrying in {wait}s (attempt {attempt+1}/5)...")
time.sleep(wait)
continue
raise
else:
raise last_err
with tarfile.open(fileobj=io.BytesIO(archive), mode='r:gz') as tar:
root_prefix = tar.getmembers()[0].name.split('/', 1)[0] + '/'
for member in tar.getmembers():
name = member.name
if name == root_prefix[:-1]:
continue
if name.startswith(root_prefix):
member.name = name[len(root_prefix):]
if member.name:
tar.extract(member, '.')
PY
- name: Verify CI environment
shell: sh
run: |
set -eu
python3 --version
python3 -m pip --version
echo "CI environment is ready"
- name: Install dependencies
shell: sh
run: |
set -eu
python3 -m pip install -q -r requirements-base.txt
python3 -m pip install -q -r requirements.txt
python3 -m pip install -q -r requirements-dev.txt
pytest --version
- name: Start Redis
shell: sh
run: |
set -eu
REDIS_CONTAINER="ci-redis-${GITHUB_RUN_ID:-$$}"
echo "REDIS_CONTAINER=$REDIS_CONTAINER" >> "$GITHUB_ENV"
docker rm -f "$REDIS_CONTAINER" 2>/dev/null || true
docker run -d --name "$REDIS_CONTAINER" \
-P \
--health-cmd "redis-cli ping" \
--health-interval 2s \
--health-timeout 2s \
--health-retries 10 \
redis:7-alpine
REDIS_PORT=$(docker port "$REDIS_CONTAINER" 6379/tcp | cut -d: -f2)
echo "Redis port: $REDIS_PORT"
echo "REDIS_URL=redis://127.0.0.1:$REDIS_PORT/0" >> "$GITHUB_ENV"
for i in $(seq 1 15); do
if docker inspect --format='{{.State.Health.Status}}' "$REDIS_CONTAINER" 2>/dev/null | grep -q healthy; then
echo "Redis is ready on port $REDIS_PORT"
break
fi
echo "Waiting for Redis... ($i/15)"
sleep 2
done
docker inspect --format='{{.State.Health.Status}}' "$REDIS_CONTAINER" | grep -q healthy
PYTHONPATH="$PWD/apps/api:$PWD" python3 -m pytest tests/unit -q \
--cov=apps --cov-report=term --cov-report=xml
- name: Start PostgreSQL for integration tests
shell: sh
@@ -391,14 +210,8 @@ jobs:
run: |
set -eu
pip install -q pytest-rerunfailures
PYTHONPATH="$PWD/apps/api:$PWD" python3 -m coverage run --append \
--source=apps/api/app,packages \
--omit="*/migrations/*,*/tests/*,*/test_*.py,*/site-packages/*" \
--branch \
-m pytest tests/integration -q --timeout=60 -x --reruns 2 --reruns-delay 1 -m "not performance"
python3 -m coverage report --show-missing
python3 -m coverage xml -o coverage.xml
python3 -m coverage report --fail-under=40 > /dev/null # 集成测试覆盖率门槛较低,核心目标是功能验证
PYTHONPATH="$PWD/apps/api:$PWD" python3 -m pytest tests/integration -q --timeout=60 -x --reruns 2 --reruns-delay 1 -m "not performance" \
--cov=apps --cov-append --cov-report=term --cov-report=xml --cov-fail-under=50
- name: Run API performance baseline tests
shell: sh
@@ -442,44 +255,25 @@ jobs:
exit 0
- name: Cleanup PostgreSQL & Redis
- name: Cleanup PostgreSQL
if: always()
shell: sh
run: |
docker rm -f "${PG_CONTAINER:-ci-pg-validate}" 2>/dev/null || true
docker rm -f "${REDIS_CONTAINER:-ci-redis-int}" 2>/dev/null || true
echo "PostgreSQL container cleaned up"
echo "Redis container cleaned up"
- name: Coverage summary
if: always()
shell: sh
env:
COVERAGE_THRESHOLD: "40"
run: |
set +e
echo "=== 覆盖率汇总 ==="
python3 scripts/ci_coverage_summary.py
- name: Notify CI failure
if: failure()
- name: Build summary
if: github.ref == 'refs/heads/develop' || github.ref == 'refs/heads/main'
shell: sh
run: |
set +e
echo "=== CI 失败通知 ==="
FAILED_JOB="Validate Code Quality And Tests" python3 scripts/ci_notify_failure.py
- name: Notify CI failure - Integration Tests
if: failure()
shell: sh
run: |
set +e
echo "=== CI 失败通知 ==="
FAILED_JOB="Integration Tests" python3 scripts/ci_notify_failure.py
set -eu
echo "Build completed successfully!"
echo "Branch: ${GITHUB_REF_NAME}"
echo "Commit: ${GITHUB_SHA}"
frontend-lint:
name: Frontend Lint
runs-on: host
runs-on: ubuntu-22.04
timeout-minutes: 10
steps:
@@ -578,22 +372,13 @@ jobs:
-w /workspace/apps/web \
docker.m.daocloud.io/library/node:20 \
sh -lc 'npx vitest run src/test'
- name: Notify CI failure
if: failure()
shell: sh
run: |
set +e
echo "=== CI 失败通知 ==="
FAILED_JOB="Frontend Lint" python3 scripts/ci_notify_failure.py
deploy-staging:
name: Build & Push Staging (Watchtower auto-deploy)
runs-on: saas
timeout-minutes: 30
needs: [validate, frontend-lint]
if: github.event_name == 'push' && (github.ref_name == 'main' || github.ref_name == 'develop')
if: github.ref_name == 'main' || github.ref_name == 'develop' || startsWith(github.ref_name, 'feature/')
steps:
- name: Checkout code
@@ -718,23 +503,6 @@ jobs:
echo "Branch: ${GITHUB_REF_NAME}"
echo "Commit: ${GITHUB_SHA}"
- name: Notify CI success
if: success()
shell: sh
run: |
set +e
echo "=== CI 成功通知 ==="
SUCCESS_JOB="Staging部署成功" python3 scripts/ci_notify_success.py
- name: Notify CI failure
if: failure()
shell: sh
run: |
set +e
echo "=== CI 失败通知 ==="
FAILED_JOB="Build & Push Staging (Watchtower auto-deploy)" python3 scripts/ci_notify_failure.py
staging-e2e:
name: Staging E2E Tests
@@ -803,15 +571,6 @@ jobs:
git.xiaoxiajianji.com/xiaoxia/base/playwright:v1.45.0-jammy \
sh -lc "npm ci && npx playwright test --reporter=line --project=chromium e2e/auth.spec.ts e2e/auth-guard.spec.ts e2e/core-upload.spec.ts e2e/core-generation.spec.ts"
- name: Notify CI failure
if: failure()
shell: sh
run: |
set +e
echo "=== CI 失败通知 ==="
FAILED_JOB="Staging E2E Tests" python3 scripts/ci_notify_failure.py
staging-api-tests:
name: Staging API Integration Tests
runs-on: saas
@@ -877,15 +636,6 @@ jobs:
git.xiaoxiajianji.com/xiaoxia/base/playwright:v1.45.0-jammy \
sh -lc 'npm ci && npx playwright test --reporter=line e2e/test_auth.spec.ts e2e/test_asset.spec.ts e2e/test_project.spec.ts'
- name: Notify CI failure
if: failure()
shell: sh
run: |
set +e
echo "=== CI 失败通知 ==="
FAILED_JOB="Staging API Integration Tests" python3 scripts/ci_notify_failure.py
build-production-runtime-images:
name: Build Production Runtime Images
@@ -966,15 +716,6 @@ jobs:
echo "Disk usage after cleanup:"
df -h / | tail -1
- name: Notify CI failure
if: failure()
shell: sh
run: |
set +e
echo "=== CI 失败通知 ==="
FAILED_JOB="Build Production Runtime Images" python3 scripts/ci_notify_failure.py
deploy-production:
name: Deploy Production
runs-on: saas
@@ -1040,23 +781,6 @@ jobs:
echo "$DEPLOY_B64" | base64 -d | ssh -p 22222 -i "$key_path" "$production_user@$production_host" "IMAGE_TAG='${GITHUB_REF_NAME}' REGISTRY_TOKEN='${REGISTRY_TOKEN}' sh"
- name: Notify CI success
if: success()
shell: sh
run: |
set +e
echo "=== CI 成功通知 ==="
SUCCESS_JOB="生产部署成功" python3 scripts/ci_notify_success.py
- name: Notify CI failure
if: failure()
shell: sh
run: |
set +e
echo "=== CI 失败通知 ==="
FAILED_JOB="Deploy Production" python3 scripts/ci_notify_failure.py
production-e2e:
name: Production Browser E2E
runs-on: saas
@@ -1125,12 +849,3 @@ jobs:
-w /workspace/apps/web \
git.xiaoxiajianji.com/xiaoxia/base/playwright:v1.45.0-jammy \
sh -lc 'npm ci && npx playwright test --reporter=line --project=chromium e2e/auth.spec.ts e2e/auth-guard.spec.ts e2e/core-upload.spec.ts e2e/core-generation.spec.ts e2e/core-titles.spec.ts'
- name: Notify CI failure
if: failure()
shell: sh
run: |
set +e
echo "=== CI 失败通知 ==="
FAILED_JOB="Production Browser E2E" python3 scripts/ci_notify_failure.py
+69
View File
@@ -0,0 +1,69 @@
name: Test SSH Secret
on:
push:
branches: [develop]
paths:
- '.gitea/workflows/test-ssh-secret.yml'
jobs:
test-ssh:
runs-on: ubuntu-22.04
steps:
- name: Install SSH client
run: |
which ssh || (apt-get update && apt-get install -y openssh-client)
ssh -V
- name: Debug environment
run: |
echo "=== Environment ==="
echo "Runner hostname: $(hostname)"
echo "Runner IP: $(hostname -i || echo 'unknown')"
echo "Current user: $(whoami)"
echo "=== Secrets check ==="
if [ -n "$STAGING_SSH_HOST" ]; then
echo "STAGING_SSH_HOST: [SET] value_length=${#STAGING_SSH_HOST}"
else
echo "STAGING_SSH_HOST: [EMPTY]"
fi
if [ -n "$STAGING_SSH_USER" ]; then
echo "STAGING_SSH_USER: [SET] value_length=${#STAGING_SSH_USER}"
else
echo "STAGING_SSH_USER: [EMPTY]"
fi
if [ -n "$STAGING_SSH_KEY" ]; then
echo "STAGING_SSH_KEY: [SET] value_length=${#STAGING_SSH_KEY}"
else
echo "STAGING_SSH_KEY: [EMPTY]"
fi
env:
STAGING_SSH_HOST: ${{ secrets.STAGING_SSH_HOST }}
STAGING_SSH_USER: ${{ secrets.STAGING_SSH_USER }}
STAGING_SSH_KEY: ${{ secrets.STAGING_SSH_KEY }}
- name: Setup SSH key
run: |
mkdir -p ~/.ssh
chmod 700 ~/.ssh
echo "$STAGING_SSH_KEY" > ~/.ssh/id_ed25519
chmod 600 ~/.ssh/id_ed25519
ssh-keygen -y -f ~/.ssh/id_ed25519 > ~/.ssh/id_ed25519.pub 2>/dev/null || echo "No public key generated"
echo "=== SSH Key fingerprint ==="
ssh-keygen -lf ~/.ssh/id_ed25519 || echo "Key fingerprint failed"
env:
STAGING_SSH_KEY: ${{ secrets.STAGING_SSH_KEY }}
- name: Test SSH connection
run: |
echo "Attempting SSH connection to $STAGING_SSH_HOST..."
ssh -i ~/.ssh/id_ed25519 \
-o StrictHostKeyChecking=no \
-o UserKnownHostsFile=/dev/null \
-o ConnectTimeout=10 \
-o BatchMode=yes \
-v \
$STAGING_SSH_USER@$STAGING_SSH_HOST "echo 'SSH_CONNECTION_SUCCESS' && hostname && whoami"
echo "=== SSH Test Complete ==="
env:
STAGING_SSH_HOST: ${{ secrets.STAGING_SSH_HOST }}
STAGING_SSH_USER: ${{ secrets.STAGING_SSH_USER }}
+163
View File
@@ -0,0 +1,163 @@
name: Tests
on:
pull_request:
branches: [ main ]
jobs:
test:
runs-on: runtime-builder
steps:
- name: Checkout code
shell: sh
run: |
set -eu
python - <<'PY'
import io
import os
import tarfile
import time
import urllib.error
import urllib.request
url = f"{os.environ['GITHUB_API_URL']}/repos/{os.environ['GITHUB_REPOSITORY']}/archive/{os.environ['GITHUB_SHA']}.tar.gz"
request = urllib.request.Request(url, headers={"Authorization": f"token {os.environ['GITHUB_TOKEN']}"})
# Retry up to 5 times with backoff for transient 5xx errors
last_err = None
for attempt in range(5):
try:
with urllib.request.urlopen(request, timeout=120) as response:
archive = response.read()
break
except urllib.error.HTTPError as e:
last_err = e
if e.code >= 500 and attempt < 4:
wait = 2 ** attempt
print(f"Checkout HTTP {e.code}, retrying in {wait}s (attempt {attempt+1}/5)...")
time.sleep(wait)
continue
raise
except Exception as e:
last_err = e
if attempt < 4:
wait = 2 ** attempt
print(f"Checkout error: {e}, retrying in {wait}s (attempt {attempt+1}/5)...")
time.sleep(wait)
continue
raise
else:
raise last_err
with tarfile.open(fileobj=io.BytesIO(archive), mode='r:gz') as tar:
root_prefix = tar.getmembers()[0].name.split('/', 1)[0] + '/'
for member in tar.getmembers():
name = member.name
if name == root_prefix[:-1]:
continue
if name.startswith(root_prefix):
member.name = name[len(root_prefix):]
if member.name:
tar.extract(member, '.')
PY
- name: Show Python version
shell: sh
run: |
set -eu
python --version
python -m pip --version
- name: Install dependencies
shell: sh
run: |
set -eu
python -m pip install --upgrade pip -i https://mirrors.aliyun.com/pypi/simple/ --trusted-host mirrors.aliyun.com
python -m pip install -r requirements.txt -r requirements-dev.txt -i https://mirrors.aliyun.com/pypi/simple/ --trusted-host mirrors.aliyun.com
- name: Run unit tests
shell: sh
run: |
set -eu
PYTHONPATH="$PWD/apps/api:$PWD" python -m pytest tests/unit -q
- name: Run integration tests
shell: sh
run: |
set -eu
PYTHONPATH="$PWD/apps/api:$PWD" python -m pytest tests/integration -q --timeout=60 -x
lint:
runs-on: runtime-builder
steps:
- name: Checkout code
shell: sh
run: |
set -eu
python - <<'PY'
import io
import os
import tarfile
import time
import urllib.error
import urllib.request
url = f"{os.environ['GITHUB_API_URL']}/repos/{os.environ['GITHUB_REPOSITORY']}/archive/{os.environ['GITHUB_SHA']}.tar.gz"
request = urllib.request.Request(url, headers={"Authorization": f"token {os.environ['GITHUB_TOKEN']}"})
# Retry up to 5 times with backoff for transient 5xx errors
last_err = None
for attempt in range(5):
try:
with urllib.request.urlopen(request, timeout=120) as response:
archive = response.read()
break
except urllib.error.HTTPError as e:
last_err = e
if e.code >= 500 and attempt < 4:
wait = 2 ** attempt
print(f"Checkout HTTP {e.code}, retrying in {wait}s (attempt {attempt+1}/5)...")
time.sleep(wait)
continue
raise
except Exception as e:
last_err = e
if attempt < 4:
wait = 2 ** attempt
print(f"Checkout error: {e}, retrying in {wait}s (attempt {attempt+1}/5)...")
time.sleep(wait)
continue
raise
else:
raise last_err
with tarfile.open(fileobj=io.BytesIO(archive), mode='r:gz') as tar:
root_prefix = tar.getmembers()[0].name.split('/', 1)[0] + '/'
for member in tar.getmembers():
name = member.name
if name == root_prefix[:-1]:
continue
if name.startswith(root_prefix):
member.name = name[len(root_prefix):]
if member.name:
tar.extract(member, '.')
PY
- name: Install dependencies
shell: sh
run: |
set -eu
python -m pip install --upgrade pip -i https://mirrors.aliyun.com/pypi/simple/ --trusted-host mirrors.aliyun.com
python -m pip install -r requirements.txt -r requirements-dev.txt -i https://mirrors.aliyun.com/pypi/simple/ --trusted-host mirrors.aliyun.com
- name: Run Black (check only)
shell: sh
run: |
set -eu
python -m black --check alembic apps packages tests scripts
- name: Run Flake8
shell: sh
run: |
set -eu
python -m flake8 apps packages tests --count --statistics
-1
View File
@@ -6,7 +6,6 @@ dist/
coverage/
# Python / backend
.cache/
.venv/
venv/
.venv-ci-root/
@@ -1,26 +0,0 @@
"""Add logs field to generation_tasks
Revision ID: 037_generation_logs
Revises: 036_expand_uuid_36
Create Date: 2026-07-10
"""
import sqlalchemy as sa
from alembic import op
revision = "037_generation_logs"
down_revision = "036_expand_uuid_36"
branch_labels = None
depends_on = None
def upgrade() -> None:
op.add_column(
"generation_tasks",
sa.Column("logs", sa.Text(), nullable=False, server_default="[]"),
)
def downgrade() -> None:
op.drop_column("generation_tasks", "logs")
+29 -10
View File
@@ -4,14 +4,17 @@ from app.api.routes.assets import router as assets_router
from app.api.routes.auth import router as auth_router
from app.api.routes.chunked_upload import router as chunked_upload_router
from app.api.routes.classification_jobs import router as classification_jobs_router
from app.api.routes.dashboard import router as dashboard_router
from app.api.routes.duplication import router as duplication_router
from app.api.routes.edit_plans import router as edit_plans_router
from app.api.routes.feature_flags import router as feature_flags_router
from app.api.routes.edit_templates import router as edit_templates_router
from app.api.routes.generated_videos import router as generated_videos_router
from app.api.routes.generation_tasks import router as generation_tasks_router
from app.api.routes.health import router as health_check_router
from app.api.routes.ingest_jobs import router as ingest_jobs_router
from app.api.routes.internal_render import router as internal_render_router
from app.api.routes.jobs import router as jobs_router
from app.api.routes.projects import router as projects_router
from app.api.routes.recipes import router as recipes_router
from app.api.routes.subscription import router as subscription_router
from app.api.routes.tags import router as tags_router
from app.api.routes.task_center import router as task_center_router
@@ -84,6 +87,15 @@ api_router.include_router(
prefix="/generation",
tags=["Generation"],
)
api_router.include_router(
jobs_router,
tags=["Job"],
)
api_router.include_router(
generated_videos_router,
prefix="/generated-videos",
tags=["GeneratedVideo"],
)
api_router.include_router(
titles_router,
prefix="/titles",
@@ -109,11 +121,26 @@ api_router.include_router(
prefix="/subscription",
tags=["Subscription"],
)
api_router.include_router(
recipes_router,
prefix="/recipes",
tags=["Recipe"],
)
api_router.include_router(
templates_router,
prefix="/templates",
tags=["Template"],
)
api_router.include_router(
dashboard_router,
prefix="/dashboard",
tags=["Dashboard"],
)
api_router.include_router(
edit_templates_router,
prefix="/edit-templates",
tags=["EditTemplate"],
)
api_router.include_router(
edit_plans_router,
prefix="/edit-plans",
@@ -124,11 +151,3 @@ api_router.include_router(
prefix="/tts",
tags=["TTS"],
)
api_router.include_router(
feature_flags_router,
tags=["Internal"],
)
api_router.include_router(
internal_render_router,
tags=["Internal"],
)
+3 -10
View File
@@ -9,19 +9,12 @@ from packages.ports.user_repository import UserRepository
def check_project_access(project_id: str, user_id: str, project_repository) -> None:
"""检查用户是否有项目访问权限。
合并自 asset_libraries.py / edit_plans.py 的同名函数。
- 空 project_id 直接放行(兼容 edit_plans 中 project_id 可选的场景)
- 错误信息使用中文,与项目其他路由保持一致
"""
if not project_id or not project_id.strip():
return
"""检查用户是否有项目访问权限。"""
project = project_repository.find_by_id(project_id)
if project is None:
raise HTTPException(status_code=404, detail="项目不存在")
raise HTTPException(status_code=404, detail=f"Project {project_id} not found")
if not project.can_access(user_id):
raise HTTPException(status_code=403, detail="无权访问该项目")
raise HTTPException(status_code=403, detail="Access denied to project")
def get_user_plan(user_id: str, user_repository: UserRepository) -> str:
+10 -3
View File
@@ -22,11 +22,18 @@ from packages.application import (
)
from packages.domain import AssetLibrary, AssetLibraryKind
from ._helpers import check_project_access
router = APIRouter()
def _check_project_access(project_id: str, user_id: str, project_repository) -> None:
"""检查用户是否有项目访问权限"""
project = project_repository.find_by_id(project_id)
if project is None:
raise HTTPException(status_code=404, detail=f"Project {project_id} not found")
if not project.can_access(user_id):
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="Access denied to project")
def _to_asset_library_response(item) -> AssetLibraryResponse:
return AssetLibraryResponse(
id=item.id,
@@ -161,7 +168,7 @@ def delete_asset_library(
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="素材库不存在")
# 权限校验:检查用户是否有项目访问权限
check_project_access(library.project_id, authenticated_user.user.id, project_repository)
_check_project_access(library.project_id, authenticated_user.user.id, project_repository)
# 删除库内所有素材(无 FK 级联,需手动清理)
assets_in_library = asset_repository.find_by_library(library_id)
+2 -2
View File
@@ -206,7 +206,7 @@ async def verify_email_post(
return _verify_email_token(request.token, user_repository)
@router.post("/forgot-password", response_model=MessageResponse, status_code=status.HTTP_202_ACCEPTED)
@router.post("/password/forgot", response_model=MessageResponse, status_code=status.HTTP_202_ACCEPTED)
async def forgot_password(
request: PasswordResetRequestModel,
user_repository: UserRepository = Depends(get_user_repository),
@@ -223,7 +223,7 @@ async def forgot_password(
return MessageResponse(message="如果账户存在,密码重置邮件已发送")
@router.post("/reset-password", response_model=MessageResponse)
@router.post("/password/reset", response_model=MessageResponse)
async def reset_password(
request: ResetPasswordModel,
user_repository: UserRepository = Depends(get_user_repository),
+92
View File
@@ -0,0 +1,92 @@
from typing import Any
from app.auth import AuthenticatedUser, get_current_user
from app.dependencies import (
get_asset_repository,
get_generation_task_repository,
get_project_repository,
get_title_library_repository,
get_voice_library_repository,
)
from app.schemas.dashboard import DashboardOverviewResponse, RecentTaskItem, SubscriptionInfo
from fastapi import APIRouter, Depends
router = APIRouter()
def _status_value(status) -> str:
return status.value if hasattr(status, "value") else str(status)
def _generation_step(status: str) -> str:
if status == "pending":
return "等待 Worker 执行"
if status == "running":
return "正在生成成片"
if status == "completed":
return "生成完成"
if status == "failed":
return "生成失败"
return status
@router.get("/overview", response_model=DashboardOverviewResponse)
def get_dashboard_overview(
authenticated_user: AuthenticatedUser = Depends(get_current_user),
project_repository: Any = Depends(get_project_repository),
asset_repository: Any = Depends(get_asset_repository),
generation_task_repository: Any = Depends(get_generation_task_repository),
title_library_repository: Any = Depends(get_title_library_repository),
voice_library_repository: Any = Depends(get_voice_library_repository),
) -> DashboardOverviewResponse:
"""Dashboard 概览:用户级汇总数据。"""
user_id = authenticated_user.user.id
# 获取用户可访问的所有 project
projects = project_repository.find_accessible_projects(user_id)
project_ids = [p.id for p in projects]
# 素材统计
total_assets = asset_repository.count_by_project_ids(project_ids)
used_storage_bytes = asset_repository.sum_storage_by_project_ids(project_ids)
# 标题库 / 配音库统计
total_titles = title_library_repository.count_by_user(user_id)
total_voices = voice_library_repository.count_by_user(user_id)
# 生成任务统计
total_tasks = generation_task_repository.count_by_user(user_id)
# 最近任务(SQL 层 LIMIT 5)
recent = generation_task_repository.list_recent_by_user(user_id, limit=5)
recent_tasks = []
for task in recent:
s = _status_value(task.status)
recent_tasks.append(
RecentTaskItem(
id=task.id,
task_type="generation",
status=s,
current_step=_generation_step(s),
error_message=task.error_message or "",
updated_at=task.completed_at or task.started_at or task.created_at,
)
)
# 订阅信息
user = authenticated_user.user
subscription = SubscriptionInfo(
plan=getattr(user, "subscription_plan", "free") or "free",
is_active=getattr(user, "subscription_status", "") == "active",
)
return DashboardOverviewResponse(
total_assets=total_assets,
used_storage_bytes=used_storage_bytes,
total_titles=total_titles,
total_voices=total_voices,
total_tasks=total_tasks,
total_products=len(projects),
subscription=subscription,
recent_tasks=recent_tasks,
)
+23 -40
View File
@@ -24,7 +24,6 @@ from typing import Any, List, Optional
from app.auth import AuthenticatedUser, get_current_user
from app.core.celery_app import celery_app
from app.core.task_enqueue import GLOBAL_PENDING_LIMIT, USER_PENDING_LIMIT
from app.dependencies import get_asset_library_repository, get_asset_repository, get_db_session, get_project_repository
from app.schemas.generation_task import GenerationTaskResponse
from app.services import EditPlanService, PlanGeneratorService
@@ -45,8 +44,6 @@ from packages.application.generation_tasks import (
CreateGenerationTaskCommand,
CreateGenerationTaskUseCase,
)
from ._helpers import check_project_access
from packages.domain.config_schemas import normalize_plan_config
from packages.domain.edit_plan import EditPlan, EditPlanStatus
@@ -243,6 +240,17 @@ class GenerateFromTemplateResponse(BaseModel):
# ── Helpers ───────────────────────────────────────────────────────────────────
def _check_project_access(project_id: str, user_id: str, project_repository: Any) -> None:
"""校验用户对项目的访问权限(参照 assets.py 的 can_access 模式)"""
if not project_id or not project_id.strip():
return
project = project_repository.find_by_id(project_id)
if project is None:
raise HTTPException(status_code=404, detail="项目不存在")
if not project.can_access(user_id):
raise HTTPException(status_code=403, detail="无权访问该项目")
def _to_response(p: EditPlan) -> EditPlanResponse:
return EditPlanResponse(
id=p.id,
@@ -296,7 +304,7 @@ def list_plans(
# 项目鉴权:如果指定了 project_id,校验用户是否有权访问
if project_id:
check_project_access(project_id, current_user.user.id, project_repository)
_check_project_access(project_id, current_user.user.id, project_repository)
skip = (page - 1) * page_size
plans = svc.list_plans(
@@ -338,7 +346,7 @@ def get_plan(
)
# 项目鉴权
if plan.project_id:
check_project_access(plan.project_id, current_user.user.id, project_repository)
_check_project_access(plan.project_id, current_user.user.id, project_repository)
return _to_response(plan)
@@ -354,7 +362,7 @@ def create_plan(
project_id = (body.project_id or "").strip()
# 项目鉴权
if project_id:
check_project_access(project_id, current_user.user.id, project_repository)
_check_project_access(project_id, current_user.user.id, project_repository)
svc = EditPlanService(db)
# 标准化 config,填充 cover/title/subtitle/bgm 默认值
normalized_config = normalize_plan_config(body.config)
@@ -396,7 +404,7 @@ def update_plan(
if existing is None:
raise HTTPException(status_code=404, detail=f"剪辑计划不存在: {plan_id}")
if existing.project_id:
check_project_access(existing.project_id, current_user.user.id, project_repository)
_check_project_access(existing.project_id, current_user.user.id, project_repository)
# 基础字段更新
try:
@@ -450,7 +458,7 @@ def delete_plan(
# 项目鉴权
existing = svc.get_plan(plan_id)
if existing and existing.project_id:
check_project_access(existing.project_id, current_user.user.id, project_repository)
_check_project_access(existing.project_id, current_user.user.id, project_repository)
deleted = svc.delete_plan(plan_id)
if not deleted:
raise HTTPException(
@@ -492,7 +500,7 @@ def generate_plan(
if plan_check is None:
raise HTTPException(status_code=404, detail=f"剪辑计划不存在: {plan_id}")
if plan_check.project_id:
check_project_access(plan_check.project_id, current_user.user.id, project_repository)
_check_project_access(plan_check.project_id, current_user.user.id, project_repository)
# ── 自动兜底 1: draft → editing ──────────────────────────────────────
if plan_check.status == EditPlanStatus.DRAFT:
@@ -630,31 +638,6 @@ def generate_plan(
# 创建 GenerationTask
gen_task_repo = SQLAlchemyGenerationTaskRepository(db)
# 队列限流预检查(repository 不支持计数时跳过)
user_id = current_user.user.id
try:
has_count = hasattr(gen_task_repo, "count_pending_by_user") and hasattr(
gen_task_repo, "count_pending_total"
)
if has_count:
user_pending = gen_task_repo.count_pending_by_user(user_id)
global_pending = gen_task_repo.count_pending_total()
if user_pending >= USER_PENDING_LIMIT:
raise HTTPException(
status_code=429,
detail=f"您的待处理任务过多(当前 {user_pending}/{USER_PENDING_LIMIT}),请等待完成后再提交",
)
if global_pending >= GLOBAL_PENDING_LIMIT:
raise HTTPException(
status_code=503,
detail="系统繁忙,请稍后再试",
)
except HTTPException:
raise
except Exception as e:
logger.warning("[队列限流] 剪辑计划限流检查失败,跳过: %s", e)
gen_task_use_case = CreateGenerationTaskUseCase(gen_task_repo)
plan = svc.get_plan_or_raise(plan_id)
gen_task = gen_task_use_case.execute(
@@ -734,7 +717,7 @@ def get_generation_status(
plan = gen_status["plan"]
# 项目鉴权
if plan.project_id:
check_project_access(plan.project_id, current_user.user.id, project_repository)
_check_project_access(plan.project_id, current_user.user.id, project_repository)
clips = gen_status["clips"]
clip_items = [
@@ -776,7 +759,7 @@ def list_plan_generations(
# 验证计划存在 + 项目鉴权
plan = svc.get_plan_or_raise(plan_id)
if plan.project_id:
check_project_access(plan.project_id, current_user.user.id, project_repository)
_check_project_access(plan.project_id, current_user.user.id, project_repository)
gen_task_repo = SQLAlchemyGenerationTaskRepository(db)
tasks = gen_task_repo.list_by_source_edit_plan(plan_id)
@@ -845,7 +828,7 @@ def ai_recommend_clips(
# 项目鉴权
if plan.project_id:
check_project_access(plan.project_id, current_user.user.id, project_repository)
_check_project_access(plan.project_id, current_user.user.id, project_repository)
# 验证状态:只允许 draft 或 editing
plan_status = plan.status.value if hasattr(plan.status, "value") else plan.status
@@ -975,7 +958,7 @@ def generate_cover(
# 项目鉴权
if plan.project_id:
check_project_access(plan.project_id, current_user.user.id, project_repository)
_check_project_access(plan.project_id, current_user.user.id, project_repository)
# 调用 AI 封面生成服务
from apps.worker.worker_app.tasks.ai_tasks import run_generate_cover
@@ -1095,7 +1078,7 @@ def get_plan_timeline(
plan = svc.get_plan_or_raise(plan_id)
# 项目鉴权
if plan.project_id:
check_project_access(plan.project_id, current_user.user.id, project_repository)
_check_project_access(plan.project_id, current_user.user.id, project_repository)
clips = svc.list_clips(plan_id=plan_id, skip=0, limit=200)
# 按 order 排序
@@ -1156,7 +1139,7 @@ def generate_from_template(
# 项目鉴权
if body.project_id:
check_project_access(body.project_id, current_user.user.id, project_repository)
_check_project_access(body.project_id, current_user.user.id, project_repository)
template_svc = EditTemplateService(db)
+289
View File
@@ -0,0 +1,289 @@
"""模板管理 API — Phase 8 模板编排引擎.
RESTful CRUD for EditTemplate:
- GET /api/v1/edit-templates 列表(分页 + 类型筛选)
- GET /api/v1/edit-templates/{id} 详情
- POST /api/v1/edit-templates 创建(管理员)
- PUT /api/v1/edit-templates/{id} 更新
- DELETE /api/v1/edit-templates/{id} 删除(软删除 → inactive)
业务逻辑委托给 EditTemplateService 服务层。
"""
from __future__ import annotations
import logging
from datetime import datetime
from typing import Any, List, Optional
from app.auth import AuthenticatedUser, get_current_user
from app.dependencies import get_db_session
from app.services import EditTemplateService
from fastapi import APIRouter, Depends, HTTPException, Query, status
from fastapi.responses import Response
from pydantic import BaseModel, Field
from sqlalchemy.orm import Session
from packages.domain.config_schemas import normalize_template_config
from packages.domain.edit_template import EditTemplate, EditTemplateStatus
logger = logging.getLogger(__name__)
router = APIRouter()
# ── Pydantic Schemas ─────────────────────────────────────────────────────────
class EditTemplateCreateRequest(BaseModel):
"""创建模板请求体"""
name: str = Field(..., min_length=1, max_length=200, description="模板名称")
description: str = Field(default="", max_length=2000, description="模板描述")
template_type: str = Field(default="default", max_length=50, description="模板类型")
editing_mode: str = Field(
default="one_take", max_length=20, description="剪辑模式: one_take/pip/voice_over/voice_pip"
)
config: dict[str, Any] = Field(default_factory=dict, description="模板配置 (JSON)")
preview_url: str = Field(default="", max_length=500, description="预览地址")
sort_weight: int = Field(default=0, ge=0, le=9999, description="排序权重")
class EditTemplateUpdateRequest(BaseModel):
"""更新模板请求体"""
name: Optional[str] = Field(default=None, min_length=1, max_length=200, description="模板名称")
description: Optional[str] = Field(default=None, max_length=2000, description="模板描述")
template_type: Optional[str] = Field(default=None, max_length=50, description="模板类型")
editing_mode: Optional[str] = Field(
default=None, max_length=20, description="剪辑模式: one_take/pip/voice_over/voice_pip"
)
config: Optional[dict[str, Any]] = Field(default=None, description="模板配置 (JSON)")
preview_url: Optional[str] = Field(default=None, max_length=500, description="预览地址")
sort_weight: Optional[int] = Field(default=None, ge=0, le=9999, description="排序权重")
status: Optional[str] = Field(default=None, description="状态: active / inactive")
class EditTemplateResponse(BaseModel):
"""模板响应体"""
id: str
name: str
description: str
template_type: str
editing_mode: str
config: dict[str, Any]
preview_url: str
sort_weight: int
status: str
created_at: datetime
updated_at: datetime
model_config = {"from_attributes": True}
class EditTemplateListResponse(BaseModel):
"""模板列表响应体"""
items: List[EditTemplateResponse]
total: int
page: int
page_size: int
# ── Helpers ───────────────────────────────────────────────────────────────────
def _require_admin(current_user: AuthenticatedUser) -> None:
"""校验当前用户是否为管理员,非管理员返回 403"""
if not getattr(current_user.user, "is_admin", False):
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail="仅管理员可执行此操作",
)
def _to_response(t: EditTemplate) -> EditTemplateResponse:
return EditTemplateResponse(
id=t.id,
name=t.name,
description=t.description,
template_type=t.template_type,
editing_mode=t.editing_mode,
config=t.config,
preview_url=t.preview_url,
sort_weight=t.sort_weight,
status=t.status.value if hasattr(t.status, "value") else t.status,
created_at=t.created_at,
updated_at=t.updated_at,
)
# ── Routes ────────────────────────────────────────────────────────────────────
@router.get("", response_model=EditTemplateListResponse)
def list_templates(
page: int = Query(default=1, ge=1, description="页码"),
page_size: int = Query(default=20, ge=1, le=100, description="每页数量"),
template_type: Optional[str] = Query(default=None, description="按类型筛选"),
status_filter: Optional[str] = Query(
default=None,
alias="status",
description="按状态筛选: active / inactive",
),
db: Session = Depends(get_db_session),
current_user: AuthenticatedUser = Depends(get_current_user),
) -> EditTemplateListResponse:
"""获取模板列表(支持分页、按类型/状态筛选)"""
svc = EditTemplateService(db)
# 解析状态筛选
status_enum: Optional[EditTemplateStatus] = None
if status_filter:
try:
status_enum = EditTemplateStatus(status_filter)
except ValueError:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail=f"无效的状态值: {status_filter},可选值: active, inactive",
)
skip = (page - 1) * page_size
templates = svc.list_templates(
template_type=template_type,
status=status_enum,
skip=skip,
limit=page_size,
)
total = svc.count_templates(
template_type=template_type,
status=status_enum,
)
return EditTemplateListResponse(
items=[_to_response(t) for t in templates],
total=total,
page=page,
page_size=page_size,
)
@router.get("/{template_id}", response_model=EditTemplateResponse)
def get_template(
template_id: str,
db: Session = Depends(get_db_session),
current_user: AuthenticatedUser = Depends(get_current_user),
) -> EditTemplateResponse:
"""获取单个模板详情"""
svc = EditTemplateService(db)
try:
template = svc.get_template_or_raise(template_id)
except ValueError as exc:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail=str(exc),
)
return _to_response(template)
@router.post("", response_model=EditTemplateResponse, status_code=status.HTTP_201_CREATED)
def create_template(
body: EditTemplateCreateRequest,
db: Session = Depends(get_db_session),
current_user: AuthenticatedUser = Depends(get_current_user),
) -> EditTemplateResponse:
"""创建模板(管理员)"""
_require_admin(current_user)
svc = EditTemplateService(db)
# 标准化 config,填充 cover/title/subtitle/bgm 默认值
normalized_config = normalize_template_config(body.config)
try:
created = svc.create_template(
name=body.name,
description=body.description,
template_type=body.template_type,
editing_mode=body.editing_mode,
config=normalized_config,
preview_url=body.preview_url,
sort_weight=body.sort_weight,
)
except ValueError as exc:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail=str(exc),
)
logger.info("创建模板: id=%s name=%s by user=%s", created.id, created.name, current_user.user.id)
return _to_response(created)
@router.put("/{template_id}", response_model=EditTemplateResponse)
def update_template(
template_id: str,
body: EditTemplateUpdateRequest,
db: Session = Depends(get_db_session),
current_user: AuthenticatedUser = Depends(get_current_user),
) -> EditTemplateResponse:
"""更新模板"""
_require_admin(current_user)
svc = EditTemplateService(db)
# 解析状态
status_enum: Optional[EditTemplateStatus] = None
if body.status is not None:
try:
status_enum = EditTemplateStatus(body.status)
except ValueError:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail=f"无效的状态值: {body.status},可选值: active, inactive",
)
# 标准化 config(如果提供了)
config_to_update = normalize_template_config(body.config) if body.config is not None else None
try:
result = svc.update_template(
template_id,
name=body.name,
description=body.description,
template_type=body.template_type,
editing_mode=body.editing_mode,
config=config_to_update,
preview_url=body.preview_url,
sort_weight=body.sort_weight,
status=status_enum,
)
except ValueError as exc:
err_msg = str(exc)
if "不存在" in err_msg:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail=err_msg,
)
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail=err_msg,
)
logger.info("更新模板: id=%s by user=%s", template_id, current_user.user.id)
return _to_response(result)
@router.delete("/{template_id}", status_code=status.HTTP_204_NO_CONTENT, response_model=None)
def delete_template(
template_id: str,
db: Session = Depends(get_db_session),
current_user: AuthenticatedUser = Depends(get_current_user),
) -> Response:
"""删除模板(软删除 → 设为 inactive)"""
_require_admin(current_user)
svc = EditTemplateService(db)
try:
svc.deactivate_template(template_id)
except ValueError as exc:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail=str(exc),
)
logger.info("删除模板(软删除): id=%s by user=%s", template_id, current_user.user.id)
return Response(status_code=204)
-195
View File
@@ -1,195 +0,0 @@
"""Feature Flag 内部管理接口。
通过内部 API Key 鉴权,支持查看和修改 Feature Flag 配置。
主要用于灰度发布期间的动态开关控制。
API:
GET /api/v1/internal/feature-flags - 列出所有 flag
GET /api/v1/internal/feature-flags/{name} - 查看单个 flag
PUT /api/v1/internal/feature-flags/{name} - 设置 flag 配置
DELETE /api/v1/internal/feature-flags/{name} - 删除 flag
鉴权:X-API-Key header,走内部 API Key 验证
"""
from __future__ import annotations
import logging
from typing import Optional
from app.api.routes.auth import _verify_internal_api_key
from app.config import settings
from fastapi import APIRouter, Depends, HTTPException, Query, status
from pydantic import BaseModel, Field
from packages.adapters.redis.feature_flag_store import (
FEATURE_FLAG_REDIS_PREFIX,
FeatureFlagConfig,
RedisFeatureFlagStore,
)
logger = logging.getLogger(__name__)
router = APIRouter(prefix="/internal/feature-flags", tags=["Internal"])
# 允许管理的 flag 白名单(防止误操作其他系统 flag)
ALLOWED_FLAGS = {
"render_engine",
}
def _get_feature_flag_store() -> RedisFeatureFlagStore:
"""获取 Feature Flag 存储实例。"""
return RedisFeatureFlagStore(redis_url=settings.REDIS_URL)
class FeatureFlagUpdateRequest(BaseModel):
"""Feature Flag 更新请求体。"""
enabled: bool = Field(..., description="是否启用")
percentage: int = Field(0, ge=0, le=100, description="灰度百分比 (0-100)")
whitelist: list[str] = Field(default_factory=list, description="白名单列表(如 user_id)")
class FeatureFlagResponse(BaseModel):
"""Feature Flag 响应。"""
name: str
enabled: bool
percentage: int
whitelist: list[str]
@classmethod
def from_config(cls, config: FeatureFlagConfig) -> "FeatureFlagResponse":
return cls(
name=config.name,
enabled=config.enabled,
percentage=config.percentage,
whitelist=sorted(config.whitelist),
)
class FeatureFlagCheckResponse(BaseModel):
"""Flag 激活检查响应。"""
name: str
active: bool
identifier: Optional[str] = None
def _validate_flag_name(name: str) -> None:
"""校验 flag 名称是否在允许列表中。"""
if name not in ALLOWED_FLAGS:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail=f"Unsupported flag: {name}. Allowed: {sorted(ALLOWED_FLAGS)}",
)
@router.get("", response_model=list[FeatureFlagResponse])
async def list_feature_flags(
_: bool = Depends(_verify_internal_api_key),
store: RedisFeatureFlagStore = Depends(_get_feature_flag_store),
):
"""列出所有 Feature Flag。"""
try:
flags = store.list_all()
# 同时返回预定义的 flag(即使未设置也显示默认值)
result = []
for name in sorted(ALLOWED_FLAGS):
config = flags.get(name) or FeatureFlagConfig(name=name, enabled=False)
result.append(FeatureFlagResponse.from_config(config))
# 加上已存在但不在白名单中的 flag(只读展示)
for name, config in flags.items():
if name not in ALLOWED_FLAGS:
result.append(FeatureFlagResponse.from_config(config))
return sorted(result, key=lambda x: x.name)
except Exception as exc:
logger.error("Failed to list feature flags: %s", exc)
raise HTTPException(status_code=500, detail=f"Failed to list flags: {exc}")
@router.get("/{name}", response_model=FeatureFlagResponse)
async def get_feature_flag(
name: str,
_: bool = Depends(_verify_internal_api_key),
store: RedisFeatureFlagStore = Depends(_get_feature_flag_store),
):
"""获取单个 Feature Flag 配置。"""
try:
config = store.get(name)
return FeatureFlagResponse.from_config(config)
except Exception as exc:
logger.error("Failed to get feature flag %s: %s", name, exc)
raise HTTPException(status_code=500, detail=f"Failed to get flag: {exc}")
@router.get("/{name}/check", response_model=FeatureFlagCheckResponse)
async def check_feature_flag(
name: str,
identifier: Optional[str] = Query(None, description="标识符,如 user_id"),
_: bool = Depends(_verify_internal_api_key),
store: RedisFeatureFlagStore = Depends(_get_feature_flag_store),
):
"""检查某个标识符是否命中 Feature Flag。"""
try:
active = store.is_active(name, identifier=identifier)
return FeatureFlagCheckResponse(name=name, active=active, identifier=identifier)
except Exception as exc:
logger.error("Failed to check feature flag %s: %s", name, exc)
raise HTTPException(status_code=500, detail=f"Failed to check flag: {exc}")
@router.put("/{name}", response_model=FeatureFlagResponse)
async def update_feature_flag(
name: str,
request: FeatureFlagUpdateRequest,
_: bool = Depends(_verify_internal_api_key),
store: RedisFeatureFlagStore = Depends(_get_feature_flag_store),
):
"""更新 Feature Flag 配置。
只允许修改 ALLOWED_FLAGS 列表中的 flag。
"""
_validate_flag_name(name)
try:
config = FeatureFlagConfig(
name=name,
enabled=request.enabled,
percentage=request.percentage,
whitelist=set(request.whitelist),
)
store.set(config)
logger.info(
"Feature flag updated: name=%s enabled=%s percentage=%d whitelist=%d",
name,
config.enabled,
config.percentage,
len(config.whitelist),
)
return FeatureFlagResponse.from_config(config)
except Exception as exc:
logger.error("Failed to update feature flag %s: %s", name, exc)
raise HTTPException(status_code=500, detail=f"Failed to update flag: {exc}")
@router.delete("/{name}", status_code=status.HTTP_204_NO_CONTENT)
async def delete_feature_flag(
name: str,
_: bool = Depends(_verify_internal_api_key),
store: RedisFeatureFlagStore = Depends(_get_feature_flag_store),
):
"""删除 Feature Flag。
只允许删除 ALLOWED_FLAGS 列表中的 flag。
"""
_validate_flag_name(name)
try:
deleted = store.delete(name)
logger.info("Feature flag deleted: name=%s deleted=%s", name, deleted)
return None
except Exception as exc:
logger.error("Failed to delete feature flag %s: %s", name, exc)
raise HTTPException(status_code=500, detail=f"Failed to delete flag: {exc}")
+123
View File
@@ -0,0 +1,123 @@
from typing import Any
from app.auth import AuthenticatedUser, get_current_user
from app.core.storage import OSSStorageService, get_storage_service
from app.dependencies import get_generated_video_repository, get_project_repository
from app.schemas.generated_video import (
GeneratedVideoDownloadUrlResponse,
GeneratedVideoResponse,
ListGeneratedVideosResponse,
UpdateGeneratedVideoReviewRequest,
)
from fastapi import APIRouter, Depends, HTTPException, Query
from packages.application import (
GetGeneratedVideoDownloadUrlUseCase,
GetGeneratedVideoUseCase,
ListGeneratedVideosUseCase,
)
router = APIRouter()
def _to_generated_video_response(item, download_url: str | None = None) -> GeneratedVideoResponse:
return GeneratedVideoResponse(
id=item.id,
project_id=item.project_id,
generation_task_id=item.generation_task_id,
name=item.name,
file_url=item.file_url,
file_size=item.file_size,
duration=item.duration,
thumbnail_url=item.thumbnail_url,
width=item.width,
height=item.height,
fps=item.fps,
status=item.status,
review_status=item.review_status,
generation_params=item.generation_params,
download_url=download_url,
)
@router.get("", response_model=ListGeneratedVideosResponse)
def list_generated_videos(
project_id: str | None = Query(None),
authenticated_user: AuthenticatedUser = Depends(get_current_user),
generated_video_repository: Any = Depends(get_generated_video_repository),
project_repository: Any = Depends(get_project_repository),
storage_service: OSSStorageService = Depends(get_storage_service),
) -> ListGeneratedVideosResponse:
user_id = authenticated_user.user.id
use_case = ListGeneratedVideosUseCase(generated_video_repository)
if project_id:
# If project_id provided, check access and filter by project
project = project_repository.find_by_id(project_id)
if project is None:
raise HTTPException(status_code=404, detail=f"Project {project_id} not found")
items = use_case.execute(project_id)
else:
# If no project_id, list all videos from accessible projects
accessible_projects = project_repository.find_accessible_projects(user_id)
all_items = []
for proj in accessible_projects:
all_items.extend(use_case.execute(proj.id))
items = all_items
# Generate download URLs for each video
responses = []
for item in items:
download_url = storage_service.get_download_url(item.file_url)
responses.append(_to_generated_video_response(item, download_url=download_url))
return ListGeneratedVideosResponse(items=responses)
@router.get("/{video_id}", response_model=GeneratedVideoResponse)
def get_generated_video(
video_id: str,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
generated_video_repository: Any = Depends(get_generated_video_repository),
storage_service: OSSStorageService = Depends(get_storage_service),
) -> GeneratedVideoResponse:
use_case = GetGeneratedVideoUseCase(generated_video_repository)
item = use_case.execute(video_id)
if item is None:
raise HTTPException(status_code=404, detail=f"GeneratedVideo {video_id} not found")
download_url = storage_service.get_download_url(item.file_url)
return _to_generated_video_response(item, download_url=download_url)
@router.patch("/{video_id}/review", response_model=GeneratedVideoResponse)
def update_generated_video_review_status(
video_id: str,
request: UpdateGeneratedVideoReviewRequest,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
generated_video_repository: Any = Depends(get_generated_video_repository),
storage_service: OSSStorageService = Depends(get_storage_service),
) -> GeneratedVideoResponse:
video = generated_video_repository.get(video_id)
if video is None:
raise HTTPException(status_code=404, detail=f"GeneratedVideo {video_id} not found")
video.review_status = request.review_status
updated = generated_video_repository.update(video)
download_url = storage_service.get_download_url(updated.file_url)
return _to_generated_video_response(updated, download_url=download_url)
@router.get("/{video_id}/download-url", response_model=GeneratedVideoDownloadUrlResponse)
def get_generated_video_download_url(
video_id: str,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
generated_video_repository: Any = Depends(get_generated_video_repository),
storage_service: OSSStorageService = Depends(get_storage_service),
) -> GeneratedVideoDownloadUrlResponse:
video = generated_video_repository.get(video_id)
if video is None:
raise HTTPException(status_code=404, detail=f"GeneratedVideo {video_id} not found")
use_case = GetGeneratedVideoDownloadUrlUseCase(generated_video_repository)
file_url = use_case.execute(video_id)
if file_url is None:
raise HTTPException(status_code=404, detail=f"GeneratedVideo {video_id} not found")
download_url = storage_service.get_download_url(file_url)
return GeneratedVideoDownloadUrlResponse(video_id=video_id, download_url=download_url)
+26 -142
View File
@@ -1,18 +1,10 @@
import logging
import random
import uuid
from typing import Any
from app.auth import AuthenticatedUser, get_current_user
from app.core.celery_app import celery_app
from app.core.storage import OSSStorageService, get_storage_service
from app.core.task_enqueue import (
GLOBAL_PENDING_LIMIT,
USER_PENDING_LIMIT,
GlobalQueueFull,
UserPendingLimitExceeded,
check_queue_limits,
safe_enqueue_generation_task,
)
from app.dependencies import (
get_asset_library_repository,
get_asset_repository,
@@ -41,10 +33,9 @@ from packages.application import (
ListGeneratedVideosByTaskUseCase,
)
logger = logging.getLogger(__name__)
router = APIRouter()
def _to_generation_task_response(task) -> GenerationTaskResponse:
return GenerationTaskResponse(
id=task.id,
@@ -59,7 +50,6 @@ def _to_generation_task_response(task) -> GenerationTaskResponse:
source_edit_plan_id=task.source_edit_plan_id or "",
asset_select_mode=getattr(task, "asset_select_mode", ""),
batch_id=getattr(task, "batch_id", ""),
logs=getattr(task, "logs", "[]"),
status=task.status,
progress=task.progress,
result_count=task.result_count,
@@ -184,37 +174,19 @@ def create_generation_task(
asset_library_repository: Any = Depends(get_asset_library_repository),
asset_repository: Any = Depends(get_asset_repository),
) -> BatchGenerationTaskResponse:
logger.info(
"[生成任务] 接收请求: user_id=%s, template_id=%s, asset_count=%d, mode=%s, count=%d",
authenticated_user.user.id,
request.template_id,
len(request.asset_ids),
request.asset_select_mode,
request.count,
project_id, asset_library_id = _resolve_project_and_library(
request, project_repository, asset_library_repository, asset_repository, authenticated_user
)
try:
project_id, asset_library_id = _resolve_project_and_library(
request, project_repository, asset_library_repository, asset_repository, authenticated_user
)
except HTTPException as e:
logger.warning("[生成任务] 校验失败: %s", e.detail)
raise
# asset_library 存在性校验(仅在提供了 asset_library_id 时)
resolved_asset_ids: list[str] = list(request.asset_ids)
if asset_library_id:
library = asset_library_repository.get(asset_library_id)
if library is None or (project_id and library.project_id != project_id):
logger.warning("[生成任务] 素材库不存在: library_id=%s", asset_library_id)
raise HTTPException(status_code=404, detail=f"AssetLibrary {asset_library_id} not found")
assets = asset_repository.find_by_library(asset_library_id)
try:
_ensure_library_has_ready_video_assets(assets)
except HTTPException as e:
logger.warning("[生成任务] 素材校验失败: %s", e.detail)
raise
_ensure_library_has_ready_video_assets(assets)
# 素材库自动匹配:当未显式指定 asset_ids 时,按模式自动选取
if not resolved_asset_ids:
@@ -227,85 +199,30 @@ def create_generation_task(
use_case = CreateGenerationTaskUseCase(generation_task_repository)
count = request.count
created_tasks = []
failed_tasks = []
user_id = authenticated_user.user.id
# 同批次任务共享 batch_id,用于视频查重时批次内比对
batch_id = uuid.uuid4().hex if count > 1 else ""
# 预检查:批量提交前先看会不会超限,避免建一半才拒
try:
user_pending = generation_task_repository.count_pending_by_user(user_id)
global_pending = generation_task_repository.count_pending_total()
if user_pending + count > USER_PENDING_LIMIT:
raise UserPendingLimitExceeded(
user_id=user_id, pending_count=user_pending + count, limit=USER_PENDING_LIMIT
for _ in range(count):
task = use_case.execute(
CreateGenerationTaskCommand(
project_id=project_id,
asset_library_id=asset_library_id,
strategy_id=request.strategy_id,
voice_library_id=request.voice_library_id,
template_id=request.template_id,
asset_ids=resolved_asset_ids,
title_ids=request.title_ids,
voice_ids=request.voice_ids,
created_by_user_id=authenticated_user.user.id,
source_edit_plan_id=request.source_edit_plan_id,
asset_select_mode=request.asset_select_mode,
batch_id=batch_id,
)
if global_pending + count > GLOBAL_PENDING_LIMIT:
raise GlobalQueueFull(pending_count=global_pending + count, limit=GLOBAL_PENDING_LIMIT)
except UserPendingLimitExceeded as e:
raise HTTPException(
status_code=429,
detail=f"您的待处理任务过多(当前 {e.pending_count - count}/{e.limit},本次提交 {count} 个),请等待完成后再提交",
) from e
except GlobalQueueFull as e:
raise HTTPException(
status_code=503,
detail="系统繁忙,请稍后再试",
) from e
)
celery_app.send_task("worker.generate_video", args=[task.id])
created_tasks.append(task)
try:
for _ in range(count):
task = use_case.execute(
CreateGenerationTaskCommand(
project_id=project_id,
asset_library_id=asset_library_id,
strategy_id=request.strategy_id,
voice_library_id=request.voice_library_id,
template_id=request.template_id,
asset_ids=resolved_asset_ids,
title_ids=request.title_ids,
voice_ids=request.voice_ids,
created_by_user_id=user_id,
source_edit_plan_id=request.source_edit_plan_id,
asset_select_mode=request.asset_select_mode,
batch_id=batch_id,
)
)
try:
if safe_enqueue_generation_task(
task,
generation_task_repository,
user_id=user_id,
log_prefix="[生成任务]",
log_task_status=True,
):
created_tasks.append(task)
else:
failed_tasks.append(task)
except UserPendingLimitExceeded:
# 兜底:如果预检查后又并发提交了,在这里也拦住
failed_tasks.append(task)
if not created_tasks:
raise HTTPException(
status_code=429,
detail="您的待处理任务过多,请等待完成后再提交",
)
break
except GlobalQueueFull:
failed_tasks.append(task)
if not created_tasks:
raise HTTPException(
status_code=503,
detail="系统繁忙,请稍后再试",
)
break
except HTTPException:
raise
except Exception as e:
logger.error("[生成任务] 创建失败: %s", e, exc_info=True)
raise HTTPException(status_code=500, detail="创建生成任务失败,请稍后重试或查看任务日志")
items = [_to_generation_task_response(t) for t in created_tasks + failed_tasks]
items = [_to_generation_task_response(t) for t in created_tasks]
return BatchGenerationTaskResponse(items=items, total=len(items))
@@ -375,21 +292,6 @@ def retry_generation_task(
if status_val != "failed":
raise HTTPException(status_code=409, detail="Only failed tasks can be retried")
user_id = authenticated_user.user.id
# 预检查:创建前判断,>= 上限就拒绝
user_pending = generation_task_repository.count_pending_by_user(user_id)
global_pending = generation_task_repository.count_pending_total()
if user_pending >= USER_PENDING_LIMIT:
raise HTTPException(
status_code=429,
detail=f"您的待处理任务过多(当前 {user_pending}/{USER_PENDING_LIMIT}),请等待完成后再提交",
)
if global_pending >= GLOBAL_PENDING_LIMIT:
raise HTTPException(
status_code=503,
detail="系统繁忙,请稍后再试",
)
use_case = CreateGenerationTaskUseCase(generation_task_repository)
retried = use_case.execute(
CreateGenerationTaskCommand(
@@ -401,28 +303,10 @@ def retry_generation_task(
asset_ids=task.asset_ids,
title_ids=task.title_ids,
voice_ids=task.voice_ids,
created_by_user_id=user_id,
created_by_user_id=authenticated_user.user.id,
source_edit_plan_id=task.source_edit_plan_id or "",
asset_select_mode=getattr(task, "asset_select_mode", ""),
)
)
try:
if not safe_enqueue_generation_task(
retried,
generation_task_repository,
user_id=user_id,
log_prefix="[生成任务]",
log_task_status=True,
):
logger.warning("[生成任务] 重试入队失败: task_id=%s", retried.id)
except UserPendingLimitExceeded:
raise HTTPException(
status_code=429,
detail="您的待处理任务过多,请等待完成后再提交",
) from None
except GlobalQueueFull:
raise HTTPException(
status_code=503,
detail="系统繁忙,请稍后再试",
) from None
celery_app.send_task("worker.generate_video", args=[retried.id])
return _to_generation_task_response(retried)
-120
View File
@@ -1,120 +0,0 @@
"""渲染结果内部下载接口。
通过内部 API Key 鉴权,为灰度对比工具等内部系统提供渲染结果下载能力。
API:
GET /api/v1/internal/render/videos/{video_id}/download-url - 获取单个视频下载URL
GET /api/v1/internal/render/tasks/{task_id}/videos - 获取任务下所有视频及下载URL
鉴权:X-API-Key header,走内部 API Key 验证
"""
from __future__ import annotations
import logging
from typing import Any
from app.api.routes.auth import _verify_internal_api_key
from app.core.storage import OSSStorageService, get_storage_service
from app.dependencies import get_generated_video_repository
from fastapi import APIRouter, Depends, HTTPException, Query
from pydantic import BaseModel
logger = logging.getLogger(__name__)
router = APIRouter(prefix="/internal/render", tags=["Internal"])
class InternalRenderVideoItem(BaseModel):
"""内部渲染视频项。"""
video_id: str
generation_task_id: str
project_id: str
name: str
file_url: str
file_size: int | None = None
duration: float | None = None
width: int | None = None
height: int | None = None
fps: float | None = None
status: str
download_url: str
class InternalRenderTaskVideosResponse(BaseModel):
"""任务下所有渲染视频响应。"""
task_id: str
count: int
videos: list[InternalRenderVideoItem]
class InternalRenderDownloadUrlResponse(BaseModel):
"""单个视频下载URL响应。"""
video_id: str
download_url: str
def _video_to_item(video: Any, download_url: str) -> InternalRenderVideoItem:
"""将 GeneratedVideo 领域对象转为响应项。"""
return InternalRenderVideoItem(
video_id=video.id,
generation_task_id=video.generation_task_id,
project_id=video.project_id,
name=video.name,
file_url=video.file_url,
file_size=getattr(video, "file_size", None),
duration=getattr(video, "duration", None),
width=getattr(video, "width", None),
height=getattr(video, "height", None),
fps=getattr(video, "fps", None),
status=video.status,
download_url=download_url,
)
@router.get("/videos/{video_id}/download-url", response_model=InternalRenderDownloadUrlResponse)
def get_render_video_download_url(
video_id: str,
_: bool = Depends(_verify_internal_api_key),
generated_video_repository: Any = Depends(get_generated_video_repository),
storage_service: OSSStorageService = Depends(get_storage_service),
) -> InternalRenderDownloadUrlResponse:
"""获取单个渲染视频的下载URL(预签名)。"""
video = generated_video_repository.get(video_id)
if video is None:
raise HTTPException(status_code=404, detail=f"GeneratedVideo {video_id} not found")
download_url = storage_service.get_download_url(video.file_url, expires_seconds=86400)
logger.info("内部渲染下载URL生成: video_id=%s", video_id)
return InternalRenderDownloadUrlResponse(video_id=video_id, download_url=download_url)
@router.get("/tasks/{task_id}/videos", response_model=InternalRenderTaskVideosResponse)
def get_render_task_videos(
task_id: str,
status: str | None = Query(None, description="按状态筛选,如 completed/failed"),
_: bool = Depends(_verify_internal_api_key),
generated_video_repository: Any = Depends(get_generated_video_repository),
storage_service: OSSStorageService = Depends(get_storage_service),
) -> InternalRenderTaskVideosResponse:
"""获取生成任务下所有渲染视频及下载URL。"""
videos = generated_video_repository.list_by_generation_task(task_id)
# 状态筛选
if status:
videos = [v for v in videos if v.status == status]
items = []
for video in videos:
download_url = storage_service.get_download_url(video.file_url, expires_seconds=86400)
items.append(_video_to_item(video, download_url))
logger.info("内部渲染任务视频查询: task_id=%s count=%d", task_id, len(items))
return InternalRenderTaskVideosResponse(
task_id=task_id,
count=len(items),
videos=items,
)
+325
View File
@@ -0,0 +1,325 @@
"""Job API 路由 — Phase 8 任务 2.10.
提供统一异步任务管理 RESTful 接口:
- POST /api/v1/jobs 创建任务
- GET /api/v1/jobs/{job_id} 任务详情
- GET /api/v1/projects/{project_id}/jobs 项目任务列表
- GET /api/v1/projects/{project_id}/jobs/stats 任务统计
- PUT /api/v1/jobs/{job_id}/progress 更新进度
- POST /api/v1/jobs/{job_id}/complete 标记完成
- POST /api/v1/jobs/{job_id}/fail 标记失败
- POST /api/v1/jobs/{job_id}/retry 重试任务
- POST /api/v1/jobs/{job_id}/cancel 取消任务
- POST /api/v1/jobs/{job_id}/submit 提交执行
"""
from __future__ import annotations
import logging
from typing import Any
from app.auth import AuthenticatedUser, get_current_user
from app.core.celery_app import celery_app
from app.dependencies import get_job_repository, get_project_repository
from app.schemas.job import (
CompleteJobRequest,
CreateJobRequest,
FailJobRequest,
JobResponse,
JobStatisticsResponse,
ListJobsResponse,
UpdateProgressRequest,
job_to_response,
)
from fastapi import APIRouter, Depends, HTTPException, Query, status
from packages.application.jobs import (
CancelJobUseCase,
CompleteJobCommand,
CompleteJobUseCase,
CreateJobCommand,
CreateJobUseCase,
FailJobCommand,
FailJobUseCase,
GetJobStatisticsUseCase,
GetJobUseCase,
ListJobsUseCase,
RetryJobUseCase,
SubmitJobUseCase,
UpdateJobProgressCommand,
UpdateJobProgressUseCase,
)
from packages.domain.job import JobType
from app.api.routes._helpers import check_project_access
logger = logging.getLogger(__name__)
router = APIRouter()
# 任务类型 → Celery task name 映射
_JOB_TYPE_TO_CELERY_TASK: dict[str, str] = {
JobType.VIDEO_COMPOSE: "worker.compose_video",
JobType.RENDER_EDIT_PLAN: "worker.render_edit_plan",
JobType.ASSET_INGEST: "worker.ingest_asset",
JobType.CLASSIFICATION: "worker.classify_asset",
JobType.VOICE_EXTRACTION: "worker.extract_voice",
JobType.GENERATION: "worker.generate_video",
}
# ── 创建任务 ──────────────────────────────────────────────────────────────────
@router.post("/jobs", response_model=JobResponse, status_code=status.HTTP_201_CREATED)
def create_job(
request: CreateJobRequest,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
job_repo: Any = Depends(get_job_repository),
project_repository: Any = Depends(get_project_repository),
) -> JobResponse:
"""创建异步任务。
创建后任务处于 pending 状态,需要调用 /submit 提交执行。
"""
check_project_access(request.project_id, authenticated_user.user.id, project_repository)
# 校验 job_type
try:
JobType(request.job_type)
except ValueError:
raise HTTPException(
status_code=400,
detail=f"不支持的任务类型: {request.job_type}," f"可选值: {[t.value for t in JobType]}",
)
use_case = CreateJobUseCase(job_repo)
job = use_case.execute(
CreateJobCommand(
project_id=request.project_id,
job_type=request.job_type,
payload=request.payload,
source_id=request.source_id,
created_by_user_id=authenticated_user.user.id,
max_retries=request.max_retries,
)
)
return job_to_response(job)
# ── 提交执行 ──────────────────────────────────────────────────────────────────
@router.post("/jobs/{job_id}/submit", response_model=JobResponse)
def submit_job(
job_id: str,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
job_repo: Any = Depends(get_job_repository),
) -> JobResponse:
"""提交任务执行。
将任务状态从 pending 切换为 running,并 dispatch Celery 异步任务。
"""
# 权限检查:先获取任务并验证权限,再执行状态变更
job = job_repo.get(job_id)
if job is None:
raise HTTPException(status_code=404, detail=f"Job {job_id} not found")
if job.created_by_user_id and job.created_by_user_id != authenticated_user.user.id:
raise HTTPException(status_code=403, detail="Access denied to this job")
use_case = SubmitJobUseCase(job_repo)
try:
job = use_case.execute(job_id)
except ValueError as e:
raise HTTPException(status_code=400, detail=str(e))
# Dispatch Celery 任务
celery_task_name = _JOB_TYPE_TO_CELERY_TASK.get(job.job_type.value)
if celery_task_name:
result = celery_app.send_task(celery_task_name, args=[job.id], kwargs=job.payload)
job.celery_task_id = result.id
job_repo.update(job)
logger.info("已提交 Celery 任务: job_id=%s celery_task_id=%s", job.id, result.id)
return job_to_response(job)
# ── 查询接口 ──────────────────────────────────────────────────────────────────
@router.get("/jobs/{job_id}", response_model=JobResponse)
def get_job(
job_id: str,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
job_repo: Any = Depends(get_job_repository),
) -> JobResponse:
"""获取任务详情。"""
use_case = GetJobUseCase(job_repo)
job = use_case.execute(job_id)
if job is None:
raise HTTPException(status_code=404, detail=f"Job {job_id} not found")
return job_to_response(job)
@router.get("/projects/{project_id}/jobs", response_model=ListJobsResponse)
def list_project_jobs(
project_id: str,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
job_repo: Any = Depends(get_job_repository),
project_repository: Any = Depends(get_project_repository),
job_type: str | None = Query(default=None, description="按任务类型过滤"),
status_filter: str | None = Query(default=None, alias="status", description="按状态过滤"),
limit: int = Query(default=50, ge=1, le=200),
offset: int = Query(default=0, ge=0),
) -> ListJobsResponse:
"""获取项目下的任务列表。"""
check_project_access(project_id, authenticated_user.user.id, project_repository)
use_case = ListJobsUseCase(job_repo)
jobs = use_case.execute(
project_id=project_id,
job_type=job_type,
status=status_filter,
limit=limit,
offset=offset,
)
items = [job_to_response(j) for j in jobs]
return ListJobsResponse(items=items, total=len(items))
@router.get("/projects/{project_id}/jobs/stats", response_model=JobStatisticsResponse)
def get_job_statistics(
project_id: str,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
job_repo: Any = Depends(get_job_repository),
project_repository: Any = Depends(get_project_repository),
) -> JobStatisticsResponse:
"""获取项目任务统计摘要。"""
check_project_access(project_id, authenticated_user.user.id, project_repository)
use_case = GetJobStatisticsUseCase(job_repo)
stats = use_case.execute(project_id)
return JobStatisticsResponse(**stats)
# ── 进度更新 ──────────────────────────────────────────────────────────────────
@router.put("/jobs/{job_id}/progress", response_model=JobResponse)
def update_job_progress(
job_id: str,
request: UpdateProgressRequest,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
job_repo: Any = Depends(get_job_repository),
) -> JobResponse:
"""更新任务进度。"""
use_case = UpdateJobProgressUseCase(job_repo)
try:
job = use_case.execute(
UpdateJobProgressCommand(
job_id=job_id,
progress=request.progress,
current_stage=request.current_stage,
)
)
except ValueError as e:
raise HTTPException(status_code=400, detail=str(e))
return job_to_response(job)
# ── 完成 / 失败 ────────────────────────────────────────────────────────────────
@router.post("/jobs/{job_id}/complete", response_model=JobResponse)
def complete_job(
job_id: str,
request: CompleteJobRequest,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
job_repo: Any = Depends(get_job_repository),
) -> JobResponse:
"""标记任务完成。"""
use_case = CompleteJobUseCase(job_repo)
try:
job = use_case.execute(CompleteJobCommand(job_id=job_id, result=request.result))
except ValueError as e:
raise HTTPException(status_code=400, detail=str(e))
return job_to_response(job)
@router.post("/jobs/{job_id}/fail", response_model=JobResponse)
def fail_job(
job_id: str,
request: FailJobRequest,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
job_repo: Any = Depends(get_job_repository),
) -> JobResponse:
"""标记任务失败。"""
use_case = FailJobUseCase(job_repo)
try:
job = use_case.execute(FailJobCommand(job_id=job_id, error_message=request.error_message))
except ValueError as e:
raise HTTPException(status_code=400, detail=str(e))
return job_to_response(job)
# ── 重试 / 取消 ────────────────────────────────────────────────────────────────
@router.post("/jobs/{job_id}/retry", response_model=JobResponse)
def retry_job(
job_id: str,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
job_repo: Any = Depends(get_job_repository),
) -> JobResponse:
"""重试失败任务。
将任务重置为 pending,retry_count + 1,但不自动 dispatch。
需要再次调用 /submit 提交执行。
"""
# 权限检查:先获取任务并验证权限,再执行状态变更
job = job_repo.get(job_id)
if job is None:
raise HTTPException(status_code=404, detail=f"Job {job_id} not found")
if job.created_by_user_id and job.created_by_user_id != authenticated_user.user.id:
raise HTTPException(status_code=403, detail="Access denied to this job")
use_case = RetryJobUseCase(job_repo)
try:
job = use_case.execute(job_id)
except ValueError as e:
raise HTTPException(status_code=400, detail=str(e))
return job_to_response(job)
@router.post("/jobs/{job_id}/cancel", response_model=JobResponse)
def cancel_job(
job_id: str,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
job_repo: Any = Depends(get_job_repository),
) -> JobResponse:
"""取消任务。"""
# 权限检查:先获取任务并验证权限,再执行状态变更
job = job_repo.get(job_id)
if job is None:
raise HTTPException(status_code=404, detail=f"Job {job_id} not found")
if job.created_by_user_id and job.created_by_user_id != authenticated_user.user.id:
raise HTTPException(status_code=403, detail="Access denied to this job")
use_case = CancelJobUseCase(job_repo)
try:
job = use_case.execute(job_id)
except ValueError as e:
raise HTTPException(status_code=400, detail=str(e))
return job_to_response(job)
+207
View File
@@ -0,0 +1,207 @@
"""Recipe CRUD + use routes."""
from __future__ import annotations
from app.auth import AuthenticatedUser, get_current_user
from app.dependencies import get_db_session, get_user_repository
from app.schemas.recipe import (
CreateRecipeRequest,
ListRecipesResponse,
RecipeItemResponse,
RecipeResponse,
UpdateRecipeRequest,
UseRecipeResponse,
)
from fastapi import APIRouter, Depends, HTTPException, Query, Response, status
from sqlalchemy.orm import Session
from packages.adapters.sqlalchemy_impl.recipe_repository import SQLAlchemyRecipeRepository
from packages.application.recipe.commands import (
CreateRecipeCommand,
RecipeItemCommand,
UpdateRecipeCommand,
)
from packages.application.recipe.use_cases import (
CreateRecipeUseCase,
DeleteRecipeUseCase,
FeatureDisabledError,
GetRecipeUseCase,
ListRecipesUseCase,
NotFoundError,
UpdateRecipeUseCase,
UseRecipeUseCase,
)
from packages.ports.user_repository import UserRepository
from app.api.routes._helpers import get_user_plan
router = APIRouter()
def _get_recipe_repository(session: Session = Depends(get_db_session)) -> SQLAlchemyRecipeRepository:
return SQLAlchemyRecipeRepository(session)
def _item_to_response(item) -> RecipeItemResponse:
return RecipeItemResponse(
id=item.id,
recipe_id=item.recipe_id,
item_type=item.item_type,
item_id=item.item_id,
position=item.position,
metadata=item.metadata_,
)
def _to_response(recipe) -> RecipeResponse:
return RecipeResponse(
id=recipe.id,
user_id=recipe.user_id,
name=recipe.name,
description=recipe.description,
template_id=recipe.template_id,
generation_params=recipe.generation_params,
items=[_item_to_response(i) for i in getattr(recipe, "items", [])],
is_active=recipe.is_active,
metadata=recipe.metadata_,
created_at=recipe.created_at,
updated_at=recipe.updated_at,
)
@router.get("", response_model=ListRecipesResponse)
def list_recipes(
skip: int = Query(0, ge=0),
limit: int = Query(50, ge=1, le=200),
authenticated_user: AuthenticatedUser = Depends(get_current_user),
recipe_repository: SQLAlchemyRecipeRepository = Depends(_get_recipe_repository),
) -> ListRecipesResponse:
user_id = authenticated_user.user.id
use_case = ListRecipesUseCase(recipe_repository)
recipes = use_case.execute(user_id, skip=skip, limit=limit)
total = recipe_repository.count_by_user(user_id)
return ListRecipesResponse(
items=[_to_response(r) for r in recipes],
total=total,
)
@router.get("/{recipe_id}", response_model=RecipeResponse)
def get_recipe(
recipe_id: str,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
recipe_repository: SQLAlchemyRecipeRepository = Depends(_get_recipe_repository),
) -> RecipeResponse:
user_id = authenticated_user.user.id
use_case = GetRecipeUseCase(recipe_repository)
recipe = use_case.execute(recipe_id, user_id)
if recipe is None:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Recipe not found")
return _to_response(recipe)
@router.post("", response_model=RecipeResponse, status_code=status.HTTP_201_CREATED)
def create_recipe(
request: CreateRecipeRequest,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
recipe_repository: SQLAlchemyRecipeRepository = Depends(_get_recipe_repository),
) -> RecipeResponse:
user_id = authenticated_user.user.id
command = CreateRecipeCommand(
user_id=user_id,
name=request.name,
description=request.description,
template_id=request.template_id,
generation_params=request.generation_params,
items=[
RecipeItemCommand(
item_type=ic.item_type,
item_id=ic.item_id,
position=ic.position,
metadata_=ic.metadata_,
)
for ic in request.items
],
metadata_=request.metadata_,
)
use_case = CreateRecipeUseCase(recipe_repository)
recipe = use_case.execute(command)
return _to_response(recipe)
@router.patch("/{recipe_id}", response_model=RecipeResponse)
def update_recipe(
recipe_id: str,
request: UpdateRecipeRequest,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
recipe_repository: SQLAlchemyRecipeRepository = Depends(_get_recipe_repository),
) -> RecipeResponse:
user_id = authenticated_user.user.id
command = UpdateRecipeCommand(
recipe_id=recipe_id,
user_id=user_id,
name=request.name,
description=request.description,
template_id=request.template_id,
generation_params=request.generation_params,
items=(
[
RecipeItemCommand(
item_type=ic.item_type,
item_id=ic.item_id,
position=ic.position,
metadata_=ic.metadata_,
)
for ic in request.items
]
if request.items is not None
else None
),
metadata_=request.metadata_,
)
use_case = UpdateRecipeUseCase(recipe_repository)
try:
recipe = use_case.execute(command)
except NotFoundError:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Recipe not found")
return _to_response(recipe)
@router.delete("/{recipe_id}", status_code=status.HTTP_204_NO_CONTENT, response_model=None)
def delete_recipe(
recipe_id: str,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
recipe_repository: SQLAlchemyRecipeRepository = Depends(_get_recipe_repository),
) -> Response:
user_id = authenticated_user.user.id
use_case = DeleteRecipeUseCase(recipe_repository)
deleted = use_case.execute(recipe_id, user_id)
if not deleted:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Recipe not found")
return Response(status_code=204)
@router.post("/{recipe_id}/use", response_model=UseRecipeResponse)
def use_recipe(
recipe_id: str,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
recipe_repository: SQLAlchemyRecipeRepository = Depends(_get_recipe_repository),
user_repository: UserRepository = Depends(get_user_repository),
) -> UseRecipeResponse:
user_id = authenticated_user.user.id
plan_name = get_user_plan(user_id, user_repository)
use_case = UseRecipeUseCase(recipe_repository)
try:
result = use_case.execute(recipe_id, user_id, user_plan=plan_name)
except FeatureDisabledError as exc:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail=str(exc),
)
except NotFoundError:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Recipe not found")
return UseRecipeResponse(
recipe=_to_response(result.recipe),
warnings=[{"item_type": w.item_type, "item_id": w.item_id, "position": w.position} for w in result.warnings],
)
+5 -73
View File
@@ -1,15 +1,7 @@
import logging
from typing import Any
from app.auth import AuthenticatedUser, get_current_user
from app.core.celery_app import celery_app
from app.core.task_enqueue import (
GLOBAL_PENDING_LIMIT,
USER_PENDING_LIMIT,
GlobalQueueFull,
UserPendingLimitExceeded,
safe_enqueue_generation_task,
)
from app.dependencies import (
get_generation_task_repository,
get_ingest_job_repository,
@@ -30,10 +22,9 @@ from packages.application import (
SubmitIngestJobUseCase,
)
logger = logging.getLogger(__name__)
router = APIRouter()
def _humanize_task_error(error_message: str) -> str:
raw = (error_message or "").strip()
if not raw:
@@ -148,21 +139,6 @@ def retry_task_by_id(
if _status_value(task.status) != "failed":
raise HTTPException(status_code=409, detail="Only failed tasks can be retried")
user_id = authenticated_user.user.id
# 预检查
user_pending = generation_task_repository.count_pending_by_user(user_id)
global_pending = generation_task_repository.count_pending_total()
if user_pending >= USER_PENDING_LIMIT:
raise HTTPException(
status_code=429,
detail=f"您的待处理任务过多(当前 {user_pending}/{USER_PENDING_LIMIT}),请等待完成后再提交",
)
if global_pending >= GLOBAL_PENDING_LIMIT:
raise HTTPException(
status_code=503,
detail="系统繁忙,请稍后再试",
)
use_case = CreateGenerationTaskUseCase(generation_task_repository)
retried = use_case.execute(
CreateGenerationTaskCommand(
@@ -174,24 +150,10 @@ def retry_task_by_id(
asset_ids=task.asset_ids,
title_ids=task.title_ids,
voice_ids=task.voice_ids,
created_by_user_id=user_id,
created_by_user_id=authenticated_user.user.id,
)
)
try:
if not safe_enqueue_generation_task(
retried, generation_task_repository, user_id=user_id, log_prefix="[任务中心]"
):
logger.warning("[任务中心] 用户级重试入队失败: task_id=%s", retried.id)
except UserPendingLimitExceeded:
raise HTTPException(
status_code=429,
detail="您的待处理任务过多,请等待完成后再提交",
) from None
except GlobalQueueFull:
raise HTTPException(
status_code=503,
detail="系统繁忙,请稍后再试",
) from None
celery_app.send_task("worker.generate_video", args=[retried.id])
return UserTaskResponse(
id=f"generation:{retried.id}",
task_type="generation",
@@ -259,22 +221,6 @@ def retry_project_task(
raise HTTPException(status_code=404, detail="Generation task not found")
if _status_value(task.status) != "failed":
raise HTTPException(status_code=409, detail="Only failed tasks can be retried")
user_id = authenticated_user.user.id
# 预检查
user_pending = generation_task_repository.count_pending_by_user(user_id)
global_pending = generation_task_repository.count_pending_total()
if user_pending >= USER_PENDING_LIMIT:
raise HTTPException(
status_code=429,
detail=f"您的待处理任务过多(当前 {user_pending}/{USER_PENDING_LIMIT}),请等待完成后再提交",
)
if global_pending >= GLOBAL_PENDING_LIMIT:
raise HTTPException(
status_code=503,
detail="系统繁忙,请稍后再试",
)
use_case = CreateGenerationTaskUseCase(generation_task_repository)
retried = use_case.execute(
CreateGenerationTaskCommand(
@@ -286,24 +232,10 @@ def retry_project_task(
asset_ids=task.asset_ids,
title_ids=task.title_ids,
voice_ids=task.voice_ids,
created_by_user_id=user_id,
created_by_user_id=authenticated_user.user.id,
)
)
try:
if not safe_enqueue_generation_task(
retried, generation_task_repository, user_id=user_id, log_prefix="[任务中心]"
):
logger.warning("[任务中心] 项目级重试入队失败: task_id=%s", retried.id)
except UserPendingLimitExceeded:
raise HTTPException(
status_code=429,
detail="您的待处理任务过多,请等待完成后再提交",
) from None
except GlobalQueueFull:
raise HTTPException(
status_code=503,
detail="系统繁忙,请稍后再试",
) from None
celery_app.send_task("worker.generate_video", args=[retried.id])
return _generation_task_to_project_response(retried)
if task_type == "ingest":
job = ingest_job_repository.get(source_id)
Executable → Regular
+6 -17
View File
@@ -7,7 +7,6 @@ from typing import Optional
from app.auth import AuthenticatedUser, get_current_user
from app.dependencies import (
get_audio_url_signer,
get_cosyvoice_service,
get_db_session,
get_user_repository,
@@ -57,10 +56,7 @@ def _get_repository(session: Session = Depends(get_db_session)) -> SQLAlchemyTTS
return SQLAlchemyTTSJobRepository(session)
def _to_response(job, sign_url=None) -> TTSJobResponse:
output_url = job.output_audio_url
if sign_url and output_url:
output_url = sign_url(output_url)
def _to_response(job) -> TTSJobResponse:
return TTSJobResponse(
id=job.id,
user_id=job.user_id,
@@ -70,7 +66,7 @@ def _to_response(job, sign_url=None) -> TTSJobResponse:
project_id=job.project_id,
voice_clone_profile_id=job.voice_clone_profile_id,
status=job.status,
output_audio_url=output_url,
output_audio_url=job.output_audio_url,
output_audio_key=job.output_audio_key,
duration=job.duration,
file_size=job.file_size,
@@ -180,7 +176,6 @@ def list_tts_jobs(
status_filter: Optional[str] = Query(None, alias="status"),
authenticated_user: AuthenticatedUser = Depends(get_current_user),
repository: SQLAlchemyTTSJobRepository = Depends(_get_repository),
sign_url=Depends(get_audio_url_signer),
) -> ListTTSJobResponse:
"""列出用户的 TTS 合成任务。"""
user_id = authenticated_user.user.id
@@ -188,7 +183,7 @@ def list_tts_jobs(
skip = (page - 1) * page_size
items, total = use_case.execute(user_id, status=status_filter, skip=skip, limit=page_size)
return ListTTSJobResponse(
items=[_to_response(j, sign_url) for j in items],
items=[_to_response(j) for j in items],
total=total,
page=page,
page_size=page_size,
@@ -200,7 +195,6 @@ def get_tts_job(
job_id: str,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
repository: SQLAlchemyTTSJobRepository = Depends(_get_repository),
sign_url=Depends(get_audio_url_signer),
) -> TTSJobResponse:
"""获取 TTS 任务详情。"""
user_id = authenticated_user.user.id
@@ -209,7 +203,7 @@ def get_tts_job(
job = use_case.execute(job_id, user_id)
except TTSJobNotFoundError:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="TTS job not found")
return _to_response(job, sign_url)
return _to_response(job)
@router.get("/jobs/{job_id}/status", response_model=TTSStatusResponse)
@@ -217,7 +211,6 @@ def get_tts_job_status(
job_id: str,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
repository: SQLAlchemyTTSJobRepository = Depends(_get_repository),
sign_url=Depends(get_audio_url_signer),
) -> TTSStatusResponse:
"""查询 TTS 合成状态(用于前端轮询)。"""
user_id = authenticated_user.user.id
@@ -226,13 +219,10 @@ def get_tts_job_status(
job = use_case.execute(job_id, user_id)
except TTSJobNotFoundError:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="TTS job not found")
output_url = job.output_audio_url
if output_url:
output_url = sign_url(output_url)
return TTSStatusResponse(
id=job.id,
status=job.status,
output_audio_url=output_url,
output_audio_url=job.output_audio_url,
error_message=job.error_message,
duration=job.duration,
retry_count=job.retry_count,
@@ -268,7 +258,6 @@ def save_tts_job_to_library(
tts_repository: SQLAlchemyTTSJobRepository = Depends(_get_repository),
voice_library_repository: SQLAlchemyVoiceLibraryRepository = Depends(get_voice_library_repository),
user_repository: UserRepository = Depends(get_user_repository),
sign_url=Depends(get_audio_url_signer),
) -> SaveToLibraryResponse:
"""将已完成的 TTS 合成结果保存到配音库。
@@ -339,7 +328,7 @@ def save_tts_job_to_library(
return SaveToLibraryResponse(
id=item.id,
name=item.name,
audio_url=sign_url(item.audio_url) if item.audio_url else "",
audio_url=item.audio_url,
duration=item.duration,
voice_id=item.voice_id,
voice_name=item.voice_name,
-1
View File
@@ -37,7 +37,6 @@ router = APIRouter()
def _to_response(profile) -> VoiceCloneProfileResponse:
# source_audio_url 是用户传入的原始 URL(可能是外部地址),不做预签名转换
return VoiceCloneProfileResponse(
id=profile.id,
user_id=profile.user_id,
+10 -22
View File
@@ -8,7 +8,7 @@ from __future__ import annotations
from typing import Literal, Optional
from app.auth import AuthenticatedUser, get_current_user
from app.dependencies import get_audio_url_signer, get_db_session, get_user_repository
from app.dependencies import get_db_session, get_user_repository
from app.schemas.voice import (
PresetVoiceItemResponse,
PresetVoiceListResponse,
@@ -52,10 +52,7 @@ def _get_clone_profile_repository(session: Session = Depends(get_db_session)) ->
return SQLAlchemyVoiceCloneProfileRepository(session)
def _to_response(item, sign_url=None) -> VoiceLibraryItemResponse:
audio = item.audio_url
if sign_url and audio:
audio = sign_url(audio)
def _to_response(item) -> VoiceLibraryItemResponse:
return VoiceLibraryItemResponse(
id=item.id,
user_id=item.user_id,
@@ -64,7 +61,7 @@ def _to_response(item, sign_url=None) -> VoiceLibraryItemResponse:
voice_provider=item.voice_provider,
voice_id=item.voice_id,
voice_name=item.voice_name,
audio_url=audio,
audio_url=item.audio_url,
duration=item.duration,
file_size=item.file_size,
status=item.status,
@@ -75,20 +72,16 @@ def _to_response(item, sign_url=None) -> VoiceLibraryItemResponse:
)
def _to_unified_response(item, profile_id_map: dict | None = None, sign_url=None) -> UnifiedVoiceItemResponse:
def _to_unified_response(item, profile_id_map: dict | None = None) -> UnifiedVoiceItemResponse:
"""将数据库音色转换为统一响应格式。
Args:
item: VoiceLibraryItem
profile_id_map: voice_id → profile_id 映射,用于填充 voice_clone_profile_id
sign_url: 音频URL预签名函数
"""
profile_id = None
if profile_id_map and item.voice_id:
profile_id = profile_id_map.get(item.voice_id)
audio = item.audio_url
if sign_url and audio:
audio = sign_url(audio)
return UnifiedVoiceItemResponse(
id=item.id,
type="clone",
@@ -98,7 +91,7 @@ def _to_unified_response(item, profile_id_map: dict | None = None, sign_url=None
language="zh-CN",
voice_id=item.voice_id,
voice_provider=item.voice_provider or "cosyvoice",
audio_url=audio,
audio_url=item.audio_url,
duration=item.duration,
file_size=item.file_size,
status=item.status,
@@ -142,7 +135,6 @@ def list_voices_unified(
authenticated_user: AuthenticatedUser = Depends(get_current_user),
voice_repository: SQLAlchemyVoiceLibraryRepository = Depends(_get_voice_repository),
clone_profile_repository: SQLAlchemyVoiceCloneProfileRepository = Depends(_get_clone_profile_repository),
sign_url=Depends(get_audio_url_signer),
) -> UnifiedVoiceListResponse:
"""获取配音列表(预置音色 + 用户克隆音色)。
@@ -170,7 +162,7 @@ def list_voices_unified(
# 批量查询 voice_id → profile_id 映射,填充 voice_clone_profile_id
voice_ids = [i.voice_id for i in clone_items_raw if i.voice_id]
profile_id_map = clone_profile_repository.find_profile_ids_by_voice_ids(voice_ids) if voice_ids else {}
clone_items = [_to_unified_response(i, profile_id_map, sign_url) for i in clone_items_raw]
clone_items = [_to_unified_response(i, profile_id_map) for i in clone_items_raw]
# 组装结果
if type == "preset":
@@ -227,7 +219,6 @@ def list_voices_legacy(
limit: int = Query(50, ge=1, le=200),
authenticated_user: AuthenticatedUser = Depends(get_current_user),
voice_repository: SQLAlchemyVoiceLibraryRepository = Depends(_get_voice_repository),
sign_url=Depends(get_audio_url_signer),
) -> ListVoiceLibraryResponse:
"""原有配音列表接口(仅返回用户克隆音色)。
@@ -237,7 +228,7 @@ def list_voices_legacy(
use_case = ListVoiceLibraryUseCase(voice_repository)
items, total = use_case.execute(user_id, status=status_filter, skip=skip, limit=limit)
return ListVoiceLibraryResponse(
items=[_to_response(i, sign_url) for i in items],
items=[_to_response(i) for i in items],
total=total,
)
@@ -247,14 +238,13 @@ def get_voice(
voice_id: str,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
voice_repository: SQLAlchemyVoiceLibraryRepository = Depends(_get_voice_repository),
sign_url=Depends(get_audio_url_signer),
) -> VoiceLibraryItemResponse:
user_id = authenticated_user.user.id
use_case = GetVoiceLibraryUseCase(voice_repository)
item = use_case.execute(voice_id, user_id)
if item is None:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Voice not found")
return _to_response(item, sign_url)
return _to_response(item)
@router.post("", response_model=VoiceLibraryItemResponse, status_code=status.HTTP_201_CREATED)
@@ -263,7 +253,6 @@ def create_voice(
authenticated_user: AuthenticatedUser = Depends(get_current_user),
voice_repository: SQLAlchemyVoiceLibraryRepository = Depends(_get_voice_repository),
user_repository: UserRepository = Depends(get_user_repository),
sign_url=Depends(get_audio_url_signer),
) -> VoiceLibraryItemResponse:
user_id = authenticated_user.user.id
plan_name = get_user_plan(user_id, user_repository)
@@ -289,7 +278,7 @@ def create_voice(
status_code=status.HTTP_429_TOO_MANY_REQUESTS,
detail=f"配音库配额已满({exc.used}/{exc.limit}),请升级套餐",
)
return _to_response(item, sign_url)
return _to_response(item)
@router.put("/{voice_id}", response_model=VoiceLibraryItemResponse)
@@ -298,7 +287,6 @@ def update_voice(
request: UpdateVoiceLibraryRequest,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
voice_repository: SQLAlchemyVoiceLibraryRepository = Depends(_get_voice_repository),
sign_url=Depends(get_audio_url_signer),
) -> VoiceLibraryItemResponse:
user_id = authenticated_user.user.id
command = UpdateVoiceLibraryCommand(
@@ -320,7 +308,7 @@ def update_voice(
item = use_case.execute(command)
except NotFoundError:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Voice not found")
return _to_response(item, sign_url)
return _to_response(item)
@router.delete("/{voice_id}", status_code=status.HTTP_204_NO_CONTENT, response_model=None)
-3
View File
@@ -114,9 +114,6 @@ class Settings(BaseSettings):
LOG_LEVEL: str = "INFO"
CORS_ORIGINS_RAW: str = "http://localhost:3000,http://localhost:5173,http://localhost:8000"
# 渲染引擎选择:legacy=旧VideoComposeService,unified=新UnifiedRenderService
RENDER_ENGINE: str = "legacy"
model_config = SettingsConfigDict(
env_file=".env",
env_file_encoding="utf-8",
-221
View File
@@ -1,221 +0,0 @@
import logging
from typing import Any
from app.core.celery_app import celery_app
logger = logging.getLogger(__name__)
# ── 限流阈值常量(全系统统一管理,不要在业务代码里硬编码) ──
USER_PENDING_LIMIT = 3 # 单用户 pending 上限
GLOBAL_PENDING_LIMIT = 20 # 全局 pending 上限
class UserPendingLimitExceeded(Exception):
"""用户 pending 任务数超限,返回 429。"""
def __init__(self, user_id: str, pending_count: int, limit: int):
self.user_id = user_id
self.pending_count = pending_count
self.limit = limit
super().__init__(f"用户 {user_id} pending 任务数 {pending_count} 超过上限 {limit}")
class GlobalQueueFull(Exception):
"""全局限流,返回 503。"""
def __init__(self, pending_count: int, limit: int):
self.pending_count = pending_count
self.limit = limit
super().__init__(f"系统 pending 任务数 {pending_count} 超过上限 {limit}")
def check_queue_limits(
user_id: str,
generation_task_repository: Any,
*,
user_pending_limit: int = USER_PENDING_LIMIT,
global_pending_limit: int = GLOBAL_PENDING_LIMIT,
) -> None:
"""检查队列限流(预检查用,任务创建前调用),超限抛对应异常。
边界语义:>= 上限即拒绝(达到上限就不能再加新任务)。
Args:
user_id: 用户 ID
generation_task_repository: 任务仓储
user_pending_limit: 单用户 pending 上限,默认 USER_PENDING_LIMIT
global_pending_limit: 全局 pending 上限,默认 GLOBAL_PENDING_LIMIT
Raises:
GlobalQueueFull: 全局超限时抛出(优先级更高,先查全局)
UserPendingLimitExceeded: 用户超限时抛出
"""
# 先查全局(系统级保护优先级更高)
global_pending = generation_task_repository.count_pending_total()
if global_pending >= global_pending_limit:
logger.warning(
"[队列限流] 全局 pending 任务数超限: %d/%d, user_id=%s",
global_pending,
global_pending_limit,
user_id,
)
raise GlobalQueueFull(pending_count=global_pending, limit=global_pending_limit)
# 再查用户级
if user_id:
user_pending = generation_task_repository.count_pending_by_user(user_id)
if user_pending >= user_pending_limit:
logger.warning(
"[队列限流] 用户 pending 任务数超限: user_id=%s, count=%d/%d",
user_id,
user_pending,
user_pending_limit,
)
raise UserPendingLimitExceeded(user_id=user_id, pending_count=user_pending, limit=user_pending_limit)
def _mark_task_failed_safely(
task: Any,
generation_task_repository: Any,
log_prefix: str,
reason: str,
) -> None:
"""安全地把任务标记为 failed,更新失败只打日志不崩溃。"""
try:
task.mark_failed(f"任务被限流拒绝: {reason}")
generation_task_repository.update(task)
except Exception as update_err:
logger.error(
"%s 限流后更新状态也失败: task_id=%s error=%s",
log_prefix,
task.id,
update_err,
exc_info=True,
)
def safe_enqueue_generation_task(
task: Any,
generation_task_repository: Any,
*,
user_id: str = "",
log_prefix: str = "[任务队列]",
log_task_status: bool = False,
user_pending_limit: int = USER_PENDING_LIMIT,
global_pending_limit: int = GLOBAL_PENDING_LIMIT,
) -> bool:
"""安全入队:入队前限流检查 → 发送 Celery 任务 → 入队后最终校验兜底。
边界说明:
入队前检查用 > 而非 >=。因为调用此函数时 task 已经是 pending 状态并计入 DB,
pending 总数包含了当前任务本身。pending > limit 等价于"其他任务数 >= limit",
与预检查的 >= 语义一致(都是达到上限就拒绝新任务)。
入队后最终校验:发送 Celery 成功后再查一次 DB 计数,处理并发竞态场景
(两个请求同时通过入队前检查,后到的那个在这里被兜住)。
Args:
task: 生成任务对象,需有 id 属性和 mark_failed 方法(状态已为 pending)
generation_task_repository: 任务仓储,用于更新状态
user_id: 用户 ID,传了才做用户级限流检查
log_prefix: 日志前缀,便于区分调用来源
log_task_status: 成功日志中是否额外打印任务状态
user_pending_limit: 单用户 pending 上限,默认 USER_PENDING_LIMIT
global_pending_limit: 全局 pending 上限,默认 GLOBAL_PENDING_LIMIT
Returns:
True 表示入队成功,False 表示入队失败(已标记为 failed)
Raises:
GlobalQueueFull: 全局 pending 超限时抛出,任务会被标记为 failed
UserPendingLimitExceeded: 用户 pending 超限时抛出,任务会被标记为 failed
"""
# ── 入队前检查:任务已是 pending,用 > 判断(包含当前任务) ──
# 全局限流检查(始终生效)
global_pending = generation_task_repository.count_pending_total()
if global_pending > global_pending_limit:
logger.warning(
"[队列限流] 全局 pending 任务数超限(入队前): %d/%d, user_id=%s",
global_pending,
global_pending_limit,
user_id or "unknown",
)
exc = GlobalQueueFull(pending_count=global_pending, limit=global_pending_limit)
_mark_task_failed_safely(task, generation_task_repository, log_prefix, str(exc))
raise exc
# 用户级限流检查(传了 user_id 才做)
if user_id:
user_pending = generation_task_repository.count_pending_by_user(user_id)
if user_pending > user_pending_limit:
logger.warning(
"[队列限流] 用户 pending 任务数超限(入队前): user_id=%s, count=%d/%d",
user_id,
user_pending,
user_pending_limit,
)
exc = UserPendingLimitExceeded(user_id=user_id, pending_count=user_pending, limit=user_pending_limit)
_mark_task_failed_safely(task, generation_task_repository, log_prefix, str(exc))
raise exc
# ── 发送 Celery 任务 ──
try:
celery_app.send_task("worker.generate_video", args=[task.id])
except Exception as e:
logger.error(
"%s 入队失败,标记为失败: task_id=%s error=%s",
log_prefix,
task.id,
e,
exc_info=True,
)
try:
task.mark_failed(f"任务入队失败: {e}")
generation_task_repository.update(task)
except Exception as update_err:
logger.error(
"%s 入队失败后更新状态也失败: task_id=%s error=%s",
log_prefix,
task.id,
update_err,
exc_info=True,
)
return False
# ── 入队后最终校验:并发竞态兜底 ──
# 发送成功后再查一次,防止两个请求同时通过入队前检查导致超限
global_after = generation_task_repository.count_pending_total()
user_after = generation_task_repository.count_pending_by_user(user_id) if user_id else 0
global_over = global_after > global_pending_limit
user_over = bool(user_id and user_after > user_pending_limit)
if global_over or user_over:
if global_over:
reason = f"全局 pending 超限(入队后): {global_after}/{global_pending_limit}"
exc: Exception = GlobalQueueFull(pending_count=global_after, limit=global_pending_limit)
else:
reason = f"用户 pending 超限(入队后): {user_after}/{user_pending_limit}"
exc = UserPendingLimitExceeded(user_id=user_id, pending_count=user_after, limit=user_pending_limit)
logger.warning(
"[队列限流] %s, task_id=%s, user_id=%s — 回滚状态为 failed",
reason,
task.id,
user_id or "unknown",
)
_mark_task_failed_safely(task, generation_task_repository, log_prefix, reason)
raise exc
# 入队成功日志
if log_task_status:
logger.info(
"%s 入队成功: task_id=%s, status=%s",
log_prefix,
task.id,
task.status,
)
else:
logger.info("%s 入队成功: task_id=%s", log_prefix, task.id)
return True
Regular → Executable
+2 -32
View File
@@ -189,37 +189,7 @@ def get_voice_clone_profile_repository(
def get_cosyvoice_service():
"""Provide the CosyVoice service instance.
注入 OSS 音频URL预签名函数,确保私有bucket下的参考音频
能被 CosyVoice 服务器下载。
"""
from app.core.storage import get_storage_service
"""Provide the CosyVoice service instance."""
from packages.application.cosyvoice_service import CosyVoiceService
storage = get_storage_service()
def _sign_audio_url(url: str) -> str:
"""对音频URL做预签名,私有bucket下 CosyVoice 服务器才能下载."""
return storage.get_download_url(url, expires_seconds=86400)
return CosyVoiceService(audio_url_signer=_sign_audio_url)
def get_audio_url_signer():
"""提供音频URL预签名函数(24小时有效期)。
用于所有 API 返回给前端的音频 URL,确保私有 OSS bucket 下可正常访问。
空 URL、非 OSS URL 直接原样返回;签名失败时回退到原始 URL。
"""
from app.core.storage import get_storage_service
storage = get_storage_service()
def sign_audio_url(url: str) -> str:
if not url:
return url
return storage.get_download_url(url, expires_seconds=86400)
return sign_audio_url
return CosyVoiceService()
+32
View File
@@ -0,0 +1,32 @@
from datetime import datetime
from pydantic import BaseModel, Field
class RecentTaskItem(BaseModel):
id: str
task_type: str = "generation"
status: str
current_step: str = ""
error_message: str = ""
updated_at: datetime | None = None
class SubscriptionInfo(BaseModel):
"""用户订阅信息。"""
plan: str = "free"
is_active: bool = False
class DashboardOverviewResponse(BaseModel):
"""Dashboard 概览数据。"""
total_assets: int = 0
used_storage_bytes: int = 0
total_titles: int = 0
total_voices: int = 0
total_tasks: int = 0
total_products: int = 0
subscription: SubscriptionInfo = Field(default_factory=SubscriptionInfo)
recent_tasks: list[RecentTaskItem] = Field(default_factory=list)
+1 -18
View File
@@ -1,6 +1,4 @@
import json
from pydantic import BaseModel, Field, field_validator, model_validator
from pydantic import BaseModel, Field, model_validator
class CreateGenerationTaskRequest(BaseModel):
@@ -64,21 +62,6 @@ class GenerationTaskResponse(BaseModel):
progress: float
result_count: int
error_message: str
logs: list[dict] = Field(default_factory=list)
@field_validator("logs", mode="before")
@classmethod
def _parse_logs(cls, v: object) -> list[dict]:
"""将 JSON 字符串解析为 list[dict]。"""
if isinstance(v, str):
try:
parsed = json.loads(v)
return parsed if isinstance(parsed, list) else []
except (json.JSONDecodeError, TypeError):
return []
if isinstance(v, list):
return v
return []
class BatchGenerationTaskResponse(BaseModel):
+109
View File
@@ -0,0 +1,109 @@
"""Job API schemas — Phase 8 任务 2.10."""
from __future__ import annotations
from datetime import datetime
from typing import Any, Optional
from pydantic import BaseModel, Field
class CreateJobRequest(BaseModel):
"""创建任务请求体。"""
project_id: str = Field(..., min_length=1, description="项目 ID")
job_type: str = Field(
...,
description="任务类型: video_compose / render_edit_plan / asset_ingest / classification / voice_extraction / generation",
)
payload: dict[str, Any] = Field(default_factory=dict, description="任务输入参数")
source_id: str = Field(default="", description="关联的业务实体 ID(如 edit_plan_id)")
max_retries: int = Field(default=3, ge=0, le=10, description="最大重试次数")
class UpdateProgressRequest(BaseModel):
"""更新任务进度请求体。"""
progress: float = Field(..., ge=0.0, le=100.0, description="进度百分比")
current_stage: str = Field(default="", description="当前阶段描述")
class CompleteJobRequest(BaseModel):
"""完成任务请求体。"""
result: dict[str, Any] = Field(default_factory=dict, description="任务结果")
class FailJobRequest(BaseModel):
"""标记任务失败请求体。"""
error_message: str = Field(..., min_length=1, description="错误信息")
class JobResponse(BaseModel):
"""任务响应体。"""
id: str
project_id: str
job_type: str
status: str
progress: float
current_stage: str
payload: dict[str, Any]
result: dict[str, Any]
error_message: str
retry_count: int
max_retries: int
celery_task_id: str
source_id: str
created_by_user_id: str
is_retryable: bool
started_at: Optional[datetime] = None
completed_at: Optional[datetime] = None
created_at: datetime
updated_at: datetime
model_config = {"from_attributes": True}
class ListJobsResponse(BaseModel):
"""任务列表响应体。"""
items: list[JobResponse]
total: int
class JobStatisticsResponse(BaseModel):
"""任务统计响应体。"""
project_id: str
total: int
pending: int
running: int
success: int
failed: int
def job_to_response(job) -> JobResponse:
"""将 Job 领域对象转换为 API 响应。"""
return JobResponse(
id=job.id,
project_id=job.project_id,
job_type=job.job_type.value if hasattr(job.job_type, "value") else str(job.job_type),
status=job.status.value if hasattr(job.status, "value") else str(job.status),
progress=job.progress,
current_stage=job.current_stage,
payload=job.payload,
result=job.result,
error_message=job.error_message,
retry_count=job.retry_count,
max_retries=job.max_retries,
celery_task_id=job.celery_task_id,
source_id=job.source_id,
created_by_user_id=job.created_by_user_id,
is_retryable=job.is_retryable,
started_at=job.started_at,
completed_at=job.completed_at,
created_at=job.created_at,
updated_at=job.updated_at,
)
+86
View File
@@ -0,0 +1,86 @@
"""Recipe API schemas."""
from __future__ import annotations
from datetime import datetime
from typing import Any, Dict, List, Optional
from pydantic import BaseModel, Field
# ── Response ──
class RecipeItemResponse(BaseModel):
id: str
recipe_id: str
item_type: str
item_id: str
position: int
metadata_: Dict[str, Any] = Field(default_factory=dict, alias="metadata")
class Config:
populate_by_name = True
class RecipeResponse(BaseModel):
id: str
user_id: str
name: str
description: str = ""
template_id: str = ""
generation_params: Dict[str, Any] = Field(default_factory=dict)
items: List[RecipeItemResponse] = Field(default_factory=list)
is_active: bool = True
metadata_: Dict[str, Any] = Field(default_factory=dict, alias="metadata")
created_at: datetime
updated_at: datetime
class Config:
populate_by_name = True
class ListRecipesResponse(BaseModel):
items: List[RecipeResponse]
total: int = 0
class UseRecipeResponse(BaseModel):
recipe: RecipeResponse
warnings: List[Dict[str, Any]] = Field(default_factory=list)
# ── Request ──
class RecipeItemRequest(BaseModel):
item_type: str
item_id: str
position: int = 0
metadata_: Dict[str, Any] = Field(default_factory=dict, alias="metadata")
class Config:
populate_by_name = True
class CreateRecipeRequest(BaseModel):
name: str
description: str = ""
template_id: str = ""
generation_params: Dict[str, Any] = Field(default_factory=dict)
items: List[RecipeItemRequest] = Field(default_factory=list)
metadata_: Dict[str, Any] = Field(default_factory=dict, alias="metadata")
class Config:
populate_by_name = True
class UpdateRecipeRequest(BaseModel):
name: Optional[str] = None
description: Optional[str] = None
template_id: Optional[str] = None
generation_params: Optional[Dict[str, Any]] = None
items: Optional[List[RecipeItemRequest]] = None
metadata_: Optional[Dict[str, Any]] = Field(default=None, alias="metadata")
class Config:
populate_by_name = True
View File
+161
View File
@@ -0,0 +1,161 @@
/**
* 账号管理 Mock API
*
* 模拟多平台账号绑定/解绑操作
* 支持平台:抖音、快手、小红书、微信视频号
*/
/* ── 类型定义 ───────────────────────────────────────────── */
/** 平台 ID */
export type PlatformId = "douyin" | "kuaishou" | "xiaohongshu" | "wechat";
/** 账号状态 */
export type AccountStatus = "active" | "expired" | "limited";
/** 已绑定的账号 */
export interface Account {
id: string;
platform_id: PlatformId;
name: string;
avatar?: string;
status: AccountStatus;
bound_at: string;
}
/** 平台信息 */
export interface Platform {
id: PlatformId;
name: string;
subName: string;
icon: string;
gradient: string;
}
/** 绑定账号请求 */
export interface BindAccountRequest {
platform_id: PlatformId;
name: string;
}
/* ── 平台配置 ───────────────────────────────────────────── */
export const PLATFORMS: Platform[] = [
{
id: "douyin",
name: "抖音",
subName: "短视频发布平台",
icon: "📱",
gradient: "linear-gradient(135deg, #fe2c55, #25f4ee)",
},
{
id: "kuaishou",
name: "快手",
subName: "短视频发布平台",
icon: "🎬",
gradient: "linear-gradient(135deg, #ff4906, #ffba00)",
},
{
id: "xiaohongshu",
name: "小红书",
subName: "种草笔记发布平台",
icon: "📕",
gradient: "linear-gradient(135deg, #ff2442, #ff6b6b)",
},
{
id: "wechat",
name: "微信视频号",
subName: "视频号发布平台",
icon: "💬",
gradient: "linear-gradient(135deg, #07c160, #4cd964)",
},
];
/* ── Mock 数据 ───────────────────────────────────────────── */
let MOCK_ACCOUNTS: Account[] = [
{
id: "acc-001",
platform_id: "douyin",
name: "小虾官方号",
avatar: "🦐",
status: "active",
bound_at: "2025-12-01T10:00:00Z",
},
{
id: "acc-002",
platform_id: "douyin",
name: "小虾日常",
avatar: "🐟",
status: "active",
bound_at: "2025-12-15T14:30:00Z",
},
{
id: "acc-003",
platform_id: "kuaishou",
name: "小虾剪辑",
avatar: "🎬",
status: "active",
bound_at: "2026-01-05T09:00:00Z",
},
{
id: "acc-004",
platform_id: "xiaohongshu",
name: "小虾种草",
avatar: "📕",
status: "limited",
bound_at: "2026-02-20T16:00:00Z",
},
];
/* ── 模拟延迟 ───────────────────────────────────────────── */
const delay = (ms: number) => new Promise((r) => setTimeout(r, ms));
/* ── API 函数 ───────────────────────────────────────────── */
/** 获取指定平台的账号列表 */
export async function getAccountsByPlatform(
platformId: PlatformId,
): Promise<Account[]> {
await delay(300);
return MOCK_ACCOUNTS.filter((a) => a.platform_id === platformId);
}
/** 获取所有平台的账号总数 */
export async function getAllAccounts(): Promise<Account[]> {
await delay(200);
return [...MOCK_ACCOUNTS];
}
/** 绑定新账号 */
export async function bindAccount(data: BindAccountRequest): Promise<Account> {
await delay(500);
const newAccount: Account = {
id: `acc-${Date.now()}`,
platform_id: data.platform_id,
name: data.name,
avatar: undefined,
status: "active",
bound_at: new Date().toISOString(),
};
MOCK_ACCOUNTS = [...MOCK_ACCOUNTS, newAccount];
return newAccount;
}
/** 解绑账号 */
export async function unbindAccount(accountId: string): Promise<void> {
await delay(400);
MOCK_ACCOUNTS = MOCK_ACCOUNTS.filter((a) => a.id !== accountId);
}
/* ── 状态配置 ───────────────────────────────────────────── */
export const ACCOUNT_STATUS_CONFIG: Record<
AccountStatus,
{ label: string; className: string }
> = {
active: { label: "正常", className: "acc-status--active" },
expired: { label: "已过期", className: "acc-status--expired" },
limited: { label: "受限", className: "acc-status--limited" },
};
+42
View File
@@ -0,0 +1,42 @@
/**
* 仪表盘 API
* Phase 1 新增:用户仪表盘概览
*/
import apiClient from "./client";
/** 仪表盘概览数据 */
export interface DashboardOverview {
/** 素材总数 */
total_assets: number;
/** 已用存储(字节) */
used_storage_bytes: number;
/** 总标题数 */
total_titles: number;
/** 总配音数 */
total_voices: number;
/** 生成任务总数 */
total_tasks: number;
/** 成品总数 */
total_products: number;
/** 最近生成任务 */
recent_tasks: Array<{
id: string;
task_type: string;
status: string;
progress: number;
user_message: string;
created_at: string;
}>;
/** 订阅信息 */
subscription: {
plan: "free" | "pro" | "enterprise";
status: "active" | "inactive" | "expired";
expires_at?: string;
};
}
/** 获取仪表盘概览数据 */
export const getDashboardOverview = async (): Promise<DashboardOverview> => {
const response = await apiClient.get("/dashboard/overview");
return response.data;
};
@@ -0,0 +1,365 @@
/* V21 业务组件统一样式 */
/* ==================== 按钮 ==================== */
.xx-primary-btn {
background: var(--gradient-primary) !important;
color: var(--text-inverse) !important;
border: none !important;
border-radius: var(--radius-md) !important;
padding: 10px 20px !important;
font-weight: var(--font-weight-bold) !important;
box-shadow: var(--shadow-primary) !important;
transition: var(--transition-all) !important;
cursor: pointer;
height: auto !important;
}
.xx-primary-btn:hover {
box-shadow: var(--shadow-hover) !important;
transform: translateY(-1px);
}
.xx-ghost-btn {
background: transparent !important;
color: var(--primary-color) !important;
border: 2px solid var(--primary-color) !important;
border-radius: var(--radius-md) !important;
padding: var(--space-sm) 18px !important;
font-weight: var(--font-weight-bold) !important;
transition: var(--transition-all) !important;
cursor: pointer;
height: auto !important;
}
.xx-ghost-btn:hover {
background: var(--primary-soft) !important;
}
/* ==================== 卡片 ==================== */
.xx-card {
background: var(--bg-elevated);
border: 1px solid var(--border-color);
border-radius: var(--radius-xl);
box-shadow: var(--shadow-card);
padding: var(--space-lg);
margin-bottom: 20px;
transition: all var(--transition-slow);
}
.xx-card:hover {
box-shadow: var(--shadow-md);
transform: translateY(-2px);
}
/* ==================== 页面结构 ==================== */
.xx-page {
max-width: 1200px;
margin: 0 auto;
padding: var(--space-lg);
}
.xx-page-head {
display: flex;
justify-content: space-between;
align-items: flex-start;
gap: 18px;
margin-bottom: 28px;
}
.xx-page-head h2 {
font-size: 26px;
font-weight: var(--font-weight-extrabold);
color: var(--text-primary);
margin: 0 0 var(--space-sm);
}
.xx-page-head p {
font-size: var(--font-size-base);
color: var(--text-secondary);
margin: 0;
}
/* ==================== 表格样式 ==================== */
.xx-table-card {
background: var(--bg-elevated);
border: 1px solid var(--border-color);
border-radius: var(--radius-xl);
box-shadow: var(--shadow-card);
padding: 20px;
overflow: hidden;
}
/* 表格包装器 */
.xx-table-wrapper {
border-radius: var(--radius-lg);
overflow: hidden;
}
/* ==================== 标签/Tag ==================== */
.xx-tag {
padding: var(--space-xs) 12px;
border-radius: var(--radius-xs);
font-size: 13px;
font-weight: var(--font-weight-medium);
}
.xx-tag-indigo {
background: var(--primary-soft);
color: var(--primary-color);
border: 1px solid var(--color-primary-200);
}
.xx-tag-success {
background: var(--success-soft);
color: var(--color-secondary-500);
border: 1px solid var(--success-border);
}
.xx-tag-warning {
background: var(--warning-soft);
color: var(--accent-dark);
border: 1px solid var(--color-accent-200);
}
.xx-tag-error {
background: var(--error-soft);
color: var(--error-color);
border: 1px solid var(--error-border);
}
/* ==================== 搜索栏 ==================== */
.xx-search-bar {
margin-bottom: 20px;
}
.xx-search-input {
width: 100%;
padding: 12px 18px;
border: 2px solid var(--border-color);
border-radius: var(--radius-md);
font-size: var(--font-size-base);
background: var(--bg-primary);
transition: var(--transition-all);
outline: none;
}
.xx-search-input:focus {
border-color: var(--primary-color);
box-shadow: 0 0 0 4px
color-mix(in srgb, var(--primary-color) 10%, transparent);
}
/* ==================== Modal ==================== */
.xx-modal .ant-modal-content {
border-radius: var(--radius-xl);
padding: var(--space-lg);
}
.xx-modal .ant-modal-header {
border-radius: var(--radius-xl) var(--radius-xl) 0 0;
padding: 20px var(--space-lg);
border-bottom: 1px solid var(--border-color);
}
.xx-modal .ant-modal-title {
font-size: var(--font-size-lg);
font-weight: var(--font-weight-bold);
color: var(--text-primary);
}
.xx-modal .ant-modal-footer {
border-top: 1px solid var(--border-color);
padding: var(--space-md) var(--space-lg);
}
/* ==================== 空状态 ==================== */
.xx-empty-state {
text-align: center;
padding: var(--space-3xl) var(--space-lg);
color: var(--text-secondary);
}
.xx-empty-state-icon {
font-size: 48px;
margin-bottom: var(--space-md);
}
/* ==================== 网格布局 ==================== */
.xx-grid-2 {
display: grid;
grid-template-columns: repeat(2, 1fr);
gap: 20px;
}
.xx-grid-3 {
display: grid;
grid-template-columns: repeat(3, 1fr);
gap: 20px;
}
.xx-grid-4 {
display: grid;
grid-template-columns: repeat(4, 1fr);
gap: 20px;
}
@media (max-width: 768px) {
.xx-grid-2,
.xx-grid-3,
.xx-grid-4 {
grid-template-columns: 1fr;
}
}
/* ==================== 配额展示 ==================== */
.xx-quota-item {
padding: 20px;
background: var(--bg-primary);
border: 1px solid var(--border-color);
border-radius: var(--radius-lg);
transition: var(--transition-all);
}
.xx-quota-item:hover {
border-color: var(--primary-color);
box-shadow: 0 8px 24px
color-mix(in srgb, var(--primary-color) 10%, transparent);
}
/* ==================== 进度条 ==================== */
.xx-progress {
margin-top: 12px;
}
/* ==================== Ant Design 覆盖样式 ==================== */
/* Table overrides */
.ant-table-wrapper .ant-table-thead > tr > th {
background: var(--bg-secondary) !important;
font-weight: var(--font-weight-bold) !important;
color: var(--text-primary) !important;
border-bottom: 2px solid var(--border-color) !important;
padding: 14px var(--space-md) !important;
}
.ant-table-wrapper .ant-table-tbody > tr > td {
padding: 14px var(--space-md) !important;
border-bottom: 1px solid var(--color-gray-100) !important;
}
.ant-table-wrapper .ant-table-tbody > tr:hover > td {
background: var(--color-gray-50) !important;
}
/* Card overrides */
.ant-card {
border-radius: var(--radius-xl) !important;
border: 1px solid var(--border-color) !important;
}
.ant-card-head {
border-bottom: 1px solid var(--border-color) !important;
min-height: 52px !important;
padding: 0 var(--space-lg) !important;
}
.ant-card-head-title {
font-weight: var(--font-weight-bold) !important;
font-size: var(--font-size-md) !important;
color: var(--text-primary) !important;
}
.ant-card-body {
padding: 20px var(--space-lg) !important;
}
/* Modal overrides */
.ant-modal-content {
border-radius: var(--radius-xl) !important;
overflow: hidden;
}
.ant-modal-header {
padding: 20px var(--space-lg) !important;
background: var(--bg-primary) !important;
}
.ant-modal-title {
font-weight: var(--font-weight-bold) !important;
font-size: var(--font-size-lg) !important;
color: var(--text-primary) !important;
}
.ant-modal-body {
padding: var(--space-lg) !important;
}
.ant-modal-footer {
padding: var(--space-md) var(--space-lg) !important;
}
/* Button overrides */
.ant-btn-primary {
background: var(--gradient-primary) !important;
border: none !important;
border-radius: var(--radius-md) !important;
box-shadow: var(--shadow-primary) !important;
height: auto !important;
padding: 10px 20px !important;
font-weight: var(--font-weight-bold) !important;
}
.ant-btn-primary:hover {
background: var(--gradient-primary) !important;
box-shadow: var(--shadow-hover) !important;
transform: translateY(-1px);
}
/* Tag overrides */
.ant-tag {
border-radius: var(--radius-xs) !important;
padding: var(--space-xs) 12px !important;
font-weight: var(--font-weight-medium) !important;
}
/* Select overrides */
.ant-select-selector {
border-radius: var(--radius-md) !important;
border-color: var(--border-color) !important;
}
.ant-select:not(.ant-select-disabled):hover .ant-select-selector {
border-color: var(--primary-color) !important;
}
.ant-select-focused .ant-select-selector {
border-color: var(--primary-color) !important;
box-shadow: 0 0 0 3px
color-mix(in srgb, var(--primary-color) 10%, transparent) !important;
}
/* Input overrides */
.ant-input {
border-radius: var(--radius-md) !important;
border-color: var(--border-color) !important;
padding: 10px 14px !important;
}
.ant-input:hover {
border-color: var(--primary-color) !important;
}
.ant-input:focus {
border-color: var(--primary-color) !important;
box-shadow: 0 0 0 3px
color-mix(in srgb, var(--primary-color) 10%, transparent) !important;
}
/* Progress overrides */
.ant-progress-inner {
background: var(--color-gray-100) !important;
border-radius: var(--radius-xs) !important;
}
.ant-progress-bg {
border-radius: var(--radius-xs) !important;
}
@@ -0,0 +1,290 @@
/**
* CloneVoiceModal — 音色克隆弹窗
*
* 三步骤状态:input → uploading → success
* 支持上传音频文件或直接录制(mock,无真实录音)
*
* V21 Design System — 零 antd 直接导入
*/
import React, { useState, useCallback, useRef } from "react";
import { Modal, Button } from "@/components/ui";
import { createVoiceClone, toVoiceClone } from "@/api/voiceClone";
import type { VoiceClone } from "@/api/voiceClone";
import { uploadAsset } from "@/api/assets";
import "./clone-voice-modal.css";
/* ── 类型定义 ───────────────────────────────────────────── */
type ModalStep = "input" | "uploading" | "success";
export interface CloneVoiceModalProps {
/** 弹窗是否可见 */
open: boolean;
/** 关闭弹窗回调 */
onClose: () => void;
/** 克隆成功回调(返回新创建的音色) */
onSuccess?: (voice: VoiceClone) => void;
}
/* ── 默认音色名称计数器 ─────────────────────────────────── */
let cloneCounter = 1;
const getNextDefaultName = (): string => {
const name = `我的声音 ${cloneCounter}`;
cloneCounter += 1;
return name;
};
/* ── 组件 ───────────────────────────────────────────────── */
const CloneVoiceModal: React.FC<CloneVoiceModalProps> = ({
open,
onClose,
onSuccess,
}) => {
const [step, setStep] = useState<ModalStep>("input");
const [voiceName, setVoiceName] = useState("");
const [isRecording, setIsRecording] = useState(false);
const [selectedFile, setSelectedFile] = useState<File | null>(null);
const [dragActive, setDragActive] = useState(false);
const fileInputRef = useRef<HTMLInputElement>(null);
/** 重置弹窗状态 */
const resetState = useCallback(() => {
setStep("input");
setVoiceName("");
setSelectedFile(null);
setIsRecording(false);
setDragActive(false);
}, []);
/** 关闭弹窗 */
const handleClose = useCallback(() => {
resetState();
onClose();
}, [resetState, onClose]);
/** 上传区域点击 */
const handleUploadClick = () => {
fileInputRef.current?.click();
};
/** 文件选择 */
const handleFileChange = (e: React.ChangeEvent<HTMLInputElement>) => {
const file = e.target.files?.[0];
if (file) {
setSelectedFile(file);
// 清除之前的录制状态
setIsRecording(false);
}
// 清空 input 以允许重复选择同一文件
e.target.value = "";
};
/** 拖拽事件 */
const handleDrag = (e: React.DragEvent) => {
e.preventDefault();
e.stopPropagation();
if (e.type === "dragenter" || e.type === "dragover") {
setDragActive(true);
} else if (e.type === "dragleave") {
setDragActive(false);
}
};
const handleDrop = (e: React.DragEvent) => {
e.preventDefault();
e.stopPropagation();
setDragActive(false);
const file = e.dataTransfer.files?.[0];
if (file) {
const ext = file.name.split(".").pop()?.toLowerCase();
if (ext === "mp3" || ext === "wav") {
setSelectedFile(file);
setIsRecording(false);
}
}
};
/** 录制按钮(mock) */
const handleRecord = () => {
setIsRecording((prev) => !prev);
if (!isRecording) {
// 开始录制 — 清除已选文件
setSelectedFile(null);
}
};
/** 开始克隆 */
const handleStartClone = async () => {
const name = voiceName.trim() || getNextDefaultName();
setStep("uploading");
try {
// 先上传音频文件获取真实 URL
let audioUrl: string;
if (selectedFile) {
const formData = new FormData();
formData.append("file", selectedFile);
formData.append("kind", "voice");
const uploadResult = await uploadAsset(formData);
audioUrl = uploadResult.url;
} else {
// 录制功能暂未实现,提示用户上传
setStep("input");
return;
}
// 提交克隆请求
const result = await createVoiceClone({
name,
audio_url: audioUrl,
});
setStep("success");
// 2秒后自动关闭
setTimeout(() => {
onSuccess?.(toVoiceClone(result));
handleClose();
}, 2000);
} catch {
setStep("input");
}
};
/** 弹窗打开时初始化默认名称 */
const handleAfterOpenChange = (visible: boolean) => {
if (visible) {
setVoiceName(getNextDefaultName());
}
};
const canStart = selectedFile || isRecording;
return (
<Modal
open={open}
onCancel={handleClose}
title="🎤 克隆新音色"
width={520}
footer={null}
destroyOnClose
afterOpenChange={handleAfterOpenChange}
>
{/* ── 输入步骤 ──────────────────────────────────── */}
{step === "input" && (
<div className="cvm-body">
{/* 音色名称 */}
<div className="cvm-field">
<label className="cvm-label">音色名称</label>
<input
type="text"
className="cvm-input"
value={voiceName}
onChange={(e) => setVoiceName(e.target.value)}
placeholder="输入音色名称"
/>
</div>
{/* 上传区域 */}
<div className="cvm-field">
<label className="cvm-label">上传音频</label>
<div
className={`cvm-upload-zone${dragActive ? " cvm-upload-zone--active" : ""}`}
onClick={handleUploadClick}
onDragEnter={handleDrag}
onDragOver={handleDrag}
onDragLeave={handleDrag}
onDrop={handleDrop}
>
<div className="cvm-upload-icon">🎵</div>
<p className="cvm-upload-title">
{selectedFile ? selectedFile.name : "拖拽音频文件到此处"}
</p>
<p className="cvm-upload-hint">支持 MP3、WAV 格式</p>
<input
ref={fileInputRef}
type="file"
accept=".mp3,.wav,audio/mpeg,audio/wav"
style={{ display: "none" }}
onChange={handleFileChange}
/>
</div>
</div>
{/* 或分隔 */}
<div className="cvm-divider">
<div className="cvm-divider-line" />
<span className="cvm-divider-text">或</span>
<div className="cvm-divider-line" />
</div>
{/* 录制区域 */}
<div className="cvm-field">
<label className="cvm-label">直接录制</label>
<div className="cvm-record-area">
<p className="cvm-record-hint">
{isRecording
? "录制中…再次点击停止"
: "点击按钮开始录制你的声音"}
</p>
<button
type="button"
className={`cvm-record-btn${isRecording ? " cvm-record-btn--recording" : ""}`}
onClick={handleRecord}
>
🎙️
</button>
</div>
</div>
{/* 提示 */}
<div className="cvm-tip">
<span className="cvm-tip-icon">💡</span>
<span>
建议上传10秒~3分钟的清晰语音,环境安静、语速均匀效果最佳
</span>
</div>
{/* 底部按钮 */}
<div className="cvm-footer">
<Button buttonType="ghost" onClick={handleClose}>
取消
</Button>
<Button
buttonType="primary"
disabled={!canStart}
onClick={handleStartClone}
>
🎤 开始克隆
</Button>
</div>
</div>
)}
{/* ── 上传中步骤 ────────────────────────────────── */}
{step === "uploading" && (
<div className="cvm-uploading">
<div className="cvm-uploading-spinner" />
<p className="cvm-uploading-text">正在克隆你的音色…</p>
<p className="cvm-uploading-sub">AI 正在分析你的声音特征,请稍候</p>
</div>
)}
{/* ── 成功步骤 ──────────────────────────────────── */}
{step === "success" && (
<div className="cvm-success">
<div className="cvm-success-icon">✅</div>
<h3 className="cvm-success-title">克隆已提交</h3>
<p className="cvm-success-desc">
音色正在生成中,完成后将出现在列表中
</p>
</div>
)}
</Modal>
);
};
export default CloneVoiceModal;
@@ -0,0 +1,325 @@
/**
* CloneVoiceModal — V21 Design System
*
* 音色克隆弹窗样式
* 三步骤状态:input → uploading → success
*/
/* ── 弹窗内容区 ─────────────────────────────────────────── */
.cvm-body {
display: flex;
flex-direction: column;
gap: 20px;
}
/* ── 表单区 ─────────────────────────────────────────────── */
.cvm-field {
display: flex;
flex-direction: column;
gap: 6px;
}
.cvm-label {
font-size: 13px;
font-weight: 600;
color: var(--text-secondary, #475467);
}
.cvm-input {
width: 100%;
padding: 10px 14px;
border: 1px solid var(--line, #e4e7ec);
border-radius: var(--radius-sm);
background: var(--bg-surface, #fff);
color: var(--text-primary, #101828);
font-size: 14px;
line-height: 1.5;
transition:
border-color 0.2s,
box-shadow 0.2s;
outline: none;
}
.cvm-input:focus {
border-color: var(--primary, #6366f1);
box-shadow: 0 0 0 3px
color-mix(in srgb, var(--primary-color) 12%, transparent);
}
.cvm-input::placeholder {
color: var(--muted, #98a2b3);
}
/* ── 上传区域 ───────────────────────────────────────────── */
.cvm-upload-zone {
border: 2px dashed var(--line, #e4e7ec);
border-radius: var(--radius-md);
padding: 28px 20px;
text-align: center;
background: var(--bg-subtle, #f8fafc);
cursor: pointer;
transition:
border-color 0.2s,
background 0.2s;
}
.cvm-upload-zone:hover {
border-color: var(--primary, #6366f1);
background: color-mix(in srgb, var(--primary-color) 4%, transparent);
}
.cvm-upload-zone.cvm-upload-zone--active {
border-color: var(--primary, #6366f1);
background: color-mix(in srgb, var(--primary-color) 6%, transparent);
}
.cvm-upload-icon {
font-size: 36px;
margin-bottom: 8px;
line-height: 1;
}
.cvm-upload-title {
font-size: 14px;
font-weight: 600;
color: var(--text-primary, #101828);
margin: 0 0 4px;
}
.cvm-upload-hint {
font-size: 13px;
color: var(--muted, #98a2b3);
margin: 0;
}
/* ── 或分隔线 ───────────────────────────────────────────── */
.cvm-divider {
display: flex;
align-items: center;
gap: 16px;
margin: 4px 0;
}
.cvm-divider-line {
flex: 1;
height: 1px;
background: var(--line, #e4e7ec);
}
.cvm-divider-text {
font-size: 13px;
color: var(--muted, #98a2b3);
flex-shrink: 0;
}
/* ── 录制区域 ───────────────────────────────────────────── */
.cvm-record-area {
border: 1px solid var(--line, #e4e7ec);
border-radius: var(--radius-md);
padding: 24px;
text-align: center;
}
.cvm-record-hint {
font-size: 13px;
color: var(--muted, #98a2b3);
margin: 0 0 14px;
}
.cvm-record-btn {
width: 80px;
height: 80px;
border-radius: 50%;
border: none;
cursor: pointer;
font-size: 32px;
line-height: 1;
padding: 0;
background: linear-gradient(
135deg,
var(--error-color, #ef4444),
var(--error-dark, #dc2626)
);
color: var(--text-inverse);
box-shadow: 0 4px 14px
color-mix(in srgb, var(--error-color, #ef4444) 35%, transparent);
transition:
transform 0.15s,
box-shadow 0.15s;
display: inline-flex;
align-items: center;
justify-content: center;
}
.cvm-record-btn:hover {
transform: scale(1.06);
box-shadow: 0 6px 20px
color-mix(in srgb, var(--error-color, #ef4444) 45%, transparent);
}
.cvm-record-btn:active {
transform: scale(0.96);
}
.cvm-record-btn--recording {
animation: cvm-pulse 1.2s ease-in-out infinite;
}
@keyframes cvm-pulse {
0%,
100% {
box-shadow: 0 4px 14px
color-mix(in srgb, var(--error-color, #ef4444) 35%, transparent);
}
50% {
box-shadow: 0 4px 28px
color-mix(in srgb, var(--error-color, #ef4444) 60%, transparent);
}
}
/* ── 提示条 ─────────────────────────────────────────────── */
.cvm-tip {
display: flex;
align-items: flex-start;
gap: 8px;
padding: 12px 16px;
background: var(--warning-soft, #fef3c7);
border-radius: var(--radius-sm);
font-size: 13px;
color: var(--warning-color, #92400e);
line-height: 1.5;
}
.cvm-tip-icon {
flex-shrink: 0;
font-size: 14px;
line-height: 1.5;
}
/* ── 底部按钮 ───────────────────────────────────────────── */
.cvm-footer {
display: flex;
gap: 12px;
margin-top: 4px;
}
.cvm-footer .xx-btn {
flex: 1;
}
/* ── 上传中状态 ─────────────────────────────────────────── */
.cvm-uploading {
display: flex;
flex-direction: column;
align-items: center;
justify-content: center;
padding: 48px 20px;
gap: 16px;
}
.cvm-uploading-spinner {
width: 48px;
height: 48px;
border: 3px solid var(--line, #e4e7ec);
border-top-color: var(--primary, #6366f1);
border-radius: 50%;
animation: cvm-spin 0.8s linear infinite;
}
@keyframes cvm-spin {
to {
transform: rotate(360deg);
}
}
.cvm-uploading-text {
font-size: 15px;
font-weight: 500;
color: var(--text-primary, #101828);
margin: 0;
}
.cvm-uploading-sub {
font-size: 13px;
color: var(--muted, #98a2b3);
margin: 0;
}
/* ── 成功状态 ───────────────────────────────────────────── */
.cvm-success {
display: flex;
flex-direction: column;
align-items: center;
justify-content: center;
padding: 48px 20px;
gap: 12px;
}
.cvm-success-icon {
font-size: 56px;
line-height: 1;
}
.cvm-success-title {
font-size: 18px;
font-weight: 700;
color: var(--text-primary, #101828);
margin: 0;
}
.cvm-success-desc {
font-size: 14px;
color: var(--muted, #98a2b3);
margin: 0;
}
/* ── 响应式 ─────────────────────────────────────────────── */
@media (max-width: 768px) {
.cvm-overlay {
padding: var(--space-md);
}
.cvm-modal {
width: 100%;
max-width: 100%;
padding: var(--space-lg);
}
}
@media (max-width: 576px) {
.cvm-upload-zone {
padding: 20px 14px;
}
.cvm-record-btn {
width: 64px;
height: 64px;
font-size: 26px;
}
.cvm-footer {
flex-direction: column;
}
}
@media (max-width: 480px) {
.cvm-record-btn {
width: 60px;
height: 60px;
}
.cvm-tip {
font-size: 12px;
padding: var(--space-sm);
}
}
-20
View File
@@ -532,23 +532,3 @@
padding: 8px 16px !important;
}
}
/* ── xx-card antd 子元素覆盖样式(从 Admin.css 迁移) ── */
/* AdminComingSoon 等页面使用 <Card className="xx-card"> 时需要 */
/* .xx-card 基础样式和 :hover 已在 global.css 中定义(V21 设计系统) */
.xx-card .ant-card-head {
border-bottom: 1px solid var(--border-color);
padding: 20px 24px;
}
.xx-card .ant-card-head-title {
font-weight: 800;
font-size: 17px;
color: var(--text-primary);
}
.xx-card .ant-card-body {
padding: 24px;
}
+211
View File
@@ -0,0 +1,211 @@
/**
* 统一导航配置
* Header 和 Sidebar 共用此数据源
*/
import React from "react";
import {
DashboardOutlined,
VideoCameraOutlined,
FileOutlined,
AudioOutlined,
FileTextOutlined,
TrophyOutlined,
AppstoreOutlined,
HistoryOutlined,
ControlOutlined,
CrownOutlined,
ScanOutlined,
EditOutlined,
FolderOutlined,
} from "@ant-design/icons";
/** 导航项定义 */
export interface NavItem {
key: string;
label: string;
path: string;
icon: React.ReactNode;
}
/** 导航分组定义 */
export interface NavGroup {
title: string;
items: NavItem[];
}
/**
* 扁平导航列表(Header 使用)
*/
export const NAV_ITEMS: NavItem[] = [
{
key: "dashboard",
label: "概览",
path: "/app/dashboard",
icon: <DashboardOutlined />,
},
{
key: "assets",
label: "素材库",
path: "/app/assets",
icon: <FileOutlined />,
},
{
key: "titles",
label: "标题库",
path: "/app/titles",
icon: <FileTextOutlined />,
},
{
key: "voices",
label: "配音库",
path: "/app/voices",
icon: <AudioOutlined />,
},
{
key: "voice-clone",
label: "我的音色",
path: "/app/voice-clone",
icon: <AudioOutlined />,
},
{
key: "voice-materials",
label: "配音素材库",
path: "/app/voice-materials",
icon: <AudioOutlined />,
},
{
key: "templates",
label: "模板库",
path: "/app/templates",
icon: <AppstoreOutlined />,
},
{
key: "editing-planner",
label: "剪辑编辑器",
path: "/app/editing-planner",
icon: <EditOutlined />,
},
{
key: "my-templates",
label: "我的模板",
path: "/app/my-templates",
icon: <FolderOutlined />,
},
{
key: "generate",
label: "一键生成",
path: "/app/generate",
icon: <VideoCameraOutlined />,
},
{
key: "history",
label: "任务历史",
path: "/app/history",
icon: <HistoryOutlined />,
},
{
key: "products",
label: "成品库",
path: "/app/products",
icon: <TrophyOutlined />,
},
{
key: "duplication",
label: "查重",
path: "/app/duplication",
icon: <ScanOutlined />,
},
];
/**
* 分组导航列表(Sidebar 使用)
*/
export const NAV_GROUPS: NavGroup[] = [
{
title: "创作工具",
items: [
{
key: "dashboard",
label: "首页",
path: "/app/dashboard",
icon: <DashboardOutlined />,
},
{
key: "generate",
label: "一键生成",
path: "/app/generate",
icon: <VideoCameraOutlined />,
},
],
},
{
title: "资源管理",
items: [
{
key: "assets",
label: "素材库",
path: "/app/assets",
icon: <FileOutlined />,
},
{
key: "voices",
label: "配音库",
path: "/app/voices",
icon: <AudioOutlined />,
},
{
key: "voice-clone",
label: "我的音色",
path: "/app/voice-clone",
icon: <AudioOutlined />,
},
{
key: "voice-materials",
label: "配音素材库",
path: "/app/voice-materials",
icon: <AudioOutlined />,
},
{
key: "titles",
label: "标题库",
path: "/app/titles",
icon: <FileTextOutlined />,
},
{
key: "products",
label: "成片库",
path: "/app/products",
icon: <TrophyOutlined />,
},
{
key: "templates",
label: "模板库",
path: "/app/templates",
icon: <AppstoreOutlined />,
},
],
},
{
title: "系统",
items: [
{
key: "history",
label: "任务历史",
path: "/app/history",
icon: <HistoryOutlined />,
},
{
key: "admin",
label: "控制台",
path: "/app/admin",
icon: <ControlOutlined />,
},
{
key: "subscription",
label: "订阅管理",
path: "/app/subscription",
icon: <CrownOutlined />,
},
],
},
];
+198 -75
View File
@@ -2,66 +2,194 @@
* 账号管理页面 — V21 Design System
*
* 展示多平台账号绑定状态(抖音/快手/小红书/微信视频号)
* 后端账号管理 API 尚未就绪,当前展示占位状态
* 支持绑定/解绑操作
*
* 零 antd 直接导入,全部使用 CSS 变量
*/
import React from "react";
import React, { useState, useCallback } from "react";
import { useQueries, useMutation, useQueryClient } from "@tanstack/react-query";
import { Button } from "@/components/ui";
import PageHead from "@/components/layout/PageHead";
import {
PLATFORMS,
getAccountsByPlatform,
unbindAccount,
bindAccount,
ACCOUNT_STATUS_CONFIG,
type Platform,
type Account,
type PlatformId,
} from "@/api/accounts";
import "./accounts.css";
/* ── 类型定义 ───────────────────────────────────────────── */
/* ── Toast 系统 ─────────────────────────────────────────── */
export type PlatformId =
| "douyin"
| "kuaishou"
| "xiaohongshu"
| "wechat";
export interface Platform {
id: PlatformId;
name: string;
subName: string;
icon: string;
gradient: string;
interface Toast {
id: number;
message: string;
type: "success" | "error";
}
/** 支持的平台列表 */
const PLATFORMS: Platform[] = [
{
id: "douyin",
name: "抖音",
subName: "短视频发布",
icon: "🎵",
gradient: "linear-gradient(135deg, #000 0%, #333 100%)",
},
{
id: "kuaishou",
name: "快手",
subName: "短视频发布",
icon: "📹",
gradient: "linear-gradient(135deg, #ff6600 0%, #ff9933 100%)",
},
{
id: "xiaohongshu",
name: "小红书",
subName: "种草笔记 + 视频",
icon: "📕",
gradient: "linear-gradient(135deg, #fe2c55 0%, #ff6680 100%)",
},
{
id: "wechat",
name: "微信视频号",
subName: "视频号发布",
icon: "💬",
gradient: "linear-gradient(135deg, #07c160 0%, #38d97a 100%)",
},
];
let toastIdCounter = 0;
/* ── 平台卡片组件 ───────────────────────────────────────── */
interface PlatformCardProps {
platform: Platform;
accounts: Account[];
isLoading: boolean;
onBind: (platformId: PlatformId) => void;
onUnbind: (accountId: string, accountName: string) => void;
}
const PlatformCard: React.FC<PlatformCardProps> = ({
platform,
accounts,
isLoading,
onBind,
onUnbind,
}) => {
return (
<div className="acc-card">
{/* 平台头部 */}
<div className="acc-card-header">
<div
className="acc-card-icon"
style={{ background: platform.gradient }}
>
{platform.icon}
</div>
<div>
<h3 className="acc-card-title">{platform.name}</h3>
<p className="acc-card-subtitle">{platform.subName}</p>
</div>
</div>
{/* 账号列表 */}
<div className="acc-account-list">
{isLoading ? (
<div className="acc-empty">
<p className="acc-empty-text">加载中…</p>
</div>
) : accounts.length > 0 ? (
accounts.map((account) => {
const statusCfg = ACCOUNT_STATUS_CONFIG[account.status];
return (
<div key={account.id} className="acc-account-row">
<div
className="acc-account-avatar"
style={{ background: platform.gradient }}
>
{account.avatar || platform.icon}
</div>
<div className="acc-account-info">
<div className="acc-account-name">{account.name}</div>
<span className={`acc-status-pill ${statusCfg.className}`}>
{statusCfg.label}
</span>
</div>
<Button
buttonType="ghost"
buttonSize="sm"
onClick={() => onUnbind(account.id, account.name)}
>
解绑
</Button>
</div>
);
})
) : (
<div className="acc-empty">
<div className="acc-empty-icon">🔓</div>
<p className="acc-empty-text">暂未绑定账号</p>
</div>
)}
</div>
{/* 绑定按钮 */}
<Button buttonType="ghost" onClick={() => onBind(platform.id)}>
+ 绑定新账号
</Button>
</div>
);
};
/* ── 主页面 ─────────────────────────────────────────────── */
const Accounts: React.FC = () => {
const queryClient = useQueryClient();
const [toasts, setToasts] = useState<Toast[]>([]);
/** 显示 toast */
const showToast = useCallback((message: string, type: Toast["type"]) => {
const id = ++toastIdCounter;
setToasts((prev) => [...prev, { id, message, type }]);
setTimeout(() => {
setToasts((prev) => prev.filter((t) => t.id !== id));
}, 3000);
}, []);
/** 查询所有平台的账号 */
const _queriesResults = useQueries({
queries: PLATFORMS.map((platform) => ({
queryKey: ["accounts", platform.id] as const,
queryFn: () => getAccountsByPlatform(platform.id),
})),
});
const accountQueries = PLATFORMS.map((platform, i) => ({
platform,
..._queriesResults[i],
}));
/** 解绑 mutation */
const unbindMutation = useMutation({
mutationFn: unbindAccount,
onSuccess: () => {
queryClient.invalidateQueries({ queryKey: ["accounts"] });
showToast("已解绑账号", "success");
},
onError: () => {
showToast("解绑失败", "error");
},
});
/** 绑定 mutation(mock) */
const bindMutation = useMutation({
mutationFn: bindAccount,
onSuccess: () => {
queryClient.invalidateQueries({ queryKey: ["accounts"] });
showToast("账号绑定成功", "success");
},
onError: () => {
showToast("绑定失败", "error");
},
});
/** 绑定新账号(mock:直接创建) */
const handleBind = (platformId: PlatformId) => {
const platform = PLATFORMS.find((p) => p.id === platformId);
if (!platform) return;
const name = window.prompt(`请输入要绑定的${platform.name}账号名称:`);
if (name && name.trim()) {
bindMutation.mutate({ platform_id: platformId, name: name.trim() });
}
};
/** 解绑账号 */
const handleUnbind = (accountId: string, accountName: string) => {
if (window.confirm(`确定解绑账号「${accountName}」吗?`)) {
unbindMutation.mutate(accountId);
}
};
/** 统计已绑定账号数 */
const totalBound = accountQueries.reduce(
(sum, q) => sum + (q.data?.length ?? 0),
0,
);
const totalPlatforms = PLATFORMS.length;
return (
<div className="acc-page">
<PageHead
@@ -69,34 +197,17 @@ const Accounts: React.FC = () => {
description="绑定您的社交平台账号,用于视频一键发布到各平台"
/>
{/* 平台卡片网格 — 占位状态 */}
{/* 平台卡片网格 */}
<div className="acc-grid">
{PLATFORMS.map((platform) => (
<div key={platform.id} className="acc-card">
<div className="acc-card-header">
<div
className="acc-card-icon"
style={{ background: platform.gradient }}
>
{platform.icon}
</div>
<div>
<h3 className="acc-card-title">{platform.name}</h3>
<p className="acc-card-subtitle">{platform.subName}</p>
</div>
</div>
<div className="acc-account-list">
<div className="acc-empty">
<div className="acc-empty-icon">🔒</div>
<p className="acc-empty-text">功能即将上线</p>
</div>
</div>
<Button buttonType="ghost" disabled>
即将支持绑定
</Button>
</div>
{accountQueries.map(({ platform, data, isLoading }) => (
<PlatformCard
key={platform.id}
platform={platform}
accounts={data ?? []}
isLoading={isLoading}
onBind={handleBind}
onUnbind={handleUnbind}
/>
))}
</div>
@@ -104,10 +215,22 @@ const Accounts: React.FC = () => {
<div className="acc-stats-bar">
<span className="acc-stats-icon">📊</span>
<span className="acc-stats-text">
支持 <span className="acc-stats-highlight">{PLATFORMS.length}</span>{" "}
个平台,功能即将上线
已绑定 <span className="acc-stats-highlight">{totalBound}</span>{" "}
个账号 / 支持{" "}
<span className="acc-stats-highlight">{totalPlatforms}</span> 个平台
</span>
</div>
{/* Toast 提示 */}
{toasts.length > 0 && (
<div className="vc-toast-container">
{toasts.map((t) => (
<div key={t.id} className={`vc-toast vc-toast--${t.type}`}>
{t.type === "success" ? "✅" : "❌"} {t.message}
</div>
))}
</div>
)}
</div>
);
};
+500 -10
View File
@@ -1,20 +1,482 @@
/* Admin 页面样式(Phase 3 精简)
*
* 原始 477 行 → 精简至仅保留实际使用的 class。
* 已迁移至 global.css / ui.css 的样式不再重复定义:
* .xx-page-head → global.css
* .xx-primary-btn → global.css
* .xx-tag / .xx-card → ui.css / global.css
*
* 以下 class 仅被 AdminComingSoon.tsx 使用。
*/
/* V21 Admin 页面样式 */
/* 页面容器 */
.dashboard-page,
.analytics-page,
.user-management-page,
.log-viewer-page,
.system-monitor-page,
.admin-coming-soon-page {
padding: 32px;
max-width: 1400px;
margin: 0 auto;
}
/* 页面头部 */
.xx-page-head {
background: rgba(255, 255, 255, 0.94);
border: 1px solid rgba(226, 232, 240, 0.95);
border-radius: var(--radius-xl);
box-shadow: 0 24px 70px rgba(15, 23, 42, 0.09);
padding: 32px 40px;
margin-bottom: 32px;
display: flex;
justify-content: space-between;
align-items: center;
}
.xx-page-head-content {
display: flex;
flex-direction: column;
gap: 8px;
}
.xx-page-head h2 {
font-size: 28px;
font-weight: 900;
color: var(--slate, #0f172a);
margin: 0;
letter-spacing: -0.02em;
}
.xx-page-head p {
font-size: 15px;
color: var(--muted, #64748b);
margin: 0;
}
.xx-page-head-actions {
display: flex;
gap: 12px;
align-items: center;
}
/* 简化页面头部(无操作按钮) */
.xx-page-head-simple {
background: rgba(255, 255, 255, 0.94);
border: 1px solid rgba(226, 232, 240, 0.95);
border-radius: var(--radius-xl);
box-shadow: 0 24px 70px rgba(15, 23, 42, 0.09);
padding: 32px 40px;
margin-bottom: 32px;
}
.xx-page-head-simple h2 {
font-size: 28px;
font-weight: 900;
color: var(--slate, #0f172a);
margin: 0 0 8px;
letter-spacing: -0.02em;
}
.xx-page-head-simple p {
font-size: 15px;
color: var(--muted, #64748b);
margin: 0;
}
/* V21 卡片 */
.xx-card {
background: rgba(255, 255, 255, 0.94);
border: 1px solid rgba(226, 232, 240, 0.95);
border-radius: var(--radius-lg);
box-shadow: 0 10px 30px rgba(15, 23, 42, 0.06);
padding: 24px;
transition: all 0.3s;
}
.xx-card:hover {
box-shadow: 0 16px 40px rgba(15, 23, 42, 0.08);
}
.xx-card .ant-card-head {
border-bottom: 1px solid rgba(226, 232, 240, 0.8);
padding: 20px 24px;
}
.xx-card .ant-card-head-title {
font-weight: 800;
font-size: 17px;
color: var(--slate, #0f172a);
}
.xx-card .ant-card-body {
padding: 24px;
}
/* 统计卡片网格 - 4列 */
.xx-grid-4 {
display: grid;
grid-template-columns: repeat(4, 1fr);
gap: 24px;
margin-bottom: 32px;
}
@media (max-width: 1200px) {
.xx-grid-4 {
grid-template-columns: repeat(2, 1fr);
}
}
@media (max-width: 768px) {
.xx-grid-4 {
grid-template-columns: 1fr;
}
}
/* 统计卡片 */
.xx-stat-card {
background: rgba(255, 255, 255, 0.94);
border: 1px solid rgba(226, 232, 240, 0.95);
border-radius: var(--radius-lg);
box-shadow: 0 10px 30px rgba(15, 23, 42, 0.06);
padding: 24px;
transition: all 0.3s;
}
.xx-stat-card:hover {
transform: translateY(-2px);
box-shadow: 0 16px 40px rgba(15, 23, 42, 0.08);
}
.xx-stat-card-header {
display: flex;
justify-content: space-between;
align-items: flex-start;
margin-bottom: 16px;
}
.xx-stat-card-icon {
width: 48px;
height: 48px;
border-radius: var(--radius-md);
display: flex;
align-items: center;
justify-content: center;
font-size: 22px;
}
.xx-stat-card-icon.primary {
background: linear-gradient(135deg, #6366f1, #4f46e5);
color: white;
}
.xx-stat-card-icon.success {
background: linear-gradient(135deg, #34d399, #10b981);
color: white;
}
.xx-stat-card-icon.warning {
background: linear-gradient(135deg, #fbbf24, #f59e0b);
color: white;
}
.xx-stat-card-icon.purple {
background: linear-gradient(135deg, #a78bfa, #8b5cf6);
color: white;
}
.xx-stat-card-icon.info {
background: linear-gradient(135deg, #60a5fa, #3b82f6);
color: white;
}
.xx-stat-card-icon.orange {
background: linear-gradient(135deg, #fb923c, #f97316);
color: white;
}
.xx-stat-card-label {
font-size: 14px;
color: var(--muted, #64748b);
font-weight: 500;
margin-bottom: 8px;
}
.xx-stat-card-value {
font-size: 32px;
font-weight: 900;
color: var(--slate, #0f172a);
line-height: 1.2;
letter-spacing: -0.02em;
}
.xx-stat-card-value.primary {
color: var(--indigo, #4f46e5);
}
.xx-stat-card-value.success {
color: var(--green, #10b981);
}
.xx-stat-card-value.warning {
color: var(--amber, #f59e0b);
}
.xx-stat-card-value.purple {
color: #8b5cf6;
}
.xx-stat-card-growth {
font-size: 13px;
color: var(--green, #10b981);
font-weight: 600;
margin-top: 8px;
display: flex;
align-items: center;
gap: 4px;
}
/* 数据表格 */
.xx-table-wrapper {
background: rgba(255, 255, 255, 0.94);
border: 1px solid rgba(226, 232, 240, 0.95);
border-radius: var(--radius-lg);
box-shadow: 0 10px 30px rgba(15, 23, 42, 0.06);
overflow: hidden;
}
.xx-table-wrapper .ant-table {
background: transparent;
}
.xx-table-wrapper .ant-table-thead > tr > th {
background: rgba(248, 250, 252, 0.8);
font-weight: 800;
font-size: 13px;
color: var(--slate, #0f172a);
text-transform: uppercase;
letter-spacing: 0.05em;
padding: 16px 20px;
}
.xx-table-wrapper .ant-table-tbody > tr > td {
padding: 16px 20px;
font-size: 14px;
}
.xx-table-wrapper .ant-table-tbody > tr:hover > td {
background: rgba(79, 70, 229, 0.03);
}
/* V21 按钮 */
.xx-primary-btn {
background: linear-gradient(135deg, #6366f1, #4f46e5) !important;
color: white !important;
border: none !important;
border-radius: var(--radius-md) !important;
font-weight: 700 !important;
box-shadow: 0 8px 20px rgba(79, 70, 229, 0.25) !important;
transition: all 0.2s !important;
}
.xx-primary-btn:hover {
box-shadow: 0 12px 28px rgba(79, 70, 229, 0.3) !important;
transform: translateY(-1px);
}
.xx-ghost-btn {
background: white !important;
border: 1px solid rgba(226, 232, 240, 0.95) !important;
color: var(--slate, #0f172a) !important;
border-radius: var(--radius-md) !important;
font-weight: 600 !important;
transition: all 0.2s !important;
}
.xx-ghost-btn:hover {
border-color: var(--indigo, #4f46e5) !important;
color: var(--indigo, #4f46e5) !important;
}
/* V21 输入框 */
.xx-search-input {
border-radius: var(--radius-md) !important;
border: 1px solid rgba(226, 232, 240, 0.95) !important;
padding: 8px 16px !important;
}
.xx-search-input:hover,
.xx-search-input:focus {
border-color: var(--indigo, #4f46e5) !important;
box-shadow: 0 0 0 3px rgba(79, 70, 229, 0.1) !important;
}
/* V21 Select */
.xx-select {
border-radius: var(--radius-md) !important;
}
.xx-select:hover,
.xx-select:focus {
border-color: var(--indigo, #4f46e5) !important;
box-shadow: 0 0 0 3px rgba(79, 70, 229, 0.1) !important;
}
/* V21 Tag */
.xx-tag {
border-radius: var(--radius-xs) !important;
font-weight: 600 !important;
font-size: 12px !important;
padding: 4px 10px !important;
}
.xx-tag.info {
background: rgba(59, 130, 246, 0.1) !important;
color: #3b82f6 !important;
border: 1px solid rgba(59, 130, 246, 0.2) !important;
}
.xx-tag.success {
background: rgba(16, 185, 129, 0.1) !important;
color: #10b981 !important;
border: 1px solid rgba(16, 185, 129, 0.2) !important;
}
.xx-tag.warning {
background: rgba(245, 158, 11, 0.1) !important;
color: #f59e0b !important;
border: 1px solid rgba(245, 158, 11, 0.2) !important;
}
.xx-tag.error {
background: rgba(239, 68, 68, 0.1) !important;
color: #ef4444 !important;
border: 1px solid rgba(239, 68, 68, 0.2) !important;
}
.xx-tag.debug {
background: rgba(100, 116, 139, 0.1) !important;
color: #64748b !important;
border: 1px solid rgba(100, 116, 139, 0.2) !important;
}
/* V21 Progress */
.xx-progress-primary .ant-progress-circle .ant-progress-text {
color: var(--indigo, #4f46e5) !important;
font-weight: 700 !important;
}
.xx-progress-success .ant-progress-circle .ant-progress-text {
color: var(--green, #10b981) !important;
font-weight: 700 !important;
}
.xx-progress-warning .ant-progress-circle .ant-progress-text {
color: var(--amber, #f59e0b) !important;
font-weight: 700 !important;
}
/* 图表容器 */
.xx-chart-container {
background: rgba(255, 255, 255, 0.94);
border: 1px solid rgba(226, 232, 240, 0.95);
border-radius: var(--radius-lg);
box-shadow: 0 10px 30px rgba(15, 23, 42, 0.06);
padding: 24px;
}
.xx-chart-container .ant-card-head {
border-bottom: 1px solid rgba(226, 232, 240, 0.8);
padding: 20px 24px;
}
.xx-chart-container .ant-card-head-title {
font-weight: 800;
font-size: 17px;
color: var(--slate, #0f172a);
}
/* 筛选器区域 */
.xx-filter-bar {
background: rgba(255, 255, 255, 0.94);
border: 1px solid rgba(226, 232, 240, 0.95);
border-radius: var(--radius-lg);
box-shadow: 0 10px 30px rgba(15, 23, 42, 0.06);
padding: 20px 24px;
margin-bottom: 24px;
display: flex;
flex-wrap: wrap;
gap: 16px;
align-items: center;
}
/* 资源监控卡片 */
.xx-resource-card {
background: rgba(255, 255, 255, 0.94);
border: 1px solid rgba(226, 232, 240, 0.95);
border-radius: var(--radius-lg);
box-shadow: 0 10px 30px rgba(15, 23, 42, 0.06);
padding: 24px;
text-align: center;
}
.xx-resource-card-icon {
width: 64px;
height: 64px;
border-radius: var(--radius-md);
display: flex;
align-items: center;
justify-content: center;
margin: 0 auto 16px;
font-size: 28px;
}
.xx-resource-card-label {
font-size: 14px;
color: var(--muted, #64748b);
font-weight: 600;
margin-bottom: 16px;
}
/* 日期选择器 */
.xx-date-picker {
border-radius: var(--radius-md) !important;
}
/* Drawer */
.xx-drawer .ant-drawer-header {
border-bottom: 1px solid rgba(226, 232, 240, 0.8);
padding: 20px 24px;
}
.xx-drawer .ant-drawer-title {
font-weight: 800;
font-size: 18px;
color: var(--slate, #0f172a);
}
/* 日志详情 */
.xx-log-detail {
padding: 4px 0;
}
.xx-log-detail-label {
font-size: 13px;
font-weight: 700;
color: var(--slate, #0f172a);
margin-bottom: 4px;
}
.xx-log-detail-value {
font-size: 14px;
color: var(--muted, #64748b);
background: rgba(248, 250, 252, 0.8);
padding: 12px;
border-radius: var(--radius-sm);
margin-bottom: 16px;
}
.xx-log-detail-code {
font-family: "JetBrains Mono", "Fira Code", monospace;
background: rgba(248, 250, 252, 0.8);
padding: 12px;
border-radius: var(--radius-sm);
white-space: pre-wrap;
word-break: break-all;
}
/* Result 页面居中 */
.xx-result-center {
display: flex;
justify-content: center;
@@ -25,3 +487,31 @@
.xx-result-center .ant-result {
padding: 48px;
}
/* 刷新时间显示 */
.xx-refresh-time {
font-size: 13px;
color: var(--muted, #64748b);
margin-right: 12px;
}
/* 图表配色覆盖 */
.recharts-text {
fill: #64748b !important;
font-size: 12px !important;
}
.recharts-cartesian-grid-horizontal line,
.recharts-cartesian-grid-vertical line {
stroke: rgba(226, 232, 240, 0.8) !important;
}
/* 图表 Legend */
.recharts-legend-wrapper {
padding-top: 16px !important;
}
.recharts-legend-item-text {
color: #64748b !important;
font-size: 13px !important;
}
+422 -56
View File
@@ -1,85 +1,451 @@
/**
* 控制台页面 — V21 设计系统
* KPI 卡片网格 + 快速入口 + 最近任务卡片列表 + 使用统计图表 + 公告
* CSS 变量,V21 组件
* 使用 mock 数据,CSS 变量,V21 组件
*/
import React from "react";
import { useNavigate } from "react-router-dom";
import { Button, Tag } from "@/components/ui";
import { DatabaseOutlined } from "@ant-design/icons";
import {
VideoCameraOutlined,
AppstoreOutlined,
ThunderboltOutlined,
DatabaseOutlined,
FileTextOutlined,
} from "@ant-design/icons";
import "./dashboard.css";
/* ── 主组件 ─────────────────────────────────────────────── */
/* ============================================================
* Mock 数据
* ============================================================ */
interface KpiItem {
key: string;
icon: string;
iconGradient: string;
value: string;
label: string;
trend: string;
trendDirection: "up" | "down" | "neutral";
accent: string;
}
const kpiData: KpiItem[] = [
{
key: "projects",
icon: "video",
iconGradient: "linear-gradient(135deg, #6366f1, #4f46e5)",
value: "12",
label: "项目总数",
trend: "↑ 2 本月新增",
trendDirection: "up",
accent: "#6366f1",
},
{
key: "assets",
icon: "appstore",
iconGradient: "linear-gradient(135deg, #0ea5e9, #0284c7)",
value: "486",
label: "素材总数",
trend: "↑ 38 本月上传",
trendDirection: "up",
accent: "#0ea5e9",
},
{
key: "generations",
icon: "thunderbolt",
iconGradient: "linear-gradient(135deg, #10b981, #059669)",
value: "156",
label: "本月生成数",
trend: "↑ 23% 较上月",
trendDirection: "up",
accent: "#10b981",
},
{
key: "storage",
icon: "database",
iconGradient: "linear-gradient(135deg, #f59e0b, #d97706)",
value: "2.4GB",
label: "存储空间",
trend: "已用 24%",
trendDirection: "neutral",
accent: "#f59e0b",
},
];
interface QuickEntry {
id: string;
icon: string;
iconGradient: string;
title: string;
description: string;
path: string;
}
const quickEntries: QuickEntry[] = [
{
id: "titles",
icon: "filetext",
iconGradient: "linear-gradient(135deg, #6366f1, #4f46e5)",
title: "标题库",
description: "24条标题 · 5个分类",
path: "/app/titles",
},
{
id: "assets",
icon: "appstore",
iconGradient: "linear-gradient(135deg, #0ea5e9, #0284c7)",
title: "素材库",
description: "486个素材 · 3个素材库",
path: "/app/assets",
},
{
id: "generate",
icon: "thunderbolt",
iconGradient: "linear-gradient(135deg, #10b981, #059669)",
title: "一键生成",
description: "开始创作新视频",
path: "/app/generate",
},
{
id: "products",
icon: "video",
iconGradient: "linear-gradient(135deg, #f59e0b, #d97706)",
title: "成片库",
description: "89个成片 · 3个待复核",
path: "/app/products",
},
];
type TaskStatus = "completed" | "processing" | "pending" | "failed";
interface RecentTask {
id: string;
name: string;
type: string;
template: string;
status: TaskStatus;
date: string;
duration?: string;
}
const statusLabel: Record<TaskStatus, string> = {
completed: "已完成",
processing: "进行中",
pending: "排队中",
failed: "失败",
};
const recentTasks: RecentTask[] = [
{
id: "t-1",
name: "产品介绍视频_春季促销",
type: "视频生成",
template: "商品展示模板",
status: "completed",
date: "2026-07-01 09:30",
duration: "2分18秒",
},
{
id: "t-2",
name: "品牌宣传片_终版",
type: "视频生成",
template: "品牌宣传模板",
status: "processing",
date: "2026-07-01 10:15",
},
{
id: "t-3",
name: "用户评价合集",
type: "视频生成",
template: "评价展示模板",
status: "completed",
date: "2026-06-30 16:42",
duration: "1分45秒",
},
{
id: "t-4",
name: "新品发布预告",
type: "视频生成",
template: "新品预告模板",
status: "pending",
date: "2026-06-30 14:20",
},
{
id: "t-5",
name: "活动回顾_618大促",
type: "视频生成",
template: "活动回顾模板",
status: "failed",
date: "2026-06-29 11:05",
},
];
interface ChartItem {
label: string;
value: number;
}
const weeklyData: ChartItem[] = [
{ label: "周一", value: 18 },
{ label: "周二", value: 25 },
{ label: "周三", value: 32 },
{ label: "周四", value: 28 },
{ label: "周五", value: 42 },
{ label: "周六", value: 15 },
{ label: "周日", value: 8 },
];
interface Announcement {
id: string;
tag: "update" | "notice" | "activity";
tagLabel: string;
title: string;
date: string;
}
const announcements: Announcement[] = [
{
id: "a-1",
tag: "update",
tagLabel: "更新",
title: "系统已升级至 v2.0,新增批量生成功能",
date: "2026-07-01",
},
{
id: "a-2",
tag: "activity",
tagLabel: "活动",
title: "7月创作挑战赛已开启,参与赢积分奖励",
date: "2026-06-28",
},
{
id: "a-3",
tag: "notice",
tagLabel: "公告",
title: "7月3日凌晨 2:00-4:00 系统维护通知",
date: "2026-06-25",
},
];
/* ============================================================
* 工具函数
* ============================================================ */
const getGreeting = () => {
const hour = new Date().getHours();
if (hour < 6) return "夜深了";
if (hour < 12) return "早上好";
if (hour < 14) return "中午好";
if (hour < 18) return "下午好";
return "晚上好";
};
const formatDate = () => {
const d = new Date();
const weekDays = ["日", "一", "二", "三", "四", "五", "六"];
return `${d.getFullYear()}年${d.getMonth() + 1}月${d.getDate()}日 星期${weekDays[d.getDay()]}`;
};
/** 图标名称 → Ant Design 组件映射 */
const iconMap: Record<string, React.ReactNode> = {
video: <VideoCameraOutlined />,
appstore: <AppstoreOutlined />,
thunderbolt: <ThunderboltOutlined />,
database: <DatabaseOutlined />,
filetext: <FileTextOutlined />,
};
/* ============================================================
* 组件
* ============================================================ */
const Dashboard: React.FC = () => {
const navigate = useNavigate();
const maxChart = Math.max(...weeklyData.map((d) => d.value));
return (
<div className="xx-dashboard-page">
{/* KPI 卡片网格 */}
{/* ── 欢迎头部 ─────────────────────────────────────────── */}
<div className="xx-dashboard-welcome">
<h2>{getGreeting()},创作者</h2>
<p>{formatDate()} — 欢迎回到小小剪辑控制台</p>
</div>
{/* ── KPI 卡片网格 ─────────────────────────────────────── */}
<div className="xx-kpi-grid">
<div className="xx-dashboard-empty">
<p>暂无统计数据</p>
{kpiData.map((item) => (
<div
key={item.key}
className="xx-kpi-card"
style={{ "--kpi-accent": item.accent } as React.CSSProperties}
>
<div
className="xx-kpi-icon"
style={{ background: item.iconGradient }}
>
{iconMap[item.icon] ?? item.icon}
</div>
<div className="xx-kpi-value">{item.value}</div>
<div className="xx-kpi-label">{item.label}</div>
<span
className={`xx-kpi-trend xx-kpi-trend--${item.trendDirection}`}
>
{item.trend}
</span>
</div>
))}
</div>
{/* ── 主内容区:左侧任务+图表 / 右侧公告 ──────────────── */}
<div className="xx-dashboard-main">
{/* 左列 */}
<div className="xx-dashboard-left-col">
{/* 最近任务 */}
<div className="xx-dashboard-section">
<div className="xx-dashboard-section-header">
<h3>最近任务</h3>
<button onClick={() => navigate("/app/history")}>查看全部</button>
</div>
<div className="xx-task-list">
{recentTasks.map((task) => (
<div key={task.id} className="xx-task-item">
<div className="xx-task-info">
<h4>{task.name}</h4>
<span>
{task.type} · 模板:{task.template}
</span>
</div>
<Tag
variant={
task.status === "completed"
? "success"
: task.status === "processing"
? "info"
: task.status === "failed"
? "error"
: "warning"
}
>
{statusLabel[task.status]}
</Tag>
<div className="xx-task-time">
<span>{task.date}</span>
{task.status === "completed"
? `耗时 ${task.duration}`
: task.status === "processing"
? "生成中..."
: task.status === "failed"
? "请重试"
: "等待中"}
</div>
<div className="xx-task-action">
<Button
buttonType="ghost"
buttonSize="sm"
onClick={() => navigate("/app/history")}
>
查看
</Button>
</div>
</div>
))}
</div>
</div>
{/* 使用统计图表 */}
<div className="xx-dashboard-section">
<div className="xx-dashboard-section-header">
<h3>本周生成趋势</h3>
<span className="xx-chart-total">
共 {weeklyData.reduce((s, d) => s + d.value, 0)} 次
</span>
</div>
<div className="xx-chart-container">
<div className="xx-chart-bars">
{weeklyData.map((d, i) => (
<div key={i} className="xx-chart-bar-wrapper">
<div
className="xx-chart-bar"
style={{
height: `${(d.value / maxChart) * 100}%`,
}}
>
<span className="xx-chart-bar-value">{d.value}</span>
</div>
</div>
))}
</div>
<div className="xx-chart-labels">
{weeklyData.map((d, i) => (
<div key={i} className="xx-chart-label">
{d.label}
</div>
))}
</div>
</div>
</div>
</div>
{/* 右列 — 公告 + 存储用量 */}
<div className="xx-dashboard-section xx-dashboard-section--start">
<div className="xx-dashboard-section-header">
<h3>系统公告</h3>
</div>
<div className="xx-announcement-list">
{announcements.map((a) => (
<div key={a.id} className="xx-announcement-item">
<span
className={`xx-announcement-tag xx-announcement-tag--${a.tag}`}
>
{a.tagLabel}
</span>
<div className="xx-announcement-content">
<h4>{a.title}</h4>
<time>{a.date}</time>
</div>
</div>
))}
</div>
{/* 存储用量 */}
<div className="xx-storage-section">
<div className="xx-storage-section-title">存储用量</div>
<div className="xx-storage-bar">
<div className="xx-storage-bar-track">
<div className="xx-storage-bar-fill" style={{ width: "24%" }} />
</div>
<div className="xx-storage-bar-label">
<span>2.4 GB 已用</span>
<span>10 GB 总量</span>
</div>
</div>
</div>
</div>
</div>
{/* 快速入口 */}
{/* ── 快速入口 ─────────────────────────────────────────── */}
<div className="xx-quick-entry-section">
<div className="xx-quick-entry-header">
<h3 className="xx-quick-entry-title">快速入口</h3>
</div>
<div className="xx-quick-grid">
<div className="xx-dashboard-empty">
<p>暂无快速入口</p>
</div>
{quickEntries.map((entry) => (
<div
key={entry.id}
className="xx-quick-card"
onClick={() => navigate(entry.path)}
>
<div
className="xx-quick-card-icon"
style={{ background: entry.iconGradient }}
>
{iconMap[entry.icon] ?? entry.icon}
</div>
<h3>{entry.title}</h3>
<p>{entry.description}</p>
</div>
))}
</div>
</div>
{/* 最近任务 */}
<section className="xx-dashboard-section">
<div className="xx-dashboard-section-header">
<h3>最近任务</h3>
<Button buttonType="ghost" buttonSize="sm" onClick={() => navigate("/app/history")}>
查看全部
</Button>
</div>
<div className="xx-task-list">
<div className="xx-dashboard-empty">
<p>暂无最近任务</p>
</div>
</div>
</section>
{/* 使用统计 */}
<section className="xx-dashboard-section" style={{ marginTop: "var(--space-md)" }}>
<div className="xx-dashboard-section-header">
<h3>使用统计</h3>
</div>
<div className="xx-chart-container">
<div className="xx-chart-bars">
<div className="xx-dashboard-empty" style={{ width: "100%" }}>
<DatabaseOutlined style={{ fontSize: 24, marginBottom: 8 }} />
<p>暂无统计数据</p>
</div>
</div>
</div>
</section>
{/* 公告 */}
<section className="xx-dashboard-section" style={{ marginTop: "var(--space-md)" }}>
<div className="xx-dashboard-section-header">
<h3>公告</h3>
</div>
<div className="xx-announcement-list">
<div className="xx-announcement-item">
<span className="xx-announcement-tag xx-announcement-tag--notice">官方</span>
<div className="xx-announcement-content">
<h4>欢迎使用小应 SaaS 平台</h4>
<time>当前为演示版本,部分功能正在开发中。</time>
</div>
</div>
</div>
</section>
</div>
);
};
+24 -12
View File
@@ -69,6 +69,22 @@ const VOICE_GENDER_ICON: Record<string, string> = {
neutral: "✨",
};
/* ── 时间线 Mock ──
* TODO: 后端暂无时间线场景数据 API,当前使用硬编码预览数据。
* 待后端提供 timeline/scene 接口后替换为真实 API 调用。
*/
interface TimelineScene {
scene: string;
time: string;
duration: number;
}
const MOCK_TIMELINE: TimelineScene[] = [
{ scene: "主讲口播 · 开场钩子", time: "0-8s", duration: 8 },
{ scene: "产品特写 · B-roll", time: "8-20s", duration: 12 },
{ scene: "用户反馈 · 结尾", time: "20-30s", duration: 10 },
];
/* ── 步骤定义 ── */
const STEPS = [
{ key: 1, label: "选择模板" },
@@ -1742,19 +1758,15 @@ const GeneratePage: React.FC = () => {
<div className="xx-preview-title">剪辑计划预览</div>
{/* 时间线列表 */}
{generated ? (
<div className="xx-preview-timeline">
<div className="xx-timeline-item">
<span className="scene-name">生成完成,可下载或分享视频</span>
<div className="xx-preview-timeline">
{MOCK_TIMELINE.map((item, idx) => (
<div key={idx} className="xx-timeline-item">
<div className="num">{idx + 1}</div>
<span className="scene-name">{item.scene}</span>
<span>{item.time}</span>
</div>
</div>
) : (
<div className="xx-preview-timeline">
<div className="xx-timeline-item">
<span className="scene-name">确认标题后自动生成剪辑计划</span>
</div>
</div>
)}
))}
</div>
{/* 生成操作按钮 */}
<div className="xx-generate-actions">
+48
View File
@@ -612,6 +612,54 @@
flex: 1;
}
/* ============================================================
按钮(匹配原型 .btn .ghost / .btn .primary)
============================================================ */
.xx-btn {
display: inline-flex;
align-items: center;
justify-content: center;
gap: 6px;
height: 42px;
padding: 0 20px;
border-radius: var(--radius-sm);
font-size: 14px;
font-weight: 600;
cursor: pointer;
transition: all 0.15s ease;
border: none;
outline: none;
white-space: nowrap;
}
.xx-btn:disabled {
opacity: 0.5;
cursor: not-allowed;
}
.xx-btn-primary {
background: var(--gradient-primary);
color: var(--text-inverse);
box-shadow: 0 14px 26px rgba(79, 70, 229, 0.22);
}
.xx-btn-primary:hover:not(:disabled) {
transform: translateY(-2px);
box-shadow: 0 18px 34px rgba(79, 70, 229, 0.28);
}
.xx-btn-ghost {
background: var(--bg-primary);
border: 1px solid var(--border-color);
color: var(--text-secondary);
}
.xx-btn-ghost:hover:not(:disabled) {
border-color: var(--info-border);
color: var(--primary-dark);
background: var(--primary-soft);
}
/* ============================================================
右侧预览区 generate-preview
============================================================ */
+118 -27
View File
@@ -36,24 +36,40 @@ type TitleType = "hot" | "normal" | "creative";
type Industry = "general" | "food" | "tech" | "beauty" | "education" | "travel";
type Frequency = "all" | "high" | "medium" | "low";
interface CategoryItem {
id: string;
name: string;
count: number;
}
interface TitleData {
id: string;
content: string;
type: TitleType;
industry: Industry;
category: string;
usageCount: number;
isFavorited: boolean;
createdAt: string;
}
/* ============================================================
* Mock 数据
* ============================================================ */
const MOCK_CATEGORIES: CategoryItem[] = [
{ id: "cat-all", name: "全部标题", count: 15 },
{ id: "cat-1", name: "美食探店", count: 4 },
{ id: "cat-2", name: "科技数码", count: 3 },
{ id: "cat-3", name: "生活日常", count: 4 },
{ id: "cat-4", name: "美妆穿搭", count: 2 },
{ id: "cat-5", name: "教育学习", count: 2 },
];
/** 后端 TitleItem → 前端 TitleData 映射 */
const toTitleData = (item: TitleItem): TitleData => ({
id: item.id,
content: item.content,
type: (item.category as TitleType) || "normal",
industry: "general",
category: item.category || "未分类",
usageCount: 0,
isFavorited: false,
createdAt: item.created_at?.slice(0, 10) || "",
@@ -234,8 +250,9 @@ const TitleCard: React.FC<{
const TitleLibrary: React.FC = () => {
const queryClient = useQueryClient();
/* 分类数据 — 从真实标题数据动态派生 */
const [activeCatId, setActiveCatId] = useState<string>("cat-all");
/* 分类数据 */
const [categories, setCategories] = useState<CategoryItem[]>(MOCK_CATEGORIES);
const [activeCatId, setActiveCatId] = useState<string>(MOCK_CATEGORIES[0].id);
/* 标题数据 — 真实 API */
const { data: apiTitles = [] } = useQuery({
@@ -248,23 +265,6 @@ const TitleLibrary: React.FC = () => {
[apiTitles],
);
/* 从真实标题数据动态派生分类(无需后端分类 API) */
const categories = useMemo(() => {
const cats = new Map<string, number>();
apiTitles.forEach((t) => {
const cat = t.category || "未分类";
cats.set(cat, (cats.get(cat) || 0) + 1);
});
return [
{ id: "cat-all", name: "全部标题", count: apiTitles.length },
...Array.from(cats.entries()).map(([name, count]) => ({
id: `cat-${name}`,
name,
count,
})),
];
}, [apiTitles]);
/* CRUD mutations */
const createMutation = useMutation({
mutationFn: (content: string) => createTitle({ content }),
@@ -301,6 +301,10 @@ const TitleLibrary: React.FC = () => {
const [editingId, setEditingId] = useState<string | null>(null);
const [editText, setEditText] = useState("");
/* 新建分类 */
const [createCatModalOpen, setCreateCatModalOpen] = useState(false);
const [newCatName, setNewCatName] = useState("");
/* 新建标题 */
const [createTitleModalOpen, setCreateTitleModalOpen] = useState(false);
const [newTitleContent, setNewTitleContent] = useState("");
@@ -318,11 +322,19 @@ const TitleLibrary: React.FC = () => {
const filteredTitles = useMemo(() => {
let list = titles;
/* 按分类过滤("全部标题" 不过滤)— 直接匹配后端 category 字段 */
/* 按分类过滤("全部标题" 不过滤) */
if (activeCatId !== "cat-all") {
const catName = activeCategory?.name || "";
if (catName) {
list = list.filter((t) => t.category === catName);
const catToIndustry: Record<string, Industry> = {
美食探店: "food",
科技数码: "tech",
生活日常: "general",
美妆穿搭: "beauty",
教育学习: "education",
};
const mappedIndustry = catToIndustry[catName];
if (mappedIndustry) {
list = list.filter((t) => t.industry === mappedIndustry);
}
}
@@ -416,6 +428,32 @@ const TitleLibrary: React.FC = () => {
[deleteMutation],
);
/* 新建分类 */
const handleCreateCategory = () => {
if (!newCatName.trim()) {
message.warning("请输入分类名称");
return;
}
const cat: CategoryItem = {
id: `cat-${Date.now()}`,
name: newCatName.trim(),
count: 0,
};
setCategories((prev) => [...prev, cat]);
setActiveCatId(cat.id);
setCreateCatModalOpen(false);
setNewCatName("");
message.success(`分类 "${cat.name}" 创建成功`);
};
/* 删除分类 */
const handleDeleteCategory = (id: string) => {
setCategories((prev) => prev.filter((c) => c.id !== id));
if (activeCatId === id) {
setActiveCatId("cat-all");
}
message.success("分类已删除");
};
/* 新建标题 */
const handleCreateTitle = () => {
@@ -503,12 +541,38 @@ const TitleLibrary: React.FC = () => {
</h4>
<span>{cat.count} 条</span>
</div>
{cat.id !== "cat-all" && (
<Popconfirm
title={`确定删除分类 "${cat.name}"?`}
onConfirm={(e) => {
e?.stopPropagation();
handleDeleteCategory(cat.id);
}}
onCancel={(e) => e?.stopPropagation()}
okText="删除"
cancelText="取消"
>
<button
className="xx-title-category-delete"
onClick={(e) => e.stopPropagation()}
title="删除分类"
>
<DeleteOutlined />
</button>
</Popconfirm>
)}
</div>
</div>
))}
{/* TODO: 新建分类功能待后端分类 API 就绪后启用 */}
{/* 新建分类 */}
<div
className="xx-title-category-add"
onClick={() => setCreateCatModalOpen(true)}
>
<PlusOutlined />
新建分类
</div>
</div>
{/* ─── 右侧:内容区 ─── */}
@@ -614,7 +678,34 @@ const TitleLibrary: React.FC = () => {
</div>
</div>
{/* ─── 新建分类弹窗 ─── */}
<AntModal
title="新建分类"
open={createCatModalOpen}
onCancel={() => setCreateCatModalOpen(false)}
onOk={handleCreateCategory}
okText="创建"
cancelText="取消"
destroyOnClose
>
<div style={{ padding: "8px 0" }}>
<div
style={{
marginBottom: 6,
fontSize: "var(--font-size-sm)",
color: "var(--text-secondary)",
}}
>
分类名称
</div>
<Input
placeholder="请输入分类名称"
value={newCatName}
onChange={(e) => setNewCatName(e.target.value)}
maxLength={30}
/>
</div>
</AntModal>
{/* ─── 新建标题弹窗 ─── */}
<AntModal
@@ -18,7 +18,7 @@ import {
CloseCircleOutlined,
} from "@ant-design/icons";
import PageHead from "@/components/layout/PageHead";
import CloneModal from "@/components/voice/CloneModal";
import CloneVoiceModal from "@/components/modals/CloneVoiceModal";
import {
getVoiceClones,
deleteVoiceClone,
@@ -356,7 +356,7 @@ const VoiceClone: React.FC = () => {
)}
{/* 克隆音色弹窗 */}
<CloneModal
<CloneVoiceModal
open={cloneModalOpen}
onClose={() => setCloneModalOpen(false)}
onSuccess={() => {
+7
View File
@@ -170,6 +170,13 @@ export const router = createBrowserRouter([
Component: m.default,
})),
},
{
path: "my-voices",
lazy: () =>
import("@/pages/my-voices/MyVoices").then((m) => ({
Component: m.default,
})),
},
{
path: "accounts",
lazy: () =>
+7 -6
View File
@@ -1,17 +1,18 @@
"""
视频处理模块
轻量工具(ffmpeg_utils / oss_helpers / dedup_helpers)顶层直接导出,
无额外依赖。渲染相关组件(UnifiedRenderService / RenderAdapter /
VideoProcessor 等)按需从子模块导入,避免 __init__ 阶段引入
packages / DB 等重依赖。
"""
# 共享工具模块(零外部依赖,供 editing_modes / generation / edit_plan_generation 等复用)
# 共享工具模块(供 editing_modes / generation / edit_plan_generation 等复用)
from . import dedup_helpers, ffmpeg_utils, oss_helpers
from .processor import VideoProcessor, VideoResult
from .unified_render_service import RenderResult, UnifiedRenderService
__all__ = [
"VideoProcessor",
"VideoResult",
"ffmpeg_utils",
"oss_helpers",
"dedup_helpers",
"UnifiedRenderService",
"RenderResult",
]
+3 -7
View File
@@ -93,16 +93,12 @@ class VideoFingerprint:
resolution: tuple[int, int]
def to_dict(self) -> dict:
# 注意:color_histograms 里的值可能是 np.float32(来自 cv2.normalize),
# 直接存进 dict 后 SQLAlchemy JSON 序列化会报 "float32 is not JSON serializable"。
# 这里统一转成 Python 原生 float。
native_histograms = [[float(v) for v in hist] for hist in self.color_histograms]
return {
"md5": self.md5,
"keyframe_phashes": self.keyframe_phashes,
"color_histograms": native_histograms,
"duration": float(self.duration),
"resolution": [int(self.resolution[0]), int(self.resolution[1])],
"color_histograms": self.color_histograms,
"duration": self.duration,
"resolution": list(self.resolution),
}
@@ -0,0 +1,657 @@
"""
视频剪辑模式处理器
支持四种剪辑模式:一镜到底、画中画、口播、口播+画中画
"""
import logging
import os
import sys
import tempfile
from dataclasses import dataclass
if sys.version_info >= (3, 11):
from enum import StrEnum
else:
from enum import Enum
class StrEnum(str, Enum):
pass
from pathlib import Path
from typing import Optional
from video_processing.ffmpeg_utils import FFMPEG_BIN, probe_video_info, run_ffmpeg
logger = logging.getLogger(__name__)
# 从 domain 层导入 EditingMode,避免重复定义
from packages.domain.editing_mode import EditingMode
class PIPPosition(StrEnum):
"""画中画位置枚举"""
TOP_LEFT = "top_left"
TOP_RIGHT = "top_right"
BOTTOM_LEFT = "bottom_left"
BOTTOM_RIGHT = "bottom_right"
@dataclass
class EditingModeConfig:
"""剪辑模式配置"""
mode: EditingMode
output_width: int = 1280
output_height: int = 720
output_fps: int = 25
pip_position: PIPPosition = PIPPosition.TOP_RIGHT
pip_scale: float = 0.25 # 画中画占主画面的比例
transition_duration: float = 0.5 # 转场时长(秒)
output_codec: str = "libx264"
output_preset: str = "medium"
output_crf: int = 23
class EditingModeProcessor:
"""剪辑模式处理器"""
def __init__(self, config: EditingModeConfig, work_dir: Optional[str] = None):
"""
初始化剪辑模式处理器
Args:
config: 剪辑模式配置
work_dir: 工作目录,默认使用系统临时目录
"""
self.config = config
self.work_dir = work_dir or tempfile.gettempdir()
def process(
self,
video_paths: list[str],
audio_path: Optional[str] = None,
output_path: Optional[str] = None,
) -> str:
"""
根据模式处理视频,返回输出文件路径
Args:
video_paths: 视频素材路径列表
audio_path: 音频路径(用于口播模式)
output_path: 输出文件路径,默认自动生成
Returns:
输出文件路径
"""
if not video_paths:
raise ValueError("video_paths cannot be empty")
self._validate_inputs(video_paths, audio_path)
if output_path is None:
output_path = self._generate_output_path()
logger.info(f"Processing videos with mode: {self.config.mode}, count: {len(video_paths)}")
try:
if self.config.mode == EditingMode.ONE_TAKE:
return self._one_take(video_paths, output_path)
elif self.config.mode == EditingMode.PIP:
return self._pip(video_paths, output_path)
elif self.config.mode == EditingMode.VOICE_OVER:
return self._voice_over(video_paths, audio_path, output_path)
elif self.config.mode == EditingMode.VOICE_PIP:
return self._voice_pip(video_paths, audio_path, output_path)
else:
raise ValueError(f"Unsupported editing mode: {self.config.mode}")
except Exception as e:
logger.error(f"Error processing videos: {e}")
raise
def _validate_inputs(self, video_paths: list[str], audio_path: Optional[str]) -> None:
"""验证输入文件"""
for path in video_paths:
if not os.path.exists(path):
raise FileNotFoundError(f"Video file not found: {path}")
if not os.path.getsize(path) > 0:
raise ValueError(f"Video file is empty: {path}")
if audio_path and not os.path.exists(audio_path):
raise FileNotFoundError(f"Audio file not found: {audio_path}")
def _generate_output_path(self) -> str:
"""生成输出文件路径"""
os.makedirs(self.work_dir, exist_ok=True)
return os.path.join(self.work_dir, f"output_{self.config.mode}_{os.getpid()}.mp4")
def _run_ffmpeg(self, command: list[str], capture_output: bool = True) -> tuple:
"""执行 FFmpeg 命令 — 委托给共享 ffmpeg_utils.run_ffmpeg"""
try:
return run_ffmpeg(command, capture_output=capture_output)
except RuntimeError as e:
logger.error(f"FFmpeg error: {e}")
raise
def _get_video_info(self, video_path: str) -> dict:
"""获取视频信息 — 委托给共享 ffmpeg_utils.probe_video_info,补充 codec/size 字段"""
try:
info = probe_video_info(video_path)
info["codec"] = "unknown"
info["size"] = os.path.getsize(video_path) if os.path.exists(video_path) else 0
return info
except Exception as e:
logger.warning(f"Failed to get video info for {video_path}: {e}")
return {"width": 0, "height": 0, "fps": 25, "duration": 0, "codec": "unknown", "size": 0}
def _get_pip_position_offset(
self, main_width: int, main_height: int, pip_width: int, pip_height: int
) -> tuple[int, int]:
"""获取画中画位置偏移量"""
margin = 10
position_offsets = {
PIPPosition.TOP_LEFT: (margin, margin),
PIPPosition.TOP_RIGHT: (main_width - pip_width - margin, margin),
PIPPosition.BOTTOM_LEFT: (margin, main_height - pip_height - margin),
PIPPosition.BOTTOM_RIGHT: (main_width - pip_width - margin, main_height - pip_height - margin),
}
return position_offsets.get(self.config.pip_position, position_offsets[PIPPosition.TOP_RIGHT])
def _normalize_video(self, input_path: str, output_path: str) -> dict:
"""标准化视频格式:先统一帧率,再缩放/填充"""
command = [
FFMPEG_BIN,
"-y",
"-i",
input_path,
"-r",
str(self.config.output_fps), # 先统一帧率
"-vf",
f"scale={self.config.output_width}:{self.config.output_height}:force_original_aspect_ratio=decrease,pad={self.config.output_width}:{self.config.output_height}:(ow-iw)/2:(oh-ih)/2,setsar=1",
"-r",
str(self.config.output_fps),
"-c:v",
self.config.output_codec,
"-preset",
self.config.output_preset,
"-crf",
str(self.config.output_crf),
"-pix_fmt",
"yuv420p",
"-movflags",
"+faststart",
"-an",
output_path,
]
run_ffmpeg(command)
return self._get_video_info(output_path)
def _one_take(self, video_paths: list[str], output_path: str) -> str:
"""一镜到底模式:顺序拼接视频,添加淡入淡出转场"""
if len(video_paths) == 1:
return self._normalize_video(video_paths[0], output_path)
normalized_paths = []
for i, path in enumerate(video_paths):
normalized = os.path.join(self.work_dir, f"normalized_{i}_{os.getpid()}.mp4")
self._normalize_video(path, normalized)
normalized_paths.append(normalized)
durations = [self._get_video_info(p)["duration"] for p in normalized_paths]
if len(normalized_paths) <= 5:
output_path = self._one_take_with_xfade(normalized_paths, durations, output_path)
else:
output_path = self._one_take_simple_concat(normalized_paths, output_path)
for p in normalized_paths:
try:
if p != output_path:
os.remove(p)
except Exception as e:
logger.warning(f"Operation failed in apps/worker/video_processing/editing_modes.py: {e}", exc_info=True)
return output_path
def _one_take_with_xfade(self, normalized_paths: list[str], durations: list[float], output_path: str) -> str:
"""使用 xfade 滤镜实现转场"""
if len(normalized_paths) == 2:
transition = self.config.transition_duration
offset1 = durations[0] - transition / 2
command = [
FFMPEG_BIN,
"-y",
"-i",
normalized_paths[0],
"-i",
normalized_paths[1],
"-filter_complex",
f"[0:v][1:v]xfade=transition=fade:duration={transition}:offset={offset1}[v]",
"-map",
"[v]",
"-c:v",
self.config.output_codec,
"-preset",
self.config.output_preset,
"-crf",
str(self.config.output_crf),
"-pix_fmt",
"yuv420p",
output_path,
]
run_ffmpeg(command)
return output_path
else:
return self._one_take_simple_concat(normalized_paths, output_path)
def _one_take_simple_concat(self, normalized_paths: list[str], output_path: str) -> str:
"""使用 concat demuxer 简单拼接"""
concat_file = os.path.join(self.work_dir, f"concat_list_{os.getpid()}.txt")
with open(concat_file, "w") as f:
for path in normalized_paths:
f.write(f"file '{os.path.abspath(path)}'\n")
command = [
FFMPEG_BIN,
"-y",
"-f",
"concat",
"-safe",
"0",
"-i",
concat_file,
"-c",
"copy",
output_path,
]
run_ffmpeg(command)
try:
os.remove(concat_file)
except Exception as e:
logger.warning(f"Operation failed in apps/worker/video_processing/editing_modes.py: {e}", exc_info=True)
return output_path
def _pip(self, video_paths: list[str], output_path: str) -> str:
"""画中画模式:主视频全屏,后续视频叠加在角落"""
if not video_paths:
raise ValueError("No video paths provided")
main_video = video_paths[0]
main_normalized = os.path.join(self.work_dir, f"main_{os.getpid()}.mp4")
main_info = self._normalize_video(main_video, main_normalized)
if len(video_paths) == 1:
os.rename(main_normalized, output_path)
return output_path
pip_width = int(self.config.output_width * self.config.pip_scale)
pip_height = int(self.config.output_height * self.config.pip_scale)
x_offset, y_offset = self._get_pip_position_offset(
self.config.output_width, self.config.output_height, pip_width, pip_height
)
pip_normalized = os.path.join(self.work_dir, f"pip_{os.getpid()}.mp4")
pip_info = self._get_video_info(video_paths[1])
if pip_info["duration"] > main_info["duration"]:
temp_pip = os.path.join(self.work_dir, f"pip_temp_{os.getpid()}.mp4")
command = [
FFMPEG_BIN,
"-y",
"-i",
video_paths[1],
"-t",
str(main_info["duration"]),
"-vf",
f"scale={pip_width}:{pip_height}",
"-c:v",
self.config.output_codec,
"-preset",
self.config.output_preset,
"-crf",
str(self.config.output_crf),
"-pix_fmt",
"yuv420p",
temp_pip,
]
run_ffmpeg(command)
pip_normalized_input = temp_pip
else:
command = [
FFMPEG_BIN,
"-y",
"-i",
video_paths[1],
"-vf",
f"scale={pip_width}:{pip_height}",
"-c:v",
self.config.output_codec,
"-preset",
self.config.output_preset,
"-crf",
str(self.config.output_crf),
"-pix_fmt",
"yuv420p",
pip_normalized,
]
run_ffmpeg(command)
pip_normalized_input = pip_normalized
if main_info["duration"] > pip_info["duration"]:
looped_pip = os.path.join(self.work_dir, f"pip_looped_{os.getpid()}.mp4")
command = [
FFMPEG_BIN,
"-y",
"-stream_loop",
"-1",
"-i",
pip_normalized_input,
"-t",
str(main_info["duration"]),
"-vf",
f"scale={pip_width}:{pip_height}",
"-c:v",
self.config.output_codec,
"-preset",
self.config.output_preset,
"-crf",
str(self.config.output_crf),
"-pix_fmt",
"yuv420p",
looped_pip,
]
run_ffmpeg(command)
pip_normalized_input = looped_pip
command = [
FFMPEG_BIN,
"-y",
"-i",
main_normalized,
"-i",
pip_normalized_input,
"-filter_complex",
f"[0:v][1:v]overlay={x_offset}:{y_offset}[v]",
"-map",
"[v]",
"-c:v",
self.config.output_codec,
"-preset",
self.config.output_preset,
"-crf",
str(self.config.output_crf),
"-pix_fmt",
"yuv420p",
output_path,
]
run_ffmpeg(command)
for temp_file in [main_normalized, pip_normalized]:
if temp_file and temp_file != output_path:
try:
os.remove(temp_file)
except Exception as e:
logger.warning(
f"Operation failed in apps/worker/video_processing/editing_modes.py: {e}", exc_info=True
)
return output_path
def _voice_over(self, video_paths: list[str], audio_path: Optional[str], output_path: str) -> str:
"""口播模式:背景画面 + 配音"""
if not audio_path:
raise ValueError("audio_path is required for VOICE_OVER mode")
if not video_paths:
raise ValueError("No background video provided")
audio_info = self._get_video_info(audio_path)
audio_duration = audio_info["duration"]
bg_normalized = os.path.join(self.work_dir, f"bg_{os.getpid()}.mp4")
bg_info = self._normalize_video(video_paths[0], bg_normalized)
if bg_info["duration"] < audio_duration:
looped_bg = os.path.join(self.work_dir, f"bg_looped_{os.getpid()}.mp4")
command = [
FFMPEG_BIN,
"-y",
"-stream_loop",
"-1",
"-i",
bg_normalized,
"-t",
str(audio_duration),
"-vf",
f"scale={self.config.output_width}:{self.config.output_height}",
"-c:v",
self.config.output_codec,
"-preset",
self.config.output_preset,
"-crf",
str(self.config.output_crf),
"-pix_fmt",
"yuv420p",
looped_bg,
]
run_ffmpeg(command)
bg_normalized = looped_bg
elif bg_info["duration"] > audio_duration:
temp_bg = os.path.join(self.work_dir, f"bg_trimmed_{os.getpid()}.mp4")
command = [
FFMPEG_BIN,
"-y",
"-i",
bg_normalized,
"-t",
str(audio_duration),
"-c:v",
"copy",
temp_bg,
]
run_ffmpeg(command)
bg_normalized = temp_bg
blurred_bg = os.path.join(self.work_dir, f"bg_blurred_{os.getpid()}.mp4")
command = [
FFMPEG_BIN,
"-y",
"-i",
bg_normalized,
"-vf",
f"boxblur=5:5,scale={self.config.output_width}:{self.config.output_height}",
"-c:v",
self.config.output_codec,
"-preset",
self.config.output_preset,
"-crf",
str(self.config.output_crf),
"-pix_fmt",
"yuv420p",
blurred_bg,
]
run_ffmpeg(command)
command = [
FFMPEG_BIN,
"-y",
"-i",
blurred_bg,
"-i",
audio_path,
"-filter_complex",
"[0:v]drawbox=x=0:y=0:w=iw:h=ih:color=black@0.3:t=fill[v]",
"-map",
"[v]",
"-map",
"1:a",
"-c:v",
self.config.output_codec,
"-preset",
self.config.output_preset,
"-crf",
str(self.config.output_crf),
"-pix_fmt",
"yuv420p",
"-shortest",
output_path,
]
run_ffmpeg(command)
for temp_file in [bg_normalized, blurred_bg]:
try:
if temp_file != output_path:
os.remove(temp_file)
except Exception as e:
logger.warning(f"Operation failed in apps/worker/video_processing/editing_modes.py: {e}", exc_info=True)
return output_path
def _voice_pip(self, video_paths: list[str], audio_path: Optional[str], output_path: str) -> str:
"""口播+画中画模式:口播视频在角落,其他视频作为背景"""
if not video_paths:
raise ValueError("No video paths provided")
if len(video_paths) == 1:
return self._normalize_video(video_paths[0], output_path)
voice_video = video_paths[0]
bg_video = video_paths[1] if len(video_paths) > 1 else video_paths[0]
voice_normalized = os.path.join(self.work_dir, f"voice_{os.getpid()}.mp4")
voice_info = self._normalize_video(voice_video, voice_normalized)
bg_normalized = os.path.join(self.work_dir, f"bg_{os.getpid()}.mp4")
bg_info = self._normalize_video(bg_video, bg_normalized)
final_duration = min(voice_info["duration"], bg_info["duration"])
pip_width = int(self.config.output_width * self.config.pip_scale)
pip_height = int(self.config.output_height * self.config.pip_scale)
x_offset, y_offset = self._get_pip_position_offset(
self.config.output_width, self.config.output_height, pip_width, pip_height
)
voice_adjusted = os.path.join(self.work_dir, f"voice_adj_{os.getpid()}.mp4")
command = [
FFMPEG_BIN,
"-y",
"-i",
voice_normalized,
"-t",
str(final_duration),
"-vf",
f"scale={pip_width}:{pip_height}",
"-c:v",
self.config.output_codec,
"-preset",
self.config.output_preset,
"-crf",
str(self.config.output_crf),
"-pix_fmt",
"yuv420p",
voice_adjusted,
]
run_ffmpeg(command)
bg_adjusted = os.path.join(self.work_dir, f"bg_adj_{os.getpid()}.mp4")
command = [
FFMPEG_BIN,
"-y",
"-i",
bg_normalized,
"-t",
str(final_duration),
"-c:v",
"copy",
bg_adjusted,
]
run_ffmpeg(command)
if audio_path:
command = [
FFMPEG_BIN,
"-y",
"-i",
bg_adjusted,
"-i",
voice_adjusted,
"-i",
audio_path,
"-filter_complex",
f"[0:v][1:v]overlay={x_offset}:{y_offset}[v]",
"-map",
"[v]",
"-map",
"2:a",
"-shortest",
"-c:v",
self.config.output_codec,
"-preset",
self.config.output_preset,
"-crf",
str(self.config.output_crf),
"-pix_fmt",
"yuv420p",
output_path,
]
else:
command = [
FFMPEG_BIN,
"-y",
"-i",
bg_adjusted,
"-i",
voice_adjusted,
"-filter_complex",
f"[0:v][1:v]overlay={x_offset}:{y_offset}[v]",
"-map",
"[v]",
"-map",
"1:a",
"-shortest",
"-c:v",
self.config.output_codec,
"-preset",
self.config.output_preset,
"-crf",
str(self.config.output_crf),
"-pix_fmt",
"yuv420p",
output_path,
]
run_ffmpeg(command)
for temp_file in [voice_normalized, voice_adjusted, bg_normalized, bg_adjusted]:
try:
if temp_file != output_path:
os.remove(temp_file)
except Exception as e:
logger.warning(f"Operation failed in apps/worker/video_processing/editing_modes.py: {e}", exc_info=True)
return output_path
def create_processor(mode: str, work_dir: Optional[str] = None, **kwargs) -> EditingModeProcessor:
"""便捷工厂函数:创建剪辑模式处理器"""
try:
editing_mode = EditingMode(mode)
except ValueError:
raise ValueError(f"Invalid editing mode: {mode}. Valid modes: {[m.value for m in EditingMode]}")
config = EditingModeConfig(
mode=editing_mode,
output_width=kwargs.get("output_width", 1280),
output_height=kwargs.get("output_height", 720),
output_fps=kwargs.get("output_fps", 25),
pip_position=PIPPosition(kwargs.get("pip_position", "top_right")),
pip_scale=kwargs.get("pip_scale", 0.25),
transition_duration=kwargs.get("transition_duration", 0.5),
)
return EditingModeProcessor(config=config, work_dir=work_dir)
+13 -85
View File
@@ -1,7 +1,8 @@
"""FFmpeg 工具函数 — 共享原语.
"""FFmpeg 工具函数 — 从 editing_modes.py / video_compose_service.py 提取的共享原语.
提供 FFmpeg / FFprobe 调用、视频信息探测、视频标准化、xfade 转场滤镜构建
等底层能力,供 UnifiedRenderService、VideoComposeService 等复用。
等底层能力,供 EditingModeProcessor、VideoComposeService、UnifiedRenderService
共同复用。
"""
from __future__ import annotations
@@ -31,10 +32,6 @@ XFADE_TRANSITION_MAP: dict[str, str] = {
"slide_left": "slideleft",
"slideright": "slideright",
"slide_right": "slideright",
"slideup": "slideup",
"slide_up": "slideup",
"slidedown": "slidedown",
"slide_down": "slidedown",
"dissolve": "dissolve",
"wipe": "wipeleft",
"wipeleft": "wipeleft",
@@ -42,10 +39,6 @@ XFADE_TRANSITION_MAP: dict[str, str] = {
DEFAULT_TRANSITION_DURATION = 0.5
# FFmpeg 执行默认超时(秒),防止 FFmpeg hang 住导致 worker 永久阻塞
# 默认 30 分钟,足够处理大部分短视频渲染;超长视频可单独传参覆盖
DEFAULT_FFMPEG_TIMEOUT = 1800
# ── FFmpeg 执行 ───────────────────────────────────────────────────────────────
@@ -54,14 +47,12 @@ def run_ffmpeg(
command: list[str],
*,
capture_output: bool = True,
timeout: int | None = DEFAULT_FFMPEG_TIMEOUT,
) -> tuple[str, str]:
"""执行 FFmpeg 命令。
Args:
command: 完整的 ffmpeg 命令列表(含 "ffmpeg" 本身)
capture_output: 是否捕获 stdout/stderr
timeout: 超时时间(秒),默认 1800s(30分钟);None 表示不设超时(不推荐)
Returns:
(stdout, stderr) 元组
@@ -69,7 +60,6 @@ def run_ffmpeg(
Raises:
subprocess.CalledProcessError: 命令执行失败时抛出,
异常信息包含完整 stderr 以便排查。
subprocess.TimeoutExpired: 超时未完成时抛出,FFmpeg 进程会被 kill。
"""
try:
result = subprocess.run( # nosec B603
@@ -78,16 +68,8 @@ def run_ffmpeg(
stdout=subprocess.PIPE if capture_output else None,
stderr=subprocess.PIPE if capture_output else None,
text=True,
timeout=timeout,
)
return (result.stdout or "", result.stderr or "")
except subprocess.TimeoutExpired as e:
logger.error(
"FFmpeg 命令超时 (%ds): command=%s",
timeout or -1,
" ".join(str(c) for c in command[:20]),
)
raise
except subprocess.CalledProcessError as e:
# 把完整 stderr 打到日志,方便排查 exit code 183 等问题
stderr_text = (e.stderr or "").strip()
@@ -100,41 +82,6 @@ def run_ffmpeg(
raise
def probe_has_audio(local_path: str | Path) -> bool:
"""探测文件是否包含音频流。
Args:
local_path: 本地文件路径
Returns:
True 表示有音频流(或探测失败保守返回),False 表示确认无音频流
"""
try:
result = subprocess.run( # nosec B603
[
FFPROBE_BIN,
"-v",
"error",
"-select_streams",
"a:0",
"-show_entries",
"stream=codec_type",
"-of",
"default=noprint_wrappers=1:nokey=1",
str(local_path),
],
check=True,
stdout=subprocess.PIPE,
stderr=subprocess.PIPE,
text=True,
timeout=10,
)
return result.stdout.strip() == "audio"
except Exception:
# 探测失败保守返回 True,让 FFmpeg 自己处理(避免误删音频)
return True
def probe_duration(local_path: str | Path) -> float:
"""用 ffprobe 获取视频时长(秒)。
@@ -163,14 +110,10 @@ def probe_duration(local_path: str | Path) -> float:
def probe_video_info(video_path: str) -> dict[str, Any]:
"""获取视频信息(宽、高、时长、fps、编码、像素格式)。
"""获取视频信息(宽、高、时长、fps)。
Returns:
{
"width": int, "height": int, "duration": float, "fps": float,
"video_codec": str, "audio_codec": str, "pix_fmt": str,
"has_audio": bool,
}
{"width": int, "height": int, "duration": float, "fps": float}
失败时返回默认值。
"""
try:
@@ -179,8 +122,10 @@ def probe_video_info(video_path: str) -> dict[str, Any]:
FFPROBE_BIN,
"-v",
"error",
"-select_streams",
"v:0",
"-show_entries",
"stream=width,height,r_frame_rate,duration,codec_name,codec_type,pix_fmt",
"stream=width,height,r_frame_rate,duration",
"-show_entries",
"format=duration",
"-of",
@@ -191,25 +136,19 @@ def probe_video_info(video_path: str) -> dict[str, Any]:
stdout=subprocess.PIPE,
stderr=subprocess.PIPE,
text=True,
timeout=15,
)
import json
info = json.loads(result.stdout)
streams = info.get("streams", [])
stream = info.get("streams", [{}])[0]
fmt = info.get("format", {})
video_stream = next((s for s in streams if s.get("codec_type") == "video"), {})
audio_stream = next((s for s in streams if s.get("codec_type") == "audio"), {})
width = int(video_stream.get("width", DEFAULT_OUTPUT_WIDTH))
height = int(video_stream.get("height", DEFAULT_OUTPUT_HEIGHT))
video_codec = video_stream.get("codec_name", "") or ""
pix_fmt = video_stream.get("pix_fmt", "") or ""
width = int(stream.get("width", DEFAULT_OUTPUT_WIDTH))
height = int(stream.get("height", DEFAULT_OUTPUT_HEIGHT))
# 解析帧率
fps_str = video_stream.get("r_frame_rate", "25/1")
fps_str = stream.get("r_frame_rate", "25/1")
if "/" in fps_str:
num, den = fps_str.split("/")
fps = float(num) / float(den) if float(den) > 0 else DEFAULT_FPS
@@ -217,20 +156,13 @@ def probe_video_info(video_path: str) -> dict[str, Any]:
fps = float(fps_str) if fps_str else DEFAULT_FPS
# 时长
duration = float(fmt.get("duration", 0)) or float(video_stream.get("duration", 0))
has_audio = bool(audio_stream)
audio_codec = audio_stream.get("codec_name", "") or ""
duration = float(fmt.get("duration", 0)) or float(stream.get("duration", 0))
return {
"width": width,
"height": height,
"duration": duration,
"fps": round(fps, 2),
"video_codec": video_codec,
"audio_codec": audio_codec,
"pix_fmt": pix_fmt,
"has_audio": has_audio,
}
except Exception as e:
logger.warning("获取视频信息失败: %s, error: %s", video_path, e)
@@ -239,10 +171,6 @@ def probe_video_info(video_path: str) -> dict[str, Any]:
"height": DEFAULT_OUTPUT_HEIGHT,
"duration": 0.0,
"fps": DEFAULT_FPS,
"video_codec": "",
"audio_codec": "",
"pix_fmt": "",
"has_audio": True,
}
+8 -110
View File
@@ -9,7 +9,6 @@ from __future__ import annotations
import hashlib
import logging
import os
import threading
from pathlib import Path
from typing import Optional
from urllib.parse import urlparse
@@ -18,13 +17,6 @@ import oss2
logger = logging.getLogger(__name__)
# OSS 上传配置
OSS_CONNECT_TIMEOUT = 10 # 连接超时(秒),防止 TCP 握手挂死
OSS_UPLOAD_TOTAL_TIMEOUT = 300 # 单文件上传总超时(秒),防止网络慢时无限卡住
OSS_MULTIPART_THRESHOLD = 100 * 1024 * 1024 # 分片上传阈值:100MB 以上走分片
OSS_PART_SIZE = 8 * 1024 * 1024 # 分片大小:8MB
OSS_MULTIPART_NUM_THREADS = 3 # 分片上传并发数
# ── OSS 配置 ──────────────────────────────────────────────────────────────────
@@ -48,12 +40,6 @@ def oss_settings() -> tuple[str, str, str, str] | None:
def oss_bucket() -> oss2.Bucket | None:
"""获取 OSS Bucket 实例。
P0-2 修复:endpoint 不带 scheme 时自动补 https:// 前缀,
确保 sign_url 等依赖 scheme 的方法返回 HTTPS URL。
P0-staging 修复:增加 connect_timeout=10s,防止网络抖动时
TCP 握手阶段无限挂死,导致 worker 进程卡死。
Returns:
oss2.Bucket 实例,配置缺失时返回 None。
"""
@@ -61,15 +47,7 @@ def oss_bucket() -> oss2.Bucket | None:
if settings is None:
return None
access_key_id, access_key_secret, endpoint, bucket_name = settings
# endpoint 无 scheme 时补 https://,与 API 端 storage.py 保持一致
if not endpoint.startswith(("http://", "https://")):
endpoint = f"https://{endpoint}"
return oss2.Bucket(
oss2.Auth(access_key_id, access_key_secret),
endpoint,
bucket_name,
connect_timeout=OSS_CONNECT_TIMEOUT,
)
return oss2.Bucket(oss2.Auth(access_key_id, access_key_secret), endpoint, bucket_name)
def normalize_storage_key(storage_key_or_url: str) -> str:
@@ -112,9 +90,6 @@ def download_asset(asset_storage_key: str, local_path: Path) -> bool:
def upload_to_oss(local_path: Path, storage_key: str) -> str | None:
"""上传文件到 OSS,返回公开 URL。
大文件(>100MB)自动走分片上传,降低内存峰值,减少 OOM 风险。
上传加总超时保护(默认 300s),防止网络异常时无限挂死。
Args:
local_path: 本地文件路径
storage_key: 目标存储键
@@ -123,94 +98,17 @@ def upload_to_oss(local_path: Path, storage_key: str) -> str | None:
公开访问 URL,上传失败或 OSS 未配置时返回 None。
"""
bucket = oss_bucket()
if bucket is None:
return None
result: dict = {"url": None, "error": None, "file_size": 0}
done = threading.Event()
def _do_upload():
try:
# 尝试获取文件大小,用于分片判断和日志;stat 失败时 fallback 走普通上传
try:
file_size = local_path.stat().st_size
result["file_size"] = file_size
use_multipart = file_size >= OSS_MULTIPART_THRESHOLD
except OSError:
use_multipart = False
file_size = 0
if use_multipart:
# 分片上传:降低内存峰值,每片 8MB,3 线程并发
logger.info(
"大文件分片上传: storage_key=%s, size=%.1fMB, part_size=%dMB, threads=%d",
storage_key[:80],
file_size / 1024 / 1024,
OSS_PART_SIZE // 1024 // 1024,
OSS_MULTIPART_NUM_THREADS,
)
oss2.resumable_upload(
bucket,
storage_key,
str(local_path),
multipart_threshold=OSS_MULTIPART_THRESHOLD,
part_size=OSS_PART_SIZE,
num_threads=OSS_MULTIPART_NUM_THREADS,
)
else:
bucket.put_object_from_file(storage_key, str(local_path))
# 构造返回 URL
settings = oss_settings()
if settings:
_, _, endpoint, bucket_name = settings
endpoint_clean = endpoint.replace("https://", "").replace("http://", "")
result["url"] = f"https://{bucket_name}.{endpoint_clean}/{storage_key}"
except Exception as e:
result["error"] = e
logger.exception("上传 OSS 失败: %s", storage_key)
finally:
done.set()
upload_thread = threading.Thread(target=_do_upload, daemon=True)
upload_thread.start()
finished = done.wait(timeout=OSS_UPLOAD_TOTAL_TIMEOUT)
if not finished:
logger.error(
"OSS 上传超时(%.0fs),强制中止: storage_key=%s, size=%.1fMB",
OSS_UPLOAD_TOTAL_TIMEOUT,
storage_key[:80],
result["file_size"] / 1024 / 1024 if result["file_size"] else 0,
)
return None
if result["error"]:
return None
return result["url"]
def get_signed_download_url(storage_key_or_url: str, expires_seconds: int = 3600) -> str | None:
"""生成预签名下载 URL(用于私有 bucket 的 URL 校验或临时下载)。
Args:
storage_key_or_url: 存储键或完整 URL(URL 会自动提取 path)
expires_seconds: 签名有效期(秒)
Returns:
预签名 URL,失败或 OSS 未配置时返回 None。
"""
bucket = oss_bucket()
if bucket is None:
return None
try:
storage_key = normalize_storage_key(storage_key_or_url)
signed = bucket.sign_url("GET", storage_key, expires_seconds)
logger.info("生成预签名URL: key=%s url_prefix=%s", storage_key[:80], signed[:60])
return signed
bucket.put_object_from_file(storage_key, str(local_path))
settings = oss_settings()
if settings:
_, _, endpoint, bucket_name = settings
return f"https://{bucket_name}.{endpoint.replace('https://', '').replace('http://', '')}/{storage_key}"
return None
except Exception:
logger.exception("生成预签名URL失败: %s", storage_key_or_url[:80])
logger.exception("上传 OSS 失败: %s", storage_key)
return None
@@ -1,299 +0,0 @@
"""统一渲染引擎适配层 — Phase 2.
将 EditPlan + EditPlanClips(来自 DB)适配为 UnifiedRenderService 的输入格式,
封装素材下载、渲染执行、结果上传的完整流程。
职责:
1. 从 DB 读取 EditPlan + EditPlanClips
2. 下载素材到本地,构建 asset_path_map
3. 调用 UnifiedRenderService 执行渲染
4. 上传渲染结果到 OSS
5. 支持进度回调(对接 JobService)
"""
from __future__ import annotations
import logging
import tempfile
from dataclasses import dataclass
from pathlib import Path
from typing import Any, Callable
from sqlalchemy.orm import Session
from video_processing.oss_helpers import download_asset, upload_to_oss
from video_processing.unified_render_service import RenderResult, UnifiedRenderService
from packages.adapters.sqlalchemy_impl.edit_plan_clip_repository import SQLAlchemyEditPlanClipRepository
from packages.adapters.sqlalchemy_impl.edit_plan_repository import SQLAlchemyEditPlanRepository
from packages.domain.edit_plan import EditPlan, EditPlanStatus
from packages.domain.edit_plan_clip import EditPlanClip, EditPlanClipStatus
logger = logging.getLogger(__name__)
# ── 数据结构 ──────────────────────────────────────────────────────────────────
@dataclass
class RenderAdapterResult:
"""渲染适配结果。"""
success: bool
output_url: str = ""
output_path: Path | None = None
duration: float = 0.0
file_size: int = 0
width: int = 0
height: int = 0
clip_count: int = 0
error_message: str = ""
ProgressCallback = Callable[[float, str], None]
"""进度回调:(progress_0_100, stage_description) → None"""
# ── 适配层主体 ────────────────────────────────────────────────────────────────
class RenderAdapter:
"""统一渲染引擎适配层。
桥接 EditPlan 领域模型与 UnifiedRenderService 图层模型。
用法::
adapter = RenderAdapter(db)
result = adapter.render_plan(
plan_id=plan_id,
job_id=job_id,
progress_cb=lambda p, s: job_service.update_progress(job_id, p, s),
)
"""
def __init__(self, db: Session) -> None:
self._db = db
self._plan_repo = SQLAlchemyEditPlanRepository(db)
self._clip_repo = SQLAlchemyEditPlanClipRepository(db)
# ── 公开方法 ──────────────────────────────────────────────────────────
def render_plan(
self,
plan_id: str,
*,
job_id: str = "",
work_dir: Path | None = None,
progress_cb: ProgressCallback | None = None,
) -> RenderAdapterResult:
"""渲染一个 EditPlan。
完整流程:
1. 加载计划与片段
2. 下载素材
3. 执行统一渲染
4. 上传结果
Args:
plan_id: EditPlan ID
job_id: 关联的 Job ID(用于结果存储路径)
work_dir: 工作目录,不传则使用临时目录
progress_cb: 进度回调函数
Returns:
RenderAdapterResult
"""
temp_dir = None
try:
# 0. 准备工作目录
if work_dir is None:
temp_dir = tempfile.mkdtemp(prefix="render_")
work_dir = Path(temp_dir)
work_dir.mkdir(parents=True, exist_ok=True)
self._report_progress(progress_cb, 5.0, "加载剪辑计划")
# 1. 加载计划与片段
plan = self._plan_repo.get(plan_id)
if plan is None:
return RenderAdapterResult(
success=False,
error_message=f"剪辑计划不存在: {plan_id}",
)
clips = self._clip_repo.list_by_plan(plan_id, skip=0, limit=10000)
ready_clips = [c for c in clips if c.status == EditPlanClipStatus.READY and c.asset_id]
ready_clips.sort(key=lambda c: c.order)
if not ready_clips:
return RenderAdapterResult(
success=False,
error_message="没有可渲染的就绪片段",
clip_count=0,
)
logger.info(
"开始渲染: plan_id=%s job_id=%s ready_clips=%d engine=unified",
plan_id,
job_id,
len(ready_clips),
)
self._report_progress(progress_cb, 15.0, f"下载素材({len(ready_clips)} 个)")
# 2. 下载素材
asset_path_map = self._download_assets(ready_clips, work_dir)
if not asset_path_map:
return RenderAdapterResult(
success=False,
error_message="所有素材下载失败",
clip_count=len(ready_clips),
)
self._report_progress(progress_cb, 40.0, "执行视频渲染")
# 3. 执行统一渲染
render_svc = UnifiedRenderService(
plan=plan,
clips=ready_clips,
asset_path_map=asset_path_map,
work_dir=work_dir,
)
result = render_svc.render()
self._report_progress(progress_cb, 80.0, "上传渲染结果")
# 4. 上传结果
storage_key = f"rendered/{plan_id}/{job_id or plan_id}.mp4"
output_url = upload_to_oss(result.output_path, storage_key)
self._report_progress(progress_cb, 100.0, "渲染完成")
logger.info(
"[render-adapter] render success: plan_id=%s job_id=%s engine=unified "
"duration=%.2fs file_size=%d resolution=%dx%d clip_count=%d",
plan_id,
job_id,
result.duration,
result.file_size,
result.width,
result.height,
len(ready_clips),
)
return RenderAdapterResult(
success=True,
output_url=output_url or "",
output_path=result.output_path,
duration=result.duration,
file_size=result.file_size,
width=result.width,
height=result.height,
clip_count=len(ready_clips),
)
except Exception as exc:
logger.exception(
"[render-adapter] render failed: plan_id=%s job_id=%s engine=unified error=%s",
plan_id,
job_id,
str(exc)[:200],
)
return RenderAdapterResult(
success=False,
error_message=str(exc)[:500],
)
finally:
# 清理临时目录
if temp_dir:
import shutil
try:
shutil.rmtree(temp_dir, ignore_errors=True)
except Exception:
pass
def validate_plan(self, plan_id: str) -> tuple[bool, list[str], list[str], int, int]:
"""校验计划是否可渲染(兼容 VideoComposeService.validate_compose 接口)。
Returns:
(valid, errors, warnings, ready_clip_count, total_clip_count)
"""
errors: list[str] = []
warnings: list[str] = []
plan = self._plan_repo.get(plan_id)
if plan is None:
return False, [f"剪辑计划不存在: {plan_id}"], [], 0, 0
if plan.status not in (EditPlanStatus.EDITING, EditPlanStatus.RENDERING):
errors.append(f"计划状态不正确,需要 editing 或 rendering,当前: {plan.status}")
clips = self._clip_repo.list_by_plan(plan_id, skip=0, limit=10000)
if not clips:
errors.append("计划没有任何片段")
return False, errors, warnings, 0, 0
clips.sort(key=lambda c: c.order)
ready_count = 0
pending_count = 0
no_asset_count = 0
for clip in clips:
if clip.status == EditPlanClipStatus.READY:
ready_count += 1
if not clip.asset_id:
errors.append(f"片段 {clip.id} (order={clip.order}) 没有分配素材")
no_asset_count += 1
elif clip.status == EditPlanClipStatus.PENDING:
pending_count += 1
elif clip.status == EditPlanClipStatus.FAILED:
warnings.append(f"片段 {clip.id} (order={clip.order}) 状态为 failed,已跳过")
if ready_count == 0:
errors.append("没有就绪(ready)的片段可以合成")
if pending_count > 0:
warnings.append(f"有 {pending_count} 个片段仍处于 pending 状态")
return len(errors) == 0, errors, warnings, ready_count, len(clips)
# ── 内部方法 ──────────────────────────────────────────────────────────
@staticmethod
def _report_progress(progress_cb: ProgressCallback | None, progress: float, stage: str) -> None:
"""上报进度。"""
if progress_cb is not None:
try:
progress_cb(progress, stage)
except Exception:
logger.exception("进度回调失败")
@staticmethod
def _download_assets(clips: list[EditPlanClip], work_dir: Path) -> dict[str, Path]:
"""下载片段素材到本地,返回 asset_id → local_path 映射。
只保留下载成功的素材。
"""
asset_dir = work_dir / "assets"
asset_dir.mkdir(exist_ok=True)
asset_path_map: dict[str, Path] = {}
for clip in clips:
asset_id = clip.asset_id
if not asset_id:
continue
# 生成安全的本地文件名
safe_name = f"clip_{clip.order:04d}_{abs(hash(asset_id)) % 100000:05d}.mp4"
local_path = asset_dir / safe_name
if download_asset(asset_id, local_path):
asset_path_map[asset_id] = local_path
logger.debug("素材下载成功: clip_id=%s asset_id=%s", clip.id, asset_id[:60])
else:
logger.warning("素材下载失败: clip_id=%s asset_id=%s", clip.id, asset_id[:60])
return asset_path_map
@@ -1,204 +0,0 @@
"""渲染引擎 Feature Flag 解析器。
封装渲染引擎选择逻辑,支持:
- 环境变量作为默认值(RENDER_ENGINE=legacy/unified)
- Redis Feature Flag 运行时覆盖(白名单 + 百分比 + 全局开关)
- 定时刷新,支持热更新不重启 worker
使用方式:
resolver = RenderEngineResolver(redis_url="redis://...", default_engine="legacy")
engine = resolver.get_engine(user_id="user123")
# engine: "legacy" 或 "unified"
"""
from __future__ import annotations
import logging
import threading
from typing import Optional
from packages.adapters.redis.feature_flag_store import (
FeatureFlagConfig,
FeatureFlagStore,
InMemoryFeatureFlagStore,
RedisFeatureFlagStore,
)
logger = logging.getLogger(__name__)
# Feature Flag 名称常量
FLAG_RENDER_ENGINE = "render_engine"
# 引擎常量
ENGINE_LEGACY = "legacy"
ENGINE_UNIFIED = "unified"
VALID_ENGINES = {ENGINE_LEGACY, ENGINE_UNIFIED}
class RenderEngineResolver:
"""渲染引擎选择器。
判定逻辑(从高到低):
1. Redis flag 白名单匹配 → unified
2. Redis flag 百分比命中 → unified
3. Redis flag 全局开启(100%)→ unified
4. 环境变量默认值 → legacy / unified
当 Redis 不可用时,自动降级到环境变量默认值,不影响业务。
"""
def __init__(
self,
default_engine: str = ENGINE_LEGACY,
redis_url: Optional[str] = None,
refresh_interval: float = 30.0,
store: Optional[FeatureFlagStore] = None,
) -> None:
"""
Args:
default_engine: 环境变量默认的引擎名(legacy / unified)
redis_url: Redis 连接 URL,传 None 时使用内存实现(测试用)
refresh_interval: Redis flag 配置刷新间隔(秒)
store: 直接传入 store 实例(测试用,优先级高于 redis_url)
"""
self._default_engine = default_engine.lower() if default_engine else ENGINE_LEGACY
if self._default_engine not in VALID_ENGINES:
logger.warning(
"Invalid default engine '%s', fallback to '%s'",
self._default_engine,
ENGINE_LEGACY,
)
self._default_engine = ENGINE_LEGACY
if store is not None:
self._store = store
elif redis_url:
self._store = RedisFeatureFlagStore(redis_url=redis_url)
else:
self._store = InMemoryFeatureFlagStore()
logger.info("No Redis configured, using in-memory feature flag store")
self._refresh_interval = refresh_interval
self._lock = threading.Lock()
self._cached_config: Optional[FeatureFlagConfig] = None
self._last_refresh: float = 0.0
def _maybe_refresh(self) -> None:
"""惰性刷新配置,超过刷新间隔时从存储重新读取。"""
import time
now = time.time()
if now - self._last_refresh < self._refresh_interval:
return
try:
config = self._store.get(FLAG_RENDER_ENGINE)
with self._lock:
self._cached_config = config
self._last_refresh = now
except Exception as exc:
logger.warning("Failed to refresh render engine flag: %s", exc)
# 刷新失败时保留旧缓存,不中断业务
if self._cached_config is None:
# 首次就读失败,设一个默认值
with self._lock:
self._cached_config = FeatureFlagConfig(name=FLAG_RENDER_ENGINE)
self._last_refresh = now
def _get_config(self) -> FeatureFlagConfig:
"""获取当前 flag 配置(带缓存)。"""
if self._cached_config is None:
self._maybe_refresh()
else:
self._maybe_refresh()
return self._cached_config or FeatureFlagConfig(name=FLAG_RENDER_ENGINE)
def get_engine(self, user_id: Optional[str] = None) -> str:
"""获取当前应该使用的渲染引擎。
Args:
user_id: 用户ID,用于白名单匹配和百分比哈希。
传 None 时只看全局开关。
Returns:
"legacy" 或 "unified"
"""
config = self._get_config()
# 全局关闭 → 用默认值
if not config.enabled:
return self._default_engine
# 白名单匹配 / 百分比命中 → unified
if config.is_active(user_id):
return ENGINE_UNIFIED
# 未命中灰度 → 用默认值
return self._default_engine
def should_use_unified(self, user_id: Optional[str] = None) -> bool:
"""便捷方法:是否应该使用统一渲染引擎。"""
return self.get_engine(user_id) == ENGINE_UNIFIED
def force_refresh(self) -> None:
"""强制立即刷新配置(用于管理接口修改后立即生效)。"""
self._last_refresh = 0.0
if isinstance(self._store, RedisFeatureFlagStore):
self._store.invalidate_cache(FLAG_RENDER_ENGINE)
self._maybe_refresh()
def get_config_snapshot(self) -> dict:
"""获取当前配置快照(用于管理接口展示)。"""
config = self._get_config()
return {
"flag_name": FLAG_RENDER_ENGINE,
"default_engine": self._default_engine,
"enabled": config.enabled,
"percentage": config.percentage,
"whitelist": sorted(config.whitelist),
"refresh_interval": self._refresh_interval,
"last_refresh": self._last_refresh,
}
def set_flag(self, config: FeatureFlagConfig) -> None:
"""设置 flag 配置(管理接口用)。"""
config.name = FLAG_RENDER_ENGINE
self._store.set(config)
self.force_refresh()
# 全局单例
_resolver: Optional[RenderEngineResolver] = None
_resolver_lock = threading.Lock()
def get_render_engine_resolver() -> RenderEngineResolver:
"""获取全局单例(基于 worker 配置)。"""
global _resolver
if _resolver is not None:
return _resolver
with _resolver_lock:
if _resolver is not None:
return _resolver
try:
from worker_app.core.config import get_settings
settings = get_settings()
redis_url = getattr(settings, "redis_url", None) or getattr(settings, "broker_url", None)
default = getattr(settings, "render_engine", ENGINE_LEGACY)
_resolver = RenderEngineResolver(
default_engine=default,
redis_url=redis_url,
)
logger.info(
"RenderEngineResolver initialized: default=%s, redis=%s",
default,
bool(redis_url),
)
except Exception as exc:
logger.warning("Failed to init RenderEngineResolver from settings: %s", exc)
_resolver = RenderEngineResolver(default_engine=ENGINE_LEGACY)
return _resolver
+29 -1050
View File
File diff suppressed because it is too large Load Diff
+821
View File
@@ -0,0 +1,821 @@
"""
视频合成服务
支持多种剪辑模式和转场效果,包含完整的安全校验
"""
import logging
import os
import subprocess
import tempfile
from dataclasses import dataclass
from enum import Enum
try:
from enum import StrEnum
except ImportError:
class StrEnum(str, Enum): # type: ignore[no-redef]
"""Python 3.10 兼容的 StrEnum 回退实现。"""
pass
from pathlib import Path
from typing import Optional
from packages.domain.editing_mode import EditingMode
logger = logging.getLogger(__name__)
# ========== 安全常量 ==========
# 允许的输出目录白名单(使用环境变量或系统临时目录,避免硬编码 /tmp)
_VIDEO_OUTPUT_DIR = os.environ.get("VIDEO_OUTPUT_DIR", os.path.join(tempfile.gettempdir(), "video_output"))
ALLOWED_OUTPUT_DIRS = [_VIDEO_OUTPUT_DIR, "/var/app/rendered"]
# 允许的输入路径前缀白名单
ALLOWED_INPUT_PREFIXES = ("s3://", "oss://", "local://", "/var/storage/")
# 允许的转场效果白名单
ALLOWED_TRANSITIONS = {
"fade",
"slideleft",
"slideright",
"dissolve",
"wipeleft",
"wiperight",
"cut",
"slideup",
"slidedown",
}
# 转场效果映射
_XFADE_TRANSITION_MAP = {
"fade": "fade",
"slideleft": "slideleft",
"slideright": "slideright",
"dissolve": "dissolve",
"wipeleft": "wipeleft",
"wiperight": "wiperight",
"cut": "cut",
"slideup": "slideup",
"slidedown": "slidedown",
}
class VideoComposeError(Exception):
"""视频合成服务异常"""
pass
class PIPPosition(StrEnum):
"""画中画位置枚举"""
TOP_LEFT = "top_left"
TOP_RIGHT = "top_right"
BOTTOM_LEFT = "bottom_left"
BOTTOM_RIGHT = "bottom_right"
@dataclass
class Clip:
"""视频片段"""
asset_id: str # 资源ID,对应输入路径
start_time: float = 0.0
duration: float = 0.0
transition: str = "fade" # 转场效果
@dataclass
class EditingModeConfig:
"""剪辑模式配置"""
mode: EditingMode
output_width: int = 1280
output_height: int = 720
output_fps: int = 25
pip_position: PIPPosition = PIPPosition.TOP_RIGHT
pip_scale: float = 0.25 # 画中画占主画面的比例
transition_duration: float = 0.5 # 转场时长(秒)
output_codec: str = "libx264"
output_preset: str = "medium"
output_crf: int = 23
class VideoComposeService:
"""视频合成服务"""
def __init__(self, config: EditingModeConfig, work_dir: Optional[str] = None):
"""
初始化视频合成服务
Args:
config: 剪辑模式配置
work_dir: 工作目录,默认使用系统临时目录
"""
self.config = config
self.work_dir = work_dir or tempfile.gettempdir()
self._ffmpeg_bin = "ffmpeg"
self._ffprobe_bin = "ffprobe"
def _validate_output_path(self, path: str) -> str:
"""
校验输出路径是否在允许范围内 (P0 修复)
防止路径穿越攻击,如 /app/config/../../../etc/passwd
Args:
path: 用户提供的输出路径
Returns:
标准化后的绝对路径
Raises:
ValueError: 路径不在允许范围内
"""
abs_path = os.path.abspath(path)
for allowed_dir in ALLOWED_OUTPUT_DIRS:
allowed_abs = os.path.abspath(allowed_dir)
if abs_path.startswith(allowed_abs):
return abs_path
raise ValueError(f"输出路径不在允许范围内: {path}")
def _validate_input_path(self, path: str) -> bool:
"""
校验输入路径格式是否合法 (P1-1 修复)
Args:
path: 输入文件路径
Returns:
是否合法
"""
return any(path.startswith(prefix) for prefix in ALLOWED_INPUT_PREFIXES)
def _validate_transition(self, transition: str) -> str:
"""
校验转场效果是否在白名单内 (P1-2 修复)
Args:
transition: 转场效果名称
Returns:
安全的转场效果名称
"""
if transition not in ALLOWED_TRANSITIONS:
logger.warning(f"未知的转场效果 '{transition}',使用默认 'fade'")
return "fade"
return transition
def _get_validated_transition(self, transition: str) -> str:
"""获取白名单校验后的转场效果名称"""
return _XFADE_TRANSITION_MAP.get(self._validate_transition(transition), "fade")
def compose(self, clips: list[Clip], output_path: Optional[str] = None) -> str:
"""
合成视频
Args:
clips: 视频片段列表,每个片段包含 asset_id 和转场配置
output_path: 输出文件路径
Returns:
输出文件路径
"""
if not clips:
raise ValueError("clips 不能为空")
# P1-1: 校验所有输入路径
for clip in clips:
if not self._validate_input_path(clip.asset_id):
raise ValueError(f"不合法的输入路径: {clip.asset_id}")
# 生成默认输出路径并校验
if output_path is None:
output_path = self._generate_output_path()
# P0: 校验输出路径
validated_output = self._validate_output_path(output_path)
logger.info(f"合成视频,片段数: {len(clips)}, 输出: {validated_output}")
# 获取输入路径列表
input_paths = [clip.asset_id for clip in clips]
try:
if self.config.mode == EditingMode.ONE_TAKE:
return self._one_take(input_paths, validated_output, clips)
elif self.config.mode == EditingMode.PIP:
return self._pip(input_paths, validated_output)
elif self.config.mode == EditingMode.VOICE_OVER:
return self._voice_over(input_paths, validated_output)
elif self.config.mode == EditingMode.VOICE_PIP:
return self._voice_pip(input_paths, validated_output)
else:
raise ValueError(f"不支持的剪辑模式: {self.config.mode}")
except Exception as e:
logger.error(f"视频合成失败: {e}")
raise VideoComposeError(f"视频合成失败: {e}") from e
def _generate_output_path(self) -> str:
"""生成输出文件路径"""
os.makedirs(self.work_dir, exist_ok=True)
return os.path.join(self.work_dir, f"output_{self.config.mode}_{os.getpid()}.mp4")
def _validate_inputs(self, video_paths: list[str], audio_path: Optional[str] = None) -> None:
"""验证输入文件存在"""
for path in video_paths:
if not os.path.exists(path):
raise FileNotFoundError(f"视频文件不存在: {path}")
if not os.path.getsize(path) > 0:
raise ValueError(f"视频文件为空: {path}")
if audio_path and not os.path.exists(audio_path):
raise FileNotFoundError(f"音频文件不存在: {audio_path}")
def _run_ffmpeg(self, command: list[str], capture_output: bool = True) -> tuple:
"""执行 FFmpeg 命令"""
logger.debug(f"Running FFmpeg: {' '.join(command)}")
try:
result = subprocess.run(
command,
check=True,
stdout=subprocess.PIPE if capture_output else None,
stderr=subprocess.PIPE if capture_output else None,
text=capture_output,
)
return result.stdout or "", result.stderr or ""
except subprocess.CalledProcessError as e:
stderr = e.stderr.decode() if e.stderr else str(e)
logger.error(f"FFmpeg error: {stderr}")
raise RuntimeError(f"FFmpeg 执行失败: {stderr}") from e
def _get_video_info(self, video_path: str) -> dict:
"""获取视频信息"""
try:
result = subprocess.run(
[
self._ffprobe_bin,
"-v",
"error",
"-show_entries",
"stream=width,height,r_frame_rate,duration,codec_name",
"-show_entries",
"format=duration,size",
"-of",
"json",
video_path,
],
check=True,
stdout=subprocess.PIPE,
stderr=subprocess.PIPE,
text=True,
)
import json
data = json.loads(result.stdout)
streams = data.get("streams", [{}])
video_stream = next((s for s in streams if s.get("codec_type") == "video"), streams[0] if streams else {})
fmt = data.get("format", {})
fps_str = video_stream.get("r_frame_rate", "25/1")
fps_parts = fps_str.split("/")
fps = float(fps_parts[0]) / float(fps_parts[1]) if len(fps_parts) == 2 else float(fps_parts[0])
return {
"width": int(video_stream.get("width", 0)),
"height": int(video_stream.get("height", 0)),
"fps": fps,
"duration": float(fmt.get("duration", 0)),
"codec": video_stream.get("codec_name", "unknown"),
"size": int(fmt.get("size", 0)),
}
except Exception as e:
logger.warning(f"获取视频信息失败 {video_path}: {e}")
return {"width": 0, "height": 0, "fps": 25, "duration": 0, "codec": "unknown", "size": 0}
def _get_pip_position_offset(
self, main_width: int, main_height: int, pip_width: int, pip_height: int
) -> tuple[int, int]:
"""获取画中画位置偏移量"""
margin = 10
position_offsets = {
PIPPosition.TOP_LEFT: (margin, margin),
PIPPosition.TOP_RIGHT: (main_width - pip_width - margin, margin),
PIPPosition.BOTTOM_LEFT: (margin, main_height - pip_height - margin),
PIPPosition.BOTTOM_RIGHT: (main_width - pip_width - margin, main_height - pip_height - margin),
}
return position_offsets.get(self.config.pip_position, position_offsets[PIPPosition.TOP_RIGHT])
def _normalize_video(self, input_path: str, output_path: str) -> dict:
"""标准化视频格式"""
command = [
self._ffmpeg_bin,
"-y",
"-i",
input_path,
"-r",
str(self.config.output_fps),
"-vf",
f"scale={self.config.output_width}:{self.config.output_height}:force_original_aspect_ratio=decrease,pad={self.config.output_width}:{self.config.output_height}:(ow-iw)/2:(oh-ih)/2,setsar=1",
"-r",
str(self.config.output_fps),
"-c:v",
self.config.output_codec,
"-preset",
self.config.output_preset,
"-crf",
str(self.config.output_crf),
"-pix_fmt",
"yuv420p",
"-movflags",
"+faststart",
"-an",
output_path,
]
self._run_ffmpeg(command)
return self._get_video_info(output_path)
def _one_take(self, video_paths: list[str], output_path: str, clips: list[Clip]) -> str:
"""一镜到底模式"""
if len(video_paths) == 1:
return self._normalize_video(video_paths[0], output_path)
normalized_paths = []
for i, path in enumerate(video_paths):
normalized = os.path.join(self.work_dir, f"normalized_{i}_{os.getpid()}.mp4")
self._normalize_video(path, normalized)
normalized_paths.append(normalized)
durations = [self._get_video_info(p)["duration"] for p in normalized_paths]
if len(normalized_paths) <= 5:
output_path = self._one_take_with_xfade(normalized_paths, durations, output_path, clips)
else:
output_path = self._one_take_simple_concat(normalized_paths, output_path)
for p in normalized_paths:
try:
if p != output_path:
os.remove(p)
except Exception as e:
logger.warning(
f"Operation failed in apps/worker/video_processing/video_compose_service.py: {e}", exc_info=True
)
return output_path
def _one_take_with_xfade(
self, normalized_paths: list[str], durations: list[float], output_path: str, clips: list[Clip]
) -> str:
"""使用 xfade 滤镜实现转场 (P1-2: 转场参数白名单校验)"""
if len(normalized_paths) == 2:
# 获取当前片段的转场效果并校验白名单
transition = "fade"
if len(clips) > 1:
transition = self._get_validated_transition(clips[1].transition)
trans_duration = self.config.transition_duration
offset1 = durations[0] - trans_duration / 2
command = [
self._ffmpeg_bin,
"-y",
"-i",
normalized_paths[0],
"-i",
normalized_paths[1],
"-filter_complex",
f"[0:v][1:v]xfade=transition={transition}:duration={trans_duration}:offset={offset1}[v]",
"-map",
"[v]",
"-c:v",
self.config.output_codec,
"-preset",
self.config.output_preset,
"-crf",
str(self.config.output_crf),
"-pix_fmt",
"yuv420p",
output_path,
]
self._run_ffmpeg(command)
return output_path
else:
return self._one_take_simple_concat(normalized_paths, output_path)
def _one_take_simple_concat(self, normalized_paths: list[str], output_path: str) -> str:
"""使用 concat demuxer 简单拼接"""
concat_file = os.path.join(self.work_dir, f"concat_list_{os.getpid()}.txt")
with open(concat_file, "w") as f:
for path in normalized_paths:
f.write(f"file '{os.path.abspath(path)}'\n")
command = [
self._ffmpeg_bin,
"-y",
"-f",
"concat",
"-safe",
"0",
"-i",
concat_file,
"-c",
"copy",
output_path,
]
self._run_ffmpeg(command)
try:
os.remove(concat_file)
except Exception as e:
logger.warning(
f"Operation failed in apps/worker/video_processing/video_compose_service.py: {e}", exc_info=True
)
return output_path
def _pip(self, video_paths: list[str], output_path: str) -> str:
"""画中画模式"""
if not video_paths:
raise ValueError("No video paths provided")
main_video = video_paths[0]
main_normalized = os.path.join(self.work_dir, f"main_{os.getpid()}.mp4")
main_info = self._normalize_video(main_video, main_normalized)
if len(video_paths) == 1:
os.rename(main_normalized, output_path)
return output_path
pip_width = int(self.config.output_width * self.config.pip_scale)
pip_height = int(self.config.output_height * self.config.pip_scale)
x_offset, y_offset = self._get_pip_position_offset(
self.config.output_width, self.config.output_height, pip_width, pip_height
)
pip_normalized = os.path.join(self.work_dir, f"pip_{os.getpid()}.mp4")
pip_info = self._get_video_info(video_paths[1])
if pip_info["duration"] > main_info["duration"]:
temp_pip = os.path.join(self.work_dir, f"pip_temp_{os.getpid()}.mp4")
command = [
self._ffmpeg_bin,
"-y",
"-i",
video_paths[1],
"-t",
str(main_info["duration"]),
"-vf",
f"scale={pip_width}:{pip_height}",
"-c:v",
self.config.output_codec,
"-preset",
self.config.output_preset,
"-crf",
str(self.config.output_crf),
"-pix_fmt",
"yuv420p",
temp_pip,
]
self._run_ffmpeg(command)
pip_normalized_input = temp_pip
else:
command = [
self._ffmpeg_bin,
"-y",
"-i",
video_paths[1],
"-vf",
f"scale={pip_width}:{pip_height}",
"-c:v",
self.config.output_codec,
"-preset",
self.config.output_preset,
"-crf",
str(self.config.output_crf),
"-pix_fmt",
"yuv420p",
pip_normalized,
]
self._run_ffmpeg(command)
pip_normalized_input = pip_normalized
if main_info["duration"] > pip_info["duration"]:
looped_pip = os.path.join(self.work_dir, f"pip_looped_{os.getpid()}.mp4")
command = [
self._ffmpeg_bin,
"-y",
"-stream_loop",
"-1",
"-i",
pip_normalized_input,
"-t",
str(main_info["duration"]),
"-vf",
f"scale={pip_width}:{pip_height}",
"-c:v",
self.config.output_codec,
"-preset",
self.config.output_preset,
"-crf",
str(self.config.output_crf),
"-pix_fmt",
"yuv420p",
looped_pip,
]
self._run_ffmpeg(command)
pip_normalized_input = looped_pip
command = [
self._ffmpeg_bin,
"-y",
"-i",
main_normalized,
"-i",
pip_normalized_input,
"-filter_complex",
f"[0:v][1:v]overlay={x_offset}:{y_offset}[v]",
"-map",
"[v]",
"-c:v",
self.config.output_codec,
"-preset",
self.config.output_preset,
"-crf",
str(self.config.output_crf),
"-pix_fmt",
"yuv420p",
output_path,
]
self._run_ffmpeg(command)
for temp_file in [main_normalized, pip_normalized]:
if temp_file and temp_file != output_path:
try:
os.remove(temp_file)
except Exception as e:
logger.warning(
f"Operation failed in apps/worker/video_processing/video_compose_service.py: {e}", exc_info=True
)
return output_path
def _voice_over(self, video_paths: list[str], audio_path: str, output_path: str) -> str:
"""口播模式"""
if not audio_path:
raise ValueError("audio_path is required for VOICE_OVER mode")
if not video_paths:
raise ValueError("No background video provided")
audio_info = self._get_video_info(audio_path)
audio_duration = audio_info["duration"]
bg_normalized = os.path.join(self.work_dir, f"bg_{os.getpid()}.mp4")
bg_info = self._normalize_video(video_paths[0], bg_normalized)
if bg_info["duration"] < audio_duration:
looped_bg = os.path.join(self.work_dir, f"bg_looped_{os.getpid()}.mp4")
command = [
self._ffmpeg_bin,
"-y",
"-stream_loop",
"-1",
"-i",
bg_normalized,
"-t",
str(audio_duration),
"-vf",
f"scale={self.config.output_width}:{self.config.output_height}",
"-c:v",
self.config.output_codec,
"-preset",
self.config.output_preset,
"-crf",
str(self.config.output_crf),
"-pix_fmt",
"yuv420p",
looped_bg,
]
self._run_ffmpeg(command)
bg_normalized = looped_bg
elif bg_info["duration"] > audio_duration:
temp_bg = os.path.join(self.work_dir, f"bg_trimmed_{os.getpid()}.mp4")
command = [
self._ffmpeg_bin,
"-y",
"-i",
bg_normalized,
"-t",
str(audio_duration),
"-c:v",
"copy",
temp_bg,
]
self._run_ffmpeg(command)
bg_normalized = temp_bg
blurred_bg = os.path.join(self.work_dir, f"bg_blurred_{os.getpid()}.mp4")
command = [
self._ffmpeg_bin,
"-y",
"-i",
bg_normalized,
"-vf",
f"boxblur=5:5,scale={self.config.output_width}:{self.config.output_height}",
"-c:v",
self.config.output_codec,
"-preset",
self.config.output_preset,
"-crf",
str(self.config.output_crf),
"-pix_fmt",
"yuv420p",
blurred_bg,
]
self._run_ffmpeg(command)
command = [
self._ffmpeg_bin,
"-y",
"-i",
blurred_bg,
"-i",
audio_path,
"-filter_complex",
"[0:v]drawbox=x=0:y=0:w=iw:h=ih:color=black@0.3:t=fill[v]",
"-map",
"[v]",
"-map",
"1:a",
"-c:v",
self.config.output_codec,
"-preset",
self.config.output_preset,
"-crf",
str(self.config.output_crf),
"-pix_fmt",
"yuv420p",
"-shortest",
output_path,
]
self._run_ffmpeg(command)
for temp_file in [bg_normalized, blurred_bg]:
try:
if temp_file != output_path:
os.remove(temp_file)
except Exception as e:
logger.warning(
f"Operation failed in apps/worker/video_processing/video_compose_service.py: {e}", exc_info=True
)
return output_path
def _voice_pip(self, video_paths: list[str], audio_path: Optional[str], output_path: str) -> str:
"""口播+画中画模式"""
if not video_paths:
raise ValueError("No video paths provided")
if len(video_paths) == 1:
return self._normalize_video(video_paths[0], output_path)
voice_video = video_paths[0]
bg_video = video_paths[1] if len(video_paths) > 1 else video_paths[0]
voice_normalized = os.path.join(self.work_dir, f"voice_{os.getpid()}.mp4")
voice_info = self._normalize_video(voice_video, voice_normalized)
bg_normalized = os.path.join(self.work_dir, f"bg_{os.getpid()}.mp4")
bg_info = self._normalize_video(bg_video, bg_normalized)
final_duration = min(voice_info["duration"], bg_info["duration"])
pip_width = int(self.config.output_width * self.config.pip_scale)
pip_height = int(self.config.output_height * self.config.pip_scale)
x_offset, y_offset = self._get_pip_position_offset(
self.config.output_width, self.config.output_height, pip_width, pip_height
)
voice_adjusted = os.path.join(self.work_dir, f"voice_adj_{os.getpid()}.mp4")
command = [
self._ffmpeg_bin,
"-y",
"-i",
voice_normalized,
"-t",
str(final_duration),
"-vf",
f"scale={pip_width}:{pip_height}",
"-c:v",
self.config.output_codec,
"-preset",
self.config.output_preset,
"-crf",
str(self.config.output_crf),
"-pix_fmt",
"yuv420p",
voice_adjusted,
]
self._run_ffmpeg(command)
bg_adjusted = os.path.join(self.work_dir, f"bg_adj_{os.getpid()}.mp4")
command = [
self._ffmpeg_bin,
"-y",
"-i",
bg_normalized,
"-t",
str(final_duration),
"-c:v",
"copy",
bg_adjusted,
]
self._run_ffmpeg(command)
if audio_path:
command = [
self._ffmpeg_bin,
"-y",
"-i",
bg_adjusted,
"-i",
voice_adjusted,
"-i",
audio_path,
"-filter_complex",
f"[0:v][1:v]overlay={x_offset}:{y_offset}[v]",
"-map",
"[v]",
"-map",
"2:a",
"-shortest",
"-c:v",
self.config.output_codec,
"-preset",
self.config.output_preset,
"-crf",
str(self.config.output_crf),
"-pix_fmt",
"yuv420p",
output_path,
]
else:
command = [
self._ffmpeg_bin,
"-y",
"-i",
bg_adjusted,
"-i",
voice_adjusted,
"-filter_complex",
f"[0:v][1:v]overlay={x_offset}:{y_offset}[v]",
"-map",
"[v]",
"-map",
"1:a",
"-shortest",
"-c:v",
self.config.output_codec,
"-preset",
self.config.output_preset,
"-crf",
str(self.config.output_crf),
"-pix_fmt",
"yuv420p",
output_path,
]
self._run_ffmpeg(command)
for temp_file in [voice_normalized, voice_adjusted, bg_normalized, bg_adjusted]:
try:
if temp_file != output_path:
os.remove(temp_file)
except Exception as e:
logger.warning(
f"Operation failed in apps/worker/video_processing/video_compose_service.py: {e}", exc_info=True
)
return output_path
def create_compose_service(mode: str, work_dir: Optional[str] = None, **kwargs) -> VideoComposeService:
"""便捷工厂函数:创建视频合成服务"""
try:
editing_mode = EditingMode(mode)
except ValueError:
raise ValueError(f"无效的剪辑模式: {mode}. 有效模式: {[m.value for m in EditingMode]}")
config = EditingModeConfig(
mode=editing_mode,
output_width=kwargs.get("output_width", 1280),
output_height=kwargs.get("output_height", 720),
output_fps=kwargs.get("output_fps", 25),
pip_position=PIPPosition(kwargs.get("pip_position", "top_right")),
pip_scale=kwargs.get("pip_scale", 0.25),
transition_duration=kwargs.get("transition_duration", 0.5),
)
return VideoComposeService(config=config, work_dir=work_dir)
-4
View File
@@ -17,10 +17,6 @@ class WorkerSettings(BaseSettings):
database_pool_recycle: int = 3600
environment: str = "development"
auto_create_schema: bool = False
redis_url: str = "redis://redis:6379/0"
# 渲染引擎选择:legacy=旧VideoComposeService,unified=新UnifiedRenderService
render_engine: str = "legacy"
model_config = SettingsConfigDict(
env_file=".env",
+65 -155
View File
@@ -39,10 +39,6 @@ def _get_job_service():
def compose_video(self, job_id: str, **kwargs):
"""视频合成任务。
根据 RENDER_ENGINE 配置选择渲染引擎:
- legacy: 旧 VideoComposeService(filter_complex 模式)
- unified: 新 UnifiedRenderService(图层架构)
Args:
job_id: JobService 中的任务 ID
**kwargs: 来自 Job.payload 的额外参数(plan_id, output_path 等)
@@ -60,18 +56,66 @@ def compose_video(self, job_id: str, **kwargs):
job_service.fail_job(job_id, "Missing plan_id in job payload")
return {"status": "error", "message": "Missing plan_id"}
# 判断使用哪个渲染引擎
# 优先级:Redis Feature Flag(白名单 > 百分比) > 环境变量默认
from video_processing.render_engine_resolver import get_render_engine_resolver
# 标记为 running
job_service.update_progress(job_id, progress=10.0, current_stage="初始化合成环境")
resolver = get_render_engine_resolver()
user_id = job.created_by_user_id or None
engine = resolver.get_engine(user_id=user_id)
# 延迟导入 VideoComposeService
from apps.api.app.services.video_compose_service import VideoComposeService
if engine == "unified":
return _compose_with_unified_engine(self, job_service, job, plan_id, db)
else:
return _compose_with_legacy_engine(self, job_service, job, plan_id, db)
compose_svc = VideoComposeService(db)
# 校验合成条件
job_service.update_progress(job_id, progress=20.0, current_stage="校验合成条件")
validation = compose_svc.validate_compose(plan_id)
if not validation.valid:
error_msg = "; ".join(validation.errors)
job_service.fail_job(job_id, f"合成校验失败: {error_msg}")
return {"status": "error", "message": error_msg}
# 构建合成命令
job_service.update_progress(job_id, progress=30.0, current_stage="构建 FFmpeg 命令")
_output_dir = os.environ.get("VIDEO_OUTPUT_DIR", os.path.join(tempfile.gettempdir(), "video_output"))
output_path = os.path.join(_output_dir, f"{job_id}.mp4")
compose_cmd = compose_svc.build_compose_command(plan_id, output_path)
# 执行 FFmpeg
job_service.update_progress(job_id, progress=50.0, current_stage="正在执行视频合成")
logger.info("Executing FFmpeg for job %s, plan %s", job_id, plan_id)
try:
subprocess.run(
compose_cmd.command,
check=True,
stdout=subprocess.PIPE,
stderr=subprocess.PIPE,
text=True,
timeout=3600,
)
except subprocess.CalledProcessError as e:
job_service.fail_job(job_id, f"FFmpeg 执行失败: {e.stderr[:500]}")
raise
# 上传结果
job_service.update_progress(job_id, progress=80.0, current_stage="上传合成结果")
storage_key = f"rendered/{plan_id}/{job_id}.mp4"
from worker_app.tasks.edit_plan_generation import _upload_to_oss
output_url = _upload_to_oss(Path(output_path), storage_key)
# 更新 Job 状态为完成
result_data = {
"plan_id": plan_id,
"output_path": output_path,
"storage_key": storage_key,
"output_url": output_url or "",
"estimated_duration": compose_cmd.estimated_duration,
"clip_count": len(compose_cmd.clip_chains),
}
job_service.complete_job(job_id, result=result_data)
logger.info("视频合成完成: job_id=%s, plan_id=%s", job_id, plan_id)
return {"status": "completed", "job_id": job_id, "result": result_data}
except self.retry_exc as exc:
logger.warning("视频合成重试中: job_id=%s, exc=%s", job_id, exc)
@@ -85,145 +129,11 @@ def compose_video(self, job_id: str, **kwargs):
raise self.retry(exc=exc, countdown=60)
finally:
db.close()
def _compose_with_legacy_engine(task, job_service, job, plan_id: str, db) -> dict:
"""旧引擎渲染路径(VideoComposeService)。"""
job_id = job.id
# 标记为 running
job_service.update_progress(job_id, progress=10.0, current_stage="初始化合成环境")
# 延迟导入 VideoComposeService
from apps.api.app.services.video_compose_service import VideoComposeService
compose_svc = VideoComposeService(db)
# 校验合成条件
job_service.update_progress(job_id, progress=20.0, current_stage="校验合成条件")
validation = compose_svc.validate_compose(plan_id)
if not validation.valid:
error_msg = "; ".join(validation.errors)
job_service.fail_job(job_id, f"合成校验失败: {error_msg}")
return {"status": "error", "message": error_msg}
# 构建合成命令
job_service.update_progress(job_id, progress=30.0, current_stage="构建 FFmpeg 命令")
_output_dir = os.environ.get("VIDEO_OUTPUT_DIR", os.path.join(tempfile.gettempdir(), "video_output"))
output_path = os.path.join(_output_dir, f"{job_id}.mp4")
compose_cmd = compose_svc.build_compose_command(plan_id, output_path)
# 执行 FFmpeg
job_service.update_progress(job_id, progress=50.0, current_stage="正在执行视频合成")
logger.info("Executing FFmpeg for job %s, plan %s", job_id, plan_id)
try:
subprocess.run(
compose_cmd.command,
check=True,
stdout=subprocess.PIPE,
stderr=subprocess.PIPE,
text=True,
timeout=3600,
)
except subprocess.CalledProcessError as e:
job_service.fail_job(job_id, f"FFmpeg 执行失败: {e.stderr[:500]}")
raise
# 上传结果
job_service.update_progress(job_id, progress=80.0, current_stage="上传合成结果")
storage_key = f"rendered/{plan_id}/{job_id}.mp4"
from worker_app.tasks.edit_plan_generation import _upload_to_oss
output_url = _upload_to_oss(Path(output_path), storage_key)
# 更新 Job 状态为完成
result_data = {
"plan_id": plan_id,
"output_path": output_path,
"storage_key": storage_key,
"output_url": output_url or "",
"estimated_duration": compose_cmd.estimated_duration,
"clip_count": len(compose_cmd.clip_chains),
"engine": "legacy",
}
job_service.complete_job(job_id, result=result_data)
logger.info("视频合成完成(legacy): job_id=%s, plan_id=%s", job_id, plan_id)
return {"status": "completed", "job_id": job_id, "result": result_data}
def _compose_with_unified_engine(task, job_service, job, plan_id: str, db) -> dict:
"""新引擎渲染路径(UnifiedRenderService + RenderAdapter)。"""
job_id = job.id
# 标记为 running
job_service.update_progress(job_id, progress=10.0, current_stage="初始化统一渲染引擎")
from video_processing.render_adapter import RenderAdapter
adapter = RenderAdapter(db)
# 校验合成条件
job_service.update_progress(job_id, progress=15.0, current_stage="校验合成条件")
valid, errors, warnings, ready_count, total_count = adapter.validate_plan(plan_id)
if not valid:
error_msg = "; ".join(errors)
job_service.fail_job(job_id, f"合成校验失败: {error_msg}")
return {"status": "error", "message": error_msg}
# 进度回调
def progress_cb(progress: float, stage: str) -> None:
# 清理临时文件
try:
job_service.update_progress(job_id, progress=progress, current_stage=stage)
except Exception:
logger.exception("更新进度失败")
# 执行渲染
job_service.update_progress(job_id, progress=20.0, current_stage="开始渲染")
logger.info("统一渲染引擎开始: job_id=%s plan_id=%s", job_id, plan_id)
result = adapter.render_plan(
plan_id=plan_id,
job_id=job_id,
progress_cb=progress_cb,
)
if not result.success:
job_service.fail_job(job_id, f"渲染失败: {result.error_message}")
raise RuntimeError(result.error_message)
# 更新 Job 状态为完成
result_data = {
"plan_id": plan_id,
"output_path": str(result.output_path) if result.output_path else "",
"storage_key": f"rendered/{plan_id}/{job_id}.mp4",
"output_url": result.output_url,
"estimated_duration": result.duration,
"clip_count": result.clip_count,
"engine": "unified",
"width": result.width,
"height": result.height,
"file_size": result.file_size,
}
job_service.complete_job(job_id, result=result_data)
logger.info(
"视频合成完成(unified): job_id=%s plan_id=%s duration=%.2fs",
job_id,
plan_id,
result.duration,
)
return {"status": "completed", "job_id": job_id, "result": result_data}
def _cleanup_output(job_id: str) -> None:
"""清理临时输出文件。"""
try:
_output_dir = os.environ.get("VIDEO_OUTPUT_DIR", os.path.join(tempfile.gettempdir(), "video_output"))
output_path = os.path.join(_output_dir, f"{job_id}.mp4")
if Path(output_path).exists():
Path(output_path).unlink()
except Exception as e:
logger.warning(f"清理输出文件失败: {e}", exc_info=True)
_output_dir = os.environ.get("VIDEO_OUTPUT_DIR", os.path.join(tempfile.gettempdir(), "video_output"))
output_path = os.path.join(_output_dir, f"{job_id}.mp4")
if Path(output_path).exists():
Path(output_path).unlink()
except Exception as e:
logger.warning(f"Operation failed in apps/worker/worker_app/tasks/compose_video.py: {e}", exc_info=True)
+101 -325
View File
@@ -1,18 +1,13 @@
"""剪辑计划渲染任务 — 支持 Feature Flag 灰度.
"""剪辑计划渲染任务 — Phase 8 任务 2.05.
Celery 任务 worker.render_edit_plan:
1. 加载 EditPlan + EditPlanClips
2. 根据 Feature Flag 选择渲染引擎(legacy / unified)
3. 下载各片段素材 + 渲染
2. 下载各片段素材
3. 使用 UnifiedRenderService 按时间线+图层渲染
4. 上传渲染结果到 OSS
5. 创建 GeneratedVideo 记录 + 查重
6. 更新 EditPlan / EditPlanClip 状态
7. 更新 GenerationTask 进度
渲染引擎灰度:
- 走 Feature Flag (render_engine) 控制
- legacy: VideoComposeService + FFmpeg filter_complex
- unified: UnifiedRenderService 图层架构
"""
from __future__ import annotations
@@ -68,268 +63,14 @@ def _get_repos():
# ── Celery Task ───────────────────────────────────────────────────────────────
def _resolve_render_engine(user_id: str) -> str:
"""根据 Feature Flag 决定使用哪个渲染引擎。
Returns:
"legacy" 或 "unified"
"""
try:
from video_processing.render_engine_resolver import get_render_engine_resolver
resolver = get_render_engine_resolver()
return resolver.get_engine(user_id=user_id)
except Exception as exc:
logger.warning("获取渲染引擎配置失败,fallback 到 legacy: %s", exc)
return "legacy"
def _mark_plan_failed(plan_repo, plan_id, gen_task_repo, generation_task_id, error_msg: str):
"""统一的计划失败标记工具。"""
plan = plan_repo.get(plan_id)
if plan and plan.status.value == "rendering":
plan.mark_failed()
plan_repo.update(plan)
if generation_task_id:
gen_task = gen_task_repo.get(generation_task_id)
if gen_task and gen_task.status.value != "failed":
gen_task.status = "failed"
gen_task.error_message = error_msg
gen_task.completed_at = datetime.now(timezone.utc)
gen_task_repo.update(gen_task)
def _finalize_render_success(
plan,
plan_repo,
clip_repo,
gen_task_repo,
db,
plan_id: str,
output_url: str,
storage_key: str,
duration: float,
file_size: int,
width: int,
height: int,
rendered_clip_ids: list[str],
failed_clip_ids: list[str],
generation_task_id: str,
output_path: Path,
engine: str,
) -> dict:
"""渲染成功后的统一收尾:查重 + 更新状态 + 返回结果。"""
# 创建 GeneratedVideo 记录 + 查重
project_id = plan.project_id or ""
batch_id = plan.config.get("batch_id", "")
mode = plan.config.get("mode", "edit_plan")
if generation_task_id and project_id:
try:
create_video_record_and_dedup(
generation_task_id=generation_task_id,
project_id=project_id,
batch_id=batch_id,
file_url=output_url or "",
file_size=file_size,
duration=duration,
video_path=str(output_path),
mode=mode,
session=db,
width=width,
height=height,
fps=OUTPUT_FPS,
)
except Exception as dedup_err:
logger.warning("查重失败(不影响渲染结果): %s", dedup_err)
# 更新片段状态为 rendered
for clip_id in rendered_clip_ids:
clip = clip_repo.get(clip_id)
if clip and clip.status.value == "ready":
clip.mark_rendered()
clip_repo.update(clip)
# 更新 EditPlan 状态为 completed
plan.config["rendered_url"] = output_url or ""
plan.config["rendered_storage_key"] = storage_key
plan.mark_completed()
plan_repo.update(plan)
# 更新 GenerationTask 状态为 completed
if generation_task_id:
gen_task = gen_task_repo.get(generation_task_id)
if gen_task:
gen_task.status = "completed"
gen_task.progress = 100.0
gen_task.result_count = len(rendered_clip_ids)
gen_task.completed_at = datetime.now(timezone.utc)
gen_task_repo.update(gen_task)
logger.info(
"剪辑计划渲染完成: plan_id=%s engine=%s rendered=%d failed=%d duration=%.1fs",
plan_id,
engine,
len(rendered_clip_ids),
len(failed_clip_ids),
duration,
)
return {
"status": "completed",
"plan_id": plan_id,
"rendered_count": len(rendered_clip_ids),
"failed_count": len(failed_clip_ids),
"output_url": output_url,
"duration": duration,
}
def _render_with_unified(
plan,
clips,
asset_path_map: dict[str, Path],
tmpdir_path: Path,
rendered_clip_ids: list[str],
plan_id: str,
generation_task_id: str,
plan_repo,
clip_repo,
gen_task_repo,
db,
) -> dict:
"""统一渲染引擎路径(UnifiedRenderService 图层架构)。"""
render_service = UnifiedRenderService(
plan=plan,
clips=clips,
asset_path_map=asset_path_map,
work_dir=tmpdir_path,
output_width=OUTPUT_WIDTH,
output_height=OUTPUT_HEIGHT,
output_fps=int(OUTPUT_FPS),
)
try:
render_result = render_service.render()
except Exception as render_err:
logger.error("渲染失败(unified): %s — %s", plan_id, render_err)
_mark_plan_failed(plan_repo, plan_id, gen_task_repo, generation_task_id, f"渲染失败: {render_err}")
return {"status": "error", "message": f"渲染失败: {render_err}"}
output_path = render_result.output_path
# 上传到 OSS
storage_key = f"rendered/{plan_id}/output.mp4"
output_url = upload_to_oss(output_path, storage_key)
failed_clip_ids: list[str] = []
return _finalize_render_success(
plan=plan,
plan_repo=plan_repo,
clip_repo=clip_repo,
gen_task_repo=gen_task_repo,
db=db,
plan_id=plan_id,
output_url=output_url or "",
storage_key=storage_key,
duration=render_result.duration,
file_size=render_result.file_size,
width=render_result.width,
height=render_result.height,
rendered_clip_ids=rendered_clip_ids,
failed_clip_ids=failed_clip_ids,
generation_task_id=generation_task_id,
output_path=output_path,
engine="unified",
)
def _render_with_legacy(
plan,
clips,
rendered_clip_ids: list[str],
failed_clip_ids: list[str],
tmpdir_path: Path,
plan_id: str,
generation_task_id: str,
plan_repo,
clip_repo,
gen_task_repo,
db,
) -> dict:
"""旧引擎路径(VideoComposeService + FFmpeg filter_complex)。"""
import os
import subprocess
from apps.api.app.services.video_compose_service import VideoComposeService
compose_svc = VideoComposeService(db)
# 校验合成条件
validation = compose_svc.validate_compose(plan_id)
if not validation.valid:
error_msg = "; ".join(validation.errors)
logger.error("合成校验失败(legacy): %s — %s", plan_id, error_msg)
_mark_plan_failed(plan_repo, plan_id, gen_task_repo, generation_task_id, f"合成校验失败: {error_msg}")
return {"status": "error", "message": error_msg}
# 构建 FFmpeg 命令
output_dir = os.environ.get("VIDEO_OUTPUT_DIR", str(tmpdir_path))
output_path = Path(output_dir) / f"{plan_id}.mp4"
compose_cmd = compose_svc.build_compose_command(plan_id, str(output_path))
logger.info("执行 FFmpeg (legacy): plan_id=%s", plan_id)
try:
subprocess.run(
compose_cmd.command,
check=True,
stdout=subprocess.PIPE,
stderr=subprocess.PIPE,
text=True,
timeout=3600,
)
except subprocess.CalledProcessError as e:
error_msg = f"FFmpeg 执行失败: {e.stderr[:500]}"
logger.error("FFmpeg 执行失败(legacy): %s — %s", plan_id, error_msg)
_mark_plan_failed(plan_repo, plan_id, gen_task_repo, generation_task_id, error_msg)
return {"status": "error", "message": error_msg}
# 获取文件大小
file_size = output_path.stat().st_size if output_path.exists() else 0
duration = compose_cmd.estimated_duration or 0.0
# 上传到 OSS
storage_key = f"rendered/{plan_id}/output.mp4"
output_url = upload_to_oss(output_path, storage_key)
return _finalize_render_success(
plan=plan,
plan_repo=plan_repo,
clip_repo=clip_repo,
gen_task_repo=gen_task_repo,
db=db,
plan_id=plan_id,
output_url=output_url or "",
storage_key=storage_key,
duration=duration,
file_size=file_size,
width=OUTPUT_WIDTH,
height=OUTPUT_HEIGHT,
rendered_clip_ids=rendered_clip_ids,
failed_clip_ids=failed_clip_ids,
generation_task_id=generation_task_id,
output_path=output_path,
engine="legacy",
)
@celery_app.task(name="worker.render_edit_plan", bind=True, max_retries=2)
def render_edit_plan(self, plan_id: str) -> dict:
"""渲染剪辑计划
流程:
1. 加载 EditPlan + EditPlanClips
2. 根据 Feature Flag 选择渲染引擎(legacy / unified)
3. 下载素材 + 渲染
2. 下载各片段素材到临时目录,构建 asset_path_map
3. 使用 UnifiedRenderService 按时间线+图层渲染
4. 上传渲染结果到 OSS
5. 创建 GeneratedVideo 记录 + 查重
6. 更新 EditPlan → completed, EditPlanClips → rendered
@@ -338,7 +79,6 @@ def render_edit_plan(self, plan_id: str) -> dict:
logger.info("开始渲染剪辑计划: plan_id=%s", plan_id)
generation_task_id = ""
engine = "legacy"
for repos in _get_repos():
plan_repo, clip_repo, gen_task_repo, db = repos
@@ -353,12 +93,7 @@ def render_edit_plan(self, plan_id: str) -> dict:
# 获取 generation_task_id(提前读取,确保 except 块可用)
generation_task_id = plan.config.get("generation_task_id", "")
# 2. 选择渲染引擎(Feature Flag 灰度控制)
user_id = plan.created_by_user_id or ""
engine = _resolve_render_engine(user_id)
logger.info("剪辑计划渲染引擎: plan_id=%s engine=%s user_id=%s", plan_id, engine, user_id)
# 3. 加载片段列表(按 order 排序)
# 2. 加载片段列表(按 order 排序)
clips = clip_repo.list_by_plan(plan_id, skip=0, limit=10000)
if not clips:
logger.warning("剪辑计划没有片段: %s", plan_id)
@@ -381,15 +116,6 @@ def render_edit_plan(self, plan_id: str) -> dict:
rendered_clip_ids: list[str] = []
failed_clip_ids: list[str] = []
# 预先批量查询所有素材的 storage_key(file_url)
from packages.adapters.sqlalchemy_impl.models import AssetModel
clip_asset_ids = [c.asset_id for c in clips if c.asset_id]
asset_storage_map: dict[str, str] = {}
if clip_asset_ids:
assets = db.query(AssetModel).filter(AssetModel.id.in_(clip_asset_ids)).all()
asset_storage_map = {a.id: a.file_url for a in assets if a.file_url}
for clip in clips:
if not clip.asset_id:
# 没有素材的片段跳过,标记为失败
@@ -403,22 +129,10 @@ def render_edit_plan(self, plan_id: str) -> dict:
rendered_clip_ids.append(clip.id)
continue
storage_key = asset_storage_map.get(clip.asset_id)
if not storage_key:
logger.warning(
"片段素材无 storage_key,跳过: clip_id=%s asset_id=%s",
clip.id,
clip.asset_id,
)
clip.mark_failed()
clip_repo.update(clip)
failed_clip_ids.append(clip.id)
continue
# 下载素材
ext = Path(storage_key).suffix or ".mp4"
ext = Path(clip.asset_id).suffix or ".mp4"
local_path = tmpdir_path / f"clip_{clip.order:04d}{ext}"
if download_asset(storage_key, local_path):
if download_asset(clip.asset_id, local_path):
asset_path_map[clip.asset_id] = local_path
rendered_clip_ids.append(clip.id)
else:
@@ -439,38 +153,100 @@ def render_edit_plan(self, plan_id: str) -> dict:
gen_task_repo.update(gen_task)
return {"status": "error", "message": "所有片段素材下载失败"}
# 4. 根据引擎选择渲染方式
if engine == "unified":
result = _render_with_unified(
plan=plan,
clips=clips,
asset_path_map=asset_path_map,
tmpdir_path=tmpdir_path,
rendered_clip_ids=rendered_clip_ids,
plan_id=plan_id,
generation_task_id=generation_task_id,
plan_repo=plan_repo,
clip_repo=clip_repo,
gen_task_repo=gen_task_repo,
db=db,
)
else:
result = _render_with_legacy(
plan=plan,
clips=clips,
rendered_clip_ids=rendered_clip_ids,
failed_clip_ids=failed_clip_ids,
tmpdir_path=tmpdir_path,
plan_id=plan_id,
generation_task_id=generation_task_id,
plan_repo=plan_repo,
clip_repo=clip_repo,
gen_task_repo=gen_task_repo,
db=db,
)
# 4. 使用 UnifiedRenderService 渲染
render_service = UnifiedRenderService(
plan=plan,
clips=clips,
asset_path_map=asset_path_map,
work_dir=tmpdir_path,
output_width=OUTPUT_WIDTH,
output_height=OUTPUT_HEIGHT,
output_fps=int(OUTPUT_FPS),
)
result["engine"] = engine
return result
try:
render_result = render_service.render()
except Exception as render_err:
logger.error("渲染失败: %s — %s", plan_id, render_err)
plan.mark_failed()
plan_repo.update(plan)
if generation_task_id:
gen_task = gen_task_repo.get(generation_task_id)
if gen_task:
gen_task.status = "failed"
gen_task.error_message = f"渲染失败: {render_err}"
gen_task.completed_at = datetime.now(timezone.utc)
gen_task_repo.update(gen_task)
return {"status": "error", "message": f"渲染失败: {render_err}"}
output_path = render_result.output_path
# 5. 上传到 OSS
storage_key = f"rendered/{plan_id}/output.mp4"
output_url = upload_to_oss(output_path, storage_key)
# 6. 创建 GeneratedVideo 记录 + 查重
project_id = plan.project_id or ""
batch_id = plan.config.get("batch_id", "")
mode = plan.config.get("mode", "edit_plan")
if generation_task_id and project_id:
try:
create_video_record_and_dedup(
generation_task_id=generation_task_id,
project_id=project_id,
batch_id=batch_id,
file_url=output_url or "",
file_size=render_result.file_size,
duration=render_result.duration,
video_path=str(output_path),
mode=mode,
session=db,
width=render_result.width,
height=render_result.height,
fps=OUTPUT_FPS,
)
except Exception as dedup_err:
logger.warning("查重失败(不影响渲染结果): %s", dedup_err)
# 7. 更新片段状态为 rendered
for clip_id in rendered_clip_ids:
clip = clip_repo.get(clip_id)
if clip and clip.status.value == "ready":
clip.mark_rendered()
clip_repo.update(clip)
# 8. 更新 EditPlan 状态为 completed
plan.config["rendered_url"] = output_url or ""
plan.config["rendered_storage_key"] = storage_key
plan.mark_completed()
plan_repo.update(plan)
# 9. 更新 GenerationTask 状态为 completed
if generation_task_id:
gen_task = gen_task_repo.get(generation_task_id)
if gen_task:
gen_task.status = "completed"
gen_task.progress = 100.0
gen_task.result_count = len(rendered_clip_ids)
gen_task.completed_at = datetime.now(timezone.utc)
gen_task_repo.update(gen_task)
logger.info(
"剪辑计划渲染完成: plan_id=%s rendered=%d failed=%d duration=%.1fs",
plan_id,
len(rendered_clip_ids),
len(failed_clip_ids),
render_result.duration,
)
return {
"status": "completed",
"plan_id": plan_id,
"rendered_count": len(rendered_clip_ids),
"failed_count": len(failed_clip_ids),
"output_url": output_url,
"duration": render_result.duration,
}
except Exception as exc:
logger.exception("渲染剪辑计划异常: %s", plan_id)
+45 -450
View File
@@ -13,13 +13,10 @@
from __future__ import annotations
import json
import logging
import os
import tempfile
import time
from dataclasses import dataclass, field
from datetime import datetime, timezone
from pathlib import Path
from typing import Any, Optional
@@ -84,36 +81,14 @@ def _update_task_status(task_id: str, status_action: str, **kwargs) -> bool:
return False
# ── 日志持久化辅助 ────────────────────────────────────────────────────────────
def _flush_logs(task_id: str, gen_task) -> None:
"""将 gen_task.logs 持久化到 DB(独立 session,失败不抛异常)。"""
try:
session = SessionLocal()
try:
from packages.adapters.sqlalchemy_impl.models import GenerationTaskModel
model = session.query(GenerationTaskModel).filter(GenerationTaskModel.id == task_id).first()
if model:
model.logs = gen_task.logs
session.commit()
finally:
session.close()
except Exception:
logger.warning("[task_id=%s] 日志持久化失败", task_id, exc_info=True)
# ── 共享工具模块导入 ──────────────────────────────────────────────────────────
from video_processing.dedup_helpers import create_video_record_and_dedup
from video_processing.ffmpeg_utils import FFMPEG_BIN, probe_duration, run_ffmpeg
from video_processing.oss_helpers import (
download_asset,
get_signed_download_url,
upload_to_oss,
)
from video_processing.render_engine_resolver import ENGINE_LEGACY, ENGINE_UNIFIED
from video_processing.unified_render_service import UnifiedRenderService
# ── 虚拟 Plan / Clip(内存中构建,不写数据库) ────────────────────────────────
@@ -125,7 +100,6 @@ class _VirtualPlan:
id: str
name: str = ""
config: dict[str, Any] = field(default_factory=dict)
@dataclass
@@ -163,15 +137,13 @@ def _build_plan_and_clips_from_task(
"""
plan = _VirtualPlan(id=task_id, name=f"Generated-{task_id[:8]}")
# 为每个下载路径生成合成 asset_id,并预探测素材时长
# 为每个下载路径生成合成 asset_id
asset_path_map: dict[str, Path] = {}
path_to_asset_id: dict[Path, str] = {}
path_duration: dict[Path, float] = {}
for i, p in enumerate(downloaded_paths):
asset_id = f"gen_{task_id[:8]}_{i:03d}{p.suffix or '.mp4'}"
asset_path_map[asset_id] = p
path_to_asset_id[p] = asset_id
path_duration[p] = probe_duration(p)
clips: list[_VirtualClip] = []
n = len(downloaded_paths)
@@ -187,7 +159,6 @@ def _build_plan_and_clips_from_task(
clip_type=clip_type,
order=i,
asset_id=path_to_asset_id[p],
duration=path_duration[p],
)
)
elif mode == "voice_over":
@@ -200,7 +171,6 @@ def _build_plan_and_clips_from_task(
clip_type="main",
order=i,
asset_id=path_to_asset_id[p],
duration=path_duration[p],
config={"role": "b_roll"},
)
)
@@ -220,7 +190,6 @@ def _build_plan_and_clips_from_task(
clip_type=clip_type,
order=i,
asset_id=path_to_asset_id[p],
duration=path_duration[p],
)
)
else:
@@ -233,7 +202,6 @@ def _build_plan_and_clips_from_task(
clip_type="main",
order=i,
asset_id=path_to_asset_id[p],
duration=path_duration[p],
)
)
@@ -351,8 +319,6 @@ def _download_library_assets(
asset_ids: list[str] | None = None,
video_extensions: tuple = (".mp4", ".mov", ".avi", ".mkv", ".webm"),
strict: bool = True,
task_id: str = "",
gen_task=None,
) -> list[Path]:
"""下载视频素材 — 同时支持素材库模式和项目级模式。
@@ -388,37 +354,31 @@ def _download_library_assets(
session = SessionLocal()
try:
# 构建查询
# 构建查询:根据模式选择不同的过滤条件
query = session.query(AssetModel).filter(
AssetModel.status == "ready",
AssetModel.file_type.in_(["video", "video/mp4", "video/quicktime"]),
)
if asset_ids:
# 明确指定了 asset_ids:直接按 ID 查,不预先按 library/project 过滤
# 避免项目级素材或跨库素材因为 library_id 不匹配而查不到
# 归属安全由后面的归属校验保证
query = query.filter(AssetModel.id.in_(asset_ids))
if asset_library_id:
# 素材库模式
query = query.filter(AssetModel.asset_library_id == asset_library_id)
logger.info(
"下载指定素材: asset_ids=%d 个, asset_library_id=%s, project_id=%s",
len(asset_ids),
asset_library_id or "none",
project_id or "none",
"下载素材库视频: asset_library_id=%s asset_ids=%s",
asset_library_id,
asset_ids or "all",
)
else:
# 未指定 asset_ids:按 library 或 project 下载全部 ready 视频
if asset_library_id:
query = query.filter(AssetModel.asset_library_id == asset_library_id)
logger.info(
"下载素材库全部视频: asset_library_id=%s",
asset_library_id,
)
else:
query = query.filter(AssetModel.project_id == project_id)
logger.info(
"下载项目全部视频: project_id=%s",
project_id,
)
# 项目级模式
query = query.filter(AssetModel.project_id == project_id)
logger.info(
"下载项目级视频: project_id=%s asset_ids=%s",
project_id,
asset_ids or "all",
)
if asset_ids:
query = query.filter(AssetModel.id.in_(asset_ids))
assets = query.order_by(AssetModel.created_at).all()
@@ -435,15 +395,13 @@ def _download_library_assets(
if missing_ids:
raise ValueError(f"素材不存在: asset_ids={sorted(missing_ids)}")
for asset in assets:
# 校验素材库归属(只要传了 asset_library_id 就校验)
if asset_library_id and asset.asset_library_id != asset_library_id:
raise ValueError(
f"素材不属于指定素材库: asset_id={asset.id}, "
f"expected_asset_library_id={asset_library_id}, "
f"actual_asset_library_id={asset.asset_library_id}"
)
# 校验项目归属(只要传了 project_id 就校验)
if project_id and asset.project_id != project_id:
if not asset_library_id and project_id and asset.project_id != project_id:
raise ValueError(
f"素材不属于指定项目: asset_id={asset.id}, "
f"expected_project_id={project_id}, "
@@ -460,65 +418,19 @@ def _download_library_assets(
storage_key = asset.file_url if asset.file_url else None
if not storage_key:
failed_assets.append(f"{asset.name}({asset.id})")
logger.warning(
"[task_id=%s] 素材缺少 file_url, 跳过: asset_id=%s name=%s", task_id, asset.id, asset.name
)
if gen_task:
gen_task.append_log(
"下载素材",
"素材缺少file_url, 跳过",
level="WARN",
asset_id=asset.id,
asset_name=asset.name,
success=False,
file_size=0,
duration=0.0,
)
logger.warning("素材缺少 file_url, 跳过: asset_id=%s name=%s", asset.id, asset.name)
if strict:
raise RuntimeError(f"素材缺少 file_url: asset_id={asset.id}, name={asset.name}")
continue
ext = Path(storage_key).suffix or ".mp4"
local_file = temp_path / f"asset_{i:03d}_{asset.id}{ext}"
asset_start = time.monotonic()
download_ok = download_asset(storage_key, local_file)
asset_elapsed = time.monotonic() - asset_start
if download_ok:
file_size = local_file.stat().st_size if local_file.exists() else 0
if download_asset(storage_key, local_file):
downloaded.append(local_file)
logger.info(
"[task_id=%s] Downloaded asset: %s -> %s (size=%d, time=%.1fs)",
task_id,
asset.name,
local_file,
file_size,
asset_elapsed,
)
if gen_task:
gen_task.append_log(
"下载素材",
f"下载成功: {asset.name}",
asset_id=asset.id,
asset_name=asset.name,
success=True,
file_size=file_size,
duration=round(asset_elapsed, 2),
)
logger.info("Downloaded asset: %s -> %s", asset.name, local_file)
else:
failed_assets.append(f"{asset.name}({asset.id})")
logger.warning("[task_id=%s] Failed to download asset: %s (id=%s)", task_id, asset.name, asset.id)
if gen_task:
gen_task.append_log(
"下载素材",
f"下载失败: {asset.name}",
level="WARN",
asset_id=asset.id,
asset_name=asset.name,
success=False,
file_size=0,
duration=round(asset_elapsed, 2),
)
logger.warning("Failed to download asset: %s (id=%s)", asset.name, asset.id)
if strict:
raise RuntimeError(f"素材下载失败: asset_id={asset.id}, name={asset.name}")
@@ -574,148 +486,6 @@ def _validate_template_exists(template_id: str) -> None:
session.close()
# ── 渲染引擎选择 ─────────────────────────────────────────────────────────────
def _resolve_render_engine(user_id: str) -> str:
"""根据 Feature Flag 决定使用哪个渲染引擎。
Returns:
"legacy" 或 "unified"
"""
try:
from video_processing.render_engine_resolver import get_render_engine_resolver
resolver = get_render_engine_resolver()
return resolver.get_engine(user_id=user_id)
except Exception as exc:
logger.warning("获取渲染引擎配置失败,fallback 到 unified: %s", exc)
return ENGINE_UNIFIED
# ── 旧引擎渲染(FFmpeg filter_complex) ────────────────────────────────────────
def _render_with_legacy_engine(
task_id: str,
virtual_clips: list[_VirtualClip],
asset_path_map: dict[str, Path],
work_dir: Path,
output_path: Path,
) -> tuple[float, int]:
"""旧引擎渲染路径:手动构建 FFmpeg filter_complex 命令。
说明:generate_video 任务使用虚拟 clips(无 EditPlan 数据库记录),
因此无法直接复用 VideoComposeService。这里手动构建等价的 filter_complex
命令,与旧引擎行为一致(scale → crop → setpts → trim → setpts,
无 fps 归一化,保持原帧率)。
支持模式:one_take / pip / voice_over / voice_pip
- 所有模式统一走 concat 滤镜(与旧引擎多片段逻辑一致)
Returns:
(duration_seconds, file_size_bytes)
"""
import subprocess
main_clips = [
c
for c in virtual_clips
if c.clip_type in ("main", "b_roll", "background")
or (c.clip_type == "main" and c.config.get("role") == "b_roll")
]
if not main_clips:
main_clips = virtual_clips[:1]
input_args: list[str] = []
video_filters: list[str] = []
audio_filters: list[str] = []
for i, clip in enumerate(main_clips):
local_path = asset_path_map.get(clip.asset_id)
if not local_path:
continue
input_args.extend(["-i", str(local_path)])
duration = clip.duration or 0.0
# 视频滤镜:scale → crop → setpts → trim → setpts(与旧引擎一致)
vf = (
f"[{i}:v]"
f"scale={OUTPUT_WIDTH}:{OUTPUT_HEIGHT}:force_original_aspect_ratio=increase,"
f"crop={OUTPUT_WIDTH}:{OUTPUT_HEIGHT},"
f"setpts=PTS-STARTPTS,"
f"trim=0:{duration:.3f},"
f"setpts=PTS-STARTPTS"
f"[v{i}]"
)
video_filters.append(vf)
# 音频滤镜:atrim → asetpts
af = f"[{i}:a]atrim=0:{duration:.3f},asetpts=PTS-STARTPTS[a{i}]"
audio_filters.append(af)
n = len(main_clips)
if n == 1:
video_label = "[v0]"
audio_label = "[a0]"
else:
# concat 视频
v_inputs = "".join(f"[v{i}]" for i in range(n))
video_filters.append(f"{v_inputs}concat=n={n}:v=1:a=0[outv]")
# concat 音频
a_inputs = "".join(f"[a{i}]" for i in range(n))
audio_filters.append(f"{a_inputs}concat=n={n}:v=0:a=1[outa]")
video_label = "[outv]"
audio_label = "[outa]"
# 组装 filter_complex
fc_parts = video_filters + audio_filters
filter_complex = ";".join(fc_parts)
command = [
FFMPEG_BIN,
"-y",
*input_args,
"-filter_complex",
filter_complex,
"-map",
video_label,
"-map",
audio_label,
"-c:v",
"libx264",
"-crf",
"23",
"-preset",
"medium",
"-c:a",
"aac",
"-b:a",
"192k",
"-movflags",
"+faststart",
str(output_path),
]
logger.info("[task_id=%s] [渲染] legacy 引擎 FFmpeg 开始: clips=%d", task_id, n)
try:
run_ffmpeg(command)
except subprocess.CalledProcessError as e:
logger.error(
"[task_id=%s] [渲染] legacy 引擎 FFmpeg 失败: %s\nfilter_complex: %s",
task_id,
e,
filter_complex[:500],
)
raise
file_size = output_path.stat().st_size if output_path.exists() else 0
duration = probe_duration(output_path)
return duration, file_size
# ── Celery Task ──────────────────────────────────────────────────────────────
@@ -740,7 +510,7 @@ def generate_video(self, task_id: str) -> dict:
"""
from packages.domain import EditingMode
logger.info("[task_id=%s] [接收任务] 开始生成视频任务", task_id)
logger.info("开始生成视频任务: task_id=%s", task_id)
# 从数据库加载任务信息
session = SessionLocal()
@@ -752,7 +522,7 @@ def generate_video(self, task_id: str) -> dict:
task_repo = SQLAlchemyGenerationTaskRepository(session)
gen_task = task_repo.get(task_id)
if gen_task is None:
logger.error("[task_id=%s] [接收任务] 任务不存在", task_id)
logger.error("生成任务不存在: task_id=%s", task_id)
return {"status": "failed", "error": f"generation task {task_id} not found"}
project_id = gen_task.project_id
asset_library_id = gen_task.asset_library_id
@@ -761,16 +531,6 @@ def generate_video(self, task_id: str) -> dict:
mode = gen_task.strategy_id or "one_take"
task_asset_ids = list(gen_task.asset_ids or [])
batch_id = getattr(gen_task, "batch_id", "") or ""
# 记录接收任务日志
gen_task.append_log(
"接收任务",
f"模式={mode}, 模板={template_id}, 素材数={len(task_asset_ids)}",
mode=mode,
template_id=template_id,
asset_count=len(task_asset_ids),
)
_flush_logs(task_id, gen_task)
finally:
session.close()
@@ -797,40 +557,12 @@ def generate_video(self, task_id: str) -> dict:
output_path = temp_path / output_name
# 1. 从素材库/项目下载视频素材
logger.info("[task_id=%s] [下载素材] 开始下载视频素材", task_id)
download_start = time.monotonic()
downloaded_videos = _download_library_assets(
temp_path,
asset_library_id=asset_library_id,
project_id=project_id,
asset_ids=task_asset_ids or None,
task_id=task_id,
gen_task=gen_task,
)
download_elapsed = time.monotonic() - download_start
logger.info(
"[task_id=%s] [下载素材] 完成: 成功=%d个, 耗时=%.1fs",
task_id,
len(downloaded_videos),
download_elapsed,
)
# 重新加载 gen_task 以追加日志(session 已关闭)
_session = SessionLocal()
try:
_repo = SQLAlchemyGenerationTaskRepository(_session)
gen_task = _repo.get(task_id)
finally:
_session.close()
if gen_task:
gen_task.append_log(
"下载素材",
f"成功下载 {len(downloaded_videos)} 个视频素材",
count=len(downloaded_videos),
duration=round(download_elapsed, 2),
)
_flush_logs(task_id, gen_task)
# 2. 下载配音(如有)
audio_path: str | None = None
@@ -838,7 +570,6 @@ def generate_video(self, task_id: str) -> dict:
local_audio = temp_path / "voice.mp3"
if _download_voice_asset(voice_library_id, local_audio):
audio_path = str(local_audio)
logger.info("[task_id=%s] [下载配音] 配音下载成功", task_id)
# 3. 渲染
if not downloaded_videos:
@@ -856,147 +587,47 @@ def generate_video(self, task_id: str) -> dict:
mode=editing_mode.value,
)
total_duration = sum(c.duration for c in virtual_clips)
logger.info(
"[task_id=%s] [剪辑计划] 片段数=%d, 总时长=%.1fs",
task_id,
len(virtual_clips),
total_duration,
# 使用 UnifiedRenderService 渲染
render_service = UnifiedRenderService(
plan=virtual_plan,
clips=virtual_clips,
asset_path_map=asset_path_map,
work_dir=temp_path,
output_width=OUTPUT_WIDTH,
output_height=OUTPUT_HEIGHT,
output_fps=int(OUTPUT_FPS),
)
if gen_task:
gen_task.append_log(
"剪辑计划",
f"片段数={len(virtual_clips)}, 总时长={total_duration:.1f}s",
segment_count=len(virtual_clips),
total_duration=round(total_duration, 2),
)
_flush_logs(task_id, gen_task)
# 3. 根据 Feature Flag 选择渲染引擎
user_id = getattr(gen_task, "created_by_user_id", "") if gen_task else ""
engine = _resolve_render_engine(user_id) if user_id else ENGINE_UNIFIED
logger.info("[task_id=%s] [渲染] 引擎选择: %s (user_id=%s)", task_id, engine, user_id)
render_start = time.monotonic()
render_output_path = temp_path / f"rendered-{task_id}.mp4"
if engine == ENGINE_LEGACY:
# 旧引擎:filter_complex + concat(保持原帧率,无 fps 归一化)
render_duration, render_file_size = _render_with_legacy_engine(
task_id=task_id,
virtual_clips=virtual_clips,
asset_path_map=asset_path_map,
work_dir=temp_path,
output_path=render_output_path,
)
render_elapsed = time.monotonic() - render_start
logger.info(
"[task_id=%s] [渲染] legacy 引擎完成: 耗时=%.1fs, 时长=%.2fs",
task_id,
render_elapsed,
render_duration,
)
else:
# 新引擎:UnifiedRenderService 图层架构
logger.info("[task_id=%s] [渲染] unified 引擎 FFmpeg 渲染开始", task_id)
render_service = UnifiedRenderService(
plan=virtual_plan,
clips=virtual_clips,
asset_path_map=asset_path_map,
work_dir=temp_path,
output_width=OUTPUT_WIDTH,
output_height=OUTPUT_HEIGHT,
output_fps=int(OUTPUT_FPS),
)
render_result = render_service.render()
render_output_path = render_result.output_path
render_duration = render_result.duration
render_file_size = render_result.file_size
render_elapsed = time.monotonic() - render_start
logger.info(
"[task_id=%s] [渲染] unified 引擎完成: 耗时=%.1fs",
task_id,
render_elapsed,
)
if gen_task:
gen_task.append_log(
"渲染",
f"引擎={engine}, 耗时={render_elapsed:.1f}s",
duration=round(render_elapsed, 2),
engine=engine,
)
_flush_logs(task_id, gen_task)
render_result = render_service.render()
# 4. 如有配音,后处理混音
if audio_path:
final_path = temp_path / f"final-{task_id}.mp4"
try:
_mux_audio_track(render_output_path, audio_path, final_path)
_mux_audio_track(render_result.output_path, audio_path, final_path)
# 混音成功,使用混音后的文件
output_path = final_path
except Exception as mux_err:
logger.warning("[task_id=%s] [混音] 音频混合失败,使用无音频版本: %s", task_id, mux_err)
output_path = render_output_path
logger.warning("音频混合失败,使用无音频版本: %s", mux_err)
output_path = render_result.output_path
else:
output_path = render_output_path
output_path = render_result.output_path
file_size = output_path.stat().st_size
duration = probe_duration(output_path)
# 5. 上传到 OSS — 失败必须抛异常,不能静默忽略
logger.info("[task_id=%s] [OSS上传] 开始上传: size=%d", task_id, file_size)
upload_start = time.monotonic()
file_url = upload_to_oss(output_path, storage_key)
upload_elapsed = time.monotonic() - upload_start
if not file_url:
# OSS 未配置或上传失败
if gen_task:
gen_task.append_log("OSS上传", "上传失败", level="ERROR")
_flush_logs(task_id, gen_task)
raise RuntimeError(
f"OSS 上传失败: task_id={task_id}, storage_key={storage_key}, " f"output_path={output_path}"
)
# P0-2 修复:私有 bucket 下裸 URL 永远 403,改用预签名 URL 校验
# 先用预签名 URL 校验,失败则降级为检查文件是否存在(object_exists)
verify_url = get_signed_download_url(file_url, expires_seconds=300) or file_url
if not _verify_url_accessible(verify_url):
# 预签名 URL 也访问失败时,退一步用 object_exists 确认上传成功
from video_processing.oss_helpers import normalize_storage_key, oss_bucket
# HEAD 校验 URL 可访问
if not _verify_url_accessible(file_url):
raise RuntimeError(f"OSS 上传后 URL 不可访问: file_url={file_url}, " f"storage_key={storage_key}")
bucket = oss_bucket()
key = normalize_storage_key(file_url)
if bucket and bucket.object_exists(key):
logger.info("URL 校验失败但 object_exists 确认文件存在,视为上传成功: storage_key=%s", key)
if gen_task:
gen_task.append_log("OSS上传", "URL校验降级: object_exists确认存在", level="WARN")
else:
if gen_task:
gen_task.append_log("OSS上传", "上传后URL不可访问", level="ERROR", file_url=file_url)
_flush_logs(task_id, gen_task)
raise RuntimeError(
f"OSS 上传后 URL 不可访问且 object_exists 失败: file_url={file_url}, "
f"storage_key={storage_key}"
)
logger.info(
"[task_id=%s] [OSS上传] 成功: 耗时=%.1fs, file_url=%s",
task_id,
upload_elapsed,
file_url,
)
if gen_task:
gen_task.append_log(
"OSS上传",
f"上传成功, 大小={file_size}, 耗时={upload_elapsed:.1f}s",
file_size=file_size,
duration=round(upload_elapsed, 2),
file_url=file_url,
)
_flush_logs(task_id, gen_task)
logger.info("OSS 上传成功: file_url=%s", file_url)
# 6. 创建 GeneratedVideo 记录 + 查重
dedup_session = SessionLocal()
@@ -1018,23 +649,7 @@ def generate_video(self, task_id: str) -> dict:
# 7. 标记任务为 completed
_update_task_status(task_id, "mark_completed", result_count=video_count or 1)
# 记录完成日志
if gen_task:
gen_task.append_log(
"任务完成",
f"视频生成完成: 时长={duration:.2f}s, 大小={file_size}",
duration=round(duration, 2),
file_size=file_size,
video_count=video_count or 1,
)
_flush_logs(task_id, gen_task)
logger.info(
"[task_id=%s] [任务完成] duration=%.2fs file_size=%d",
task_id,
duration,
file_size,
)
logger.info("视频生成完成: task_id=%s duration=%.2fs file_size=%d", task_id, duration, file_size)
return {
"status": "completed",
@@ -1047,27 +662,7 @@ def generate_video(self, task_id: str) -> dict:
"mode": editing_mode.value,
}
except Exception as error:
logger.error("[task_id=%s] [任务失败] %s", task_id, error, exc_info=True)
# 记录失败日志
try:
_session = SessionLocal()
try:
_repo = SQLAlchemyGenerationTaskRepository(_session)
gen_task = _repo.get(task_id)
if gen_task:
gen_task.append_log(
"任务失败",
str(error),
level="ERROR",
error_type=type(error).__name__,
)
_flush_logs(task_id, gen_task)
finally:
_session.close()
except Exception:
logger.warning("[task_id=%s] 记录失败日志异常", task_id, exc_info=True)
logger.error("Video generation failed: %s", error, exc_info=True)
_update_task_status(task_id, "mark_failed", error_message=str(error))
return {
"status": "failed",
+1 -4
View File
@@ -4,7 +4,6 @@ import logging
from celery import Task
from celery.exceptions import Retry
from video_processing.oss_helpers import get_signed_download_url
from worker_app.celery_app import celery_app
from worker_app.db import SessionLocal
@@ -49,9 +48,7 @@ def process_voice_clone(self: Task, profile_id: str) -> dict:
repo = SQLAlchemyVoiceCloneProfileRepository(session)
workflow = VoiceCloneWorkflowService(
repository=repo,
cosyvoice_service=CosyVoiceService(
audio_url_signer=lambda url: get_signed_download_url(url, expires_seconds=86400) or url
),
cosyvoice_service=CosyVoiceService(),
)
updated_profile = workflow.poll_and_process_clone(profile_id, timeout=300)
+1 -1
View File
@@ -112,7 +112,7 @@
| 变量名 | 用途说明 | 默认值 |
|--------|---------|--------|
| `OSS_ENDPOINT` | OSS Endpoint | `oss-cn-hangzhou.aliyuncs.com` |
| `OSS_ENDPOINT` | OSS Endpoint | `oss-cn-hangzhou.aliiyuncs.com` |
| `OSS_ACCESS_KEY_ID` | OSS Access Key ID | `""`(空) |
| `OSS_ACCESS_KEY_SECRET` | OSS Access Key Secret | `""`(空) |
| `OSS_BUCKET_NAME` | OSS Bucket 名称 | `xiaoxia-autocut` |
-8
View File
@@ -1549,14 +1549,6 @@
"type": "JSON",
"unique": false
},
{
"index": false,
"name": "logs",
"nullable": false,
"primary_key": false,
"type": "TEXT",
"unique": false
},
{
"index": false,
"name": "created_at",
+13 -10
View File
@@ -6,14 +6,14 @@
# 基础镜像:Python 3.12
FROM git.xiaoxiajianji.com/xiaoxia/base/python:3.12-slim
# 构建参数:版本号(CI 传入 commit hash)
ARG APP_VERSION=dev
# 使用阿里云镜像加速
RUN sed -i 's|deb.debian.org|mirrors.aliyun.com|g' /etc/apt/sources.list.d/debian.sources 2>/dev/null || sed -i 's|deb.debian.org|mirrors.aliyun.com|g' /etc/apt/sources.list 2>/dev/null || true
RUN sed -i 's|deb.debian.org|mirrors.aliyun.com|g' /etc/apt/sources.list.d/debian.sources 2>/dev/null || \
sed -i 's|deb.debian.org|mirrors.aliyun.com|g' /etc/apt/sources.list 2>/dev/null || true
# 安装系统依赖
RUN apt-get update && apt-get install -y --no-install-recommends libpq-dev && rm -rf /var/lib/apt/lists/*
RUN apt-get update && apt-get install -y --no-install-recommends \
libpq-dev \
&& rm -rf /var/lib/apt/lists/*
# 设置工作目录
WORKDIR /app
@@ -21,12 +21,15 @@ WORKDIR /app
# ---- 依赖分层:基础依赖(变化少,缓存命中率高)----
COPY requirements-base.txt /tmp/requirements-base.txt
RUN python -m venv /opt/venv && /opt/venv/bin/pip install --no-cache-dir -i https://mirrors.aliyun.com/pypi/simple/ --trusted-host mirrors.aliyun.com -r /tmp/requirements-base.txt && rm /tmp/requirements-base.txt
RUN python -m venv /opt/venv \
&& /opt/venv/bin/pip install --no-cache-dir -i https://mirrors.aliyun.com/pypi/simple/ --trusted-host mirrors.aliyun.com -r /tmp/requirements-base.txt \
&& rm /tmp/requirements-base.txt
# ---- 依赖分层:业务依赖(变化频繁)----
COPY requirements.txt /tmp/requirements.txt
RUN /opt/venv/bin/pip install --no-cache-dir -i https://mirrors.aliyun.com/pypi/simple/ --trusted-host mirrors.aliyun.com -r /tmp/requirements.txt && rm /tmp/requirements.txt
RUN /opt/venv/bin/pip install --no-cache-dir -i https://mirrors.aliyun.com/pypi/simple/ --trusted-host mirrors.aliyun.com -r /tmp/requirements.txt \
&& rm /tmp/requirements.txt
# 复制应用代码
COPY apps/api/ /app/apps/api/
@@ -37,13 +40,13 @@ COPY alembic/ /app/alembic/
COPY scripts/ /app/scripts/
# 设置环境变量
ENV PATH="/opt/venv/bin:/usr/local/sbin:/usr/local/bin:/usr/sbin:/usr/bin:/sbin:/bin"
ENV PATH="/opt/venv/bin:$PATH"
ENV PYTHONPATH=/app
ENV PYTHONUNBUFFERED=1
ENV APP_VERSION=$APP_VERSION
# 健康检查
HEALTHCHECK --interval=30s --timeout=10s --start-period=5s --retries=3 CMD python -c "import urllib.request; urllib.request.urlopen('http://localhost:8000/health', timeout=5)"
HEALTHCHECK --interval=30s --timeout=10s --start-period=5s --retries=3 \
CMD python -c "import urllib.request; urllib.request.urlopen('http://localhost:8000/health', timeout=5)"
# API 入口点
WORKDIR /app/apps/api
-13
View File
@@ -1,13 +0,0 @@
#!/bin/bash
# Worker 启动脚本 — 支持 WORKER_CONCURRENCY 环境变量
# 未设置时默认 2(保持向后兼容)
set -e
CONCURRENCY="${WORKER_CONCURRENCY:-2}"
exec celery \
-A worker_app.celery_app \
worker \
--loglevel=info \
"--concurrency=${CONCURRENCY}"
+1 -9
View File
@@ -6,9 +6,6 @@
# 基础镜像:Python 3.12 + ffmpeg
FROM git.xiaoxiajianji.com/xiaoxia/base/python:3.12-slim
# 构建参数:版本号(CI 传入 commit hash)
ARG APP_VERSION=dev
# 使用阿里云镜像加速
RUN sed -i 's|deb.debian.org|mirrors.aliyun.com|g' /etc/apt/sources.list.d/debian.sources 2>/dev/null || \
sed -i 's|deb.debian.org|mirrors.aliyun.com|g' /etc/apt/sources.list 2>/dev/null || true
@@ -51,15 +48,10 @@ COPY packages/ /app/packages/
COPY alembic.ini /app/alembic.ini
COPY migrations/ /app/migrations/
# 复制 Worker 启动脚本(支持 WORKER_CONCURRENCY 环境变量)
COPY infra/docker/entrypoint-worker.sh /usr/local/bin/entrypoint-worker.sh
RUN chmod +x /usr/local/bin/entrypoint-worker.sh
# 设置 Python 路径
ENV PATH="/opt/venv/bin:$PATH"
ENV PYTHONPATH=/app
ENV PYTHONUNBUFFERED=1
ENV APP_VERSION=$APP_VERSION
# 创建非 root 用户运行 Worker
RUN groupadd -r celery && useradd -r -g celery -d /app -s /sbin/nologin celery \
@@ -69,4 +61,4 @@ USER celery
# Worker 入口点
WORKDIR /app/apps/worker
CMD ["/usr/local/bin/entrypoint-worker.sh"]
CMD ["celery", "-A", "worker_app.celery_app", "worker", "--loglevel=info", "--concurrency=2"]
+1 -16
View File
@@ -1,9 +1,3 @@
from packages.adapters.redis.feature_flag_store import (
FeatureFlagConfig,
FeatureFlagStore,
InMemoryFeatureFlagStore,
RedisFeatureFlagStore,
)
from packages.adapters.redis.session_store import (
NoopSessionStore,
RedisConfig,
@@ -11,13 +5,4 @@ from packages.adapters.redis.session_store import (
get_session_store,
)
__all__ = [
"FeatureFlagConfig",
"FeatureFlagStore",
"InMemoryFeatureFlagStore",
"NoopSessionStore",
"RedisConfig",
"RedisFeatureFlagStore",
"SessionStore",
"get_session_store",
]
__all__ = ["NoopSessionStore", "RedisConfig", "SessionStore", "get_session_store"]
@@ -1,259 +0,0 @@
"""Feature Flag 存储实现。
支持两种后端:
- RedisFeatureFlagStore:生产环境使用,支持多实例共享、热更新
- InMemoryFeatureFlagStore:测试/开发环境使用,纯内存
支持的 Flag 类型:
- 全局开关(enabled: bool)
- 白名单(whitelist: Set[str],如 user_id 列表)
- 百分比切流(percentage: 0-100,基于标识符哈希取模)
判定优先级:白名单 > 百分比 > 全局开关
"""
from __future__ import annotations
import hashlib
import json
import logging
import threading
import time
from abc import ABC, abstractmethod
from dataclasses import dataclass, field
from typing import Optional, Set
logger = logging.getLogger(__name__)
# Redis key 前缀
FEATURE_FLAG_REDIS_PREFIX = "feature_flag:"
@dataclass
class FeatureFlagConfig:
"""单个 Feature Flag 的配置。"""
name: str
enabled: bool = False
percentage: int = 0 # 0-100
whitelist: Set[str] = field(default_factory=set)
def to_dict(self) -> dict:
return {
"name": self.name,
"enabled": self.enabled,
"percentage": self.percentage,
"whitelist": sorted(self.whitelist),
}
@classmethod
def from_dict(cls, data: dict) -> "FeatureFlagConfig":
return cls(
name=data["name"],
enabled=bool(data.get("enabled", False)),
percentage=int(data.get("percentage", 0)),
whitelist=set(data.get("whitelist", [])),
)
def is_active(self, identifier: Optional[str] = None) -> bool:
"""判断当前 flag 是否激活。
判定优先级:
1. 全局关闭 → False
2. 白名单匹配 → True
3. 百分比命中 → True
4. 其他 → False
Args:
identifier: 用于白名单匹配和百分比哈希的标识符(如 user_id)。
传 None 时只看全局开关 + 百分比(百分比用随机值)。
"""
if not self.enabled:
return False
# 白名单:精确匹配
if identifier and identifier in self.whitelist:
return True
# 百分比:0 直接 False,100 直接 True
if self.percentage <= 0:
# 没有白名单且百分比为0 → 未启用
return False
if self.percentage >= 100:
return True
# 基于 identifier 做哈希取模,确保同一用户始终落在同一侧
if identifier:
hash_val = int(
hashlib.md5(f"{self.name}:{identifier}".encode("utf-8")).hexdigest(), 16 # nosec B324
) # nosec B324 - 用于哈希取模做百分比切流,非安全用途
return (hash_val % 100) < self.percentage
# 无 identifier 且百分比在 0-100 之间 → 按比例随机(不保证一致性)
import random
return random.randint(0, 99) < self.percentage
class FeatureFlagStore(ABC):
"""Feature Flag 存储抽象接口。"""
@abstractmethod
def get(self, name: str) -> FeatureFlagConfig:
"""获取指定 flag 的配置,不存在则返回默认配置(关闭状态)。"""
...
@abstractmethod
def set(self, config: FeatureFlagConfig) -> None:
"""设置 flag 配置。"""
...
@abstractmethod
def delete(self, name: str) -> bool:
"""删除 flag,返回是否成功删除。"""
...
@abstractmethod
def list_all(self) -> dict[str, FeatureFlagConfig]:
"""列出所有 flag。"""
...
def is_active(self, name: str, identifier: Optional[str] = None) -> bool:
"""便捷方法:判断 flag 是否激活。"""
return self.get(name).is_active(identifier)
class InMemoryFeatureFlagStore(FeatureFlagStore):
"""内存实现,用于测试和本地开发。"""
def __init__(self) -> None:
self._flags: dict[str, FeatureFlagConfig] = {}
self._lock = threading.Lock()
def get(self, name: str) -> FeatureFlagConfig:
with self._lock:
return self._flags.get(name, FeatureFlagConfig(name=name, enabled=False))
def set(self, config: FeatureFlagConfig) -> None:
with self._lock:
self._flags[config.name] = config
def delete(self, name: str) -> bool:
with self._lock:
if name in self._flags:
del self._flags[name]
return True
return False
def list_all(self) -> dict[str, FeatureFlagConfig]:
with self._lock:
return dict(self._flags)
class RedisFeatureFlagStore(FeatureFlagStore):
"""Redis 实现,支持多实例共享配置。
每个 flag 存在一个独立的 Redis hash key 中:
Key: feature_flag:{name}
Fields: enabled, percentage, whitelist(JSON array)
"""
def __init__(self, redis_url: str, key_prefix: str = FEATURE_FLAG_REDIS_PREFIX) -> None:
import redis as redis_lib
self._redis = redis_lib.from_url(redis_url, decode_responses=True)
self._key_prefix = key_prefix
# 本地缓存 + TTL,减少 Redis 调用
self._cache: dict[str, tuple[FeatureFlagConfig, float]] = {}
self._cache_ttl = 5.0 # 秒,默认5秒本地缓存
self._lock = threading.Lock()
def _redis_key(self, name: str) -> str:
return f"{self._key_prefix}{name}"
def _parse_whitelist(self, raw: Optional[str]) -> Set[str]:
if not raw:
return set()
try:
data = json.loads(raw)
return set(data) if isinstance(data, list) else set()
except (json.JSONDecodeError, TypeError):
return set()
def get(self, name: str) -> FeatureFlagConfig:
now = time.time()
# 先查本地缓存
with self._lock:
cached = self._cache.get(name)
if cached and now - cached[1] < self._cache_ttl:
return cached[0]
# 从 Redis 读取
try:
key = self._redis_key(name)
data = self._redis.hgetall(key)
if not data:
config = FeatureFlagConfig(name=name, enabled=False)
else:
config = FeatureFlagConfig(
name=name,
enabled=(data.get("enabled", "0") in ("1", "true", "True")),
percentage=int(data.get("percentage", 0)),
whitelist=self._parse_whitelist(data.get("whitelist")),
)
# 写入本地缓存
with self._lock:
self._cache[name] = (config, now)
return config
except Exception as exc:
logger.warning("Failed to get feature flag %s from Redis: %s", name, exc)
# Redis 不可用时返回默认值(关闭),不影响业务
return FeatureFlagConfig(name=name, enabled=False)
def set(self, config: FeatureFlagConfig) -> None:
key = self._redis_key(config.name)
self._redis.hset(
key,
mapping={
"enabled": "1" if config.enabled else "0",
"percentage": str(config.percentage),
"whitelist": json.dumps(sorted(config.whitelist), ensure_ascii=False),
},
)
# 失效本地缓存
with self._lock:
self._cache.pop(config.name, None)
def delete(self, name: str) -> bool:
key = self._redis_key(name)
result = self._redis.delete(key)
with self._lock:
self._cache.pop(name, None)
return bool(result)
def list_all(self) -> dict[str, FeatureFlagConfig]:
pattern = f"{self._key_prefix}*"
result: dict[str, FeatureFlagConfig] = {}
try:
cursor = 0
while True:
cursor, keys = self._redis.scan(cursor=cursor, match=pattern, count=100)
for key in keys:
name = key[len(self._key_prefix) :]
result[name] = self.get(name)
if cursor == 0:
break
except Exception as exc:
logger.warning("Failed to list feature flags from Redis: %s", exc)
return result
def invalidate_cache(self, name: Optional[str] = None) -> None:
"""手动失效本地缓存。"""
with self._lock:
if name:
self._cache.pop(name, None)
else:
self._cache.clear()
-20
View File
@@ -27,7 +27,6 @@ def _to_domain(model: GenerationTaskModel) -> GenerationTask:
source_edit_plan_id=model.source_edit_plan_id or "",
asset_select_mode=model.asset_select_mode or "",
batch_id=model.batch_id or "",
logs=model.logs or "[]",
created_at=model.created_at,
)
@@ -57,7 +56,6 @@ class SQLAlchemyGenerationTaskRepository:
source_edit_plan_id=task.source_edit_plan_id or None,
asset_select_mode=task.asset_select_mode or "",
batch_id=task.batch_id or "",
logs=task.logs,
created_at=task.created_at,
)
self.session.add(model)
@@ -91,23 +89,6 @@ class SQLAlchemyGenerationTaskRepository:
def count_by_user(self, user_id: str) -> int:
return self.session.query(GenerationTaskModel).filter(GenerationTaskModel.created_by_user_id == user_id).count()
def count_pending_by_user(self, user_id: str) -> int:
return (
self.session.query(GenerationTaskModel)
.filter(
GenerationTaskModel.created_by_user_id == user_id,
GenerationTaskModel.status == GenerationTaskStatus.PENDING.value,
)
.count()
)
def count_pending_total(self) -> int:
return (
self.session.query(GenerationTaskModel)
.filter(GenerationTaskModel.status == GenerationTaskStatus.PENDING.value)
.count()
)
def list_recent_by_user(self, user_id: str, limit: int = 5) -> list[GenerationTask]:
models = (
self.session.query(GenerationTaskModel)
@@ -148,6 +129,5 @@ class SQLAlchemyGenerationTaskRepository:
model.source_edit_plan_id = task.source_edit_plan_id or None
model.asset_select_mode = task.asset_select_mode or ""
model.batch_id = task.batch_id or ""
model.logs = task.logs
self.session.commit()
return task
@@ -257,7 +257,6 @@ class GenerationTaskModel(Base):
asset_select_mode = Column(String(20), nullable=False, default="")
batch_id = Column(String(36), nullable=False, default="", index=True)
extra_meta = Column("metadata", JSON, nullable=False, default=dict)
logs = Column(Text, nullable=False, default="[]", server_default="[]")
created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc))
-2
View File
@@ -27,7 +27,6 @@ from .generated_videos import (
from .generation_tasks import (
CreateGenerationTaskCommand,
CreateGenerationTaskUseCase,
GetGenerationTaskUseCase,
)
from .ingest_jobs import SubmitIngestJobCommand, SubmitIngestJobUseCase
from .jobs import (
@@ -64,7 +63,6 @@ __all__ = [
"CreateAssetUseCase",
"CreateGenerationTaskCommand",
"CreateGenerationTaskUseCase",
"GetGenerationTaskUseCase",
"CreateJobCommand",
"CreateJobUseCase",
"CreateProjectCommand",
+304 -311
View File
@@ -1,13 +1,11 @@
"""CosyVoice 语音服务 — 适配阿里云百炼 DashScope API.
"""CosyVoice 语音服务 — Phase 3.
封装阿里云百炼 CosyVoice 语音合成 API,提供:
封装阿里云 CosyVoice 语音合成 API,提供:
- 预置音色列表查询
- 音色克隆(提交 + 轮询状态)
- 语音合成(同步非流式调用)
- 音色克隆(提交任务 + 轮询状态)
- 语音合成(提交任务 + 轮询状态)
API 文档:
- 音色克隆: https://help.aliyun.com/document_detail/3027318.html
- 语音合成: https://help.aliyun.com/zh/model-studio/cosyvoice-tts-http-api
API 文档: https://help.aliyun.com/zh/model-studio/cosyvoice
"""
from __future__ import annotations
@@ -62,34 +60,33 @@ class SynthesizeResult:
class CosyVoiceService:
"""CosyVoice 语音服务.
"""CosyVoice 语音服务。
封装阿里云百炼 CosyVoice API,提供音色克隆和语音合成功能.
接口总览:
- 音色克隆: POST /services/audio/tts/customization (model=voice-enrollment)
- action=create_voice: 创建克隆音色,返回 voice_id(状态 DEPLOYING)
- action=query_voice: 查询音色状态(DEPLOYING / OK / UNDEPLOYED)
- 语音合成: POST /services/audio/tts/SpeechSynthesizer (model=cosyvoice-v3-flash)
- 非流式: 同步返回音频 URL
封装阿里云 CosyVoice API,提供音色克隆和语音合成功能。
支持同步和异步两种模式:
- 同步:API 直接返回结果
- 异步:API 返回 task_id,需要轮询状态
使用示例:
service = CosyVoiceService(
api_key="your-api-key",
base_url="https://dashscope.aliyuncs.com/api/v1",
model="cosyvoice-v3-flash",
base_url="https://dashscope.aliyuncs.com/api/v1/services/aigc/text2audio",
model="cosyvoice-v1",
)
# 获取预置音色
voices = service.list_preset_voices()
# 音色克隆
result = service.clone_voice(audio_url="https://example.com/audio.mp3")
# 语音合成
result = service.synthesize_speech(text="你好世界", voice_id="longxiaochun_v3")
result = service.synthesize_speech(text="你好世界", voice_id="longxiaochun")
"""
# 音色状态轮询配置
CLONE_POLL_INTERVAL = 5.0 # 秒
CLONE_MAX_POLL_ATTEMPTS = 60 # 最多轮询 60 次(5分钟)
# 轮询配置
POLL_INTERVAL = 2.0 # 秒
MAX_POLL_ATTEMPTS = 60 # 最多轮询 60 次(2分钟)
# 重试配置
MAX_RETRIES = 3
@@ -100,70 +97,27 @@ class CosyVoiceService:
api_key: str = "",
base_url: str = "",
model: str = "",
clone_model: str = "",
http_client: Optional[httpx.Client] = None,
audio_url_signer: Optional[callable] = None,
) -> None:
"""初始化 CosyVoice 服务.
"""初始化 CosyVoice 服务。
Args:
api_key: DashScope API Key,为空时从配置读取
base_url: DashScope API Base URL,为空时从配置读取
model: 语音合成模型名称,为空时从配置读取
clone_model: 音色克隆模型名称,为空时从配置读取
api_key: CosyVoice API Key,为空时从配置读取
base_url: CosyVoice API Base URL,为空时从配置读取
model: CosyVoice 模型名称,为空时从配置读取
http_client: 可选的 HTTP 客户端(用于测试注入)
audio_url_signer: 可选的音频URL预签名函数,签名式 fn(url) -> str.
用于私有 bucket 下,将裸 URL 转为预签名 URL,
确保 CosyVoice 服务器能下载参考音频.
"""
settings = get_shared_settings()
self._api_key = api_key or settings.cosyvoice_api_key
self._base_url = base_url or settings.cosyvoice_base_url
self._model = model or settings.cosyvoice_model
self._clone_model = clone_model or getattr(settings, "cosyvoice_clone_model", "voice-enrollment")
self._audio_url_signer = audio_url_signer
# base_url 规范化:去掉末尾的路径残留(兼容旧版配置)
# 旧版 .env 模板中 base_url 包含 /services/aigc/text2audio 完整路径,
# 新版只需 /api/v1,具体路径由代码拼接。这里自动修正,避免配置滞后导致418。
if "/services/aigc/text2audio" in self._base_url:
old_url = self._base_url
# 截取到 /api/v1 为止
idx = self._base_url.find("/api/v1")
if idx >= 0:
self._base_url = self._base_url[: idx + len("/api/v1")]
logger.warning(
"[CosyVoice Config] base_url包含旧版text2audio路径,已自动修正: " "%s -> %s",
old_url,
self._base_url,
)
self._client = http_client or httpx.Client(
timeout=httpx.Timeout(60.0, connect=10.0),
timeout=httpx.Timeout(30.0, connect=10.0),
)
self._owns_client = http_client is None
# 启动时打印配置(脱敏),方便排查环境变量覆盖问题
if self._owns_client:
masked_key = ""
if self._api_key:
if len(self._api_key) > 8:
masked_key = f"{self._api_key[:4]}...{self._api_key[-4:]}"
else:
masked_key = "***"
logger.info(
"[CosyVoice Config] 初始化配置: "
"model=%s, base_url=%s, default_voice=%s, "
"sample_rate=%d, format=%s, api_key=%s",
self._model,
self._base_url,
getattr(settings, "cosyvoice_voice", "(unset)"),
settings.cosyvoice_sample_rate,
settings.cosyvoice_format,
masked_key or "(empty)",
)
def __enter__(self) -> CosyVoiceService:
return self
@@ -178,7 +132,7 @@ class CosyVoiceService:
# ── 预置音色 ─────────────────────────────────────────
def list_preset_voices(self) -> list[PresetVoice]:
"""获取预置音色列表.
"""获取预置音色列表。
Returns:
预置音色列表
@@ -192,22 +146,20 @@ class CosyVoiceService:
audio_url: str,
voice_name: str = "",
language: str = "zh-CN",
target_model: str = "",
) -> dict:
"""提交音色克隆任务(非阻塞).
"""提交音色克隆任务(非阻塞)。
调用百炼 voice-enrollment API 创建克隆音色.
创建后音色状态为 DEPLOYING,需通过 query_voice_status 轮询直到 OK.
只提交任务到 CosyVoice API,不轮询结果。
返回的 dict 包含 task_id(异步)或 voice_id(同步)。
Args:
audio_url: 参考音频 URL(必须公网可访问)
voice_name: 音色名称前缀(字母数字,最多10字符)
language: 语言代码(zh-CN 会转换为 zh)
target_model: 目标合成模型,默认使用当前 model
audio_url: 参考音频 URL
voice_name: 音色名称(可选)
language: 语言代码
Returns:
dict: {"voice_id": str, "status": str, "request_id": str}
voice_id 非空,status 通常为 DEPLOYING
dict: {"task_id": str, "voice_id": str, "request_id": str}
task_id 和 voice_id 至少有一个非空
Raises:
CosyVoiceError: API 调用失败
@@ -219,67 +171,48 @@ class CosyVoiceService:
if not self._api_key:
raise CosyVoiceAuthError("CosyVoice API Key 未配置")
# voice_name 作为 prefix,限制字母数字,最多10字符
# 不符合要求的做清洗
prefix = self._sanitize_prefix(voice_name) if voice_name else "clone"
# 语言转换:zh-CN → zh,保留 ISO 639-1 格式
lang_code = language.split("-")[0].lower() if language else "zh"
target = target_model or self._model
# 如果配置了 audio_url_signer,对音频URL做预签名
# (私有 bucket 下 CosyVoice 服务器无法直接访问裸 URL)
signed_audio_url = audio_url
if self._audio_url_signer:
try:
signed_audio_url = self._audio_url_signer(audio_url)
logger.info("音频URL已预签名: original=%s signed_prefix=%s", audio_url[:80], signed_audio_url[:80])
except Exception as e:
logger.warning("音频URL预签名失败,使用原始URL: %s", e)
payload = {
"model": self._clone_model,
"model": self._model,
"input": {
"action": "create_voice",
"target_model": target,
"prefix": prefix,
"url": signed_audio_url,
"language_hints": [lang_code],
"audio_url": audio_url,
},
"parameters": {
"language": language,
},
}
if voice_name:
payload["parameters"]["voice_name"] = voice_name
response = self._call_api(
method="POST",
path="/services/audio/tts/customization",
path="/services/audio/voice-clone",
json=payload,
timeout=60.0,
)
output = response.get("output", {})
task_id = output.get("task_id", "")
voice_id = output.get("voice_id", "")
status = output.get("status", "DEPLOYING")
request_id = response.get("request_id", "")
if not voice_id:
raise CosyVoiceError(f"CosyVoice API 未返回 voice_id: {response}")
if not task_id and not voice_id:
raise CosyVoiceError(f"CosyVoice API 未返回 task_id 或 voice_id: {response}")
return {
"task_id": task_id,
"voice_id": voice_id,
"status": status,
"request_id": request_id,
}
def query_voice_status(self, voice_id: str) -> dict:
"""查询音色状态(单次查询,不轮询).
def check_task_status(self, task_id: str) -> dict:
"""查询克隆任务状态(单次查询,不轮询)。
Args:
voice_id: 音色 ID
task_id: 任务 ID
Returns:
dict: {"status": str, "target_model": str, "gmt_create": str,
"gmt_modified": str, "resource_link": str}
status 为 DEPLOYING / OK / UNDEPLOYED
dict: {"status": str, "voice_id": str, "message": str}
status 为 SUCCEEDED/FAILED/PENDING/RUNNING
Raises:
CosyVoiceError: API 调用失败
@@ -288,92 +221,40 @@ class CosyVoiceService:
if not self._api_key:
raise CosyVoiceAuthError("CosyVoice API Key 未配置")
if not voice_id:
raise ValueError("voice_id 不能为空")
payload = {
"model": self._clone_model,
"input": {
"action": "query_voice",
"voice_id": voice_id,
},
}
response = self._call_api(
method="POST",
path="/services/audio/tts/customization",
json=payload,
method="GET",
path=f"/tasks/{task_id}",
timeout=30.0,
)
output = response.get("output", {})
status = output.get("task_status", "").upper()
voice_id = output.get("voice_id", "")
message = output.get("message", "")
return {
"status": output.get("status", ""),
"target_model": output.get("target_model", ""),
"gmt_create": output.get("gmt_create", ""),
"gmt_modified": output.get("gmt_modified", ""),
"resource_link": output.get("resource_link", ""),
"status": status,
"voice_id": voice_id,
"message": message,
}
def check_task_status(self, task_id: str) -> dict:
"""查询克隆任务状态(兼容旧接口,实际用 voice_id 查询).
def poll_clone_task(self, task_id: str, timeout: float = 300.0) -> dict:
"""轮询音色克隆任务状态(公开方法)。
为了兼容旧代码,task_id 参数名保留,但实际传的是 voice_id.
供 Celery 后台任务调用,轮询直到完成或超时。
Args:
task_id: 音色 ID(兼容旧接口名)
Returns:
dict: {"status": str, "voice_id": str, "message": str}
"""
result = self.query_voice_status(task_id)
return {
"status": result["status"],
"voice_id": task_id,
"message": "",
}
def poll_clone_task(self, voice_id: str, timeout: float = 300.0) -> dict:
"""轮询音色克隆状态直到完成或超时.
供 Celery 后台任务调用,轮询直到状态变为 OK 或 UNDEPLOYED.
Args:
voice_id: 音色 ID
task_id: CosyVoice 任务 ID
timeout: 超时时间(秒),默认 300
Returns:
dict: {"voice_id": str}
Raises:
CosyVoiceError: 任务失败(状态 UNDEPLOYED)
CosyVoiceError: 任务失败
CosyVoiceTimeoutError: 超时
"""
start_time = time.time()
attempts = 0
while attempts < self.CLONE_MAX_POLL_ATTEMPTS:
elapsed = time.time() - start_time
if elapsed > timeout:
raise CosyVoiceTimeoutError(f"音色克隆任务超时({timeout}秒): voice_id={voice_id}")
result = self.query_voice_status(voice_id)
status = result.get("status", "").upper()
if status == "OK":
return {"voice_id": voice_id}
elif status == "UNDEPLOYED":
raise CosyVoiceError(f"音色克隆任务失败(审核未通过): voice_id={voice_id}")
elif status in ("DEPLOYING", "PENDING", "PROCESSING", ""):
# 继续轮询
time.sleep(self.CLONE_POLL_INTERVAL)
attempts += 1
else:
logger.warning("未知的音色状态: %s (voice_id=%s)", status, voice_id)
time.sleep(self.CLONE_POLL_INTERVAL)
attempts += 1
raise CosyVoiceTimeoutError(f"音色克隆任务轮询次数超限: voice_id={voice_id}")
return self._poll_clone_task(task_id, timeout=timeout)
def clone_voice(
self,
@@ -381,45 +262,122 @@ class CosyVoiceService:
voice_name: str = "",
language: str = "zh-CN",
timeout: float = 300.0,
target_model: str = "",
) -> CloneResult:
"""克隆音色(阻塞,直到完成或超时).
"""克隆音色。
提交音色克隆到百炼 API,并轮询直到状态变为 OK 或超时.
提交音色克隆任务到 CosyVoice API,并轮询直到完成或超时。
Args:
audio_url: 参考音频 URL(必须公网可访问)
voice_name: 音色名称前缀
audio_url: 参考音频 URL
voice_name: 音色名称(可选)
language: 语言代码
timeout: 超时时间(秒)
target_model: 目标合成模型
Returns:
CloneResult: 克隆结果,包含 voice_id
Raises:
CosyVoiceError: API 调用失败或克隆失败
CosyVoiceError: API 调用失败
CosyVoiceTimeoutError: 超时
CosyVoiceAuthError: 认证失败
ValueError: 参数无效
"""
submit_result = self.submit_clone_task(
audio_url=audio_url,
voice_name=voice_name,
language=language,
target_model=target_model,
if not audio_url:
raise ValueError("audio_url 不能为空")
if not self._api_key:
raise CosyVoiceAuthError("CosyVoice API Key 未配置")
# 构建请求
payload = {
"model": self._model,
"input": {
"audio_url": audio_url,
},
"parameters": {
"language": language,
},
}
if voice_name:
payload["parameters"]["voice_name"] = voice_name
# 调用 API
response = self._call_api(
method="POST",
path="/services/audio/voice-clone",
json=payload,
timeout=timeout,
)
voice_id = submit_result["voice_id"]
request_id = submit_result["request_id"]
# 解析响应
output = response.get("output", {})
# 如果创建时已经是 OK 状态,直接返回
if submit_result.get("status", "").upper() == "OK":
return CloneResult(voice_id=voice_id, request_id=request_id)
# 检查是否有 task_id(异步模式)
task_id = output.get("task_id")
voice_id = output.get("voice_id")
# 否则轮询
result = self.poll_clone_task(voice_id, timeout=timeout)
return CloneResult(voice_id=result["voice_id"], request_id=request_id)
if task_id:
# 异步模式:轮询任务状态
result = self._poll_clone_task(task_id, timeout)
return CloneResult(
voice_id=result["voice_id"],
request_id=response.get("request_id", ""),
)
elif voice_id:
# 同步模式:直接返回结果
return CloneResult(
voice_id=voice_id,
request_id=response.get("request_id", ""),
)
else:
raise CosyVoiceError(f"CosyVoice API 未返回 task_id 或 voice_id: {response}")
def _poll_clone_task(self, task_id: str, timeout: float) -> dict:
"""轮询音色克隆任务状态。
Args:
task_id: 任务 ID
timeout: 超时时间(秒)
Returns:
任务结果字典
Raises:
CosyVoiceError: 任务失败
CosyVoiceTimeoutError: 超时
"""
start_time = time.time()
attempts = 0
while attempts < self.MAX_POLL_ATTEMPTS:
elapsed = time.time() - start_time
if elapsed > timeout:
raise CosyVoiceTimeoutError(f"音色克隆任务超时({timeout}秒): task_id={task_id}")
response = self._call_api(
method="GET",
path=f"/tasks/{task_id}",
timeout=30.0,
)
output = response.get("output", {})
status = output.get("task_status", "").upper()
if status == "SUCCEEDED":
voice_id = output.get("voice_id", "")
if not voice_id:
raise CosyVoiceError(f"音色克隆任务成功但未返回 voice_id: {response}")
return {"voice_id": voice_id}
elif status == "FAILED":
error_msg = output.get("message", "未知错误")
raise CosyVoiceError(f"音色克隆任务失败: {error_msg}")
elif status in ("PENDING", "RUNNING"):
# 继续轮询
time.sleep(self.POLL_INTERVAL)
attempts += 1
else:
raise CosyVoiceError(f"未知的任务状态: {status}")
raise CosyVoiceTimeoutError(f"音色克隆任务轮询次数超限: task_id={task_id}")
# ── 语音合成 ─────────────────────────────────────────
@@ -430,12 +388,11 @@ class CosyVoiceService:
sample_rate: int = 0,
format: str = "",
speed: float = 1.0,
volume: int = 50,
) -> dict:
"""提交语音合成任务(同步非流式,直接返回结果).
"""提交语音合成任务(非阻塞)。
CosyVoice SpeechSynthesizer 非流式接口是同步的,
调用后直接返回音频 URL. 此方法保持与旧接口兼容.
只提交任务到 CosyVoice API,不轮询结果。
返回的 dict 包含 task_id(异步)或 audio_url(同步)。
Args:
text: 要合成的文本
@@ -443,11 +400,10 @@ class CosyVoiceService:
sample_rate: 采样率(Hz),0 表示使用配置默认值
format: 输出格式(mp3/wav/pcm),空表示使用配置默认值
speed: 语速(0.5-2.0),1.0 为正常速度
volume: 音量(0-100),默认 50
Returns:
dict: {"audio_url": str, "request_id": str,
"duration": float, "file_size": int}
dict: {"task_id": str, "audio_url": str, "request_id": str}
task_id 和 audio_url 至少有一个非空
Raises:
CosyVoiceError: API 调用失败
@@ -467,47 +423,55 @@ class CosyVoiceService:
"model": self._model,
"input": {
"text": text,
},
"parameters": {
"voice": voice_id,
"format": format or settings.cosyvoice_format,
"sample_rate": sample_rate or settings.cosyvoice_sample_rate,
"format": format or settings.cosyvoice_format,
"rate": speed,
"volume": volume,
},
}
response = self._call_api(
method="POST",
path="/services/audio/tts/SpeechSynthesizer",
path="/services/aigc/text2audio/generation",
json=payload,
timeout=120.0,
timeout=60.0,
)
output = response.get("output", {})
audio = output.get("audio", {})
audio_url = audio.get("url", "")
task_id = output.get("task_id", "")
audio_url = output.get("audio_url", "")
request_id = response.get("request_id", "")
if not audio_url:
raise CosyVoiceError(f"CosyVoice API 未返回 audio_url: {response}")
if not task_id and not audio_url:
raise CosyVoiceError(f"CosyVoice API 未返回 task_id 或 audio_url: {response}")
return {
"task_id": "", # 同步接口无 task_id,兼容旧接口
"task_id": task_id,
"audio_url": audio_url,
"duration": 0.0, # 同步接口不返回 duration
"file_size": 0, # 同步接口不返回 file_size
"duration": output.get("duration", 0.0),
"file_size": output.get("file_size", 0),
"request_id": request_id,
}
def poll_synthesize_task(self, task_id: str, timeout: float = 120.0) -> dict:
"""轮询合成任务(同步接口无需轮询,保留兼容).
"""轮询语音合成任务状态(公开方法)。
CosyVoice SpeechSynthesizer 非流式接口是同步的,
此方法仅为保持接口兼容,实际调用时 task_id 应该为空.
供 Celery 后台任务调用,轮询直到完成或超时。
Args:
task_id: CosyVoice 任务 ID
timeout: 超时时间(秒),默认 120
Returns:
dict: {"audio_url": str, "duration": float, "file_size": int}
Raises:
CosyVoiceError: 同步接口无需轮询
CosyVoiceError: 任务失败
CosyVoiceTimeoutError: 超时
"""
raise CosyVoiceError("CosyVoice 非流式合成接口是同步的,无需轮询. " "请直接使用 submit_synthesize_task().")
return self._poll_synthesize_task(task_id, timeout=timeout)
def synthesize_speech(
self,
@@ -516,13 +480,11 @@ class CosyVoiceService:
sample_rate: int = 0,
format: str = "",
speed: float = 1.0,
volume: int = 50,
timeout: float = 120.0,
) -> SynthesizeResult:
"""语音合成(同步非流式).
"""语音合成。
调用百炼 CosyVoice SpeechSynthesizer 非流式接口,
直接返回合成音频 URL.
提交语音合成任务到 CosyVoice API,并轮询直到完成或超时。
Args:
text: 要合成的文本
@@ -530,52 +492,128 @@ class CosyVoiceService:
sample_rate: 采样率(Hz),0 表示使用配置默认值
format: 输出格式(mp3/wav/pcm),空表示使用配置默认值
speed: 语速(0.5-2.0),1.0 为正常速度
volume: 音量(0-100),默认 50
timeout: 超时时间(秒),保留参数兼容
timeout: 超时时间(秒)
Returns:
SynthesizeResult: 合成结果,包含 audio_url
Raises:
CosyVoiceError: API 调用失败
CosyVoiceTimeoutError: 超时
CosyVoiceAuthError: 认证失败
ValueError: 参数无效
"""
result = self.submit_synthesize_task(
text=text,
voice_id=voice_id,
sample_rate=sample_rate,
format=format,
speed=speed,
volume=volume,
if not text:
raise ValueError("text 不能为空")
if not voice_id:
raise ValueError("voice_id 不能为空")
if not self._api_key:
raise CosyVoiceAuthError("CosyVoice API Key 未配置")
settings = get_shared_settings()
# 构建请求
payload = {
"model": self._model,
"input": {
"text": text,
},
"parameters": {
"voice": voice_id,
"sample_rate": sample_rate or settings.cosyvoice_sample_rate,
"format": format or settings.cosyvoice_format,
"rate": speed,
},
}
# 调用 API
response = self._call_api(
method="POST",
path="/services/aigc/text2audio/generation",
json=payload,
timeout=timeout,
)
return SynthesizeResult(
audio_url=result["audio_url"],
duration=result.get("duration", 0.0),
file_size=result.get("file_size", 0),
request_id=result.get("request_id", ""),
)
# 解析响应
output = response.get("output", {})
# ── 内部方法 ─────────────────────────────────────────
# 检查是否有 task_id(异步模式)
task_id = output.get("task_id")
audio_url = output.get("audio_url")
def _sanitize_prefix(self, name: str) -> str:
"""清洗音色名称为合法的 prefix(字母数字,最多10字符).
if task_id:
# 异步模式:轮询任务状态
result = self._poll_synthesize_task(task_id, timeout)
return SynthesizeResult(
audio_url=result["audio_url"],
duration=result.get("duration", 0.0),
file_size=result.get("file_size", 0),
request_id=response.get("request_id", ""),
)
elif audio_url:
# 同步模式:直接返回结果
return SynthesizeResult(
audio_url=audio_url,
duration=output.get("duration", 0.0),
file_size=output.get("file_size", 0),
request_id=response.get("request_id", ""),
)
else:
raise CosyVoiceError(f"CosyVoice API 未返回 audio_url 或 task_id: {response}")
def _poll_synthesize_task(self, task_id: str, timeout: float) -> dict:
"""轮询语音合成任务状态。
Args:
name: 原始音色名称
task_id: 任务 ID
timeout: 超时时间(秒)
Returns:
清洗后的 prefix
任务结果字典
Raises:
CosyVoiceError: 任务失败
CosyVoiceTimeoutError: 超时
"""
# 只保留字母和数字
cleaned = "".join(c for c in name if c.isalnum())
# 最多10字符
cleaned = cleaned[:10]
# 如果清洗后为空,用默认值
if not cleaned:
cleaned = "clone"
return cleaned
start_time = time.time()
attempts = 0
while attempts < self.MAX_POLL_ATTEMPTS:
elapsed = time.time() - start_time
if elapsed > timeout:
raise CosyVoiceTimeoutError(f"语音合成任务超时({timeout}秒): task_id={task_id}")
response = self._call_api(
method="GET",
path=f"/tasks/{task_id}",
timeout=30.0,
)
output = response.get("output", {})
status = output.get("task_status", "").upper()
if status == "SUCCEEDED":
audio_url = output.get("audio_url", "")
if not audio_url:
raise CosyVoiceError(f"语音合成任务成功但未返回 audio_url: {response}")
return {
"audio_url": audio_url,
"duration": output.get("duration", 0.0),
"file_size": output.get("file_size", 0),
}
elif status == "FAILED":
error_msg = output.get("message", "未知错误")
raise CosyVoiceError(f"语音合成任务失败: {error_msg}")
elif status in ("PENDING", "RUNNING"):
# 继续轮询
time.sleep(self.POLL_INTERVAL)
attempts += 1
else:
raise CosyVoiceError(f"未知的任务状态: {status}")
raise CosyVoiceTimeoutError(f"语音合成任务轮询次数超限: task_id={task_id}")
# ── 内部方法 ─────────────────────────────────────────
def _call_api(
self,
@@ -584,13 +622,13 @@ class CosyVoiceService:
json: Optional[dict] = None,
timeout: float = 30.0,
) -> dict:
"""调用 DashScope API.
"""调用 CosyVoice API。
支持重试和错误处理.
支持重试和错误处理。
Args:
method: HTTP 方法(GET/POST)
path: API 路径(以 / 开头)
path: API 路径
json: 请求体
timeout: 超时时间(秒)
@@ -608,22 +646,6 @@ class CosyVoiceService:
"Content-Type": "application/json",
}
# DEBUG: 打印完整请求信息,用于排查418错误
import json as json_lib
safe_headers = {k: v for k, v in headers.items()}
if "Authorization" in safe_headers:
token = safe_headers["Authorization"]
if len(token) > 20:
safe_headers["Authorization"] = token[:13] + "..." + token[-4:]
logger.info(
"[CosyVoice Debug] 请求详情: " "method=%s, url=%s, headers=%s, body=%s",
method,
url,
safe_headers,
json_lib.dumps(json, ensure_ascii=False) if json else "None",
)
last_error: Optional[Exception] = None
for attempt in range(self.MAX_RETRIES):
@@ -636,58 +658,29 @@ class CosyVoiceService:
timeout=timeout,
)
# DEBUG: 打印响应状态和完整响应体
logger.info(
"[CosyVoice Debug] 响应详情: " "status=%d, body=%s",
response.status_code,
response.text[:2000], # 最多2000字符,避免日志过大
)
# 处理响应
if response.status_code == 200:
return response.json()
elif response.status_code in (401, 403):
raise CosyVoiceAuthError(f"CosyVoice API 认证失败: HTTP {response.status_code}")
elif response.status_code == 400:
# 客户端错误,不重试
body_text = response.text
try:
body = response.json()
code = body.get("code", "")
message = body.get("message", "")
raise CosyVoiceError(f"CosyVoice API 参数错误: HTTP 400, " f"code={code}, message={message}")
except ValueError:
raise CosyVoiceError(f"CosyVoice API 调用失败: HTTP 400, body={body_text}")
elif response.status_code >= 500:
# 服务端错误,可重试
last_error = CosyVoiceError(f"CosyVoice API 服务端错误: HTTP {response.status_code}")
logger.warning(
"CosyVoice API 失败 (尝试 %d/%d): HTTP %d",
attempt + 1,
self.MAX_RETRIES,
response.status_code,
f"CosyVoice API 失败 (尝试 {attempt + 1}/{self.MAX_RETRIES}): " f"HTTP {response.status_code}"
)
else:
# 其他客户端错误,不重试
# 客户端错误,不重试
raise CosyVoiceError(
f"CosyVoice API 调用失败: HTTP {response.status_code}, " f"body={response.text}"
)
except httpx.TimeoutException as e:
last_error = CosyVoiceTimeoutError(f"请求超时: {e}")
logger.warning(
"CosyVoice API 超时 (尝试 %d/%d)",
attempt + 1,
self.MAX_RETRIES,
)
logger.warning(f"CosyVoice API 超时 (尝试 {attempt + 1}/{self.MAX_RETRIES})")
except httpx.RequestError as e:
last_error = CosyVoiceError(f"请求错误: {e}")
logger.warning(
"CosyVoice API 请求错误 (尝试 %d/%d): %s",
attempt + 1,
self.MAX_RETRIES,
e,
)
logger.warning(f"CosyVoice API 请求错误 (尝试 {attempt + 1}/{self.MAX_RETRIES}): {e}")
# 指数退避
if attempt < self.MAX_RETRIES - 1:
+61 -139
View File
@@ -14,6 +14,7 @@ import logging
import os
import shutil
import tempfile
import time
from concurrent.futures import ThreadPoolExecutor, as_completed
from typing import Optional
@@ -26,7 +27,7 @@ from packages.application.cosyvoice_service import (
)
from packages.application.tts_job.audio_merger import AudioMerger
from packages.application.tts_job.text_splitter import split_text
from packages.domain.tts_job import TTSJob, TTSJobStatus
from packages.domain.tts_job import TTSJob
from packages.ports.tts_job_repository import TTSJobRepository
from packages.shared.storage import SharedStorageService, get_shared_storage_service
@@ -189,13 +190,10 @@ class TTSWorkflowService:
return job
def poll_and_process_synthesis(self, job_id: str, timeout: float = 120.0) -> TTSJob:
"""轮询/检查 CosyVoice 合成任务并处理结果.
"""轮询 CosyVoice 合成任务并处理结果。
新 CosyVoice SpeechSynthesizer 非流式接口是同步的,
start_synthesis 阶段通常已经完成. 此方法用于:
1. job 已 completed → 直接返回(同步路径已处理)
2. job 仍在 processing → 重新提交合成(兜底)
3. 分段任务 → 检查分段状态
从 job.metadata 获取 task_id,调用 CosyVoiceService.poll_synthesize_task()
轮询状态,然后通过 process_synthesis_result / process_synthesis_failure 更新 job。
供 Celery 后台任务调用。
"""
@@ -203,38 +201,22 @@ class TTSWorkflowService:
if job is None:
raise TTSJobNotFoundError(f"TTS job {job_id} not found")
# 已完成直接返回(同步路径在 start_synthesis 里已处理)
if job.status == TTSJobStatus.COMPLETED.value:
logger.info(f"TTS 任务已完成,跳过轮询: job_id={job_id}")
return job
# 检查是否为分段合成任务
segment_task_ids = (job.metadata or {}).get("segment_task_ids", [])
if segment_task_ids:
return self._poll_segment_tasks(job)
# 单段模式:同步接口下通常不会走到这里,
# 但如果因为异常导致仍在 processing,重新提交一次
task_id = (job.metadata or {}).get("cosyvoice_task_id", "")
# 新接口(同步):没有 task_id,重新合成
if not task_id:
logger.info(f"TTS 任务无 task_id,重新同步合成: job_id={job_id}")
return self._resynthesize_and_complete(job)
raise ValueError(f"TTSJob {job_id} has no cosyvoice_task_id in metadata")
# 旧接口遗留的 task_id,尝试轮询(兼容过渡)
try:
result = self.cosyvoice_service.poll_synthesize_task(task_id, timeout=timeout)
return self.process_synthesis_result(
job_id,
audio_url=result["audio_url"],
duration=result.get("duration", 0.0),
file_size=result.get("file_size", 0),
)
except CosyVoiceError:
# 旧接口轮询失败,重新同步合成
logger.warning(f"旧 task_id 轮询失败,重新同步合成: job_id={job_id}, task_id={task_id}")
return self._resynthesize_and_complete(job)
result = self.cosyvoice_service.poll_synthesize_task(task_id, timeout=timeout)
return self.process_synthesis_result(
job_id,
audio_url=result["audio_url"],
duration=result.get("duration", 0.0),
file_size=result.get("file_size", 0),
)
def process_synthesis_result(
self,
@@ -275,40 +257,6 @@ class TTSWorkflowService:
logger.info(f"TTS 合成成功: job_id={job_id}, audio_url={permanent_url}")
return job
def _resynthesize_and_complete(self, job: TTSJob) -> TTSJob:
"""重新同步合成并完成任务(兜底路径).
当 poll_and_process_synthesis 发现 job 仍在 processing 且无 task_id 时,
重新调用同步合成接口,转存 OSS 后标记完成。
"""
try:
# 从 metadata 读取合成参数(兼容旧数据,无则用默认值)
job_metadata = job.metadata or {}
speed = float(job_metadata.get("speed", 1.0))
volume = int(job_metadata.get("volume", 50))
result = self.cosyvoice_service.submit_synthesize_task(
text=job.input_text,
voice_id=job.voice_id,
sample_rate=job.sample_rate,
format=job.format,
speed=speed,
volume=volume,
)
audio_url = result.get("audio_url", "")
if not audio_url:
raise CosyVoiceError("重新合成未返回 audio_url")
return self.process_synthesis_result(
job.id,
audio_url=audio_url,
duration=result.get("duration", 0.0),
file_size=result.get("file_size", 0),
)
except Exception as e:
logger.error(f"重新同步合成失败: job_id={job.id}, error={e}")
return self.process_synthesis_failure(job.id, str(e))
def process_synthesis_failure(self, job_id: str, error_message: str) -> TTSJob:
"""处理合成失败结果。
@@ -481,95 +429,69 @@ class TTSWorkflowService:
shutil.rmtree(temp_dir, ignore_errors=True)
def _poll_segment_tasks(self, job: TTSJob) -> TTSJob:
"""分段任务完成检查(适配新同步接口).
新 CosyVoice SpeechSynthesizer 非流式接口为同步接口,
分段任务在提交时应已同步返回 audio_url。
若历史任务处于 processing 且有 segment_task_ids 但缺少 audio_url,
则对缺失分段重新同步合成,全部完成后合并音频。
"""
"""轮询所有分段异步任务,全部完成后合并音频。"""
segment_task_ids: list[str] = (job.metadata or {}).get("segment_task_ids", [])
segment_audio_urls: list[str] = (job.metadata or {}).get("segment_audio_urls", [])
segment_count = len(segment_task_ids)
if segment_count == 0:
logger.warning(f"分段任务无 task_id: job_id={job.id}")
self._handle_segment_failure(job, "分段任务数据异常:无分段信息")
return self.repository.get(job.id)
poll_start = time.monotonic()
poll_timeout = 300.0 # 分段任务超时更长
poll_interval = 2.0
# 从 metadata 读取合成参数
job_metadata = job.metadata or {}
speed = float(job_metadata.get("speed", 1.0))
volume = int(job_metadata.get("volume", 50))
while time.monotonic() - poll_start < poll_timeout:
all_done = True
results: list[dict | None] = [None] * segment_count
# 分段文本(用于缺失段重新合成)
segments = split_text(job.input_text, max_chars=_SEGMENT_THRESHOLD)
for idx, task_id in enumerate(segment_task_ids):
# 已经有音频的分段跳过轮询
if idx < len(segment_audio_urls) and segment_audio_urls[idx]:
results[idx] = {
"audio_url": segment_audio_urls[idx],
"duration": 0.0,
"file_size": 0,
}
continue
results: list[dict | None] = [None] * segment_count
try:
result = self.cosyvoice_service.poll_synthesize_task(task_id, timeout=poll_timeout)
results[idx] = result
except Exception as e:
logger.error(f"分段任务轮询失败: job_id={job.id}, " f"segment={idx}, error={e}")
self._handle_segment_failure(job, f"分段 {idx + 1} 轮询失败: {e}")
return self.repository.get(job.id)
# 已有音频的分段直接用
for idx in range(segment_count):
if idx < len(segment_audio_urls) and segment_audio_urls[idx]:
results[idx] = {
"audio_url": segment_audio_urls[idx],
"duration": 0.0,
"file_size": 0,
}
if results[idx] is None:
all_done = False
# 找出缺失音频的分段索引
missing_indices = [i for i in range(segment_count) if results[i] is None]
if all_done and all(r is not None for r in results):
# 所有分段完成,下载合并
try:
merged_data, total_duration = self._download_and_merge_segments(results, job)
if missing_indices:
logger.info(f"分段任务重新合成缺失段: job_id={job.id}, " f"缺失={len(missing_indices)}/{segment_count}")
# 并发重新合成缺失分段
max_workers = min(len(missing_indices), _MAX_SEGMENT_WORKERS)
with ThreadPoolExecutor(max_workers=max_workers) as executor:
future_to_idx = {}
for idx in missing_indices:
segment_text = segments[idx] if idx < len(segments) else ""
future = executor.submit(
self.cosyvoice_service.submit_synthesize_task,
text=segment_text,
voice_id=job.voice_id,
sample_rate=job.sample_rate,
format=job.format,
speed=speed,
volume=volume,
# 转存 OSS
permanent_url, storage_key = self._upload_merged_to_oss(
merged_data, job.user_id, job.id, job.format
)
future_to_idx[future] = idx
for future in as_completed(future_to_idx):
idx = future_to_idx[future]
try:
results[idx] = future.result()
except Exception as e:
logger.error(f"分段重新合成失败: job_id={job.id}, " f"segment={idx}, error={e}")
self._handle_segment_failure(job, f"分段 {idx + 1} 重新合成失败: {e}")
return self.repository.get(job.id)
job.mark_completed(
output_audio_url=permanent_url,
output_audio_key=storage_key,
duration=total_duration,
file_size=len(merged_data),
)
job = self.repository.update(job)
logger.info(f"分段合成轮询完成: job_id={job.id}, " f"merged_size={len(merged_data)}")
return job
# 所有分段完成,下载合并
if all(r is not None for r in results):
try:
merged_data, total_duration = self._download_and_merge_segments(results, job)
except Exception as e:
self._handle_segment_failure(job, f"分段合并失败: {e}")
return self.repository.get(job.id)
permanent_url, storage_key = self._upload_merged_to_oss(merged_data, job.user_id, job.id, job.format)
# 等待后重试
time.sleep(poll_interval)
job.mark_completed(
output_audio_url=permanent_url,
output_audio_key=storage_key,
duration=total_duration,
file_size=len(merged_data),
)
job = self.repository.update(job)
logger.info(f"分段合成完成(重新合成路径): job_id={job.id}, " f"merged_size={len(merged_data)}")
return job
except Exception as e:
self._handle_segment_failure(job, f"分段合并失败: {e}")
return self.repository.get(job.id)
# 理论上不会到这里(全部重新合成要么成功要么失败)
self._handle_segment_failure(job, "分段合成结果不完整")
# 超时
self._handle_segment_failure(job, "分段合成轮询超时(300 秒)")
return self.repository.get(job.id)
def _handle_segment_failure(self, job: TTSJob, error_message: str) -> None:
+7 -12
View File
@@ -114,16 +114,14 @@ class VoiceCloneWorkflowService:
language=language,
)
# 4. 保存 voice_id / request_id 到 metadata
# 注意:key 保留 cosyvoice_task_id 以兼容旧数据,实际存的是 voice_id
# 4. 保存 task_id / voice_id 到 metadata
task_metadata = dict(profile.metadata)
task_metadata["cosyvoice_task_id"] = submit_result.get("voice_id", "")
task_metadata["cosyvoice_task_id"] = submit_result.get("task_id", "")
task_metadata["cosyvoice_request_id"] = submit_result.get("request_id", "")
# 如果 CosyVoice 直接返回了 OK 状态,直接标记 ready
# 如果 CosyVoice 同步返回了 voice_id,直接标记 ready
voice_id = submit_result.get("voice_id", "")
status = submit_result.get("status", "").upper()
if voice_id and status == "OK":
if voice_id:
profile.mark_ready(voice_id)
profile.metadata = task_metadata
profile = self.repository.update(profile)
@@ -132,9 +130,7 @@ class VoiceCloneWorkflowService:
profile.metadata = task_metadata
profile = self.repository.update(profile)
logger.info(
f"音色克隆任务已提交: profile_id={profile.id}, " f"voice_id={submit_result.get('voice_id')}"
)
logger.info(f"音色克隆任务已提交: profile_id={profile.id}, " f"task_id={submit_result.get('task_id')}")
except (CosyVoiceError, CosyVoiceAuthError) as e:
# CosyVoice 提交失败,标记为 failed
@@ -251,12 +247,11 @@ class VoiceCloneWorkflowService:
)
task_metadata = dict(profile.metadata)
task_metadata["cosyvoice_task_id"] = submit_result.get("voice_id", "")
task_metadata["cosyvoice_task_id"] = submit_result.get("task_id", "")
task_metadata["cosyvoice_request_id"] = submit_result.get("request_id", "")
voice_id = submit_result.get("voice_id", "")
status = submit_result.get("status", "").upper()
if voice_id and status == "OK":
if voice_id:
profile.mark_ready(voice_id)
profile.metadata = task_metadata
profile = self.repository.update(profile)
Executable → Regular
-36
View File
@@ -134,25 +134,6 @@ class AssetStatus(StrEnum):
PROCESSING = "processing"
ERROR = "error"
@classmethod
def _missing_(cls, value: object) -> "AssetStatus":
"""兼容历史数据,避免枚举转换失败导致500。
- uploaded → READY(早期版本用 uploaded 表示上传完成)
- 其他未知值 → READY(兜底,不阻塞业务)
"""
if isinstance(value, str):
normalized = value.strip().lower()
if normalized in ("uploaded", "success", "ok", "done", "complete"):
return cls.READY
if normalized in ("upload", "uploading_start", "upload_start"):
return cls.UPLOADING
if normalized in ("failed", "fail", "err"):
return cls.ERROR
if normalized in ("process", "processing", "running", "run"):
return cls.PROCESSING
return cls.READY
class ClassificationStatus(StrEnum):
PENDING = "pending"
@@ -160,23 +141,6 @@ class ClassificationStatus(StrEnum):
COMPLETED = "completed"
FAILED = "failed"
@classmethod
def _missing_(cls, value: object) -> "ClassificationStatus":
"""兼容历史数据,避免枚举转换失败导致500。
- done → COMPLETED(早期版本用 done 表示完成)
- 其他未知值 → PENDING(兜底,不阻塞业务)
"""
if isinstance(value, str):
normalized = value.strip().lower()
if normalized in ("done", "success", "finished", "complete"):
return cls.COMPLETED
if normalized in ("fail", "error", "err"):
return cls.FAILED
if normalized in ("process", "processing", "running", "run"):
return cls.PROCESSING
return cls.PENDING
@dataclass(slots=True)
class Asset:
-39
View File
@@ -8,7 +8,6 @@
from __future__ import annotations
import json
import sys
from dataclasses import dataclass, field
from datetime import datetime, timezone
@@ -86,7 +85,6 @@ class GenerationTask:
created_by_user_id: str = ""
asset_select_mode: str = ""
batch_id: str = ""
logs: str = "[]"
created_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc))
@classmethod
@@ -229,43 +227,6 @@ class GenerationTask:
self.transition_to(GenerationTaskStatus.CANCELLED)
self.completed_at = datetime.now(timezone.utc)
# ── 日志辅助 ────────────────────────────────────────────────────────────
_MAX_LOGS = 200
def append_log(self, stage: str, message: str, level: str = "INFO", **kwargs) -> None:
"""追加一条结构化日志到 logs 字段。
Args:
stage: 阶段名称(如 "接收任务"、"下载素材"、"渲染")
message: 日志消息
level: 日志级别(INFO / WARN / ERROR)
**kwargs: 额外字段(如 asset_id、duration 等)
"""
try:
entries = json.loads(self.logs) if self.logs else []
except (json.JSONDecodeError, TypeError):
entries = []
entry = {
"ts": datetime.now(timezone.utc).isoformat(),
"level": level,
"stage": stage,
"message": message,
**kwargs,
}
entries.append(entry)
# 限制最多保留 _MAX_LOGS 条,防止字段过大
if len(entries) > self._MAX_LOGS:
entries = entries[-self._MAX_LOGS :]
self.logs = json.dumps(entries, ensure_ascii=False)
def get_logs(self) -> list[dict]:
"""解析 logs 字段为 list[dict]。"""
try:
return json.loads(self.logs) if self.logs else []
except (json.JSONDecodeError, TypeError):
return []
def mark_pending_from_failed(self) -> None:
"""从失败状态重置为待处理(用于重试)。
Executable → Regular
+9 -9
View File
@@ -16,7 +16,7 @@ class PresetVoice:
"""预置音色定义。
Attributes:
voice_id: CosyVoice 模型音色名(如 longxiaochun_v3)
voice_id: CosyVoice 模型音色名(如 longxiaochun)
name: 中文展示名
description: 音色描述
gender: 性别(male/female)
@@ -49,7 +49,7 @@ class PresetVoice:
# 预置音色列表(阿里云 CosyVoice 真实可用音色)
PRESET_VOICES: list[PresetVoice] = [
PresetVoice(
voice_id="longxiaochun_v3",
voice_id="longxiaochun",
name="龙小淳",
description="温柔女声,适合情感类内容",
gender="female",
@@ -57,7 +57,7 @@ PRESET_VOICES: list[PresetVoice] = [
tags=["温柔", "女声", "情感"],
),
PresetVoice(
voice_id="longxiaoxia_v3",
voice_id="longxiaoxia",
name="龙小夏",
description="知性女声,适合新闻播报",
gender="female",
@@ -65,7 +65,7 @@ PRESET_VOICES: list[PresetVoice] = [
tags=["知性", "女声", "播报"],
),
PresetVoice(
voice_id="longxiaochen_v3",
voice_id="longxiaochen",
name="龙小晨",
description="磁性男声,适合有声书",
gender="male",
@@ -73,7 +73,7 @@ PRESET_VOICES: list[PresetVoice] = [
tags=["磁性", "男声", "有声书"],
),
PresetVoice(
voice_id="longyue_v3",
voice_id="longyue",
name="龙悦",
description="甜美女声,适合广告配音",
gender="female",
@@ -81,7 +81,7 @@ PRESET_VOICES: list[PresetVoice] = [
tags=["甜美", "女声", "广告"],
),
PresetVoice(
voice_id="longshu_v3",
voice_id="longshu",
name="龙书",
description="沉稳男声,适合教育讲解",
gender="male",
@@ -89,7 +89,7 @@ PRESET_VOICES: list[PresetVoice] = [
tags=["沉稳", "男声", "教育"],
),
PresetVoice(
voice_id="longjing_v3",
voice_id="longjing",
name="龙静",
description="优雅女声,适合纪录片解说",
gender="female",
@@ -97,7 +97,7 @@ PRESET_VOICES: list[PresetVoice] = [
tags=["优雅", "女声", "纪录片"],
),
PresetVoice(
voice_id="longbo_v3",
voice_id="longbo",
name="龙博",
description="浑厚男声,适合科技类内容",
gender="male",
@@ -105,7 +105,7 @@ PRESET_VOICES: list[PresetVoice] = [
tags=["浑厚", "男声", "科技"],
),
PresetVoice(
voice_id="longtian_v3",
voice_id="longtian",
name="龙甜",
description="活泼女声,适合短视频配音",
gender="female",
-4
View File
@@ -16,10 +16,6 @@ class GenerationTaskRepository(Protocol):
def count_by_user(self, user_id: str) -> int: ...
def count_pending_by_user(self, user_id: str) -> int: ...
def count_pending_total(self) -> int: ...
def list_recent_by_user(self, user_id: str, limit: int = 5) -> list[GenerationTask]: ...
def list_by_source_edit_plan(self, plan_id: str) -> list[GenerationTask]: ...
Executable → Regular
+4 -6
View File
@@ -29,15 +29,13 @@ class SharedSettings(BaseSettings):
oss_access_key_secret: str = ""
oss_bucket_name: str = "xiaoxia-autocut"
# CosyVoice (阿里云百炼语音合成)
# CosyVoice (阿里云语音合成)
cosyvoice_api_key: str = ""
cosyvoice_base_url: str = "https://dashscope.aliyuncs.com/api/v1"
cosyvoice_model: str = "cosyvoice-v3-flash"
cosyvoice_voice: str = "longxiaochun_v3" # 默认音色(v3 系列系统音色带 _v3 后缀)
cosyvoice_base_url: str = "https://dashscope.aliyuncs.com/api/v1/services/aigc/text2audio"
cosyvoice_model: str = "cosyvoice-v1"
cosyvoice_voice: str = "longxiaochun" # 默认音色
cosyvoice_sample_rate: int = 22050
cosyvoice_format: str = "mp3" # 输出格式:mp3/wav/pcm
# 音色克隆模型名(固定为 voice-enrollment)
cosyvoice_clone_model: str = "voice-enrollment"
# Environment
environment: str = "development"
Executable → Regular
-63
View File
@@ -1,70 +1,7 @@
[tool.black]
line-length = 120
target-version = ["py312"]
extend-exclude = '''
(
\.git
| \.cache
| \.pytest_cache
| \.mypy_cache
| __pycache__
| node_modules
| \.venv
| venv
| build
| dist
| \.next
| out
| coverage
)
'''
[tool.isort]
profile = "black"
line_length = 120
extend_skip_glob = [
".git/**",
".cache/**",
".pytest_cache/**",
".mypy_cache/**",
"__pycache__/**",
"node_modules/**",
".venv/**",
"venv/**",
"build/**",
"dist/**",
".next/**",
"out/**",
"coverage/**",
]
[tool.coverage.run]
source = ["apps/api/app", "packages"]
omit = [
"*/migrations/*",
"*/tests/*",
"*/test_*.py",
"*/site-packages/*",
]
branch = true
[tool.coverage.report]
exclude_lines = [
"pragma: no cover",
"def __repr__",
"if __name__ == .__main__.:",
"raise NotImplementedError",
"pass",
"if TYPE_CHECKING:",
"class .*Protocol",
"@abstractmethod",
"raise AssertionError",
"raise RuntimeError",
"if 0:",
"if __debug__:",
]
show_missing = true
skip_covered = false
[tool.coverage.xml]
output = "coverage.xml"
+5 -65
View File
@@ -30,8 +30,7 @@ fi
# ---- Registry 配置 ----
REGISTRY="${REGISTRY:-git.xiaoxiajianji.com/xiaoxia/xiaoxia-saas}"
CACHE_REGISTRY="${CACHE_REGISTRY:-git.xiaoxiajianji.com/xiaoxia/xiaoxia-saas}"
# 主缓存 tag:develop 分支构建时写入,所有分支读取
CACHE_TAG_PRIMARY="${CACHE_TAG:-develop}"
CACHE_TAG="${CACHE_TAG:-release}"
API_IMAGE="xiaoxia-saas-api:$VERSION"
WORKER_IMAGE="xiaoxia-saas-worker:$VERSION"
@@ -46,7 +45,6 @@ REGISTRY_WEB="${REGISTRY}/xiaoxia-saas-web:$VERSION"
USE_CACHE=0
USE_PUSH=0
CACHE_WRITE=0
# 检查 buildx 和 Registry 认证
if docker buildx version >/dev/null 2>&1; then
@@ -56,13 +54,9 @@ if docker buildx version >/dev/null 2>&1; then
docker buildx use default 2>/dev/null || true
fi
# ---- 缓存读写策略(按分支隔离)----
# 默认只读不写,防止 feature 分支污染主缓存
# 只有 develop/main 分支才写回缓存
BRANCH_NAME="${GITHUB_REF_NAME:-${CI_COMMIT_BRANCH:-unknown}}"
echo "=== Building API image ==="
if [ "$USE_CACHE" -eq 1 ]; then
docker buildx build \
--build-arg APP_VERSION="$VERSION" \
--cache-from "type=registry,ref=${CACHE_REGISTRY}/api-cache:${CACHE_TAG},ignore-error=true" \
--cache-to "type=registry,ref=${CACHE_REGISTRY}/api-cache:${CACHE_TAG},mode=max" \
-f infra/docker/api.Dockerfile \
@@ -70,59 +64,12 @@ if [ "$USE_CACHE" -eq 1 ]; then
--load \
.
else
docker build --pull=false --build-arg APP_VERSION="$VERSION" -f infra/docker/api.Dockerfile -t "$API_IMAGE" -t "$API_LATEST" .
docker build --pull=false -f infra/docker/api.Dockerfile -t "$API_IMAGE" -t "$API_LATEST" .
fi
build_with_cache() {
# usage: build_with_cache <image_name> <dockerfile> <extra_args...>
IMG_NAME="$1"
DOCKERFILE="$2"
shift 2
EXTRA_ARGS="$*"
CACHE_FROM="type=registry,ref=${CACHE_REGISTRY}/${IMG_NAME}-cache:${CACHE_TAG_PRIMARY},ignore-error=true"
if [ "$CACHE_WRITE" -eq 1 ]; then
CACHE_TO="type=registry,ref=${CACHE_REGISTRY}/${IMG_NAME}-cache:${CACHE_TAG_PRIMARY},mode=max"
echo " cache: read+write from ${CACHE_REGISTRY}/${IMG_NAME}-cache:${CACHE_TAG_PRIMARY}"
else
CACHE_TO=""
echo " cache: read-only from ${CACHE_REGISTRY}/${IMG_NAME}-cache:${CACHE_TAG_PRIMARY}"
fi
if [ "$USE_CACHE" -eq 1 ]; then
if [ -n "$CACHE_TO" ]; then
docker buildx build \
$EXTRA_ARGS \
--cache-from "$CACHE_FROM" \
--cache-to "$CACHE_TO" \
-f "$DOCKERFILE" \
-t "$IMG_NAME:$VERSION" \
--load \
.
else
docker buildx build \
$EXTRA_ARGS \
--cache-from "$CACHE_FROM" \
-f "$DOCKERFILE" \
-t "$IMG_NAME:$VERSION" \
--load \
.
fi
else
docker build --pull=false $EXTRA_ARGS -f "$DOCKERFILE" -t "$IMG_NAME:$VERSION" .
fi
}
echo "=== Building API image ==="
build_with_cache "api" "infra/docker/api.Dockerfile" \
"--build-arg APP_VERSION=$VERSION"
docker tag "$API_IMAGE" "$API_LATEST"
echo "=== Building Worker image ==="
if [ "$USE_CACHE" -eq 1 ]; then
docker buildx build \
--build-arg APP_VERSION="$VERSION" \
--cache-from "type=registry,ref=${CACHE_REGISTRY}/worker-cache:${CACHE_TAG},ignore-error=true" \
--cache-to "type=registry,ref=${CACHE_REGISTRY}/worker-cache:${CACHE_TAG},mode=max" \
-f infra/docker/worker.Dockerfile \
@@ -130,20 +77,13 @@ if [ "$USE_CACHE" -eq 1 ]; then
--load \
.
else
docker build --pull=false --build-arg APP_VERSION="$VERSION" -f infra/docker/worker.Dockerfile -t "$WORKER_IMAGE" -t "$WORKER_LATEST" .
docker build --pull=false -f infra/docker/worker.Dockerfile -t "$WORKER_IMAGE" -t "$WORKER_LATEST" .
fi
echo "=== Building Web image (with buildx cache) ==="
# 先构建前端产物(使用持久化 npm 缓存卷)
NPM_CACHE_VOLUME="xiaoxia-npm-cache"
if ! docker volume inspect "$NPM_CACHE_VOLUME" >/dev/null 2>&1; then
docker volume create "$NPM_CACHE_VOLUME" >/dev/null
echo " Created npm cache volume: $NPM_CACHE_VOLUME"
fi
# 先构建前端产物
docker run --rm \
-v "$PWD:/workspace" \
-v "$NPM_CACHE_VOLUME:/workspace/apps/web/node_modules" \
-w /workspace/apps/web \
docker.m.daocloud.io/library/node:20 \
sh -lc "npm ci && npm run build"
-34
View File
@@ -1,34 +0,0 @@
#!/usr/bin/env python3
"""解析 coverage.xml 并输出覆盖率汇总。"""
import os
import sys
import xml.etree.ElementTree as ET
THRESHOLD = int(os.environ.get("COVERAGE_THRESHOLD", 65)) # 行覆盖率门槛,百分比,可通过环境变量覆盖
def main() -> int:
try:
tree = ET.parse("coverage.xml")
except FileNotFoundError:
print("coverage.xml 不存在,跳过汇总")
return 0
root = tree.getroot()
line_rate = float(root.get("line-rate", 0)) * 100
branch_rate = float(root.get("branch-rate", 0)) * 100
lines_covered = int(root.get("lines-covered", 0))
lines_valid = int(root.get("lines-valid", 0))
print(f"行覆盖率: {line_rate:.2f}% ({lines_covered}/{lines_valid})")
print(f"分支覆盖率: {branch_rate:.2f}%")
print(f"门槛: {THRESHOLD}%")
status = "PASS ✅" if line_rate >= THRESHOLD else "FAIL ❌"
print(f"状态: {status}")
return 0 if line_rate >= THRESHOLD else 1
if __name__ == "__main__":
sys.exit(main())
-83
View File
@@ -1,83 +0,0 @@
#!/usr/bin/env python3
"""发送 CI 失败通知到飞书/项目群 webhook。"""
import json
import os
import sys
import urllib.request
def main() -> int:
webhook = os.environ.get("CI_NOTIFY_WEBHOOK", "")
if not webhook:
print("未配置 CI_NOTIFY_WEBHOOK,跳过通知")
print("如需启用,请在仓库 Settings -> Secrets and variables -> Actions 中添加 CI_NOTIFY_WEBHOOK")
return 0
failed_job = os.environ.get("FAILED_JOB", "Unknown Job")
branch = os.environ.get("GITHUB_REF_NAME", "unknown")
commit = os.environ.get("GITHUB_SHA", "unknown")[:8]
actor = os.environ.get("GITHUB_ACTOR", "unknown")
run_id = os.environ.get("GITHUB_RUN_ID", "unknown")
repo = os.environ.get("GITHUB_REPOSITORY", "unknown")
run_url = f"https://git.xiaoxiajianji.com/{repo}/actions/runs/{run_id}"
payload = {
"msg_type": "interactive",
"card": {
"header": {
"title": {
"tag": "plain_text",
"content": "❌ CI 构建失败",
},
"status": "red",
},
"elements": [
{
"tag": "div",
"text": {
"tag": "lark_md",
"content": (
f"**任务**: {failed_job}\n"
f"**分支**: {branch}\n"
f"**提交**: {commit}\n"
f"**提交者**: {actor}\n"
f"**Run ID**: {run_id}"
),
},
},
{
"tag": "action",
"actions": [
{
"tag": "button",
"text": {"tag": "plain_text", "content": "查看失败日志"},
"url": run_url,
"type": "danger",
}
],
},
],
},
}
data = json.dumps(payload).encode("utf-8")
req = urllib.request.Request(
webhook,
data=data,
headers={"Content-Type": "application/json"},
method="POST",
)
try:
with urllib.request.urlopen(req, timeout=10) as resp:
resp.read()
print("通知已发送")
except Exception as e:
print(f"通知发送失败: {e}", file=sys.stderr)
return 1
return 0
if __name__ == "__main__":
sys.exit(main())
-83
View File
@@ -1,83 +0,0 @@
#!/usr/bin/env python3
"""发送 CI 成功通知到飞书/项目群 webhook。"""
import json
import os
import sys
import urllib.request
def main() -> int:
webhook = os.environ.get("CI_NOTIFY_WEBHOOK", "")
if not webhook:
print("未配置 CI_NOTIFY_WEBHOOK,跳过成功通知")
print("如需启用,请在仓库 Settings -> Secrets and variables -> Actions 中添加 CI_NOTIFY_WEBHOOK")
return 0
success_job = os.environ.get("SUCCESS_JOB", "Unknown Job")
branch = os.environ.get("GITHUB_REF_NAME", "unknown")
commit = os.environ.get("GITHUB_SHA", "unknown")[:8]
actor = os.environ.get("GITHUB_ACTOR", "unknown")
run_id = os.environ.get("GITHUB_RUN_ID", "unknown")
repo = os.environ.get("GITHUB_REPOSITORY", "unknown")
run_url = f"https://git.xiaoxiajianji.com/{repo}/actions/runs/{run_id}"
payload = {
"msg_type": "interactive",
"card": {
"header": {
"title": {
"tag": "plain_text",
"content": "✅ CI 构建成功",
},
"status": "green",
},
"elements": [
{
"tag": "div",
"text": {
"tag": "lark_md",
"content": (
f"**任务**: {success_job}\n"
f"**分支**: {branch}\n"
f"**提交**: {commit}\n"
f"**提交者**: {actor}\n"
f"**Run ID**: {run_id}"
),
},
},
{
"tag": "action",
"actions": [
{
"tag": "button",
"text": {"tag": "plain_text", "content": "查看构建详情"},
"url": run_url,
"type": "primary",
}
],
},
],
},
}
data = json.dumps(payload).encode("utf-8")
req = urllib.request.Request(
webhook,
data=data,
headers={"Content-Type": "application/json"},
method="POST",
)
try:
with urllib.request.urlopen(req, timeout=10) as resp:
resp.read()
print("成功通知已发送")
except Exception as e:
print(f"成功通知发送失败: {e}", file=sys.stderr)
return 1
return 0
if __name__ == "__main__":
sys.exit(main())
-1
View File
@@ -3,7 +3,6 @@ max-line-length = 120
extend-ignore = E203,W503,E501,E302,E402,E722,W291,W293,F401,F403,F405,F841
exclude =
.git,
.cache,
__pycache__,
.venv,
.venv-ci-root,
+496
View File
@@ -0,0 +1,496 @@
"""
仪表盘 API 集成测试。
覆盖端点:
- GET /dashboard/overview — 仪表盘概览
验证返回数据结构、空数据场景、数据汇总正确性。
"""
from __future__ import annotations
import os
import sys
from datetime import datetime, timezone
# ── 环境变量 & sys.path(必须在导入 app.* 之前设置) ──────────────────────────
os.environ.setdefault("JWT_SECRET_KEY", "unit-test-secret-key-for-testing")
os.environ.setdefault("DATABASE_URL", "sqlite:///test.db")
import pytest
from fastapi import FastAPI
from fastapi.testclient import TestClient
sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..", "..", "apps", "api"))
from app.api.routes.dashboard import router
from app.auth import AuthenticatedUser, get_current_user
from app.dependencies import (
get_asset_repository,
get_generation_task_repository,
get_project_repository,
get_title_library_repository,
get_voice_library_repository,
)
from packages.domain.entities import Project, User
from packages.domain.generation_task import GenerationTask, GenerationTaskStatus
# ---------------------------------------------------------------------------
# 1. 内存 Repository
# ---------------------------------------------------------------------------
class InMemoryProjectRepository:
def __init__(self):
self._projects: dict[str, Project] = {}
def save(self, project: Project) -> None:
self._projects[project.id] = project
def find_by_id(self, project_id: str):
return self._projects.get(project_id)
def find_by_owner_user_id(self, owner_user_id: str):
return [p for p in self._projects.values() if p.owner_user_id == owner_user_id]
def find_accessible_projects(self, user_id: str):
return [p for p in self._projects.values() if p.owner_user_id == user_id]
def count_by_owner(self, owner_user_id: str) -> int:
return len(self.find_by_owner_user_id(owner_user_id))
def delete(self, project_id: str) -> bool:
if project_id in self._projects:
del self._projects[project_id]
return True
return False
class InMemoryAssetRepository:
def __init__(self):
self._assets = []
def add_asset(self, project_id: str, storage_size: int = 0):
self._assets.append({"project_id": project_id, "storage_size": storage_size})
def count_by_project_ids(self, project_ids: list[str]) -> int:
return sum(1 for a in self._assets if a["project_id"] in project_ids)
def sum_storage_by_project_ids(self, project_ids: list[str]) -> int:
return sum(a["storage_size"] for a in self._assets if a["project_id"] in project_ids)
# 其他方法占位
def create(self, asset):
return asset
def find_by_id(self, asset_id):
return None
def find_by_project(self, project_id, **kwargs):
return []
def find_by_library(self, library_id, **kwargs):
return []
def update(self, asset):
return asset
def delete(self, asset_id):
return False
def batch_delete(self, asset_ids):
return 0
def search_candidates(self, **kwargs):
return []
def find_by_tag_ids(self, tag_ids):
return []
def count_by_project(self, project_id):
return 0
def find_by_library_and_file_type(self, library_id, file_type):
return []
def find_by_library_and_file_hash(self, library_id, file_hash):
return None
class InMemoryGenerationTaskRepository:
def __init__(self):
self._tasks = {}
def add_task(self, task: GenerationTask):
self._tasks[task.id] = task
def count_by_user(self, user_id: str) -> int:
return len([t for t in self._tasks.values() if t.created_by_user_id == user_id])
def list_recent_by_user(self, user_id: str, limit: int = 5) -> list:
user_tasks = [t for t in self._tasks.values() if t.created_by_user_id == user_id]
# 按 created_at 倒序
user_tasks.sort(key=lambda t: t.created_at, reverse=True)
return user_tasks[:limit]
# 其他方法占位
def create(self, task):
return task
def get(self, task_id):
return None
def list_by_project(self, project_id):
return []
def list_by_user(self, user_id):
return []
def list_by_source_edit_plan(self, plan_id):
return []
def update(self, task):
return task
class InMemoryTitleLibraryRepository:
def __init__(self):
self._items = {}
def add_item(self, user_id: str):
from uuid import uuid4
item_id = uuid4().hex
self._items[item_id] = {"id": item_id, "user_id": user_id}
return item_id
def count_by_user(self, user_id: str, is_active: bool = True) -> int:
return len([i for i in self._items.values() if i["user_id"] == user_id])
# 其他方法占位
def list_by_user(self, user_id, **kwargs):
return []
def get(self, title_id, user_id):
return None
def create(self, item):
return item
def update(self, item):
return item
def delete(self, title_id, user_id):
return False
class InMemoryVoiceLibraryRepository:
def __init__(self):
self._items = {}
def add_item(self, user_id: str):
from uuid import uuid4
item_id = uuid4().hex
self._items[item_id] = {"id": item_id, "user_id": user_id}
return item_id
def count_by_user(self, user_id: str) -> int:
return len([i for i in self._items.values() if i["user_id"] == user_id])
# 其他方法占位
def list_by_user(self, user_id, **kwargs):
return []
def get(self, voice_id, user_id):
return None
def create(self, item):
return item
def update(self, item):
return item
def delete(self, voice_id, user_id):
return False
# ---------------------------------------------------------------------------
# 2. 辅助函数
# ---------------------------------------------------------------------------
def _make_user(**overrides) -> User:
defaults = dict(
id="user-test-001",
email="test@example.com",
display_name="Test User",
username="testuser",
subscription_plan="free",
subscription_status="active",
max_projects=3,
max_storage_gb=10,
created_at=datetime(2026, 1, 1, tzinfo=timezone.utc),
)
defaults.update(overrides)
return User(**defaults)
def _make_project(project_id: str, owner_user_id: str = "user-test-001") -> Project:
return Project(
id=project_id,
name=f"Project {project_id}",
owner_user_id=owner_user_id,
description="",
created_at=datetime(2026, 1, 1, tzinfo=timezone.utc),
)
def _make_generation_task(
task_id: str,
user_id: str = "user-test-001",
status: GenerationTaskStatus = GenerationTaskStatus.COMPLETED,
created_at: datetime | None = None,
) -> GenerationTask:
return GenerationTask(
id=task_id,
project_id="proj-1",
asset_library_id="lib-1",
created_by_user_id=user_id,
status=status,
error_message="",
created_at=created_at or datetime.now(timezone.utc),
started_at=datetime.now(timezone.utc) if status != GenerationTaskStatus.PENDING else None,
completed_at=datetime.now(timezone.utc) if status == GenerationTaskStatus.COMPLETED else None,
)
# ---------------------------------------------------------------------------
# 3. Fixtures
# ---------------------------------------------------------------------------
@pytest.fixture
def project_repo():
repo = InMemoryProjectRepository()
repo.save(_make_project("proj-1", "user-test-001"))
repo.save(_make_project("proj-2", "user-test-001"))
repo.save(_make_project("proj-other", "other-user"))
return repo
@pytest.fixture
def asset_repo():
return InMemoryAssetRepository()
@pytest.fixture
def generation_task_repo():
return InMemoryGenerationTaskRepository()
@pytest.fixture
def title_library_repo():
return InMemoryTitleLibraryRepository()
@pytest.fixture
def voice_library_repo():
return InMemoryVoiceLibraryRepository()
@pytest.fixture
def client(project_repo, asset_repo, generation_task_repo, title_library_repo, voice_library_repo):
"""创建带有依赖覆盖的 TestClient。"""
test_app = FastAPI()
test_app.include_router(router, prefix="/dashboard")
def _override_current_user():
return AuthenticatedUser(user=_make_user())
test_app.dependency_overrides[get_current_user] = _override_current_user
test_app.dependency_overrides[get_project_repository] = lambda: project_repo
test_app.dependency_overrides[get_asset_repository] = lambda: asset_repo
test_app.dependency_overrides[get_generation_task_repository] = lambda: generation_task_repo
test_app.dependency_overrides[get_title_library_repository] = lambda: title_library_repo
test_app.dependency_overrides[get_voice_library_repository] = lambda: voice_library_repo
yield TestClient(test_app)
test_app.dependency_overrides.clear()
# ---------------------------------------------------------------------------
# 4. GET /overview — 仪表盘概览
# ---------------------------------------------------------------------------
class TestDashboardOverview:
"""仪表盘概览端点测试。"""
def test_empty_data_returns_zeros(self, client):
"""空数据时所有计数为 0。"""
resp = client.get("/dashboard/overview")
assert resp.status_code == 200
data = resp.json()
assert data["total_assets"] == 0
assert data["used_storage_bytes"] == 0
assert data["total_titles"] == 0
assert data["total_voices"] == 0
assert data["total_tasks"] == 0
assert data["total_products"] == 2 # fixture 中有 2 个项目
assert data["recent_tasks"] == []
def test_assets_count_and_storage(self, client, asset_repo):
"""素材统计正确。"""
asset_repo.add_asset("proj-1", 1024)
asset_repo.add_asset("proj-1", 2048)
asset_repo.add_asset("proj-2", 4096)
# 其他用户的不计入
asset_repo.add_asset("proj-other", 9999)
resp = client.get("/dashboard/overview")
data = resp.json()
assert data["total_assets"] == 3
assert data["used_storage_bytes"] == 1024 + 2048 + 4096
def test_title_library_count(self, client, title_library_repo):
"""标题库统计正确。"""
title_library_repo.add_item("user-test-001")
title_library_repo.add_item("user-test-001")
title_library_repo.add_item("user-test-001")
title_library_repo.add_item("other-user")
resp = client.get("/dashboard/overview")
data = resp.json()
assert data["total_titles"] == 3
def test_voice_library_count(self, client, voice_library_repo):
"""配音库统计正确。"""
voice_library_repo.add_item("user-test-001")
voice_library_repo.add_item("other-user")
resp = client.get("/dashboard/overview")
data = resp.json()
assert data["total_voices"] == 1
def test_generation_tasks_count(self, client, generation_task_repo):
"""生成任务统计正确。"""
generation_task_repo.add_task(_make_generation_task("task-1"))
generation_task_repo.add_task(_make_generation_task("task-2"))
generation_task_repo.add_task(_make_generation_task("task-other", user_id="other-user"))
resp = client.get("/dashboard/overview")
data = resp.json()
assert data["total_tasks"] == 2
def test_recent_tasks_limited_to_5(self, client, generation_task_repo):
"""最近任务最多返回 5 个。"""
for i in range(10):
task = _make_generation_task(f"task-{i}")
generation_task_repo.add_task(task)
resp = client.get("/dashboard/overview")
data = resp.json()
assert len(data["recent_tasks"]) <= 5
def test_recent_tasks_have_correct_fields(self, client, generation_task_repo):
"""最近任务包含正确字段。"""
task = _make_generation_task("task-1", status=GenerationTaskStatus.COMPLETED)
generation_task_repo.add_task(task)
resp = client.get("/dashboard/overview")
data = resp.json()
assert len(data["recent_tasks"]) == 1
item = data["recent_tasks"][0]
for field in ["id", "task_type", "status", "current_step", "error_message", "updated_at"]:
assert field in item, f"缺少字段: {field}"
assert item["task_type"] == "generation"
def test_subscription_info(self, client):
"""订阅信息正确。"""
resp = client.get("/dashboard/overview")
data = resp.json()
assert "subscription" in data
sub = data["subscription"]
assert "plan" in sub
assert "is_active" in sub
assert sub["plan"] == "free"
assert sub["is_active"] is True
def test_pro_user_subscription(
self, project_repo, asset_repo, generation_task_repo, title_library_repo, voice_library_repo
):
"""Pro 用户订阅信息正确。"""
test_app = FastAPI()
test_app.include_router(router, prefix="/dashboard")
test_app.dependency_overrides[get_current_user] = lambda: AuthenticatedUser(
user=_make_user(subscription_plan="pro", subscription_status="active")
)
test_app.dependency_overrides[get_project_repository] = lambda: project_repo
test_app.dependency_overrides[get_asset_repository] = lambda: asset_repo
test_app.dependency_overrides[get_generation_task_repository] = lambda: generation_task_repo
test_app.dependency_overrides[get_title_library_repository] = lambda: title_library_repo
test_app.dependency_overrides[get_voice_library_repository] = lambda: voice_library_repo
c = TestClient(test_app)
resp = c.get("/dashboard/overview")
assert resp.status_code == 200
assert resp.json()["subscription"]["plan"] == "pro"
assert resp.json()["subscription"]["is_active"] is True
test_app.dependency_overrides.clear()
def test_total_products_count(self, client, project_repo):
"""项目(产品)数量正确。"""
resp = client.get("/dashboard/overview")
data = resp.json()
assert data["total_products"] == 2
# 新增一个项目后
project_repo.save(_make_project("proj-3", "user-test-001"))
resp2 = client.get("/dashboard/overview")
assert resp2.json()["total_products"] == 3
def test_unauthorized_returns_401(
self, project_repo, asset_repo, generation_task_repo, title_library_repo, voice_library_repo
):
"""未授权访问返回 401/403。"""
test_app = FastAPI()
test_app.include_router(router, prefix="/dashboard")
test_app.dependency_overrides[get_project_repository] = lambda: project_repo
test_app.dependency_overrides[get_asset_repository] = lambda: asset_repo
test_app.dependency_overrides[get_generation_task_repository] = lambda: generation_task_repo
test_app.dependency_overrides[get_title_library_repository] = lambda: title_library_repo
test_app.dependency_overrides[get_voice_library_repository] = lambda: voice_library_repo
c = TestClient(test_app)
resp = c.get("/dashboard/overview")
assert resp.status_code in (401, 403)
test_app.dependency_overrides.clear()
def test_recent_tasks_status_mapping(self, client, generation_task_repo):
"""不同状态的任务显示正确的当前步骤。"""
# 已完成任务
completed_task = _make_generation_task("task-completed", status=GenerationTaskStatus.COMPLETED)
generation_task_repo.add_task(completed_task)
resp = client.get("/dashboard/overview")
tasks = resp.json()["recent_tasks"]
completed = [t for t in tasks if t["id"] == "task-completed"][0]
assert completed["status"] == "completed"
assert "完成" in completed["current_step"] or "completed" in completed["current_step"].lower()
if __name__ == "__main__":
pytest.main([__file__, "-v"])
@@ -0,0 +1,554 @@
"""
生成视频管理 API 集成测试。
覆盖端点:
- GET /generated-videos — 列出生成视频
- GET /generated-videos/{video_id} — 获取生成视频详情
- PATCH /generated-videos/{video_id}/review — 更新审核状态
- GET /generated-videos/{video_id}/download-url — 获取下载地址
使用 FastAPI TestClient + dependency_overrides 模式,
导入真实路由模块,mock 所有外部依赖。
"""
from __future__ import annotations
import os
import sys
from dataclasses import replace
from datetime import datetime, timezone
from unittest.mock import MagicMock
# ── 环境变量 & sys.path(必须在导入 app.* 之前设置) ──────────────────────────
os.environ.setdefault("JWT_SECRET_KEY", "unit-test-secret-key-for-testing")
os.environ.setdefault("DATABASE_URL", "sqlite:///test.db")
import pytest
from fastapi import FastAPI
from fastapi.testclient import TestClient
sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..", "..", "apps", "api"))
from app.api.routes.generated_videos import router
from app.auth import AuthenticatedUser, get_current_user
from app.core.storage import get_storage_service
from app.dependencies import get_generated_video_repository, get_project_repository
from packages.domain.entities import Project, User
from packages.domain.generated_video import GeneratedVideo
# ---------------------------------------------------------------------------
# 1. 内存 Repository + 辅助函数
# ---------------------------------------------------------------------------
class InMemoryGeneratedVideoRepository:
"""内存中的生成视频 Repository。"""
def __init__(self):
self._items: dict[str, GeneratedVideo] = {}
def create(self, video: GeneratedVideo) -> GeneratedVideo:
self._items[video.id] = video
return video
def get(self, video_id: str) -> GeneratedVideo | None:
return self._items.get(video_id)
def update(self, video: GeneratedVideo) -> GeneratedVideo:
self._items[video.id] = video
return video
def list_by_project(self, project_id: str) -> list[GeneratedVideo]:
return [v for v in self._items.values() if v.project_id == project_id]
def list_by_generation_task(self, generation_task_id: str) -> list[GeneratedVideo]:
return [v for v in self._items.values() if v.generation_task_id == generation_task_id]
def list_by_batch(self, batch_id: str) -> list[GeneratedVideo]:
return []
class InMemoryProjectRepository:
"""内存中的项目 Repository。"""
def __init__(self):
self._projects: dict[str, Project] = {}
def save(self, project: Project) -> None:
self._projects[project.id] = project
def find_by_id(self, project_id: str) -> Project | None:
return self._projects.get(project_id)
def find_by_owner_user_id(self, owner_user_id: str) -> list[Project]:
return [p for p in self._projects.values() if p.owner_user_id == owner_user_id]
def find_accessible_projects(self, user_id: str) -> list[Project]:
return [p for p in self._projects.values() if p.owner_user_id == user_id]
def count_by_owner(self, owner_user_id: str) -> int:
return len(self.find_by_owner_user_id(owner_user_id))
def delete(self, project_id: str) -> bool:
if project_id in self._projects:
del self._projects[project_id]
return True
return False
class MockStorageService:
"""Mock OSS 存储服务。"""
def get_download_url(self, file_url: str) -> str:
return f"https://cdn.example.com/download/{file_url}?token=abc123"
def _make_user(**overrides) -> User:
defaults = dict(
id="user-test-001",
email="test@example.com",
display_name="Test User",
username="testuser",
subscription_plan="free",
subscription_status="active",
max_projects=3,
max_storage_gb=10,
created_at=datetime(2026, 1, 1, tzinfo=timezone.utc),
)
defaults.update(overrides)
return User(**defaults)
def _make_project(project_id: str = "proj-1", owner_user_id: str = "user-test-001") -> Project:
return Project(
id=project_id,
name=f"Project {project_id}",
owner_user_id=owner_user_id,
description="",
created_at=datetime(2026, 1, 1, tzinfo=timezone.utc),
)
def _make_video(
project_id: str = "proj-1",
name: str = "output.mp4",
status: str = "completed",
review_status: str = "pending_review",
**kwargs,
) -> GeneratedVideo:
return GeneratedVideo.create(
project_id=project_id,
generation_task_id=kwargs.pop("generation_task_id", "task-1"),
name=name,
file_url=kwargs.pop("file_url", f"generated/{name}"),
file_size=kwargs.pop("file_size", 1024000),
duration=kwargs.pop("duration", 30.5),
width=kwargs.pop("width", 1920),
height=kwargs.pop("height", 1080),
fps=kwargs.pop("fps", 30.0),
thumbnail_url=kwargs.pop("thumbnail_url", None),
generation_params=kwargs.pop("generation_params", {"resolution": "1080p"}),
)
# ---------------------------------------------------------------------------
# 2. Fixtures
# ---------------------------------------------------------------------------
@pytest.fixture
def video_repo():
return InMemoryGeneratedVideoRepository()
@pytest.fixture
def project_repo():
repo = InMemoryProjectRepository()
# 默认创建一个项目
repo.save(_make_project("proj-1", "user-test-001"))
repo.save(_make_project("proj-2", "user-test-001"))
repo.save(_make_project("proj-other", "other-user"))
return repo
@pytest.fixture
def storage_service():
return MockStorageService()
@pytest.fixture
def client(video_repo, project_repo, storage_service):
"""创建带有依赖覆盖的 TestClient。"""
test_app = FastAPI()
test_app.include_router(router, prefix="/generated-videos")
def _override_current_user():
return AuthenticatedUser(user=_make_user())
def _override_video_repo():
return video_repo
def _override_project_repo():
return project_repo
def _override_storage():
return storage_service
test_app.dependency_overrides[get_current_user] = _override_current_user
test_app.dependency_overrides[get_generated_video_repository] = _override_video_repo
test_app.dependency_overrides[get_project_repository] = _override_project_repo
test_app.dependency_overrides[get_storage_service] = _override_storage
yield TestClient(test_app)
test_app.dependency_overrides.clear()
# ---------------------------------------------------------------------------
# 3. GET / — 列出生成视频
# ---------------------------------------------------------------------------
class TestListGeneratedVideos:
"""列出生成视频端点测试。"""
def test_empty_list(self, client):
"""无视频时返回空列表。"""
resp = client.get("/generated-videos")
assert resp.status_code == 200
data = resp.json()
assert data["items"] == []
def test_list_all_user_videos(self, client, video_repo, project_repo):
"""列出当前用户所有项目的视频。"""
v1 = _make_video(project_id="proj-1", name="video1.mp4")
v2 = _make_video(project_id="proj-2", name="video2.mp4")
v3 = _make_video(project_id="proj-other", name="other.mp4") # 其他用户
video_repo.create(v1)
video_repo.create(v2)
video_repo.create(v3)
resp = client.get("/generated-videos")
assert resp.status_code == 200
data = resp.json()
assert len(data["items"]) == 2
names = {item["name"] for item in data["items"]}
assert names == {"video1.mp4", "video2.mp4"}
def test_filter_by_project_id(self, client, video_repo):
"""按 project_id 筛选视频。"""
v1 = _make_video(project_id="proj-1", name="a.mp4")
v2 = _make_video(project_id="proj-2", name="b.mp4")
video_repo.create(v1)
video_repo.create(v2)
resp = client.get("/generated-videos?project_id=proj-1")
assert resp.status_code == 200
data = resp.json()
assert len(data["items"]) == 1
assert data["items"][0]["name"] == "a.mp4"
def test_filter_by_nonexistent_project_returns_404(self, client):
"""筛选不存在的项目返回 404。"""
resp = client.get("/generated-videos?project_id=nonexistent")
assert resp.status_code == 404
def test_list_includes_download_url(self, client, video_repo):
"""列表响应应包含下载地址。"""
v = _make_video(file_url="generated/test.mp4")
video_repo.create(v)
resp = client.get("/generated-videos")
assert resp.status_code == 200
item = resp.json()["items"][0]
assert "download_url" in item
assert item["download_url"] is not None
assert "cdn.example.com" in item["download_url"]
def test_list_response_fields(self, client, video_repo):
"""列表响应包含所有必需字段。"""
v = _make_video()
video_repo.create(v)
resp = client.get("/generated-videos")
item = resp.json()["items"][0]
for field in [
"id",
"project_id",
"generation_task_id",
"name",
"file_url",
"file_size",
"duration",
"width",
"height",
"fps",
"status",
"review_status",
"generation_params",
"download_url",
]:
assert field in item, f"缺少字段: {field}"
def test_unauthorized_returns_401(self, video_repo, project_repo, storage_service):
"""未授权访问返回 401/403。"""
test_app = FastAPI()
test_app.include_router(router, prefix="/generated-videos")
# 不覆盖 get_current_user,使用默认(会拒绝无 token 请求)
test_app.dependency_overrides[get_generated_video_repository] = lambda: video_repo
test_app.dependency_overrides[get_project_repository] = lambda: project_repo
test_app.dependency_overrides[get_storage_service] = lambda: storage_service
c = TestClient(test_app)
resp = c.get("/generated-videos")
# 无 token 时 fastapi HTTPBearer auto_error=False 会返回 None,
# get_current_user 会抛 401
assert resp.status_code in (401, 403)
test_app.dependency_overrides.clear()
# ---------------------------------------------------------------------------
# 4. GET /{video_id} — 获取生成视频详情
# ---------------------------------------------------------------------------
class TestGetGeneratedVideo:
"""获取生成视频详情端点测试。"""
def test_get_existing_video(self, client, video_repo):
"""获取存在的视频返回详情。"""
v = _make_video(name="detail.mp4", duration=45.0)
video_repo.create(v)
resp = client.get(f"/generated-videos/{v.id}")
assert resp.status_code == 200
data = resp.json()
assert data["id"] == v.id
assert data["name"] == "detail.mp4"
assert data["duration"] == 45.0
assert data["status"] == "completed"
def test_get_includes_download_url(self, client, video_repo):
"""详情响应包含下载地址。"""
v = _make_video(file_url="generated/detail.mp4")
video_repo.create(v)
resp = client.get(f"/generated-videos/{v.id}")
data = resp.json()
assert "download_url" in data
assert "cdn.example.com" in data["download_url"]
def test_get_nonexistent_returns_404(self, client):
"""获取不存在的视频返回 404。"""
resp = client.get("/generated-videos/nonexistent-video-id")
assert resp.status_code == 404
assert "not found" in resp.json()["detail"].lower()
def test_get_thumbnail_url(self, client, video_repo):
"""有缩略图时返回缩略图 URL。"""
v = _make_video(thumbnail_url="thumbs/test.jpg")
video_repo.create(v)
resp = client.get(f"/generated-videos/{v.id}")
data = resp.json()
assert data["thumbnail_url"] == "thumbs/test.jpg"
def test_get_generation_params(self, client, video_repo):
"""返回生成参数。"""
params = {"resolution": "4k", "style": "cinematic"}
v = _make_video(generation_params=params)
video_repo.create(v)
resp = client.get(f"/generated-videos/{v.id}")
data = resp.json()
assert data["generation_params"]["resolution"] == "4k"
assert data["generation_params"]["style"] == "cinematic"
# ---------------------------------------------------------------------------
# 5. PATCH /{video_id}/review — 更新审核状态
# ---------------------------------------------------------------------------
class TestUpdateReviewStatus:
"""更新审核状态端点测试。"""
def test_approve_video(self, client, video_repo):
"""审核通过。"""
v = _make_video(review_status="pending_review")
video_repo.create(v)
resp = client.patch(
f"/generated-videos/{v.id}/review",
json={"review_status": "approved"},
)
assert resp.status_code == 200
data = resp.json()
assert data["review_status"] == "approved"
# 验证 repository 已更新
updated = video_repo.get(v.id)
assert updated.review_status == "approved"
def test_reject_video(self, client, video_repo):
"""审核拒绝。"""
v = _make_video(review_status="pending_review")
video_repo.create(v)
resp = client.patch(
f"/generated-videos/{v.id}/review",
json={"review_status": "rejected"},
)
assert resp.status_code == 200
assert resp.json()["review_status"] == "rejected"
def test_set_pending_review(self, client, video_repo):
"""设置为待审核。"""
v = _make_video(review_status="approved")
video_repo.create(v)
resp = client.patch(
f"/generated-videos/{v.id}/review",
json={"review_status": "pending_review"},
)
assert resp.status_code == 200
assert resp.json()["review_status"] == "pending_review"
def test_nonexistent_video_returns_404(self, client):
"""更新不存在的视频返回 404。"""
resp = client.patch(
"/nonexistent-id/review",
json={"review_status": "approved"},
)
assert resp.status_code == 404
def test_invalid_status_returns_422(self, client, video_repo):
"""无效审核状态返回 422。"""
v = _make_video()
video_repo.create(v)
resp = client.patch(
f"/generated-videos/{v.id}/review",
json={"review_status": "invalid_status"},
)
assert resp.status_code == 422
def test_missing_status_returns_422(self, client, video_repo):
"""缺少 review_status 字段返回 422。"""
v = _make_video()
video_repo.create(v)
resp = client.patch(f"/generated-videos/{v.id}/review", json={})
assert resp.status_code == 422
def test_update_returns_updated_fields(self, client, video_repo):
"""更新后返回完整的视频信息。"""
v = _make_video(name="review_test.mp4")
video_repo.create(v)
resp = client.patch(
f"/generated-videos/{v.id}/review",
json={"review_status": "approved"},
)
data = resp.json()
assert data["name"] == "review_test.mp4"
assert "id" in data
assert "download_url" in data
# ---------------------------------------------------------------------------
# 6. GET /{video_id}/download-url — 获取下载地址
# ---------------------------------------------------------------------------
class TestGetDownloadUrl:
"""获取下载地址端点测试。"""
def test_get_download_url_success(self, client, video_repo):
"""获取下载地址成功。"""
v = _make_video(file_url="generated/video.mp4")
video_repo.create(v)
resp = client.get(f"/generated-videos/{v.id}/download-url")
assert resp.status_code == 200
data = resp.json()
assert data["video_id"] == v.id
assert "download_url" in data
assert "cdn.example.com" in data["download_url"]
def test_nonexistent_video_returns_404(self, client):
"""获取不存在视频的下载地址返回 404。"""
resp = client.get("/generated-videos/nonexistent-id/download-url")
assert resp.status_code == 404
def test_download_url_format(self, client, video_repo):
"""下载地址格式正确。"""
v = _make_video(file_url="my-video.mp4")
video_repo.create(v)
resp = client.get(f"/generated-videos/{v.id}/download-url")
url = resp.json()["download_url"]
assert url.startswith("https://")
assert "token=" in url
# ---------------------------------------------------------------------------
# 7. 跨端点场景
# ---------------------------------------------------------------------------
class TestCrossEndpointScenarios:
"""跨端点集成场景。"""
def test_create_list_detail_review_flow(self, client, video_repo):
"""列表 → 详情 → 审核 完整流程。"""
# 准备数据
v = _make_video(name="flow.mp4", review_status="pending_review")
video_repo.create(v)
# 1. 列表
list_resp = client.get("/generated-videos")
assert list_resp.status_code == 200
assert len(list_resp.json()["items"]) == 1
# 2. 详情
detail_resp = client.get(f"/generated-videos/{v.id}")
assert detail_resp.status_code == 200
assert detail_resp.json()["name"] == "flow.mp4"
assert detail_resp.json()["review_status"] == "pending_review"
# 3. 审核通过
review_resp = client.patch(
f"/generated-videos/{v.id}/review",
json={"review_status": "approved"},
)
assert review_resp.status_code == 200
assert review_resp.json()["review_status"] == "approved"
# 4. 再次查看详情确认
detail_resp2 = client.get(f"/generated-videos/{v.id}")
assert detail_resp2.json()["review_status"] == "approved"
# 5. 获取下载地址
dl_resp = client.get(f"/generated-videos/{v.id}/download-url")
assert dl_resp.status_code == 200
assert dl_resp.json()["video_id"] == v.id
def test_multiple_videos_pagination_simulation(self, client, video_repo):
"""多个视频时列表正确返回所有视频。"""
for i in range(5):
v = _make_video(project_id="proj-1", name=f"video_{i}.mp4")
video_repo.create(v)
resp = client.get("/generated-videos")
assert resp.status_code == 200
items = resp.json()["items"]
assert len(items) == 5
names = {item["name"] for item in items}
assert len(names) == 5 # 全部不同
if __name__ == "__main__":
pytest.main([__file__, "-v"])
+10 -22
View File
@@ -118,18 +118,6 @@ class StubGenerationTaskRepository:
def count_by_user(self, user_id: str) -> int:
return len([t for t in self._tasks.values() if t.created_by_user_id == user_id])
def count_pending_by_user(self, user_id: str) -> int:
return len(
[
t
for t in self._tasks.values()
if t.created_by_user_id == user_id and t.status == GenerationTaskStatus.PENDING
]
)
def count_pending_total(self) -> int:
return len([t for t in self._tasks.values() if t.status == GenerationTaskStatus.PENDING])
def list_recent_by_user(self, user_id: str, limit: int = 5) -> list[GenerationTask]:
items = [t for t in self._tasks.values() if t.created_by_user_id == user_id]
items.sort(key=lambda t: t.created_at, reverse=True)
@@ -253,7 +241,7 @@ def client():
class TestCreateGenerationTask:
"""创建生成任务端点测试。"""
@patch("app.core.task_enqueue.celery_app")
@patch("app.api.routes.generation_tasks.celery_app")
def test_create_task_success(self, mock_celery, client):
"""正常创建生成任务成功。"""
mock_celery.send_task = MagicMock()
@@ -282,7 +270,7 @@ class TestCreateGenerationTask:
assert mock_celery.send_task.called
assert mock_celery.send_task.call_args[0][0] == "worker.generate_video"
@patch("app.core.task_enqueue.celery_app")
@patch("app.api.routes.generation_tasks.celery_app")
def test_create_batch_tasks(self, mock_celery, client):
"""批量创建多个生成任务。"""
mock_celery.send_task = MagicMock()
@@ -359,7 +347,7 @@ class TestListGenerationTasks:
def _create_task(self, client, task_suffix: str = "1"):
"""辅助方法:创建一个生成任务。"""
with patch("app.core.task_enqueue.celery_app") as mock_celery:
with patch("app.api.routes.generation_tasks.celery_app") as mock_celery:
mock_celery.send_task = MagicMock()
resp = client.post(
"/api/v1/generation/tasks",
@@ -380,7 +368,7 @@ class TestListGenerationTasks:
assert "items" in data
assert data["items"] == []
@patch("app.core.task_enqueue.celery_app")
@patch("app.api.routes.generation_tasks.celery_app")
def test_list_returns_user_tasks(self, mock_celery, client):
"""返回当前用户的生成任务列表。"""
mock_celery.send_task = MagicMock()
@@ -418,7 +406,7 @@ class TestGetGenerationTask:
"""获取生成任务详情端点测试。"""
def _create_task(self, client) -> str:
with patch("app.core.task_enqueue.celery_app") as mock_celery:
with patch("app.api.routes.generation_tasks.celery_app") as mock_celery:
mock_celery.send_task = MagicMock()
resp = client.post(
"/api/v1/generation/tasks",
@@ -461,7 +449,7 @@ class TestListGenerationResults:
"""列出生成结果端点测试。"""
def _create_task(self, client) -> str:
with patch("app.core.task_enqueue.celery_app") as mock_celery:
with patch("app.api.routes.generation_tasks.celery_app") as mock_celery:
mock_celery.send_task = MagicMock()
resp = client.post(
"/api/v1/generation/tasks",
@@ -501,7 +489,7 @@ class TestRetryGenerationTask:
def _create_failed_task(self, client) -> str:
"""创建一个失败状态的任务。"""
with patch("app.core.task_enqueue.celery_app") as mock_celery:
with patch("app.api.routes.generation_tasks.celery_app") as mock_celery:
mock_celery.send_task = MagicMock()
resp = client.post(
"/api/v1/generation/tasks",
@@ -521,7 +509,7 @@ class TestRetryGenerationTask:
# 让我们直接通过 retry 测试来验证
return task_id
@patch("app.core.task_enqueue.celery_app")
@patch("app.api.routes.generation_tasks.celery_app")
def test_retry_failed_task(self, mock_celery, client):
"""重试失败的任务成功。"""
mock_celery.send_task = MagicMock()
@@ -551,7 +539,7 @@ class TestRetryGenerationTask:
assert resp.status_code == 404
assert "not found" in resp.json()["detail"].lower()
@patch("app.core.task_enqueue.celery_app")
@patch("app.api.routes.generation_tasks.celery_app")
def test_retry_completed_task_returns_409(self, mock_celery, client):
"""重试已完成的任务返回 409。"""
mock_celery.send_task = MagicMock()
@@ -580,7 +568,7 @@ class TestRetryGenerationTask:
class TestGenerationTaskFlow:
"""生成任务完整流程集成测试。"""
@patch("app.core.task_enqueue.celery_app")
@patch("app.api.routes.generation_tasks.celery_app")
def test_create_list_detail_results_flow(self, mock_celery, client):
"""测试创建 → 列表 → 详情 → 结果 完整流程。"""
mock_celery.send_task = MagicMock()
+2 -14
View File
@@ -83,18 +83,6 @@ class StubGenerationTaskRepository:
def count_by_user(self, user_id: str) -> int:
return len([t for t in self._tasks.values() if t.created_by_user_id == user_id])
def count_pending_by_user(self, user_id: str) -> int:
return len(
[
t
for t in self._tasks.values()
if t.created_by_user_id == user_id and t.status == GenerationTaskStatus.PENDING
]
)
def count_pending_total(self) -> int:
return len([t for t in self._tasks.values() if t.status == GenerationTaskStatus.PENDING])
def list_recent_by_user(self, user_id: str, limit: int = 5) -> list[GenerationTask]:
items = [t for t in self._tasks.values() if t.created_by_user_id == user_id]
items.sort(key=lambda t: t.created_at, reverse=True)
@@ -440,7 +428,7 @@ class TestRetryProjectTask:
assert resp.status_code == 400
assert "Unsupported" in resp.json()["detail"]
@patch("app.core.task_enqueue.celery_app")
@patch("app.api.routes.task_center.celery_app")
def test_retry_failed_generation_task(self, mock_celery, client):
"""重试失败的 generation 任务成功。"""
mock_celery.send_task = MagicMock()
@@ -593,7 +581,7 @@ class TestRetryProjectTask:
class TestTaskCenterCrossEndpoint:
"""任务中心跨端点集成测试。"""
@patch("app.core.task_enqueue.celery_app")
@patch("app.api.routes.task_center.celery_app")
def test_list_then_retry_then_list(self, mock_celery, client):
"""列出任务 → 重试失败任务 → 再列出验证新任务。"""
mock_celery.send_task = MagicMock()
+17 -18
View File
@@ -32,11 +32,7 @@ sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..", "..", "apps", "
from app.api.routes.voice_clones import router
from app.auth import AuthenticatedUser, get_current_user
from app.dependencies import (
get_audio_url_signer,
get_cosyvoice_service,
get_voice_clone_profile_repository,
)
from app.dependencies import get_cosyvoice_service, get_voice_clone_profile_repository
from packages.domain.entities import User
from packages.domain.voice_clone_profile import (
@@ -230,7 +226,7 @@ def clone_repo():
@pytest.fixture
def cosyvoice_service():
return MockCosyVoiceService(async_mode=True) # 异步模式,匹配真实 CosyVoice API 行为
return MockCosyVoiceService(async_mode=False) # 同步模式,简化测试
@pytest.fixture
@@ -245,7 +241,6 @@ def client(clone_repo, cosyvoice_service):
test_app.dependency_overrides[get_current_user] = _override_current_user
test_app.dependency_overrides[get_voice_clone_profile_repository] = lambda: clone_repo
test_app.dependency_overrides[get_cosyvoice_service] = lambda: cosyvoice_service
test_app.dependency_overrides[get_audio_url_signer] = lambda: (lambda url: url)
yield TestClient(test_app)
@@ -261,7 +256,7 @@ class TestCreateVoiceClone:
"""创建声音克隆端点测试。"""
def test_create_with_source_audio(self, client, cosyvoice_service):
"""提供源音频时创建克隆,异步提交后状态为 processing。"""
"""提供源音频时创建克隆,同步模式下直接 ready。"""
resp = client.post(
"/voice-clones",
json={
@@ -282,9 +277,9 @@ class TestCreateVoiceClone:
assert "id" in data
assert len(data["id"]) > 0
# 异步模式下提交后状态为 processing,voice_id 为空
assert data["status"] == "processing"
assert data["voice_id"] == ""
# 同步模式下应直接 ready
assert data["status"] == "ready"
assert data["voice_id"] == "mock-voice-789"
assert data["error_message"] == ""
def test_create_without_source_audio(self, client):
@@ -559,15 +554,16 @@ class TestRetryVoiceClone:
"""重试克隆端点测试。"""
def test_retry_failed_clone(self, client, clone_repo, cosyvoice_service):
"""重试失败的克隆,重新提交后期望 processing。"""
"""重试失败的克隆应成功。"""
cosyvoice_service.async_mode = False
p = _make_clone_profile("重试测试", status=VoiceCloneStatus.FAILED)
clone_repo.create(p)
resp = client.post(f"/voice-clones/{p.id}/retry")
assert resp.status_code == 200
data = resp.json()
# 异步模式下重试后状态为 processing,等待 CosyVoice 完成
assert data["status"] == "processing"
# 同步模式下重试后应变为 ready
assert data["status"] == "ready"
assert data["retry_count"] >= 1
def test_retry_nonexistent_returns_404(self, client):
@@ -594,6 +590,7 @@ class TestRetryVoiceClone:
def test_retry_increments_retry_count(self, client, clone_repo, cosyvoice_service):
"""重试后重试次数增加。"""
cosyvoice_service.async_mode = False
p = _make_clone_profile("重试计数", status=VoiceCloneStatus.FAILED)
clone_repo.create(p)
@@ -683,7 +680,7 @@ class TestVoiceCloneLifecycle:
# 4. 状态
status_resp = client.get(f"/voice-clones/{clone_id}/status")
assert status_resp.status_code == 200
assert status_resp.json()["status"] == "processing"
assert status_resp.json()["status"] == "ready"
# 5. 删除
del_resp = client.delete(f"/voice-clones/{clone_id}")
@@ -694,7 +691,7 @@ class TestVoiceCloneLifecycle:
assert list_resp2.json()["total"] == 0
def test_failed_retry_flow(self, client, clone_repo, cosyvoice_service):
"""失败 → 重试 → processing(等待异步完成) 流程。"""
"""失败 → 重试 → 成功 流程。"""
# 创建一个失败的克隆
p = _make_clone_profile("失败重试", status=VoiceCloneStatus.FAILED)
clone_repo.create(p)
@@ -704,13 +701,15 @@ class TestVoiceCloneLifecycle:
assert status_resp.json()["status"] == "failed"
# 重试
cosyvoice_service.async_mode = False
retry_resp = client.post(f"/voice-clones/{p.id}/retry")
assert retry_resp.status_code == 200
assert retry_resp.json()["status"] == "processing"
assert retry_resp.json()["status"] == "ready"
# 再次确认状态
status_resp2 = client.get(f"/voice-clones/{p.id}/status")
assert status_resp2.json()["status"] == "processing"
assert status_resp2.json()["status"] == "ready"
assert status_resp2.json()["voice_id"] != ""
if __name__ == "__main__":
-136
View File
@@ -1,136 +0,0 @@
# 灰度对比测试工具
用于统一渲染引擎灰度发布期间的新旧引擎对比验证。
## 能力
- **像素对比**:基于 FFmpeg SSIM + PSNR 双指标,评估视频画质差异
- **音频对比**:基于差值音频 RMS,评估音频波形差异
- **批量对比**:10个预设场景覆盖 P0/P1/P2 优先级
- **HTML 报告**:可视化对比结果,包含画质、音频、性能三维度
- **两种切换方式**:支持 engine 参数直传 或 Feature Flag 白名单切换
## 目录结构
```
tests/render_compare/
├── __init__.py # 包导出
├── README.md # 本文档
├── video_diff.py # 视频像素对比(SSIM + PSNR)
├── audio_diff.py # 音频对比(差值 RMS)
├── scenarios.py # 预定义对比场景(10个)
└── runner.py # 批量对比执行器 + HTML 报告生成
```
## 快速开始
### 环境要求
- FFmpeg 4.4+(需带 ssim 和 psnr 滤镜)
- Python 3.10+
- httpx(API 调用)
### 配置环境变量
```bash
export STAGING_API_URL=https://api.staging.example.com
export STAGING_API_KEY=your_api_key
export STAGING_INTERNAL_API_KEY=your_internal_key # 可选,Feature Flag 模式需要
```
### 运行对比
```bash
# 运行所有 P0 场景(最核心的5个)
python -m tests.render_compare.runner --priority P0 --output ./report/
# 运行 P0 + P1 场景
python -m tests.render_compare.runner --priority P1 --output ./report/
# 只跑指定场景
python -m tests.render_compare.runner --scenarios simple_pass_through,subtitle_rendering
# 使用 Feature Flag 方式切换引擎(需要 internal key)
python -m tests.render_compare.runner --priority P0 --flag-mode
# 自定义阈值
python -m tests.render_compare.runner --priority P0 --ssim-threshold 0.95 --psnr-threshold 30
```
## 对比场景
| ID | 名称 | 优先级 | 验证点 |
|----|------|--------|--------|
| simple_pass_through | 简单直通 | P0 | 直通优化路径正确性 |
| multi_clip_transition | 多clip转场 | P0 | 转场效果 + concat |
| subtitle_rendering | 字幕渲染 | P0 | ASS字幕渲染 |
| independent_audio_track | 独立音频轨 | P0 | 音频混音(amix) |
| no_audio_video | 无音轨视频 | P0 | 无音轨防御逻辑 |
| picture_in_picture | 画中画 | P1 | overlay 图层 |
| multi_layer_mix | 多图层混合 | P1 | 多图层复杂场景 |
| image_background | 图片背景 | P1 | background 层 + 无音频 |
| long_video_stress | 长视频压力 | P2 | 多clip性能 |
| vertical_portrait | 竖屏9:16 | P2 | scale 策略(铺满裁剪) |
## 验收标准(建议)
### 视频质量
- **平均 SSIM >= 0.90**:通过(有微小差异但视觉可接受)
- **平均 SSIM >= 0.95**:优秀(视觉几乎无差异)
- **平均 PSNR >= 25 dB**:通过
- **分辨率一致 + 时长差 < 0.1s**:通过
### 音频质量
- **相似度 >= 0.85**:通过
- **采样率/声道数一致**:通过
### 性能
- **平均性能差异在 ±10% 以内**:可接受
- **直通场景新引擎更快**(预期 +30%)
## API 约定
Runner 默认假设渲染 API 支持以下接口:
### 提交任务
```
POST /api/v1/render/compose
Authorization: Bearer {api_key}
Body: { ...plan_payload, "engine": "legacy" | "unified" }
Response: { "task_id": "xxx" }
```
### 查询状态
```
GET /api/v1/tasks/{task_id}
Response: { "status": "completed", "output_url": "...", "duration_sec": 5.2 }
```
### Feature Flag(flag-mode)
```
PUT /api/v1/internal/feature-flags/render_engine
X-API-Key: {internal_key}
Body: { "enabled": true, "percentage": 100 }
```
如果你的 API 接口不同,请修改 `StagingAPI` 类中的对应方法。
## 故障排查
### 对比失败定位指南
1. **像素差异大(SSIM < 0.90)**
- 检查分辨率是否一致
- 检查帧率是否一致
- 用 `save_diff_frame` 生成差异帧可视化
- 检查转场效果(slideup/slidedown 是新引擎独有)
2. **音频不一致**
- 检查音频编码参数(码率、采样率)
- 检查主音频源优先级(main > broll)
- 用 ffprobe 对比两视频音频流参数
3. **渲染失败**
- 检查日志:`[unified-render] render failed`
- 检查素材是否完整下载
- 检查 FFmpeg 命令是否正确
-26
View File
@@ -1,26 +0,0 @@
"""灰度对比测试工具包.
用于新旧渲染引擎的批量对比测试,包含:
- video_diff: 视频像素对比(SSIM + PSNR)
- audio_diff: 音频对比(差值 RMS)
- scenarios: 预定义对比场景
- runner: 批量对比执行器 + HTML 报告
"""
from .audio_diff import AudioDiffResult, compute_audio_diff, extract_audio, probe_duration, probe_has_audio
from .scenarios import SCENARIOS, CompareScenario, get_scenarios_by_priority
from .video_diff import VideoDiffResult, compute_video_diff, save_diff_frame
__all__ = [
"VideoDiffResult",
"compute_video_diff",
"save_diff_frame",
"AudioDiffResult",
"compute_audio_diff",
"extract_audio",
"probe_has_audio",
"probe_duration",
"SCENARIOS",
"CompareScenario",
"get_scenarios_by_priority",
]
-322
View File
@@ -1,322 +0,0 @@
"""音频对比工具 — 基于 FFmpeg 的音频质量对比.
使用以下指标评估两段音频的相似度:
1. 波形差异(RMS 差值)
2. 频谱相似度(FFT 分帧比较)
3. 时长差异
对比方式:
- 直接对两个音频做 `ametadata=select='gt(scene\\,0.3)'` 过于复杂
- 简化方案:用 `amerge` + `astats` 计算差值音频的 RMS
更精确的方案(已实现):
- 将两轨音频做差(amix=0:weights='1 -1' → 实际上用 pan 更简单)
- 对差值音频做 astats,获取差值的 RMS、峰值等指标
"""
from __future__ import annotations
import json
import re
import shutil
import subprocess # nosec B404
from dataclasses import asdict, dataclass
from pathlib import Path
from typing import Any
FFMPEG_BIN: str = shutil.which("ffmpeg") or "ffmpeg"
FFPROBE_BIN: str = shutil.which("ffprobe") or "ffprobe"
@dataclass
class AudioDiffResult:
"""音频对比结果."""
audio_a: str
audio_b: str
duration_a: float
duration_b: float
duration_diff: float
sample_rate_match: bool
channels_match: bool
diff_rms_db: float # 差值音频的 RMS(dB,越低越相似)
diff_peak_db: float # 差值音频的峰值(dB,越低越相似)
similarity_score: float # 综合相似度评分 [0, 1],1 = 完全一致
passed: bool
def to_dict(self) -> dict[str, Any]:
return asdict(self)
def probe_duration(file_path: str) -> float:
"""探测文件时长(秒),失败返回 0."""
try:
result = subprocess.run( # nosec B603
[
FFPROBE_BIN,
"-v",
"error",
"-show_entries",
"format=duration",
"-of",
"default=noprint_wrappers=1:nokey=1",
str(file_path),
],
check=True,
stdout=subprocess.PIPE,
stderr=subprocess.PIPE,
text=True,
timeout=10,
)
return round(float(result.stdout.strip()), 3)
except Exception:
return 0.0
def probe_has_audio(file_path: str | Path) -> bool:
"""探测文件是否包含音频流."""
try:
result = subprocess.run( # nosec B603
[
FFPROBE_BIN,
"-v",
"error",
"-select_streams",
"a:0",
"-show_entries",
"stream=codec_type",
"-of",
"default=noprint_wrappers=1:nokey=1",
str(file_path),
],
check=True,
stdout=subprocess.PIPE,
stderr=subprocess.PIPE,
text=True,
timeout=10,
)
return result.stdout.strip() == "audio"
except Exception:
return False # 探测失败保守返回 False,避免误判有音频
def compute_audio_diff(
audio_a: str | Path,
audio_b: str | Path,
*,
similarity_threshold: float = 0.90,
duration_tolerance: float = 0.1,
) -> AudioDiffResult:
"""计算两段音频的差异.
方案:用 pan 滤镜将两轨相减,对差值音频做 astats 分析。
Args:
audio_a: 音频A(基线)
audio_b: 音频B(对比)
similarity_threshold: 相似度合格阈值
duration_tolerance: 时长容忍度(秒)
Returns:
AudioDiffResult 对比结果
"""
dur_a = probe_duration(str(audio_a))
dur_b = probe_duration(str(audio_b))
duration_diff = abs(dur_a - dur_b)
# 获取音频元信息
info_a = _probe_audio_info(str(audio_a))
info_b = _probe_audio_info(str(audio_b))
sample_rate_match = info_a["sample_rate"] == info_b["sample_rate"]
channels_match = info_a["channels"] == info_b["channels"]
# 相减后分析差值
# 取较短时长做对比
min_dur = min(dur_a, dur_b)
if min_dur <= 0:
return AudioDiffResult(
audio_a=str(audio_a),
audio_b=str(audio_b),
duration_a=dur_a,
duration_b=dur_b,
duration_diff=duration_diff,
sample_rate_match=sample_rate_match,
channels_match=channels_match,
diff_rms_db=-999.0,
diff_peak_db=-999.0,
similarity_score=0.0,
passed=False,
)
# 做差值音频:a - b
# 注意:amix 会自动按输入数归一化音量(除以N),
# 所以 a + (-1)*b 经过 amix=inputs=2 后整体音量会减半(-6dB)。
# 加 volume=2 补偿回来,确保差值 RMS 反映真实差异幅度。
command = [
FFMPEG_BIN,
"-i",
str(audio_a),
"-i",
str(audio_b),
"-filter_complex",
# 第2轨反相 → amix混合 → volume=2补偿amix的自动缩放
"[1:a]volume=-1[inv];[0:a][inv]amix=inputs=2:duration=shortest:dropout_transition=0,volume=2[diff]",
"-map",
"[diff]",
"-f",
"null",
"-af",
"astats=metadata=1:reset=0",
"-",
]
try:
result = subprocess.run( # nosec B603
command,
check=True,
stdout=subprocess.PIPE,
stderr=subprocess.PIPE,
text=True,
timeout=120,
)
stderr = result.stderr or ""
except subprocess.CalledProcessError as e:
# 如果音频格式不兼容,返回失败
return AudioDiffResult(
audio_a=str(audio_a),
audio_b=str(audio_b),
duration_a=dur_a,
duration_b=dur_b,
duration_diff=duration_diff,
sample_rate_match=sample_rate_match,
channels_match=channels_match,
diff_rms_db=999.0,
diff_peak_db=999.0,
similarity_score=0.0,
passed=False,
)
diff_rms_db, diff_peak_db = _parse_astats(stderr)
# 相似度评分:基于差值 RMS
# 差值 RMS -60dB → 相似度 ~1.0(几乎无声差)
# 差值 RMS -20dB → 相似度 ~0.5(有明显差异)
# 差值 RMS 0dB → 相似度 ~0.0(完全相反)
if diff_rms_db <= -60:
similarity_score = 1.0
elif diff_rms_db >= 0:
similarity_score = 0.0
else:
# 线性映射:-60dB → 1.0, 0dB → 0.0
similarity_score = max(0.0, min(1.0, 1.0 + diff_rms_db / 60.0))
passed = (
duration_diff <= duration_tolerance
and sample_rate_match
and channels_match
and similarity_score >= similarity_threshold
)
return AudioDiffResult(
audio_a=str(audio_a),
audio_b=str(audio_b),
duration_a=round(dur_a, 3),
duration_b=round(dur_b, 3),
duration_diff=round(duration_diff, 3),
sample_rate_match=sample_rate_match,
channels_match=channels_match,
diff_rms_db=round(diff_rms_db, 2),
diff_peak_db=round(diff_peak_db, 2),
similarity_score=round(similarity_score, 4),
passed=passed,
)
def _probe_audio_info(file_path: str) -> dict[str, int]:
"""探测音频元信息."""
try:
result = subprocess.run( # nosec B603
[
FFPROBE_BIN,
"-v",
"error",
"-select_streams",
"a:0",
"-show_entries",
"stream=sample_rate,channels",
"-of",
"json",
file_path,
],
check=True,
stdout=subprocess.PIPE,
stderr=subprocess.PIPE,
text=True,
timeout=10,
)
info = json.loads(result.stdout)
stream = info.get("streams", [{}])[0]
return {
"sample_rate": int(stream.get("sample_rate", 44100)),
"channels": int(stream.get("channels", 2)),
}
except Exception:
return {"sample_rate": 0, "channels": 0}
def _parse_astats(stderr: str) -> tuple[float, float]:
"""从 astats 输出中解析 RMS 和峰值.
astats 输出格式(在 stderr 中):
[Parsed_astats_1 @ 0x...] Channel: 1
[Parsed_astats_1 @ 0x...] ...
[Parsed_astats_1 @ 0x...] Overall
[Parsed_astats_1 @ 0x...] DC offset: 0.000000
[Parsed_astats_1 @ 0x...] Min level: -0.123456
[Parsed_astats_1 @ 0x...] Max level: 0.789012
[Parsed_astats_1 @ 0x...] Peak level dB: -2.01
[Parsed_astats_1 @ 0x...] RMS level dB: -10.56
...
"""
lines = stderr.split("\n")
rms_db = -999.0
peak_db = -999.0
for line in lines:
# 找 Overall 部分的统计(双声道时取整体值)
rms_match = re.search(r"RMS level dB:\s*(-?\d+\.?\d*)", line)
peak_match = re.search(r"Peak level dB:\s*(-?\d+\.?\d*)", line)
if rms_match:
rms_db = float(rms_match.group(1))
if peak_match:
peak_db = float(peak_match.group(1))
return rms_db, peak_db
def extract_audio(video_path: str | Path, output_path: str | Path) -> Path:
"""从视频中提取音频(AAC 格式).
Args:
video_path: 视频文件路径
output_path: 输出音频路径
Returns:
输出音频文件路径
"""
command = [
FFMPEG_BIN,
"-y",
"-i",
str(video_path),
"-vn",
"-acodec",
"aac",
"-b:a",
"128k",
str(output_path),
]
subprocess.run(command, check=True, capture_output=True, timeout=120) # nosec B603
return Path(output_path)
-628
View File
@@ -1,628 +0,0 @@
"""灰度对比测试 Runner — 新旧引擎批量对比 + 报告生成.
使用方法:
# 配置环境变量
export STAGING_API_URL=https://api.staging.example.com
export STAGING_API_KEY=your_key
# 运行全部 P0 场景
python -m tests.render_compare.runner --priority P0 --output ./report/
# 只跑指定场景
python -m tests.render_compare.runner --scenario simple_pass_through,subtitle_rendering
对比流程:
1. 对每个场景,分别提交到 legacy 和 unified 引擎(通过 Feature Flag 白名单/百分比控制)
- 方式A:通过内部 API 临时切换 flag(需要 admin key)
- 方式B:提交任务时指定 engine 参数(如果 API 支持)
2. 等待任务完成,下载输出视频
3. 像素对比(SSIM + PSNR)+ 音频对比(差值RMS)
4. 生成 HTML 对比报告
注意:默认假设 API 支持 `engine` 参数来指定渲染引擎。
如果不支持,需要先通过内部 API 切换 Feature Flag,然后提交任务。
"""
from __future__ import annotations
import argparse
import json
import os
import sys
import time
from dataclasses import dataclass, field
from datetime import datetime
from pathlib import Path
from typing import Any
import httpx
# 确保项目根目录在 path 中
sys.path.insert(0, str(Path(__file__).parent.parent.parent))
from .audio_diff import AudioDiffResult, compute_audio_diff
from .scenarios import SCENARIOS, CompareScenario, get_scenarios_by_priority
from .video_diff import VideoDiffResult, compute_video_diff
@dataclass
class ScenarioResult:
"""单个场景的对比结果."""
scenario: CompareScenario
legacy_task_id: str = ""
unified_task_id: str = ""
legacy_video_path: str = ""
unified_video_path: str = ""
legacy_duration_sec: float = 0.0
unified_duration_sec: float = 0.0
video_diff: VideoDiffResult | None = None
audio_diff: AudioDiffResult | None = None
legacy_success: bool = False
unified_success: bool = False
error: str = ""
@property
def passed(self) -> bool:
if not (self.legacy_success and self.unified_success):
return False
if self.video_diff and not self.video_diff.passed:
return False
if self.audio_diff and not self.audio_diff.passed:
return False
return True
class StagingAPI:
"""Staging 环境 API 客户端."""
def __init__(self, base_url: str, api_key: str, internal_api_key: str = ""):
self.base_url = base_url.rstrip("/")
self.api_key = api_key
self.internal_api_key = internal_api_key
self.client = httpx.Client(timeout=30.0)
def _headers(self, internal: bool = False) -> dict[str, str]:
headers = {"Authorization": f"Bearer {self.api_key}"}
if internal and self.internal_api_key:
headers["X-API-Key"] = self.internal_api_key
return headers
def submit_render_task(self, plan_payload: dict[str, Any], engine: str = "") -> str:
"""提交渲染任务,返回 task_id.
Args:
plan_payload: EditPlan payload
engine: 可选,指定引擎("legacy" / "unified")
Returns:
task_id
"""
url = f"{self.base_url}/api/v1/render/compose"
payload = dict(plan_payload)
if engine:
payload["engine"] = engine
resp = self.client.post(url, json=payload, headers=self._headers())
resp.raise_for_status()
data = resp.json()
return data.get("task_id") or data.get("id", "")
def get_task_status(self, task_id: str) -> dict[str, Any]:
"""获取任务状态."""
url = f"{self.base_url}/api/v1/tasks/{task_id}"
resp = self.client.get(url, headers=self._headers())
resp.raise_for_status()
return resp.json()
def wait_for_task(self, task_id: str, timeout: float = 300.0, poll_interval: float = 3.0) -> dict[str, Any]:
"""等待任务完成.
Returns:
最终任务状态
Raises:
TimeoutError: 超时
"""
start = time.time()
while time.time() - start < timeout:
status = self.get_task_status(task_id)
state = status.get("status", "")
if state in ("completed", "success", "done", "failed", "error"):
return status
time.sleep(poll_interval)
raise TimeoutError(f"Task {task_id} timed out after {timeout}s")
def set_feature_flag(self, flag_name: str, enabled: bool, percentage: int = 0, whitelist: list[str] | None = None):
"""通过内部 API 设置 Feature Flag.
用于不支持 engine 参数的场景,切换全局灰度比例。
"""
if not self.internal_api_key:
raise ValueError("internal_api_key is required for feature flag operations")
url = f"{self.base_url}/api/v1/internal/feature-flags/{flag_name}"
body: dict[str, Any] = {"enabled": enabled, "percentage": percentage}
if whitelist is not None:
body["whitelist"] = whitelist
resp = self.client.put(url, json=body, headers=self._headers(internal=True))
resp.raise_for_status()
return resp.json()
def get_feature_flag(self, flag_name: str) -> dict[str, Any]:
"""获取 Feature Flag 配置."""
if not self.internal_api_key:
raise ValueError("internal_api_key is required")
url = f"{self.base_url}/api/v1/internal/feature-flags/{flag_name}"
resp = self.client.get(url, headers=self._headers(internal=True))
resp.raise_for_status()
return resp.json()
def download_video(self, video_url: str, output_path: str | Path) -> Path:
"""下载视频文件."""
output_path = Path(output_path)
output_path.parent.mkdir(parents=True, exist_ok=True)
with self.client.stream("GET", video_url, timeout=60.0) as resp:
resp.raise_for_status()
with open(output_path, "wb") as f:
for chunk in resp.iter_bytes():
f.write(chunk)
return output_path
class CompareRunner:
"""新旧引擎对比 Runner."""
# 全局默认阈值(唯一真实来源,所有入口统一引用)
DEFAULT_SSIM_THRESHOLD: float = 0.95
DEFAULT_PSNR_THRESHOLD: float = 28.0
DEFAULT_AUDIO_SIMILARITY_THRESHOLD: float = 0.90
DEFAULT_DURATION_TOLERANCE: float = 0.1
DEFAULT_TASK_TIMEOUT: float = 300.0
def __init__(
self,
api: StagingAPI,
output_dir: Path,
*,
ssim_threshold: float | None = None,
psnr_threshold: float | None = None,
audio_similarity_threshold: float | None = None,
task_timeout: float | None = None,
flag_mode: bool = False, # 是否使用 Feature Flag 方式切换引擎
duration_tolerance: float | None = None,
):
self.api = api
self.output_dir = output_dir
self.ssim_threshold = ssim_threshold if ssim_threshold is not None else self.DEFAULT_SSIM_THRESHOLD
self.psnr_threshold = psnr_threshold if psnr_threshold is not None else self.DEFAULT_PSNR_THRESHOLD
self.audio_similarity_threshold = (
audio_similarity_threshold
if audio_similarity_threshold is not None
else self.DEFAULT_AUDIO_SIMILARITY_THRESHOLD
)
self.duration_tolerance = (
duration_tolerance if duration_tolerance is not None else self.DEFAULT_DURATION_TOLERANCE
)
self.task_timeout = task_timeout if task_timeout is not None else self.DEFAULT_TASK_TIMEOUT
self.flag_mode = flag_mode
self.results: list[ScenarioResult] = []
# flag_mode 下保存原始配置,测试结束后恢复(防污染线上)
self._original_flag_config: dict[str, Any] | None = None
def run_scenario(self, scenario: CompareScenario) -> ScenarioResult:
"""运行单个场景对比."""
print(f"\n{'='*60}")
print(f"[{scenario.priority}] {scenario.id}: {scenario.name}")
print(f" {scenario.description}")
result = ScenarioResult(scenario=scenario)
scenario_dir = self.output_dir / scenario.id
scenario_dir.mkdir(parents=True, exist_ok=True)
try:
# 1. 提交两个引擎的任务
legacy_task_id = self._submit_with_engine(scenario, "legacy")
unified_task_id = self._submit_with_engine(scenario, "unified")
result.legacy_task_id = legacy_task_id
result.unified_task_id = unified_task_id
print(f" legacy task: {legacy_task_id}")
print(f" unified task: {unified_task_id}")
# 2. 等待完成
print(" waiting for legacy...", end="", flush=True)
legacy_status = self.api.wait_for_task(legacy_task_id, timeout=self.task_timeout)
result.legacy_success = legacy_status.get("status") in ("completed", "success", "done")
legacy_video_url = legacy_status.get("output_url", "") or legacy_status.get("video_url", "")
print(f" {'✅' if result.legacy_success else '❌'} ({legacy_status.get('duration_sec', '?')}s)")
print(" waiting for unified...", end="", flush=True)
unified_status = self.api.wait_for_task(unified_task_id, timeout=self.task_timeout)
result.unified_success = unified_status.get("status") in ("completed", "success", "done")
unified_video_url = unified_status.get("output_url", "") or unified_status.get("video_url", "")
print(f" {'✅' if result.unified_success else '❌'} ({unified_status.get('duration_sec', '?')}s)")
result.legacy_duration_sec = float(legacy_status.get("duration_sec", 0))
result.unified_duration_sec = float(unified_status.get("duration_sec", 0))
if not (result.legacy_success and result.unified_success):
result.error = f"Legacy success={result.legacy_success}, Unified success={result.unified_success}"
print(" ⚠️ 任务未全部成功,跳过对比")
return result
# 3. 下载视频
print(" downloading...", end="", flush=True)
legacy_path = self.api.download_video(legacy_video_url, scenario_dir / "legacy.mp4")
unified_path = self.api.download_video(unified_video_url, scenario_dir / "unified.mp4")
result.legacy_video_path = str(legacy_path)
result.unified_video_path = str(unified_path)
print(" ✅")
# 4. 像素对比
print(" computing video diff...", end="", flush=True)
result.video_diff = compute_video_diff(
legacy_path,
unified_path,
ssim_threshold=self.ssim_threshold,
psnr_threshold=self.psnr_threshold,
duration_tolerance=self.duration_tolerance,
)
print(
f" SSIM={result.video_diff.avg_ssim:.4f} PSNR={result.video_diff.avg_psnr:.2f}dB {'✅' if result.video_diff.passed else '❌'}"
)
# 5. 音频对比(仅当都有音频时)
from .audio_diff import probe_has_audio
legacy_has_audio = probe_has_audio(legacy_path)
unified_has_audio = probe_has_audio(unified_path)
if legacy_has_audio and unified_has_audio:
print(" computing audio diff...", end="", flush=True)
result.audio_diff = compute_audio_diff(
legacy_path,
unified_path,
similarity_threshold=self.audio_similarity_threshold,
)
print(
f" similarity={result.audio_diff.similarity_score:.4f} {'✅' if result.audio_diff.passed else '❌'}"
)
elif legacy_has_audio != unified_has_audio:
result.error = f"音频不一致: legacy_has_audio={legacy_has_audio}, unified_has_audio={unified_has_audio}"
print(f" ⚠️ 音频不一致: legacy={legacy_has_audio}, unified={unified_has_audio}")
else:
print(" audio: both silent (skip)")
except Exception as e:
result.error = str(e)
print(f" ❌ 错误: {e}")
self.results.append(result)
return result
def _submit_with_engine(self, scenario: CompareScenario, engine: str) -> str:
"""提交指定引擎的任务.
如果 flag_mode=True,通过 Feature Flag 切换,否则通过 engine 参数。
"""
if self.flag_mode:
# 先设置 flag(用白名单方式,确保只有当前测试用户命中)
percentage = 0 if engine == "legacy" else 100
self.api.set_feature_flag("render_engine", enabled=True, percentage=percentage)
time.sleep(1) # 给 worker 一点时间刷新配置
return self.api.submit_render_task(scenario.plan_payload)
else:
return self.api.submit_render_task(scenario.plan_payload, engine=engine)
def run_all(self, scenarios: list[CompareScenario]) -> list[ScenarioResult]:
"""运行所有场景.
flag_mode=True 时,测试开始前保存原始 Feature Flag 配置,
结束后(无论成功失败)自动恢复,避免污染线上环境。
"""
print(f"\n灰度对比测试开始 - {len(scenarios)} 个场景")
print(f"输出目录: {self.output_dir}")
print(f"视频阈值: SSIM>={self.ssim_threshold}, PSNR>={self.psnr_threshold}dB")
print(f"音频阈值: similarity>={self.audio_similarity_threshold}")
# flag_mode:保存原始配置,测试结束后恢复(防污染)
if self.flag_mode:
try:
self._original_flag_config = self.api.get_feature_flag("render_engine")
print(f" [flag_mode] 已保存原始配置: {self._original_flag_config}")
except Exception as e:
print(f" ⚠️ [flag_mode] 保存原始配置失败: {e}")
print(" 为避免污染线上,将中止测试。请检查 internal_api_key 配置。")
return self.results
try:
for i, scenario in enumerate(scenarios):
print(f"\n进度: {i+1}/{len(scenarios)}")
self.run_scenario(scenario)
finally:
# 始终恢复原始 flag 配置
if self.flag_mode and self._original_flag_config:
try:
orig = self._original_flag_config
self.api.set_feature_flag(
"render_engine",
enabled=orig.get("enabled", False),
percentage=orig.get("percentage", 0),
whitelist=orig.get("whitelist"),
)
print("\n[flag_mode] ✅ 已恢复原始 Feature Flag 配置")
except Exception as e:
print(f"\n[flag_mode] ❌ 恢复 Feature Flag 失败: {e}")
print(" 请手动检查并恢复 render_engine flag 配置!")
return self.results
def summary(self) -> dict[str, Any]:
"""生成汇总统计."""
total = len(self.results)
passed = sum(1 for r in self.results if r.passed)
failed = total - passed
# 性能对比
perf_diffs = []
for r in self.results:
if r.legacy_success and r.unified_success and r.legacy_duration_sec > 0:
diff_pct = (r.unified_duration_sec - r.legacy_duration_sec) / r.legacy_duration_sec * 100
perf_diffs.append(diff_pct)
avg_perf_diff = sum(perf_diffs) / len(perf_diffs) if perf_diffs else 0.0
return {
"total": total,
"passed": passed,
"failed": failed,
"pass_rate": f"{passed/total*100:.1f}%" if total > 0 else "0%",
"avg_perf_diff_pct": round(avg_perf_diff, 2),
"scenarios": [self._result_to_dict(r) for r in self.results],
"timestamp": datetime.now().isoformat(),
"ssim_threshold": self.ssim_threshold,
"psnr_threshold": self.psnr_threshold,
"audio_threshold": self.audio_similarity_threshold,
}
def _result_to_dict(self, r: ScenarioResult) -> dict[str, Any]:
return {
"id": r.scenario.id,
"name": r.scenario.name,
"priority": r.scenario.priority,
"passed": r.passed,
"legacy_success": r.legacy_success,
"unified_success": r.unified_success,
"legacy_duration_sec": r.legacy_duration_sec,
"unified_duration_sec": r.unified_duration_sec,
"video_diff": r.video_diff.to_dict() if r.video_diff else None,
"audio_diff": r.audio_diff.to_dict() if r.audio_diff else None,
"error": r.error,
}
def generate_html_report(summary: dict[str, Any], output_path: Path):
"""生成 HTML 对比报告."""
scenarios = summary["scenarios"]
# 按通过/失败分组
passed_list = [s for s in scenarios if s["passed"]]
failed_list = [s for s in scenarios if not s["passed"]]
# 构建场景卡片
scenario_cards = ""
for s in scenarios:
status_class = "pass" if s["passed"] else "fail"
status_text = "✅ 通过" if s["passed"] else "❌ 失败"
vdiff = s.get("video_diff") or {}
adiff = s.get("audio_diff") or {}
video_info = ""
if vdiff:
video_info = f"""
<div class="metric-row">
<span>SSIM:</span>
<span class="{'good' if vdiff.get('avg_ssim', 0) >= 0.95 else 'warn'}">{vdiff.get('avg_ssim', 0):.4f}</span>
</div>
<div class="metric-row">
<span>PSNR:</span>
<span>{vdiff.get('avg_psnr', 0):.2f} dB</span>
</div>
<div class="metric-row">
<span>时长差:</span>
<span>{vdiff.get('duration_diff', 0):.3f}s</span>
</div>
"""
audio_info = ""
if adiff:
audio_info = f"""
<div class="metric-row">
<span>音频相似度:</span>
<span class="{'good' if adiff.get('similarity_score', 0) >= 0.9 else 'warn'}">{adiff.get('similarity_score', 0):.4f}</span>
</div>
<div class="metric-row">
<span>差值 RMS:</span>
<span>{adiff.get('diff_rms_db', 0):.2f} dB</span>
</div>
"""
perf_info = ""
if s["legacy_duration_sec"] and s["unified_duration_sec"]:
diff = s["unified_duration_sec"] - s["legacy_duration_sec"]
pct = diff / s["legacy_duration_sec"] * 100 if s["legacy_duration_sec"] else 0
trend = "🔴" if pct > 10 else ("🟡" if pct > 0 else "🟢")
perf_info = f"""
<div class="perf-row">
<span>Legacy: {s['legacy_duration_sec']:.2f}s</span>
<span>Unified: {s['unified_duration_sec']:.2f}s</span>
<span>{trend} {pct:+.1f}%</span>
</div>
"""
error_info = f'<div class="error-box">{s["error"]}</div>' if s["error"] else ""
scenario_cards += f"""
<div class="card {status_class}">
<div class="card-header">
<span class="badge">{s['priority']}</span>
<span class="scenario-name">{s['name']}</span>
<span class="status {status_class}">{status_text}</span>
</div>
<div class="card-body">
<div class="grid-2">
<div>
<h4>视频质量</h4>
{video_info or '<p class="muted">无数据</p>'}
</div>
<div>
<h4>音频质量</h4>
{audio_info or '<p class="muted">无音频或跳过</p>'}
</div>
</div>
<div>
<h4>性能对比</h4>
{perf_info or '<p class="muted">无数据</p>'}
</div>
{error_info}
</div>
</div>
"""
html = f"""<!DOCTYPE html>
<html lang="zh-CN">
<head>
<meta charset="UTF-8">
<meta name="viewport" content="width=device-width, initial-scale=1.0">
<title>统一渲染引擎灰度对比报告</title>
<style>
* {{ box-sizing: border-box; margin: 0; padding: 0; }}
body {{ font-family: -apple-system, BlinkMacSystemFont, 'Segoe UI', Roboto, sans-serif; background: #f5f5f5; color: #333; padding: 20px; }}
.container {{ max-width: 1200px; margin: 0 auto; }}
h1 {{ margin-bottom: 20px; font-size: 24px; }}
.summary {{ background: white; border-radius: 12px; padding: 24px; margin-bottom: 24px; display: flex; gap: 32px; flex-wrap: wrap; }}
.summary-item {{ text-align: center; }}
.summary-item .value {{ font-size: 32px; font-weight: bold; margin-bottom: 4px; }}
.summary-item .label {{ color: #666; font-size: 14px; }}
.pass .value {{ color: #10b981; }}
.fail .value {{ color: #ef4444; }}
.card {{ background: white; border-radius: 12px; margin-bottom: 16px; overflow: hidden; border-left: 4px solid #10b981; }}
.card.fail {{ border-left-color: #ef4444; }}
.card-header {{ padding: 16px 20px; background: #fafafa; display: flex; align-items: center; gap: 12px; border-bottom: 1px solid #eee; }}
.badge {{ background: #e5e7eb; color: #374151; padding: 2px 8px; border-radius: 4px; font-size: 12px; font-weight: 600; }}
.scenario-name {{ flex: 1; font-weight: 600; }}
.status {{ font-weight: 600; }}
.status.pass {{ color: #10b981; }}
.status.fail {{ color: #ef4444; }}
.card-body {{ padding: 20px; }}
.grid-2 {{ display: grid; grid-template-columns: 1fr 1fr; gap: 24px; margin-bottom: 16px; }}
h4 {{ margin-bottom: 12px; color: #374151; font-size: 14px; }}
.metric-row {{ display: flex; justify-content: space-between; padding: 6px 0; font-size: 14px; }}
.metric-row .good {{ color: #10b981; font-weight: 600; }}
.metric-row .warn {{ color: #f59e0b; font-weight: 600; }}
.perf-row {{ display: flex; gap: 24px; padding: 8px 0; font-size: 14px; background: #f9fafb; padding: 12px; border-radius: 8px; }}
.error-box {{ background: #fef2f2; color: #dc2626; padding: 12px; border-radius: 8px; margin-top: 12px; font-size: 13px; }}
.muted {{ color: #9ca3af; font-size: 14px; }}
.timestamp {{ text-align: center; color: #9ca3af; font-size: 12px; margin-top: 24px; }}
</style>
</head>
<body>
<div class="container">
<h1>🎬 统一渲染引擎灰度对比报告</h1>
<div class="summary">
<div class="summary-item">
<div class="value">{summary['total']}</div>
<div class="label">总场景数</div>
</div>
<div class="summary-item pass">
<div class="value">{summary['passed']}</div>
<div class="label">通过</div>
</div>
<div class="summary-item fail">
<div class="value">{summary['failed']}</div>
<div class="label">失败</div>
</div>
<div class="summary-item">
<div class="value">{summary['pass_rate']}</div>
<div class="label">通过率</div>
</div>
<div class="summary-item">
<div class="value {'good' if summary['avg_perf_diff_pct'] <= 0 else 'warn'}" style="font-size: 24px; color: {'#10b981' if summary['avg_perf_diff_pct'] <= 0 else '#f59e0b'}">{summary['avg_perf_diff_pct']:+.1f}%</div>
<div class="label">平均性能差异</div>
</div>
</div>
{scenario_cards}
<div class="timestamp">生成时间: {summary['timestamp']}</div>
</div>
</body>
</html>"""
output_path.parent.mkdir(parents=True, exist_ok=True)
output_path.write_text(html, encoding="utf-8")
return output_path
def main():
parser = argparse.ArgumentParser(description="统一渲染引擎灰度对比测试")
parser.add_argument("--priority", default="P0", choices=["P0", "P1", "P2"], help="最低优先级")
parser.add_argument("--scenarios", default="", help="指定场景ID,逗号分隔")
parser.add_argument("--output", default="./gray_compare_report", help="输出目录")
parser.add_argument("--ssim-threshold", type=float, default=None, help="SSIM阈值(默认0.95)")
parser.add_argument("--psnr-threshold", type=float, default=None, help="PSNR阈值(dB)(默认28.0)")
parser.add_argument("--audio-threshold", type=float, default=None, help="音频相似度阈值(默认0.90)")
parser.add_argument("--flag-mode", action="store_true", help="使用Feature Flag方式切换引擎")
parser.add_argument("--task-timeout", type=float, default=300.0, help="单任务超时时间(秒)")
args = parser.parse_args()
base_url = os.environ.get("STAGING_API_URL", "")
api_key = os.environ.get("STAGING_API_KEY", "")
internal_key = os.environ.get("STAGING_INTERNAL_API_KEY", "")
if not base_url or not api_key:
print("❌ 请设置环境变量 STAGING_API_URL 和 STAGING_API_KEY")
sys.exit(1)
# 选择场景
if args.scenarios:
scenario_ids = [s.strip() for s in args.scenarios.split(",")]
selected = [s for s in SCENARIOS if s.id in scenario_ids]
if not selected:
print(f"❌ 未找到匹配的场景: {scenario_ids}")
print(f"可用场景: {[s.id for s in SCENARIOS]}")
sys.exit(1)
else:
selected = get_scenarios_by_priority(args.priority)
output_dir = Path(args.output).resolve()
output_dir.mkdir(parents=True, exist_ok=True)
api = StagingAPI(base_url, api_key, internal_key)
runner = CompareRunner(
api,
output_dir,
ssim_threshold=args.ssim_threshold,
psnr_threshold=args.psnr_threshold,
audio_similarity_threshold=args.audio_threshold,
flag_mode=args.flag_mode,
task_timeout=args.task_timeout,
)
runner.run_all(selected)
# 生成报告
summary = runner.summary()
# JSON 报告
json_path = output_dir / "report.json"
json_path.write_text(json.dumps(summary, indent=2, ensure_ascii=False), encoding="utf-8")
# HTML 报告
html_path = output_dir / "report.html"
generate_html_report(summary, html_path)
print(f"\n{'='*60}")
print(f"对比完成: {summary['passed']}/{summary['total']} 通过 ({summary['pass_rate']})")
print(f"报告: {html_path}")
print(f"JSON: {json_path}")
if __name__ == "__main__":
main()
-257
View File
@@ -1,257 +0,0 @@
"""灰度对比测试场景定义 — 覆盖典型渲染场景.
每个场景对应一个 EditPlan,用于新旧引擎对比。
覆盖场景:
1. 简单直通(单clip无特效)
2. 多clip转场(fade + slide)
3. 画中画(main + overlay)
4. 字幕渲染(ASS字幕)
5. 独立音频轨(主视频 + BGM)
6. 多图层混合(main + broll + overlay + audio)
7. 背景图片 + 主视频(图片背景无音频)
8. 无音频视频(纯画面,验证无音轨防御)
9. 长视频(10+ clip,压力测试)
10. 分辨率非标(竖屏9:16,验证scale策略)
"""
from __future__ import annotations
from dataclasses import dataclass, field
from typing import Any
@dataclass
class CompareScenario:
"""对比测试场景."""
id: str
name: str
description: str
priority: str # P0 / P1 / P2
plan_payload: dict[str, Any] # EditPlan JSON payload(提交给 API 的数据)
expected: dict[str, Any] = field(default_factory=dict) # 预期结果
SCENARIOS: list[CompareScenario] = [
CompareScenario(
id="simple_pass_through",
name="简单直通",
description="单主clip,无转场无特效,验证直通优化路径",
priority="P0",
plan_payload={
"width": 1280,
"height": 720,
"fps": 25,
"clips": [
{
"clip_type": "main",
"asset_id": "sample_5s.mp4",
"duration": 5.0,
"order": 0,
}
],
},
),
CompareScenario(
id="multi_clip_transition",
name="多clip转场",
description="3个clip,fade + slideleft 转场",
priority="P0",
plan_payload={
"width": 1280,
"height": 720,
"fps": 25,
"clips": [
{
"clip_type": "main",
"asset_id": "sample_5s.mp4",
"duration": 3.0,
"order": 0,
"transition_effect": "cut",
},
{
"clip_type": "main",
"asset_id": "sample_5s.mp4",
"duration": 3.0,
"order": 1,
"transition_effect": "fade",
},
{
"clip_type": "main",
"asset_id": "sample_5s.mp4",
"duration": 3.0,
"order": 2,
"transition_effect": "slideleft",
},
],
},
),
CompareScenario(
id="picture_in_picture",
name="画中画",
description="主视频 + 角落小窗(corner_voice)",
priority="P1",
plan_payload={
"width": 1280,
"height": 720,
"fps": 25,
"clips": [
{"clip_type": "main", "asset_id": "sample_5s.mp4", "duration": 5.0, "order": 0},
{"clip_type": "corner_voice", "asset_id": "sample_5s.mp4", "duration": 5.0, "order": 0},
],
},
),
CompareScenario(
id="subtitle_rendering",
name="字幕渲染",
description="主视频 + ASS字幕",
priority="P0",
plan_payload={
"width": 1280,
"height": 720,
"fps": 25,
"clips": [
{
"clip_type": "main",
"asset_id": "sample_5s.mp4",
"duration": 5.0,
"order": 0,
"config": {"subtitles": [{"text": "测试字幕 Test Subtitle", "start_time": 0, "end_time": 5.0}]},
}
],
},
),
CompareScenario(
id="independent_audio_track",
name="独立音频轨",
description="主视频(带音频)+ 独立BGM轨,验证音频混音",
priority="P0",
plan_payload={
"width": 1280,
"height": 720,
"fps": 25,
"clips": [
{"clip_type": "main", "asset_id": "sample_5s.mp4", "duration": 5.0, "order": 0},
{
"clip_type": "main",
"asset_id": "sample_bgm.mp3",
"duration": 5.0,
"order": 0,
"config": {"role": "audio", "volume": 0.5},
},
],
},
),
CompareScenario(
id="multi_layer_mix",
name="多图层混合",
description="main + broll + overlay + audio 四图层",
priority="P1",
plan_payload={
"width": 1280,
"height": 720,
"fps": 25,
"clips": [
{
"clip_type": "main",
"asset_id": "sample_5s.mp4",
"duration": 4.0,
"order": 0,
"transition_effect": "fade",
},
{
"clip_type": "main",
"asset_id": "sample_5s.mp4",
"duration": 4.0,
"order": 1,
"transition_effect": "slideup",
},
{"clip_type": "broll", "asset_id": "sample_broll.mp4", "duration": 8.0, "order": 0},
{"clip_type": "overlay", "asset_id": "sample_overlay.png", "duration": 8.0, "order": 0},
{
"clip_type": "main",
"asset_id": "sample_bgm.mp3",
"duration": 8.0,
"order": 0,
"config": {"role": "audio", "volume": 0.3},
},
],
},
),
CompareScenario(
id="image_background",
name="图片背景",
description="background图片层 + 主视频,验证背景层无音频",
priority="P1",
plan_payload={
"width": 1280,
"height": 720,
"fps": 25,
"clips": [
{"clip_type": "background", "asset_id": "sample_bg.jpg", "duration": 5.0, "order": 0},
{"clip_type": "main", "asset_id": "sample_5s.mp4", "duration": 5.0, "order": 0},
],
},
),
CompareScenario(
id="no_audio_video",
name="无音轨视频",
description="源视频无音频流,验证无音轨防御逻辑",
priority="P0",
plan_payload={
"width": 1280,
"height": 720,
"fps": 25,
"clips": [
{"clip_type": "main", "asset_id": "sample_silent_5s.mp4", "duration": 5.0, "order": 0},
],
},
),
CompareScenario(
id="long_video_stress",
name="长视频压力",
description="10个clip + 多种转场,性能压力测试",
priority="P2",
plan_payload={
"width": 1280,
"height": 720,
"fps": 25,
"clips": [
{
"clip_type": "main",
"asset_id": "sample_5s.mp4",
"duration": 3.0,
"order": i,
"transition_effect": ["cut", "fade", "slideleft", "slidedown", "dissolve"][i % 5],
}
for i in range(10)
],
},
),
CompareScenario(
id="vertical_portrait",
name="竖屏9:16",
description="竖屏分辨率,验证scale策略(铺满裁剪)",
priority="P2",
plan_payload={
"width": 720,
"height": 1280,
"fps": 25,
"clips": [
{"clip_type": "main", "asset_id": "sample_5s.mp4", "duration": 5.0, "order": 0},
],
},
),
]
def get_scenarios_by_priority(min_priority: str = "P2") -> list[CompareScenario]:
"""按优先级过滤场景.
P0 包含 P0
P1 包含 P0 + P1
P2 包含全部
"""
priority_order = {"P0": 0, "P1": 1, "P2": 2}
threshold = priority_order.get(min_priority, 2)
return [s for s in SCENARIOS if priority_order.get(s.priority, 2) <= threshold]
-283
View File
@@ -1,283 +0,0 @@
"""视频对比工具 — 基于 FFmpeg 的像素级质量对比.
使用 SSIM + PSNR 双指标评估两个视频的相似度:
- SSIM (Structural Similarity): 结构相似性,范围 [0, 1],越接近 1 越相似
- PSNR (Peak Signal-to-Noise Ratio): 峰值信噪比,单位 dB,越高越好
灰度验收标准:
- 平均 SSIM >= 0.95 → 视觉上几乎无差异(P0 场景必达)
- 最低 SSIM >= 0.90 → 最严重帧差异可接受
- 平均 PSNR >= 28dB → 质量达标
"""
from __future__ import annotations
import json
import re
import shutil
import subprocess # nosec B404
from dataclasses import asdict, dataclass
from pathlib import Path
from typing import Any
FFMPEG_BIN: str = shutil.which("ffmpeg") or "ffmpeg"
FFPROBE_BIN: str = shutil.which("ffprobe") or "ffprobe"
@dataclass
class VideoDiffResult:
"""视频对比结果."""
video_a: str
video_b: str
width: int
height: int
duration_a: float
duration_b: float
avg_ssim: float
min_ssim: float
avg_psnr: float # dB
min_psnr: float
frame_count: int
duration_diff: float # 时长差(秒)
resolution_match: bool
passed: bool # 是否通过阈值
def to_dict(self) -> dict[str, Any]:
return asdict(self)
def probe_video_info(video_path: str) -> dict[str, Any]:
"""获取视频信息(宽、高、时长、fps)."""
try:
result = subprocess.run( # nosec B603
[
FFPROBE_BIN,
"-v",
"error",
"-select_streams",
"v:0",
"-show_entries",
"stream=width,height,r_frame_rate,duration",
"-show_entries",
"format=duration",
"-of",
"json",
video_path,
],
check=True,
stdout=subprocess.PIPE,
stderr=subprocess.PIPE,
text=True,
timeout=10,
)
info = json.loads(result.stdout)
stream = info.get("streams", [{}])[0]
fmt = info.get("format", {})
width = int(stream.get("width", 1280))
height = int(stream.get("height", 720))
fps_str = stream.get("r_frame_rate", "25/1")
if "/" in fps_str:
num, den = fps_str.split("/")
fps = float(num) / float(den) if float(den) > 0 else 25.0
else:
fps = float(fps_str) if fps_str else 25.0
duration = float(fmt.get("duration", 0)) or float(stream.get("duration", 0))
return {"width": width, "height": height, "duration": duration, "fps": round(fps, 2)}
except Exception:
return {"width": 1280, "height": 720, "duration": 0.0, "fps": 25.0}
def compute_video_diff(
video_a: str | Path,
video_b: str | Path,
*,
ssim_threshold: float = 0.95,
psnr_threshold: float = 28.0,
duration_tolerance: float = 0.1,
) -> VideoDiffResult:
"""计算两个视频的像素差异.
使用 FFmpeg ssim + psnr 滤镜一次性计算两个指标。
Args:
video_a: 视频A路径(基线)
video_b: 视频B路径(对比)
ssim_threshold: SSIM 合格阈值(默认 0.90)
psnr_threshold: PSNR 合格阈值(默认 25dB)
duration_tolerance: 时长容忍度(秒,默认 0.1s)
Returns:
VideoDiffResult 对比结果
Raises:
subprocess.CalledProcessError: FFmpeg 执行失败
"""
info_a = probe_video_info(str(video_a))
info_b = probe_video_info(str(video_b))
duration_diff = abs(info_a["duration"] - info_b["duration"])
resolution_match = info_a["width"] == info_b["width"] and info_a["height"] == info_b["height"]
# ssim 和 psnr 的 stats_file 都输出到 stdout
# 用行格式区分:SSIM 行含 "All:",PSNR 行含 "psnr_avg:"
command = [
FFMPEG_BIN,
"-i",
str(video_a),
"-i",
str(video_b),
"-lavfi",
"[0:v][1:v]ssim=stats_file=-[out1];[0:v][1:v]psnr=stats_file=-[out2]",
"-f",
"null",
"-",
]
result = subprocess.run( # nosec B603
command,
check=True,
stdout=subprocess.PIPE,
stderr=subprocess.PIPE,
text=True,
timeout=300,
)
# 逐帧统计在 stdout(stats_file=-),汇总日志在 stderr
stats_stdout = result.stdout or ""
avg_ssim, min_ssim = _parse_ssim_stats(stats_stdout)
avg_psnr, min_psnr = _parse_psnr_stats(stats_stdout)
frame_count = _count_frames(result.stderr or "")
passed = (
resolution_match
and duration_diff <= duration_tolerance
and avg_ssim >= ssim_threshold
and avg_psnr >= psnr_threshold
)
return VideoDiffResult(
video_a=str(video_a),
video_b=str(video_b),
width=info_a["width"],
height=info_a["height"],
duration_a=round(info_a["duration"], 3),
duration_b=round(info_b["duration"], 3),
avg_ssim=round(avg_ssim, 6),
min_ssim=round(min_ssim, 6),
avg_psnr=round(avg_psnr, 3),
min_psnr=round(min_psnr, 3),
frame_count=frame_count,
duration_diff=round(duration_diff, 3),
resolution_match=resolution_match,
passed=passed,
)
def _parse_ssim_stats(stats_output: str) -> tuple[float, float]:
"""从 SSIM stats_file 输出中解析逐帧 SSIM.
FFmpeg ssim 滤镜 stats_file 输出格式(每行一帧):
n:1 Y:0.987654 U:0.991234 V:0.990000 All:0.989000 (19.585642)
n:2 Y:0.986543 U:0.990123 V:0.988888 All:0.987654 (19.123456)
...
Returns:
(avg_ssim, min_ssim)
"""
ssim_values: list[float] = []
for line in stats_output.split("\n"):
# 匹配 stats_file 格式:n:数字 ... All:数字
if not line.startswith("n:"):
continue
match = re.search(r"All:(\d+\.\d+)", line)
if match:
ssim_values.append(float(match.group(1)))
if not ssim_values:
return 0.0, 0.0
avg_ssim = sum(ssim_values) / len(ssim_values)
min_ssim = min(ssim_values)
return avg_ssim, min_ssim
def _parse_psnr_stats(stats_output: str) -> tuple[float, float]:
"""从 PSNR stats_file 输出中解析逐帧 PSNR.
FFmpeg psnr 滤镜 stats_file 输出格式(每行一帧):
n:1 mse_avg:100.23 mse_y:150.12 mse_u:50.34 mse_v:80.56 psnr_avg:28.12 psnr_y:26.34 psnr_u:31.12 psnr_v:29.08
n:2 ...
Returns:
(avg_psnr, min_psnr) — avg_psnr 是逐帧 psnr_avg 的均值,min_psnr 是逐帧最小值
"""
psnr_values: list[float] = []
for line in stats_output.split("\n"):
if not line.startswith("n:"):
continue
match = re.search(r"psnr_avg:(\d+\.\d+)", line)
if match:
psnr_values.append(float(match.group(1)))
if not psnr_values:
return 0.0, 0.0
avg_psnr = sum(psnr_values) / len(psnr_values)
min_psnr = min(psnr_values)
return avg_psnr, min_psnr
def _count_frames(stderr: str) -> int:
"""从 FFmpeg 输出中统计帧数."""
match = re.search(r"frame=\s*(\d+)", stderr)
return int(match.group(1)) if match else 0
def save_diff_frame(
video_a: str | Path,
video_b: str | Path,
output_path: str | Path,
*,
timestamp: float = 1.0,
) -> Path:
"""生成差异帧可视化图(红绿色差).
使用 blend 滤镜生成差异可视化图,差异越大越亮。
Args:
video_a: 视频A
video_b: 视频B
output_path: 输出图片路径
timestamp: 截取的时间点(秒)
Returns:
输出图片路径
"""
command = [
FFMPEG_BIN,
"-y",
"-ss",
str(timestamp),
"-i",
str(video_a),
"-ss",
str(timestamp),
"-i",
str(video_b),
"-lavfi",
"[0:v][1:v]blend=all_mode=difference,eq=contrast=5:brightness=0.5[diff]",
"-map",
"[diff]",
"-vframes",
"1",
str(output_path),
]
subprocess.run(command, check=True, capture_output=True, timeout=60) # nosec B603
return Path(output_path)
-72
View File
@@ -1,72 +0,0 @@
"""AssetStatus 枚举兼容性测试。
验证历史脏数据(如 'uploaded')不会导致枚举转换失败。
"""
import pytest
from packages.domain.entities import AssetStatus
class TestAssetStatusNormalValues:
"""正常值应该正确映射。"""
def test_uploading(self):
assert AssetStatus("uploading") == AssetStatus.UPLOADING
def test_ready(self):
assert AssetStatus("ready") == AssetStatus.READY
def test_processing(self):
assert AssetStatus("processing") == AssetStatus.PROCESSING
def test_error(self):
assert AssetStatus("error") == AssetStatus.ERROR
class TestAssetStatusHistoricalValues:
"""历史脏数据应该正确映射到对应状态,不抛异常。"""
@pytest.mark.parametrize("value", ["uploaded", "Uploaded", "UPLOADED", " uploaded "])
def test_uploaded_maps_to_ready(self, value):
"""生产环境发现的 'uploaded' 历史值应映射为 READY。"""
assert AssetStatus(value) == AssetStatus.READY
@pytest.mark.parametrize("value", ["success", "ok", "done", "complete"])
def test_other_ready_like_values_map_to_ready(self, value):
assert AssetStatus(value) == AssetStatus.READY
@pytest.mark.parametrize("value", ["upload", "uploading_start", "upload_start"])
def test_upload_like_values_map_to_uploading(self, value):
assert AssetStatus(value) == AssetStatus.UPLOADING
@pytest.mark.parametrize("value", ["failed", "fail", "err"])
def test_error_like_values_map_to_error(self, value):
assert AssetStatus(value) == AssetStatus.ERROR
@pytest.mark.parametrize("value", ["process", "running", "run"])
def test_processing_like_values_map_to_processing(self, value):
assert AssetStatus(value) == AssetStatus.PROCESSING
class TestAssetStatusFallback:
"""完全未知的值兜底为 READY,不抛500。"""
@pytest.mark.parametrize("value", ["unknown", "foo_bar", ""])
def test_unknown_value_falls_back_to_ready(self, value):
assert AssetStatus(value) == AssetStatus.READY
def test_none_value_falls_back_to_ready(self):
assert AssetStatus(None) == AssetStatus.READY # type: ignore[arg-type]
def test_int_value_falls_back_to_ready(self):
assert AssetStatus(123) == AssetStatus.READY # type: ignore[arg-type]
class TestAssetStatusStrValue:
"""枚举值仍为字符串类型,不影响序列化。"""
def test_value_unchanged(self):
assert AssetStatus.READY.value == "ready"
assert AssetStatus.ERROR.value == "error"
assert isinstance(AssetStatus.READY, str)
-93
View File
@@ -1,93 +0,0 @@
"""测试音频URL预签名逻辑。
验证所有 API 返回的音频 URL 都会经过 OSS 预签名(24小时有效期),
确保私有 bucket 下的音频文件前端可正常访问。
"""
from __future__ import annotations
from unittest.mock import MagicMock, patch
import pytest
class TestAudioUrlSigner:
"""测试音频URL签名函数的行为。"""
def _make_signer(self, mock_storage):
"""构造一个签名函数(模拟 get_audio_url_signer 的逻辑)。"""
def sign_audio_url(url: str) -> str:
if not url:
return url
return mock_storage.get_download_url(url, expires_seconds=86400)
return sign_audio_url
def test_empty_url_returns_empty(self):
"""空URL直接返回,不调用签名。"""
mock_storage = MagicMock()
signer = self._make_signer(mock_storage)
result = signer("")
assert result == ""
mock_storage.get_download_url.assert_not_called()
def test_none_url_returns_none(self):
"""None URL直接返回(有些字段可能为None)。"""
mock_storage = MagicMock()
signer = self._make_signer(mock_storage)
result = signer(None) # type: ignore
assert result is None
mock_storage.get_download_url.assert_not_called()
def test_valid_url_gets_signed_24h(self):
"""有效URL会调用 storage.get_download_url,有效期24小时(86400秒)。"""
mock_storage = MagicMock()
mock_storage.get_download_url.return_value = (
"https://bucket.oss-cn-hangzhou.aliyuncs.com/audio/test.mp3?signature=xxx"
)
signer = self._make_signer(mock_storage)
result = signer("https://bucket.oss-cn-hangzhou.aliyuncs.com/audio/test.mp3")
assert "signature=xxx" in result
mock_storage.get_download_url.assert_called_once_with(
"https://bucket.oss-cn-hangzhou.aliyuncs.com/audio/test.mp3",
expires_seconds=86400,
)
def test_storage_key_format_also_works(self):
"""纯 storage key 格式也能正常签名(storage内部会处理)。"""
mock_storage = MagicMock()
mock_storage.get_download_url.return_value = "https://signed-url/audio.mp3?sig=xxx"
signer = self._make_signer(mock_storage)
result = signer("audio/test.mp3")
assert result == "https://signed-url/audio.mp3?sig=xxx"
mock_storage.get_download_url.assert_called_once_with(
"audio/test.mp3",
expires_seconds=86400,
)
def test_signer_via_dependencies_module(self):
"""通过 dependencies 模块获取 signer,验证集成正确。"""
from app.core.storage import OSSStorageService
mock_svc = MagicMock(spec=OSSStorageService)
mock_svc.get_download_url.return_value = "https://signed/a.mp3?sig=123"
# 替换全局单例
with patch("app.core.storage._storage_service", mock_svc):
from app.dependencies import get_audio_url_signer
signer = get_audio_url_signer()
result = signer("test/audio.mp3")
assert result == "https://signed/a.mp3?sig=123"
mock_svc.get_download_url.assert_called_once_with(
"test/audio.mp3",
expires_seconds=86400,
)

Some files were not shown because too many files have changed in this diff Show More