Compare commits

..

1 Commits

Author SHA1 Message Date
用户CI Test 7ed0ffd5a8 fix: 生成任务入队失败时标记为failed,避免pending僵尸任务
CI/CD Pipeline / Validate Code Quality And Tests (pull_request) Failing after 8s
CI/CD Pipeline / Frontend Lint (pull_request) Successful in 1m0s
CI/CD Pipeline / Build & Push Staging (Watchtower auto-deploy) (pull_request) Has been skipped
CI/CD Pipeline / Build Production Runtime Images (pull_request) Has been skipped
CI/CD Pipeline / Staging E2E Tests (pull_request) Has been skipped
CI/CD Pipeline / Staging API Integration Tests (pull_request) Has been skipped
CI/CD Pipeline / Deploy Production (pull_request) Has been skipped
CI/CD Pipeline / Production Browser E2E (pull_request) Has been skipped
根因:任务创建(DB commit)和入队(Celery send_task)是两个独立操作,
send_task失败时任务卡在pending状态永远不会执行。

修复:
- 新增_safe_enqueue_generation_task安全入队函数
- send_task失败时自动标记任务为failed并记录错误
- 覆盖4处入口:批量创建、generation重试、task_center两级重试
2026-07-11 01:15:21 +08:00
82 changed files with 2653 additions and 9078 deletions
+3 -8
View File
@@ -44,14 +44,9 @@ 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_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
+14 -284
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,180 +158,14 @@ jobs:
python3 scripts/check_migration_safety.py --allow-medium-risk
fi
- name: Debug coverage paths
shell: sh
run: |
set +e
echo "=== PWD ==="
pwd
echo "=== check source dirs ==="
ls -d apps/api/app packages
echo "=== python import check ==="
python3 - <<'PY'
import sys, os
os.environ["PYTHONPATH"] = f"{os.getcwd()}/apps/api:{os.getcwd()}"
sys.path.insert(0, f"{os.getcwd()}/apps/api")
sys.path.insert(0, os.getcwd())
print(f"cwd: {os.getcwd()}")
print(f"sys.path[:5]: {sys.path[:5]}")
try:
import app
print(f"app.__file__: {app.__file__}")
except Exception as e:
print(f"import app failed: {e}")
try:
import packages
print(f"packages.__file__: {packages.__file__}")
except Exception as e:
print(f"import packages failed: {e}")
PY
echo "=== coverage debug ==="
python3 - <<'PY'
import os, sys
sys.path.insert(0, f"{os.getcwd()}/apps/api")
sys.path.insert(0, os.getcwd())
import coverage
cov = coverage.Coverage(source=["apps/api/app", "packages"])
print(f"source: {cov.config.source}")
for src in cov.config.source or []:
abspath = os.path.abspath(src)
print(f" {src} -> {abspath} exists={os.path.exists(src)}")
if os.path.isdir(src):
pyfiles = []
for root, dirs, files in os.walk(src):
for f in files:
if f.endswith('.py'):
pyfiles.append(os.path.join(root, f))
print(f" .py files: {len(pyfiles)}")
PY
- name: Run unit tests
shell: sh
env:
USE_IN_MEMORY_DB: "true"
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: Build summary
if: github.ref == 'refs/heads/develop' || github.ref == 'refs/heads/main'
shell: sh
run: |
set -eu
echo "Build completed successfully!"
echo "Branch: ${GITHUB_REF_NAME}"
echo "Commit: ${GITHUB_SHA}"
# 输出最终覆盖率
python3 scripts/ci_coverage_summary.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
@@ -376,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
@@ -427,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:
@@ -563,15 +372,6 @@ 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
@@ -703,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
@@ -788,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
@@ -862,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
@@ -951,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
@@ -1025,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
@@ -1110,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/
-10
View File
@@ -8,12 +8,10 @@ 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.edit_templates import router as edit_templates_router
from app.api.routes.feature_flags import router as feature_flags_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
@@ -153,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"],
)
-26
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
@@ -645,31 +644,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(
-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}")
+46 -95
View File
@@ -4,15 +4,8 @@ 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,
@@ -43,6 +36,44 @@ logger = logging.getLogger(__name__)
router = APIRouter()
def _safe_enqueue_generation_task(
task: Any,
generation_task_repository: Any,
) -> bool:
"""安全入队:send_task 失败时自动把任务标记为 failed,避免留下 pending 僵尸任务。
Returns:
True 表示入队成功,False 表示入队失败(已标记为 failed)
"""
try:
celery_app.send_task("worker.generate_video", args=[task.id])
logger.info(
"[生成任务] 入队成功: task_id=%s, status=%s",
task.id,
task.status,
)
return True
except Exception as e:
logger.error(
"[生成任务] 入队失败,标记为失败: task_id=%s error=%s",
task.id,
e,
exc_info=True,
)
try:
task.mark_failed(f"任务入队失败: {e}")
generation_task_repository.update(task)
except Exception as update_err:
logger.error(
"[生成任务] 入队失败后更新状态也失败: task_id=%s error=%s",
task.id,
update_err,
exc_info=True,
)
return False
def _check_project_access(project_id: str, user_id: str, project_repository) -> None:
"""检查用户是否有项目访问权限"""
project = project_repository.find_by_id(project_id)
@@ -235,31 +266,9 @@ def create_generation_task(
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
)
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
try:
for _ in range(count):
task = use_case.execute(
@@ -272,42 +281,16 @@ def create_generation_task(
asset_ids=resolved_asset_ids,
title_ids=request.title_ids,
voice_ids=request.voice_ids,
created_by_user_id=user_id,
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,
)
)
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:
# 兜底:如果预检查后又并发提交了,在这里也拦住
if _safe_enqueue_generation_task(task, generation_task_repository):
created_tasks.append(task)
else:
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="创建生成任务失败,请稍后重试或查看任务日志")
@@ -382,21 +365,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(
@@ -408,28 +376,11 @@ 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
if not _safe_enqueue_generation_task(retried, generation_task_repository):
logger.warning("[生成任务] 重试入队失败: task_id=%s", 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,
)
+36 -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,38 @@ from packages.application import (
SubmitIngestJobUseCase,
)
logger = logging.getLogger(__name__)
router = APIRouter()
def _safe_enqueue_generation_task(
task: Any,
generation_task_repository: Any,
) -> bool:
"""安全入队:send_task 失败时自动把任务标记为 failed,避免留下 pending 僵尸任务。"""
try:
celery_app.send_task("worker.generate_video", args=[task.id])
logger.info("[任务中心] 生成任务入队成功: task_id=%s", task.id)
return True
except Exception as e:
logger.error(
"[任务中心] 生成任务入队失败,标记为失败: task_id=%s error=%s",
task.id,
e,
exc_info=True,
)
try:
task.mark_failed(f"任务入队失败: {e}")
generation_task_repository.update(task)
except Exception as update_err:
logger.error(
"[任务中心] 入队失败后更新状态也失败: task_id=%s error=%s",
task.id,
update_err,
exc_info=True,
)
return False
def _humanize_task_error(error_message: str) -> str:
raw = (error_message or "").strip()
if not raw:
@@ -148,21 +168,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 +179,11 @@ 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
if not _safe_enqueue_generation_task(retried, generation_task_repository):
logger.warning("[任务中心] 用户级重试入队失败: task_id=%s", retried.id)
return UserTaskResponse(
id=f"generation:{retried.id}",
task_type="generation",
@@ -259,22 +251,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 +262,11 @@ 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
if not _safe_enqueue_generation_task(retried, generation_task_repository):
logger.warning("[任务中心] 项目级重试用队失败: task_id=%s", 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
@@ -38,7 +38,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,
@@ -50,10 +50,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,
@@ -62,7 +59,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,
@@ -73,20 +70,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",
@@ -96,7 +89,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,
@@ -147,7 +140,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:
"""获取配音列表(预置音色 + 用户克隆音色)。
@@ -175,7 +167,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":
@@ -232,7 +224,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:
"""原有配音列表接口(仅返回用户克隆音色)。
@@ -242,7 +233,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,
)
@@ -252,14 +243,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)
@@ -268,7 +258,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)
@@ -294,7 +283,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)
@@ -303,7 +292,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(
@@ -325,7 +313,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
@@ -109,9 +109,6 @@ class Settings(BaseSettings):
LOG_LEVEL: str = "INFO"
CORS_ORIGINS_RAW: str = "http://localhost:3000,http://localhost:5173,http://localhost:8000"
# 渲染引擎选择:legacy=旧VideoComposeServiceunified=新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
-19
View File
@@ -207,7 +207,6 @@ def get_cosyvoice_service():
能被 CosyVoice 服务器下载。
"""
from app.core.storage import get_storage_service
from packages.application.cosyvoice_service import CosyVoiceService
storage = get_storage_service()
@@ -217,21 +216,3 @@ def get_cosyvoice_service():
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
View File
+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,
}
+10 -82
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 配置 ──────────────────────────────────────────────────────────────────
@@ -51,9 +43,6 @@ def oss_bucket() -> oss2.Bucket | None:
P0-2 修复:endpoint 不带 scheme 时自动补 https:// 前缀,
确保 sign_url 等依赖 scheme 的方法返回 HTTPS URL。
P0-staging 修复:增加 connect_timeout=10s,防止网络抖动时
TCP 握手阶段无限挂死,导致 worker 进程卡死。
Returns:
oss2.Bucket 实例,配置缺失时返回 None。
"""
@@ -64,12 +53,7 @@ def oss_bucket() -> oss2.Bucket | None:
# 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 +96,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: 目标存储键
@@ -125,71 +106,18 @@ def upload_to_oss(local_path: Path, storage_key: str) -> str | 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,
)
try:
bucket.put_object_from_file(storage_key, str(local_path))
settings = oss_settings()
if settings:
_, _, endpoint, bucket_name = settings
endpoint_clean = endpoint.replace("https://", "").replace("http://", "")
return f"https://{bucket_name}.{endpoint_clean}/{storage_key}"
return None
if result["error"]:
except Exception:
logger.exception("上传 OSS 失败: %s", storage_key)
return None
return result["url"]
def get_signed_download_url(storage_key_or_url: str, expires_seconds: int = 3600) -> str | None:
"""生成预签名下载 URL(用于私有 bucket 的 URL 校验或临时下载)。
@@ -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
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)
-3
View File
@@ -18,9 +18,6 @@ class WorkerSettings(BaseSettings):
environment: str = "development"
auto_create_schema: bool = False
# 渲染引擎选择:legacy=旧VideoComposeServiceunified=新UnifiedRenderService
render_engine: str = "legacy"
model_config = SettingsConfigDict(
env_file=".env",
env_file_encoding="utf-8",
+65 -155
View File
@@ -39,10 +39,6 @@ def _get_job_service():
def compose_video(self, job_id: str, **kwargs):
"""视频合成任务。
根据 RENDER_ENGINE 配置选择渲染引擎:
- legacy: 旧 VideoComposeServicefilter_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_keyfile_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)
+23 -36
View File
@@ -124,7 +124,6 @@ class _VirtualPlan:
id: str
name: str = ""
config: dict[str, Any] = field(default_factory=dict)
@dataclass
@@ -162,15 +161,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)
@@ -186,7 +183,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":
@@ -199,7 +195,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"},
)
)
@@ -219,7 +214,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:
@@ -232,7 +226,6 @@ def _build_plan_and_clips_from_task(
clip_type="main",
order=i,
asset_id=path_to_asset_id[p],
duration=path_duration[p],
)
)
@@ -387,37 +380,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()
@@ -434,15 +421,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}, "
@@ -793,12 +778,14 @@ def generate_video(self, task_id: str) -> dict:
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
from video_processing.oss_helpers import oss_bucket, normalize_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)
logger.info(
"URL 校验失败但 object_exists 确认文件存在,视为上传成功: storage_key=%s", key
)
if gen_task:
gen_task.append_log("OSS上传", "URL校验降级: object_exists确认存在", level="WARN")
else:
+1 -1
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
@@ -17,6 +16,7 @@ from packages.application.cosyvoice_service import (
CosyVoiceTimeoutError,
)
from packages.application.voice_clone.workflow import VoiceCloneWorkflowService
from video_processing.oss_helpers import get_signed_download_url
logger = logging.getLogger(__name__)
+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 直接 False100 直接 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()
-17
View File
@@ -91,23 +91,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)
+42 -74
View File
@@ -70,21 +70,21 @@ class CosyVoiceService:
- 音色克隆: 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)
- 语音合成: POST /services/audio/tts/SpeechSynthesizer (model=cosyvoice-v3.5-plus)
- 非流式: 同步返回音频 URL
使用示例:
service = CosyVoiceService(
api_key="your-api-key",
base_url="https://dashscope.aliyuncs.com/api/v1",
model="cosyvoice-v3-flash",
model="cosyvoice-v3.5-plus",
)
# 音色克隆
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")
"""
# 音色状态轮询配置
@@ -121,49 +121,16 @@ class CosyVoiceService:
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._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),
)
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
@@ -234,7 +201,8 @@ class CosyVoiceService:
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])
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)
@@ -355,7 +323,9 @@ class CosyVoiceService:
while attempts < self.CLONE_MAX_POLL_ATTEMPTS:
elapsed = time.time() - start_time
if elapsed > timeout:
raise CosyVoiceTimeoutError(f"音色克隆任务超时({timeout}秒): voice_id={voice_id}")
raise CosyVoiceTimeoutError(
f"音色克隆任务超时({timeout}秒): voice_id={voice_id}"
)
result = self.query_voice_status(voice_id)
status = result.get("status", "").upper()
@@ -363,7 +333,9 @@ class CosyVoiceService:
if status == "OK":
return {"voice_id": voice_id}
elif status == "UNDEPLOYED":
raise CosyVoiceError(f"音色克隆任务失败(审核未通过): voice_id={voice_id}")
raise CosyVoiceError(
f"音色克隆任务失败(审核未通过): voice_id={voice_id}"
)
elif status in ("DEPLOYING", "PENDING", "PROCESSING", ""):
# 继续轮询
time.sleep(self.CLONE_POLL_INTERVAL)
@@ -373,7 +345,9 @@ class CosyVoiceService:
time.sleep(self.CLONE_POLL_INTERVAL)
attempts += 1
raise CosyVoiceTimeoutError(f"音色克隆任务轮询次数超限: voice_id={voice_id}")
raise CosyVoiceTimeoutError(
f"音色克隆任务轮询次数超限: voice_id={voice_id}"
)
def clone_voice(
self,
@@ -488,7 +462,9 @@ class CosyVoiceService:
request_id = response.get("request_id", "")
if not audio_url:
raise CosyVoiceError(f"CosyVoice API 未返回 audio_url: {response}")
raise CosyVoiceError(
f"CosyVoice API 未返回 audio_url: {response}"
)
return {
"task_id": "", # 同步接口无 task_id,兼容旧接口
@@ -498,7 +474,9 @@ class CosyVoiceService:
"request_id": request_id,
}
def poll_synthesize_task(self, task_id: str, timeout: float = 120.0) -> dict:
def poll_synthesize_task(
self, task_id: str, timeout: float = 120.0
) -> dict:
"""轮询合成任务(同步接口无需轮询,保留兼容).
CosyVoice SpeechSynthesizer 非流式接口是同步的,
@@ -507,7 +485,10 @@ class CosyVoiceService:
Raises:
CosyVoiceError: 同步接口无需轮询
"""
raise CosyVoiceError("CosyVoice 非流式合成接口是同步的,无需轮询. " "请直接使用 submit_synthesize_task().")
raise CosyVoiceError(
"CosyVoice 非流式合成接口是同步的,无需轮询. "
"请直接使用 submit_synthesize_task()."
)
def synthesize_speech(
self,
@@ -608,22 +589,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,18 +601,13 @@ 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}")
raise CosyVoiceAuthError(
f"CosyVoice API 认证失败: HTTP {response.status_code}"
)
elif response.status_code == 400:
# 客户端错误,不重试
body_text = response.text
@@ -655,12 +615,19 @@ class CosyVoiceService:
body = response.json()
code = body.get("code", "")
message = body.get("message", "")
raise CosyVoiceError(f"CosyVoice API 参数错误: HTTP 400, " f"code={code}, message={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}")
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}")
last_error = CosyVoiceError(
f"CosyVoice API 服务端错误: HTTP {response.status_code}"
)
logger.warning(
"CosyVoice API 失败 (尝试 %d/%d): HTTP %d",
attempt + 1,
@@ -670,7 +637,8 @@ class CosyVoiceService:
else:
# 其他客户端错误,不重试
raise CosyVoiceError(
f"CosyVoice API 调用失败: HTTP {response.status_code}, " f"body={response.text}"
f"CosyVoice API 调用失败: HTTP {response.status_code}, "
f"body={response.text}"
)
except httpx.TimeoutException as e:
+24 -7
View File
@@ -219,7 +219,9 @@ class TTSWorkflowService:
# 新接口(同步):没有 task_id,重新合成
if not task_id:
logger.info(f"TTS 任务无 task_id,重新同步合成: job_id={job_id}")
logger.info(
f"TTS 任务无 task_id,重新同步合成: job_id={job_id}"
)
return self._resynthesize_and_complete(job)
# 旧接口遗留的 task_id,尝试轮询(兼容过渡)
@@ -233,7 +235,9 @@ class TTSWorkflowService:
)
except CosyVoiceError:
# 旧接口轮询失败,重新同步合成
logger.warning(f"旧 task_id 轮询失败,重新同步合成: job_id={job_id}, task_id={task_id}")
logger.warning(
f"旧 task_id 轮询失败,重新同步合成: job_id={job_id}, task_id={task_id}"
)
return self._resynthesize_and_complete(job)
def process_synthesis_result(
@@ -520,7 +524,10 @@ class TTSWorkflowService:
missing_indices = [i for i in range(segment_count) if results[i] is None]
if missing_indices:
logger.info(f"分段任务重新合成缺失段: job_id={job.id}, " f"缺失={len(missing_indices)}/{segment_count}")
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:
@@ -543,8 +550,13 @@ class TTSWorkflowService:
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}")
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)
# 所有分段完成,下载合并
@@ -552,7 +564,9 @@ class TTSWorkflowService:
try:
merged_data, total_duration = self._download_and_merge_segments(results, job)
permanent_url, storage_key = self._upload_merged_to_oss(merged_data, job.user_id, job.id, job.format)
permanent_url, storage_key = self._upload_merged_to_oss(
merged_data, job.user_id, job.id, job.format
)
job.mark_completed(
output_audio_url=permanent_url,
@@ -561,7 +575,10 @@ class TTSWorkflowService:
file_size=len(merged_data),
)
job = self.repository.update(job)
logger.info(f"分段合成完成(重新合成路径): job_id={job.id}, " f"merged_size={len(merged_data)}")
logger.info(
f"分段合成完成(重新合成路径): job_id={job.id}, "
f"merged_size={len(merged_data)}"
)
return job
except Exception as e:
+1 -3
View File
@@ -133,9 +133,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"voice_id={submit_result.get('voice_id')}")
except (CosyVoiceError, CosyVoiceAuthError) as e:
# CosyVoice 提交失败,标记为 failed
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:
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
+2 -2
View File
@@ -32,8 +32,8 @@ class SharedSettings(BaseSettings):
# 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_model: str = "cosyvoice-v3.5-plus"
cosyvoice_voice: str = "longxiaochun" # 默认音色
cosyvoice_sample_rate: int = 22050
cosyvoice_format: str = "mp3" # 输出格式:mp3/wav/pcm
# 音色克隆模型名(固定为 voice-enrollment
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"
-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,
+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
# 步模式下提交后状态为 processingvoice_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+
- httpxAPI 调用)
### 配置环境变量
```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 Flagflag-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个clipfade + 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,
)
# 逐帧统计在 stdoutstats_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,
)
@@ -1,68 +0,0 @@
"""ClassificationStatus 枚举兼容性测试。
验证历史脏数据 'done'不会导致枚举转换失败
"""
import pytest
from packages.domain.entities import ClassificationStatus
class TestClassificationStatusNormalValues:
"""正常值应该正确映射。"""
def test_pending(self):
assert ClassificationStatus("pending") == ClassificationStatus.PENDING
def test_processing(self):
assert ClassificationStatus("processing") == ClassificationStatus.PROCESSING
def test_completed(self):
assert ClassificationStatus("completed") == ClassificationStatus.COMPLETED
def test_failed(self):
assert ClassificationStatus("failed") == ClassificationStatus.FAILED
class TestClassificationStatusHistoricalValues:
"""历史脏数据应该正确映射到对应状态,不抛异常。"""
@pytest.mark.parametrize("value", ["done", "Done", "DONE", " done "])
def test_done_maps_to_completed(self, value):
"""生产环境发现的 'done' 历史值应映射为 COMPLETED。"""
assert ClassificationStatus(value) == ClassificationStatus.COMPLETED
@pytest.mark.parametrize("value", ["success", "finished", "complete"])
def test_other_done_like_values_map_to_completed(self, value):
assert ClassificationStatus(value) == ClassificationStatus.COMPLETED
@pytest.mark.parametrize("value", ["fail", "error", "err"])
def test_error_like_values_map_to_failed(self, value):
assert ClassificationStatus(value) == ClassificationStatus.FAILED
@pytest.mark.parametrize("value", ["process", "running", "run"])
def test_processing_like_values_map_to_processing(self, value):
assert ClassificationStatus(value) == ClassificationStatus.PROCESSING
class TestClassificationStatusFallback:
"""完全未知的值兜底为 PENDING,不抛500。"""
@pytest.mark.parametrize("value", ["unknown", "foo_bar", ""])
def test_unknown_value_falls_back_to_pending(self, value):
assert ClassificationStatus(value) == ClassificationStatus.PENDING
def test_none_value_falls_back_to_pending(self):
assert ClassificationStatus(None) == ClassificationStatus.PENDING # type: ignore[arg-type]
def test_int_value_falls_back_to_pending(self):
assert ClassificationStatus(123) == ClassificationStatus.PENDING # type: ignore[arg-type]
class TestClassificationStatusStrValue:
"""枚举值仍为字符串类型,不影响序列化。"""
def test_value_unchanged(self):
assert ClassificationStatus.COMPLETED.value == "completed"
assert ClassificationStatus.PENDING.value == "pending"
assert isinstance(ClassificationStatus.COMPLETED, str)
+29 -65
View File
@@ -23,7 +23,7 @@ def _make_service(
*,
api_key: str = "test-api-key",
base_url: str = "https://dashscope.aliyuncs.com/api/v1",
model: str = "cosyvoice-v3-flash",
model: str = "cosyvoice-v3.5-plus",
clone_model: str = "voice-enrollment",
http_client: httpx.Client | None = None,
audio_url_signer=None,
@@ -78,7 +78,7 @@ class TestSubmitCloneTask:
200,
{
"output": {
"voice_id": "cosyvoice-v3-flash-clone-abc123",
"voice_id": "cosyvoice-v3.5-plus-clone-abc123",
"status": "DEPLOYING",
},
"usage": {"count": 1},
@@ -92,7 +92,7 @@ class TestSubmitCloneTask:
voice_name="myvoice",
)
assert result["voice_id"] == "cosyvoice-v3-flash-clone-abc123"
assert result["voice_id"] == "cosyvoice-v3.5-plus-clone-abc123"
assert result["status"] == "DEPLOYING"
assert result["request_id"] == "req-001"
@@ -104,7 +104,7 @@ class TestSubmitCloneTask:
payload = call_args.kwargs["json"]
assert payload["model"] == "voice-enrollment"
assert payload["input"]["action"] == "create_voice"
assert payload["input"]["target_model"] == "cosyvoice-v3-flash"
assert payload["input"]["target_model"] == "cosyvoice-v3.5-plus"
assert payload["input"]["prefix"] == "myvoice"
assert payload["input"]["url"] == "https://example.com/audio.wav"
assert payload["input"]["language_hints"] == ["zh"]
@@ -197,11 +197,7 @@ class TestSubmitCloneTask:
payload = mock_client.request.call_args.kwargs["json"]
# 中文和特殊字符被过滤,剩下字母数字
assert (
payload["input"]["prefix"] == "2024"
or payload["input"]["prefix"] == "clone"
or len(payload["input"]["prefix"]) <= 10
)
assert payload["input"]["prefix"] == "2024" or payload["input"]["prefix"] == "clone" or len(payload["input"]["prefix"]) <= 10
def test_submit_auth_401_raises(self) -> None:
mock_client = MagicMock()
@@ -233,7 +229,7 @@ class TestQueryVoiceStatus:
{
"output": {
"status": "DEPLOYING",
"target_model": "cosyvoice-v3-flash",
"target_model": "cosyvoice-v3.5-plus",
"gmt_create": "2026-01-01T00:00:00Z",
"gmt_modified": "2026-01-01T00:01:00Z",
"resource_link": "https://...",
@@ -246,7 +242,7 @@ class TestQueryVoiceStatus:
result = service.query_voice_status("voice-123")
assert result["status"] == "DEPLOYING"
assert result["target_model"] == "cosyvoice-v3-flash"
assert result["target_model"] == "cosyvoice-v3.5-plus"
# 验证请求
payload = mock_client.request.call_args.kwargs["json"]
@@ -257,7 +253,7 @@ class TestQueryVoiceStatus:
def test_query_ok_status(self) -> None:
mock_client = MagicMock()
mock_client.request.return_value = _mock_response(
200, {"output": {"status": "OK", "target_model": "cosyvoice-v3-flash"}}
200, {"output": {"status": "OK", "target_model": "cosyvoice-v3.5-plus"}}
)
service = _make_service(http_client=mock_client)
@@ -282,7 +278,7 @@ class TestPollCloneTask:
def test_poll_ok_on_first_check(self) -> None:
mock_client = MagicMock()
mock_client.request.return_value = _mock_response(
200, {"output": {"status": "OK", "target_model": "cosyvoice-v3-flash"}}
200, {"output": {"status": "OK", "target_model": "cosyvoice-v3.5-plus"}}
)
service = _make_service(http_client=mock_client)
@@ -309,7 +305,9 @@ class TestPollCloneTask:
def test_poll_undeployed_raises_error(self) -> None:
mock_client = MagicMock()
mock_client.request.return_value = _mock_response(200, {"output": {"status": "UNDEPLOYED"}})
mock_client.request.return_value = _mock_response(
200, {"output": {"status": "UNDEPLOYED"}}
)
service = _make_service(http_client=mock_client)
service.CLONE_POLL_INTERVAL = 0.01
@@ -319,7 +317,9 @@ class TestPollCloneTask:
def test_poll_timeout_raises(self) -> None:
mock_client = MagicMock()
mock_client.request.return_value = _mock_response(200, {"output": {"status": "DEPLOYING"}})
mock_client.request.return_value = _mock_response(
200, {"output": {"status": "DEPLOYING"}}
)
service = _make_service(http_client=mock_client)
service.CLONE_POLL_INTERVAL = 0.01
@@ -396,7 +396,9 @@ class TestSynthesizeSpeech:
)
service = _make_service(http_client=mock_client)
result = service.synthesize_speech(text="你好世界", voice_id="longxiaochun_v3")
result = service.synthesize_speech(
text="你好世界", voice_id="longxiaochun"
)
assert isinstance(result, SynthesizeResult)
assert result.audio_url == "https://dashscope-result.oss.com/output.mp3"
@@ -407,9 +409,9 @@ class TestSynthesizeSpeech:
assert "/services/audio/tts/SpeechSynthesizer" in call_args.kwargs["url"]
payload = call_args.kwargs["json"]
assert payload["model"] == "cosyvoice-v3-flash"
assert payload["model"] == "cosyvoice-v3.5-plus"
assert payload["input"]["text"] == "你好世界"
assert payload["input"]["voice"] == "longxiaochun_v3"
assert payload["input"]["voice"] == "longxiaochun"
assert payload["input"]["format"] == "mp3"
assert payload["input"]["sample_rate"] == 22050
assert payload["input"]["rate"] == 1.0
@@ -424,12 +426,8 @@ class TestSynthesizeSpeech:
service = _make_service(http_client=mock_client)
service.synthesize_speech(
text="test",
voice_id="v1",
sample_rate=44100,
format="wav",
speed=1.5,
volume=80,
text="test", voice_id="v1", sample_rate=44100,
format="wav", speed=1.5, volume=80,
)
payload = mock_client.request.call_args.kwargs["json"]
@@ -465,8 +463,7 @@ class TestSynthesizeSpeech:
"""同步接口的 submit_synthesize_task 返回空 task_id 字段(兼容旧接口)."""
mock_client = MagicMock()
mock_client.request.return_value = _mock_response(
200,
{"output": {"audio": {"url": "https://e.com/a.mp3"}}},
200, {"output": {"audio": {"url": "https://e.com/a.mp3"}}},
)
service = _make_service(http_client=mock_client)
@@ -493,7 +490,9 @@ class TestRetryLogic:
mock_client.request.side_effect = [
_mock_response(500, text="Server Error"),
_mock_response(502, text="Bad Gateway"),
_mock_response(200, {"output": {"audio": {"url": "https://e.com/a.mp3"}}}),
_mock_response(
200, {"output": {"audio": {"url": "https://e.com/a.mp3"}}}
),
]
service = _make_service(http_client=mock_client)
@@ -548,7 +547,9 @@ class TestSanitizePrefix:
class TestCheckTaskStatus:
def test_check_task_status_uses_query_voice(self) -> None:
mock_client = MagicMock()
mock_client.request.return_value = _mock_response(200, {"output": {"status": "OK"}})
mock_client.request.return_value = _mock_response(
200, {"output": {"status": "OK"}}
)
service = _make_service(http_client=mock_client)
result = service.check_task_status("voice-123")
@@ -559,40 +560,3 @@ class TestCheckTaskStatus:
# 验证走的是 query_voice 路径
payload = mock_client.request.call_args.kwargs["json"]
assert payload["input"]["action"] == "query_voice"
# ── 配置与初始化 ─────────────────────────────────────────
class TestServiceConfiguration:
"""测试 CosyVoiceService 配置与初始化逻辑."""
def test_base_url_old_text2audio_path_auto_fixed(self) -> None:
"""旧版 base_url 包含 text2audio 路径时,应自动修正为 /api/v1."""
service = CosyVoiceService(
api_key="test-key",
base_url="https://dashscope.aliyuncs.com/api/v1/services/aigc/text2audio",
model="cosyvoice-v3-flash",
)
# 应自动去掉 text2audio 后缀,保留到 /api/v1
assert service._base_url == "https://dashscope.aliyuncs.com/api/v1"
def test_base_url_normal_unchanged(self) -> None:
"""正常的 base_url 不应被修改."""
url = "https://dashscope.aliyuncs.com/api/v1"
service = CosyVoiceService(
api_key="test-key",
base_url=url,
model="cosyvoice-v3-flash",
)
assert service._base_url == url
def test_base_url_workspace_domain_unchanged(self) -> None:
"""工作空间专属域名的 base_url 不应被修改."""
url = "https://workspace-xxx.cn-beijing.maas.aliyuncs.com/api/v1"
service = CosyVoiceService(
api_key="test-key",
base_url=url,
model="cosyvoice-v3-flash",
)
assert service._base_url == url
-12
View File
@@ -177,18 +177,6 @@ class StubGenerationTaskRepository:
def count_by_user(self, user_id: str) -> int:
return len([t for t in self._store.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._store.values()
if t.created_by_user_id == user_id and getattr(t, "status", "") == "pending"
]
)
def count_pending_total(self) -> int:
return len([t for t in self._store.values() if getattr(t, "status", "") == "pending"])
def list_recent_by_user(self, user_id: str, limit: int = 5) -> list[Any]:
items = [t for t in self._store.values() if t.created_by_user_id == user_id]
items.sort(key=lambda t: t.created_at, reverse=True)
-6
View File
@@ -194,12 +194,6 @@ class StubGenerationTaskRepository:
self._tasks[task.id] = task
return task
def count_pending_by_user(self, user_id: str) -> int:
return 0
def count_pending_total(self) -> int:
return 0
# ---------------------------------------------------------------------------
# Service factory
-2
View File
@@ -63,8 +63,6 @@ class StubEditPlan:
template_id: str = "tmpl-001"
status: Any = None
config: dict = field(default_factory=dict)
project_id: str = ""
created_by_user_id: str = "user-001"
def mark_failed(self):
self.status = _StubStatus("failed")
-424
View File
@@ -1,424 +0,0 @@
"""Feature Flag 单元测试。
测试 FeatureFlagConfigInMemoryFeatureFlagStoreRenderEngineResolver 的核心逻辑
"""
from __future__ import annotations
import time
from unittest.mock import MagicMock, patch
import pytest
from packages.adapters.redis.feature_flag_store import (
FeatureFlagConfig,
InMemoryFeatureFlagStore,
)
# ── FeatureFlagConfig 测试 ──────────────────────────────────────────────────
class TestFeatureFlagConfig:
"""FeatureFlagConfig 核心逻辑测试。"""
def test_default_disabled(self):
"""默认配置为关闭状态。"""
config = FeatureFlagConfig(name="test_flag")
assert config.enabled is False
assert config.percentage == 0
assert config.whitelist == set()
assert config.is_active() is False
assert config.is_active("user1") is False
def test_global_enabled_100_percent(self):
"""100% + 启用 = 全部命中。"""
config = FeatureFlagConfig(name="test_flag", enabled=True, percentage=100)
assert config.is_active() is True
assert config.is_active("user1") is True
assert config.is_active("any_user") is True
def test_global_enabled_0_percent_no_whitelist(self):
"""启用但 0% 且无白名单 = 不命中。"""
config = FeatureFlagConfig(name="test_flag", enabled=True, percentage=0)
assert config.is_active() is False
assert config.is_active("user1") is False
def test_whitelist_takes_priority(self):
"""白名单优先级高于百分比。"""
config = FeatureFlagConfig(
name="test_flag",
enabled=True,
percentage=0,
whitelist={"user1", "user2"},
)
assert config.is_active("user1") is True
assert config.is_active("user2") is True
assert config.is_active("user3") is False
def test_whitelist_with_percentage(self):
"""白名单用户即使百分比为0也命中,非白名单按百分比。"""
config = FeatureFlagConfig(
name="test_flag",
enabled=True,
percentage=100, # 100% 所有人命中
whitelist={"user1"},
)
assert config.is_active("user1") is True
assert config.is_active("user999") is True # 100% 命中
def test_percentage_consistency_same_user(self):
"""同一用户多次调用结果一致(哈希确定性)。"""
config = FeatureFlagConfig(name="test_flag", enabled=True, percentage=50)
results = [config.is_active("user_fixed") for _ in range(100)]
assert all(r == results[0] for r in results)
def test_percentage_different_users_distributed(self):
"""不同用户分布大致符合百分比(统计检验,宽松阈值)。"""
config = FeatureFlagConfig(name="test_flag", enabled=True, percentage=50)
active_count = sum(1 for i in range(1000) if config.is_active(f"user_{i}"))
# 50% 上下浮动 10% 都算合理
assert 400 <= active_count <= 600, f"Expected ~500, got {active_count}"
def test_percentage_boundary_0_and_100(self):
"""0% 和 100% 的边界情况。"""
config_0 = FeatureFlagConfig(name="test", enabled=True, percentage=0)
config_100 = FeatureFlagConfig(name="test", enabled=True, percentage=100)
for i in range(100):
assert config_0.is_active(f"user_{i}") is False
assert config_100.is_active(f"user_{i}") is True
def test_disabled_ignores_all_other_settings(self):
"""关闭时忽略白名单和百分比。"""
config = FeatureFlagConfig(
name="test_flag",
enabled=False,
percentage=100,
whitelist={"user1"},
)
assert config.is_active("user1") is False
assert config.is_active() is False
def test_none_identifier_with_percentage(self):
"""无 identifier 时按随机比例(0% 和 100% 是确定的)。"""
config_0 = FeatureFlagConfig(name="test", enabled=True, percentage=0)
config_100 = FeatureFlagConfig(name="test", enabled=True, percentage=100)
assert config_0.is_active(None) is False
assert config_100.is_active(None) is True
def test_to_dict_and_from_dict(self):
"""序列化和反序列化对称。"""
original = FeatureFlagConfig(
name="test_flag",
enabled=True,
percentage=30,
whitelist={"user_a", "user_b", "user_c"},
)
data = original.to_dict()
restored = FeatureFlagConfig.from_dict(data)
assert restored.name == original.name
assert restored.enabled == original.enabled
assert restored.percentage == original.percentage
assert restored.whitelist == original.whitelist
def test_from_dict_with_missing_fields(self):
"""from_dict 缺失字段时使用默认值。"""
config = FeatureFlagConfig.from_dict({"name": "minimal"})
assert config.name == "minimal"
assert config.enabled is False
assert config.percentage == 0
assert config.whitelist == set()
# ── InMemoryFeatureFlagStore 测试 ───────────────────────────────────────────
class TestInMemoryFeatureFlagStore:
"""内存存储实现测试。"""
def test_get_nonexistent_returns_default(self):
"""获取不存在的 flag 返回默认配置(关闭)。"""
store = InMemoryFeatureFlagStore()
config = store.get("nonexistent")
assert config.name == "nonexistent"
assert config.enabled is False
def test_set_and_get(self):
"""设置后可以读取。"""
store = InMemoryFeatureFlagStore()
config = FeatureFlagConfig(name="test", enabled=True, percentage=50, whitelist={"u1"})
store.set(config)
got = store.get("test")
assert got.enabled is True
assert got.percentage == 50
assert got.whitelist == {"u1"}
def test_delete_existing(self):
"""删除存在的 flag 返回 True。"""
store = InMemoryFeatureFlagStore()
store.set(FeatureFlagConfig(name="test", enabled=True))
assert store.delete("test") is True
assert store.get("test").enabled is False
def test_delete_nonexistent(self):
"""删除不存在的 flag 返回 False。"""
store = InMemoryFeatureFlagStore()
assert store.delete("nonexistent") is False
def test_list_all(self):
"""列出所有 flag。"""
store = InMemoryFeatureFlagStore()
store.set(FeatureFlagConfig(name="flag_a", enabled=True))
store.set(FeatureFlagConfig(name="flag_b", percentage=10))
all_flags = store.list_all()
assert len(all_flags) == 2
assert "flag_a" in all_flags
assert "flag_b" in all_flags
assert all_flags["flag_a"].enabled is True
def test_is_active_convenience(self):
"""is_active 便捷方法。"""
store = InMemoryFeatureFlagStore()
store.set(FeatureFlagConfig(name="render", enabled=True, percentage=0, whitelist={"vip_user"}))
assert store.is_active("render", "vip_user") is True
assert store.is_active("render", "normal_user") is False
assert store.is_active("nonexistent") is False
# ── RenderEngineResolver 测试 ───────────────────────────────────────────────
class TestRenderEngineResolver:
"""渲染引擎选择器测试。"""
def test_default_legacy_when_flag_disabled(self):
"""flag 关闭时使用默认引擎(legacy)。"""
store = InMemoryFeatureFlagStore()
resolver = self._make_resolver(store=store, default="legacy")
assert resolver.get_engine() == "legacy"
assert resolver.get_engine("user1") == "legacy"
def test_default_unified_when_flag_disabled(self):
"""flag 关闭但默认值是 unified 时返回 unified。"""
store = InMemoryFeatureFlagStore()
resolver = self._make_resolver(store=store, default="unified")
assert resolver.get_engine() == "unified"
def test_whitelist_user_uses_unified(self):
"""白名单用户走新引擎。"""
store = InMemoryFeatureFlagStore()
store.set(
FeatureFlagConfig(
name="render_engine",
enabled=True,
percentage=0,
whitelist={"beta_tester"},
)
)
resolver = self._make_resolver(store=store, default="legacy")
assert resolver.get_engine("beta_tester") == "unified"
assert resolver.get_engine("normal_user") == "legacy"
def test_100_percent_all_unified(self):
"""100% 时所有用户走新引擎。"""
store = InMemoryFeatureFlagStore()
store.set(FeatureFlagConfig(name="render_engine", enabled=True, percentage=100))
resolver = self._make_resolver(store=store, default="legacy")
for i in range(50):
assert resolver.get_engine(f"user_{i}") == "unified"
def test_invalid_default_engine_fallback(self):
"""无效默认值回退到 legacy。"""
store = InMemoryFeatureFlagStore()
resolver = self._make_resolver(store=store, default="invalid_value")
assert resolver.get_engine() == "legacy"
def test_should_use_unified_helper(self):
"""should_use_unified 便捷方法。"""
store = InMemoryFeatureFlagStore()
store.set(
FeatureFlagConfig(
name="render_engine",
enabled=True,
percentage=0,
whitelist={"user_a"},
)
)
resolver = self._make_resolver(store=store)
assert resolver.should_use_unified("user_a") is True
assert resolver.should_use_unified("user_b") is False
def test_config_snapshot(self):
"""配置快照。"""
store = InMemoryFeatureFlagStore()
store.set(
FeatureFlagConfig(
name="render_engine",
enabled=True,
percentage=30,
whitelist={"u1", "u2"},
)
)
resolver = self._make_resolver(store=store)
snapshot = resolver.get_config_snapshot()
assert snapshot["flag_name"] == "render_engine"
assert snapshot["enabled"] is True
assert snapshot["percentage"] == 30
assert snapshot["whitelist"] == ["u1", "u2"]
def test_set_flag_updates_config(self):
"""通过 set_flag 修改后立即生效。"""
store = InMemoryFeatureFlagStore()
resolver = self._make_resolver(store=store, default="legacy")
# 初始:关闭
assert resolver.get_engine("user1") == "legacy"
# 开启 100%
resolver.set_flag(FeatureFlagConfig(name="render_engine", enabled=True, percentage=100))
assert resolver.get_engine("user1") == "unified"
# 关闭
resolver.set_flag(FeatureFlagConfig(name="render_engine", enabled=False))
assert resolver.get_engine("user1") == "legacy"
def test_force_refresh(self):
"""强制刷新不报错。"""
store = InMemoryFeatureFlagStore()
resolver = self._make_resolver(store=store)
resolver.force_refresh() # 不抛异常即可
def test_does_not_affect_in_flight_tasks(self):
"""
热更新不影响在途任务验证
任务开始时确定引擎中途配置变更不改变当前任务的引擎选择
这是通过"每次调用 get_engine 时读取当前配置"来保证的
任务开始时调用一次拿到结果之后不再变化
"""
store = InMemoryFeatureFlagStore()
store.set(FeatureFlagConfig(name="render_engine", enabled=True, percentage=100))
resolver = self._make_resolver(store=store, default="legacy")
# 模拟任务开始时获取引擎
engine_at_start = resolver.get_engine("user1")
assert engine_at_start == "unified"
# 任务进行中关闭 flag
store.set(FeatureFlagConfig(name="render_engine", enabled=False))
resolver.force_refresh()
# 在途任务持有的 engine_at_start 仍然是 unified(不随配置变化)
assert engine_at_start == "unified"
# 新任务会拿到 legacy
assert resolver.get_engine("user1") == "legacy"
# ── 辅助方法 ──
@staticmethod
def _make_resolver(store=None, default="legacy"):
from apps.worker.video_processing.render_engine_resolver import (
RenderEngineResolver,
)
return RenderEngineResolver(
default_engine=default,
store=store or InMemoryFeatureFlagStore(),
refresh_interval=9999, # 测试时禁用自动刷新
)
# ── RedisFeatureFlagStore 降级测试(无 Redis 环境) ───────────────────────
class TestRedisStoreDegradation:
"""Redis 不可用时的降级行为测试。"""
def test_get_returns_default_when_redis_unavailable(self):
"""Redis 连接失败时返回默认关闭配置,不抛异常。"""
import importlib
from packages.adapters.redis import feature_flag_store as ff_module
# 模拟 redis 模块不存在的场景不好做,这里直接测试异常捕获逻辑
store = ff_module.RedisFeatureFlagStore.__new__(ff_module.RedisFeatureFlagStore)
store._redis = MagicMock()
store._redis.hgetall.side_effect = ConnectionError("Redis down")
store._key_prefix = ff_module.FEATURE_FLAG_REDIS_PREFIX
store._cache = {}
store._cache_ttl = 5.0
import threading
store._lock = threading.Lock()
config = store.get("render_engine")
assert config.enabled is False
assert config.name == "render_engine"
def test_list_all_returns_empty_on_redis_error(self):
"""Redis 错误时 list_all 返回空字典。"""
import importlib
from packages.adapters.redis import feature_flag_store as ff_module
store = ff_module.RedisFeatureFlagStore.__new__(ff_module.RedisFeatureFlagStore)
store._redis = MagicMock()
store._redis.scan.side_effect = ConnectionError("Redis down")
store._key_prefix = ff_module.FEATURE_FLAG_REDIS_PREFIX
store._cache = {}
store._cache_ttl = 5.0
import threading
store._lock = threading.Lock()
result = store.list_all()
assert result == {}
class TestRedisStoreListAll:
"""RedisFeatureFlagStore list_all 正常路径测试。"""
def _make_store(self):
from packages.adapters.redis import feature_flag_store as ff_module
store = ff_module.RedisFeatureFlagStore.__new__(ff_module.RedisFeatureFlagStore)
store._redis = MagicMock()
store._key_prefix = ff_module.FEATURE_FLAG_REDIS_PREFIX
store._cache = {}
store._cache_ttl = 5.0
import threading
store._lock = threading.Lock()
return store
def test_list_all_scan_with_match_param(self):
"""list_all 调用 redis.scan 时使用正确的 match 参数名。"""
store = self._make_store()
prefix = store._key_prefix
# 模拟 scan 返回 2 个 key,分 2 次游标
store._redis.scan.side_effect = [
(10, [f"{prefix}render_engine", f"{prefix}other_flag"]),
(0, []),
]
# 模拟 hgetall 返回配置
store._redis.hgetall.return_value = {
b"enabled": b"true",
b"percentage": b"50",
b"whitelist": b'["user1","user2"]',
}
result = store.list_all()
# 验证 scan 被调用了 2 次(游标遍历)
assert store._redis.scan.call_count == 2
# 验证参数名是 match(不是 match_pattern
first_call_kwargs = store._redis.scan.call_args_list[0][1]
assert "match" in first_call_kwargs
assert "match_pattern" not in first_call_kwargs
assert first_call_kwargs["match"] == f"{prefix}*"
# 验证返回了 2 个 flag
assert len(result) == 2
assert "render_engine" in result
assert "other_flag" in result
@@ -1,93 +0,0 @@
"""FFmpeg 超时保护测试。
验证 run_ffmpeg / probe_video_info 的超时保护机制
防止 FFmpeg hang 住导致 worker 永久阻塞
"""
from __future__ import annotations
import subprocess
from unittest.mock import MagicMock, patch
import pytest
from video_processing.ffmpeg_utils import (
DEFAULT_FFMPEG_TIMEOUT,
probe_video_info,
run_ffmpeg,
)
# ── run_ffmpeg 超时保护 ──────────────────────────────────────────────────────
class TestRunFFmpegTimeout:
"""run_ffmpeg 超时保护测试。"""
def test_default_timeout_is_set(self):
"""默认超时应为 1800 秒(30分钟)。"""
assert DEFAULT_FFMPEG_TIMEOUT == 1800
def test_timeout_expired_is_raised(self):
"""超时未完成时 TimeoutExpired 异常被传播。"""
with patch("video_processing.ffmpeg_utils.subprocess.run") as mock_run:
mock_run.side_effect = subprocess.TimeoutExpired(cmd=["ffmpeg", "test"], timeout=1)
with pytest.raises(subprocess.TimeoutExpired):
run_ffmpeg(["ffmpeg", "test"])
def test_custom_timeout(self):
"""支持自定义超时时间。"""
with patch("video_processing.ffmpeg_utils.subprocess.run") as mock_run:
mock_run.side_effect = subprocess.TimeoutExpired(cmd=["ffmpeg"], timeout=5)
with pytest.raises(subprocess.TimeoutExpired):
run_ffmpeg(["ffmpeg", "test"], timeout=5)
def test_none_timeout_disables_protection(self):
"""timeout=None 可以禁用超时保护(不推荐)。"""
with patch("video_processing.ffmpeg_utils.subprocess.run") as mock_run:
mock_result = MagicMock()
mock_result.stdout = ""
mock_result.stderr = ""
mock_run.return_value = mock_result
run_ffmpeg(["ffmpeg", "test"], timeout=None)
# 验证 timeout=None 被传递
call_kwargs = mock_run.call_args.kwargs
assert call_kwargs["timeout"] is None
def test_called_process_error_still_raised(self):
"""超时异常不影响原有 CalledProcessError 的抛出。"""
with patch("video_processing.ffmpeg_utils.subprocess.run") as mock_run:
mock_run.side_effect = subprocess.CalledProcessError(returncode=1, cmd=["ffmpeg"], stderr="error msg")
with pytest.raises(subprocess.CalledProcessError):
run_ffmpeg(["ffmpeg", "test"])
# ── probe_video_info 超时保护 ────────────────────────────────────────────────
class TestProbeVideoInfoTimeout:
"""probe_video_info 超时保护测试。"""
def test_probe_uses_timeout(self):
"""probe_video_info 调用 ffprobe 时应设置 timeout=15。"""
with patch("video_processing.ffmpeg_utils.subprocess.run") as mock_run:
mock_run.side_effect = subprocess.TimeoutExpired(cmd=["ffprobe"], timeout=15)
# 超时异常被捕获,返回默认值
result = probe_video_info("/tmp/test.mp4")
assert result["width"] == 1280 # DEFAULT_OUTPUT_WIDTH
assert result["height"] == 720 # DEFAULT_OUTPUT_HEIGHT
def test_probe_success(self):
"""正常情况应解析 ffprobe JSON 输出。"""
fake_output = """
{
"streams": [{"width": 1920, "height": 1080, "r_frame_rate": "30/1", "duration": "10.5"}],
"format": {"duration": "10.5"}
}
"""
with patch("video_processing.ffmpeg_utils.subprocess.run") as mock_run:
mock_result = MagicMock()
mock_result.stdout = fake_output
mock_run.return_value = mock_result
result = probe_video_info("/tmp/test.mp4")
assert result["width"] == 1920
assert result["height"] == 1080
assert abs(result["duration"] - 10.5) < 0.01
-6
View File
@@ -58,12 +58,6 @@ class StubGenerationTaskRepository:
def get(self, task_id):
return self._tasks.get(task_id)
def count_pending_by_user(self, user_id):
return 0
def count_pending_total(self):
return 0
class StubGeneratedVideoRepository:
def __init__(self, videos=None):
-190
View File
@@ -1,190 +0,0 @@
"""渲染结果内部下载接口单元测试。
测试 internal_render 路由的核心逻辑mock repository storage 依赖
"""
from __future__ import annotations
from unittest.mock import MagicMock
import pytest
from app.api.routes.internal_render import (
InternalRenderDownloadUrlResponse,
InternalRenderTaskVideosResponse,
_video_to_item,
get_render_task_videos,
get_render_video_download_url,
)
# ── Helpers ────────────────────────────────────────────────────────────────
class MockVideo:
"""模拟 GeneratedVideo 领域对象。"""
def __init__(self, **kwargs):
self.id = kwargs.get("id", "video-1")
self.generation_task_id = kwargs.get("generation_task_id", "task-1")
self.project_id = kwargs.get("project_id", "proj-1")
self.name = kwargs.get("name", "test_video.mp4")
self.file_url = kwargs.get("file_url", "videos/test/output.mp4")
self.file_size = kwargs.get("file_size", 1024000)
self.duration = kwargs.get("duration", 30.5)
self.width = kwargs.get("width", 1080)
self.height = kwargs.get("height", 1920)
self.fps = kwargs.get("fps", 30.0)
self.status = kwargs.get("status", "completed")
# ── _video_to_item 测试 ────────────────────────────────────────────────────
class TestVideoToItem:
"""测试视频对象转响应项。"""
def test_basic_conversion(self):
video = MockVideo(id="v1", generation_task_id="t1", status="completed")
item = _video_to_item(video, "https://oss.example.com/download?v1")
assert item.video_id == "v1"
assert item.generation_task_id == "t1"
assert item.status == "completed"
assert item.download_url == "https://oss.example.com/download?v1"
def test_missing_optional_fields(self):
"""缺可选字段时返回 None。"""
video = MockVideo()
# 去掉可选字段
del video.file_size
del video.duration
item = _video_to_item(video, "https://example.com/dl")
assert item.file_size is None
assert item.duration is None
assert item.width == 1080 # 还在
# ── 路由函数测试 ────────────────────────────────────────────────────────────
class TestGetRenderVideoDownloadUrl:
"""测试单个视频下载URL接口。"""
def test_video_exists(self):
video = MockVideo(id="v-abc", file_url="videos/abc/out.mp4")
mock_repo = MagicMock()
mock_repo.get.return_value = video
mock_storage = MagicMock()
mock_storage.get_download_url.return_value = "https://oss.test/signed?v=abc"
result = get_render_video_download_url(
video_id="v-abc",
_=True,
generated_video_repository=mock_repo,
storage_service=mock_storage,
)
assert isinstance(result, InternalRenderDownloadUrlResponse)
assert result.video_id == "v-abc"
assert result.download_url == "https://oss.test/signed?v=abc"
mock_repo.get.assert_called_once_with("v-abc")
mock_storage.get_download_url.assert_called_once()
def test_video_not_found_raises_404(self):
from fastapi import HTTPException
mock_repo = MagicMock()
mock_repo.get.return_value = None
mock_storage = MagicMock()
with pytest.raises(HTTPException) as exc_info:
get_render_video_download_url(
video_id="nonexistent",
_=True,
generated_video_repository=mock_repo,
storage_service=mock_storage,
)
assert exc_info.value.status_code == 404
def test_download_url_long_expiry(self):
"""过期时间应为 24 小时(86400s)。"""
video = MockVideo(id="v1")
mock_repo = MagicMock()
mock_repo.get.return_value = video
mock_storage = MagicMock()
mock_storage.get_download_url.return_value = "https://oss.test/signed"
get_render_video_download_url(
video_id="v1",
_=True,
generated_video_repository=mock_repo,
storage_service=mock_storage,
)
# 验证 expires_seconds=86400
call_kwargs = mock_storage.get_download_url.call_args
assert call_kwargs.kwargs.get("expires_seconds") == 86400 or call_kwargs[1].get("expires_seconds") == 86400
class TestGetRenderTaskVideos:
"""测试任务视频列表接口。"""
def test_list_multiple_videos(self):
videos = [
MockVideo(id="v1", status="completed"),
MockVideo(id="v2", status="completed"),
MockVideo(id="v3", status="failed"),
]
mock_repo = MagicMock()
mock_repo.list_by_generation_task.return_value = videos
mock_storage = MagicMock()
mock_storage.get_download_url.return_value = "https://oss.test/signed"
result = get_render_task_videos(
task_id="task-1",
status=None,
_=True,
generated_video_repository=mock_repo,
storage_service=mock_storage,
)
assert isinstance(result, InternalRenderTaskVideosResponse)
assert result.task_id == "task-1"
assert result.count == 3
assert len(result.videos) == 3
def test_filter_by_status(self):
videos = [
MockVideo(id="v1", status="completed"),
MockVideo(id="v2", status="completed"),
MockVideo(id="v3", status="failed"),
]
mock_repo = MagicMock()
mock_repo.list_by_generation_task.return_value = videos
mock_storage = MagicMock()
mock_storage.get_download_url.return_value = "https://oss.test/signed"
result = get_render_task_videos(
task_id="task-1",
status="completed",
_=True,
generated_video_repository=mock_repo,
storage_service=mock_storage,
)
assert result.count == 2
assert all(v.status == "completed" for v in result.videos)
def test_empty_task(self):
mock_repo = MagicMock()
mock_repo.list_by_generation_task.return_value = []
mock_storage = MagicMock()
result = get_render_task_videos(
task_id="empty-task",
status=None,
_=True,
generated_video_repository=mock_repo,
storage_service=mock_storage,
)
assert result.count == 0
assert result.videos == []
+3 -1
View File
@@ -132,8 +132,10 @@ class TestDownloadLibraryAssets:
session.query.return_value = query
filter_result = MagicMock()
query.filter.return_value = filter_result
in_filter = MagicMock()
filter_result.filter.return_value = in_filter
id_filter = MagicMock()
filter_result.filter.return_value = id_filter
in_filter.filter.return_value = id_filter
assets = [self._make_asset("a1", "video/a1.mp4")]
id_filter.order_by.return_value.all.return_value = assets
-236
View File
@@ -1,236 +0,0 @@
"""P0-stagingOSS 上传崩溃修复测试.
测试
1. oss_bucket() 传递 connect_timeout 参数
2. upload_to_oss() 小文件走 put_object_from_file大文件走分片上传
3. upload_to_oss() 超时保护超过总超时返回 None
4. upload_to_oss() 异常时返回 None
"""
from __future__ import annotations
import os
import tempfile
import time
from pathlib import Path
from unittest.mock import MagicMock, patch
import pytest
# ── oss_bucket connect_timeout 测试 ───────────────────────────────────────────
class TestOSSBucketConnectTimeout:
"""测试 oss_bucket() 传递 connect_timeout 参数."""
def test_oss_bucket_has_connect_timeout(self):
"""oss_bucket 应传递 connect_timeout=10s 参数."""
from video_processing.oss_helpers import oss_bucket
mock_bucket_instance = MagicMock()
with (
patch.dict(
os.environ,
{
"OSS_ACCESS_KEY_ID": "test-key",
"OSS_ACCESS_KEY_SECRET": "test-secret",
"OSS_ENDPOINT": "oss-cn-hangzhou.aliyuncs.com",
"OSS_BUCKET_NAME": "test-bucket",
},
),
patch("video_processing.oss_helpers.oss2.Auth"),
patch("video_processing.oss_helpers.oss2.Bucket", return_value=mock_bucket_instance) as mock_bucket_cls,
):
bucket = oss_bucket()
assert bucket is mock_bucket_instance
# 验证 connect_timeout 关键字参数
call_kwargs = mock_bucket_cls.call_args[1]
assert "connect_timeout" in call_kwargs, "oss_bucket 应传递 connect_timeout 参数"
assert (
call_kwargs["connect_timeout"] == 10
), f"connect_timeout 应为 10,实际为 {call_kwargs['connect_timeout']}"
def test_oss_bucket_no_config_returns_none(self):
"""OSS 配置缺失时返回 None."""
from video_processing.oss_helpers import oss_bucket
with patch.dict(os.environ, {}, clear=True):
bucket = oss_bucket()
assert bucket is None
# ── upload_to_oss 分片上传测试 ────────────────────────────────────────────────
class TestUploadToOSSMultipart:
"""测试 upload_to_oss() 根据文件大小选择上传方式."""
def _create_temp_file(self, size_bytes: int) -> Path:
"""创建指定大小的临时文件."""
tmp = tempfile.NamedTemporaryFile(delete=False, suffix=".mp4")
tmp.write(b"x" * size_bytes)
tmp.close()
return Path(tmp.name)
def test_small_file_uses_put_object(self):
"""小文件(<100MB)走 put_object_from_file."""
from video_processing.oss_helpers import upload_to_oss
small_file = self._create_temp_file(10 * 1024 * 1024) # 10MB
try:
mock_bucket = MagicMock()
with (
patch.dict(
os.environ,
{
"OSS_ACCESS_KEY_ID": "test-key",
"OSS_ACCESS_KEY_SECRET": "test-secret",
"OSS_ENDPOINT": "oss-cn-hangzhou.aliyuncs.com",
"OSS_BUCKET_NAME": "test-bucket",
},
),
patch("video_processing.oss_helpers.oss2.Auth"),
patch("video_processing.oss_helpers.oss2.Bucket", return_value=mock_bucket),
patch("video_processing.oss_helpers.oss2.resumable_upload") as mock_resumable,
):
url = upload_to_oss(small_file, "test/small.mp4")
# 验证调用了 put_object_from_file
mock_bucket.put_object_from_file.assert_called_once()
# 验证没调用分片上传
mock_resumable.assert_not_called()
# 验证返回 URL
assert url == "https://test-bucket.oss-cn-hangzhou.aliyuncs.com/test/small.mp4"
finally:
small_file.unlink()
def test_large_file_uses_resumable_upload(self):
"""大文件(>=100MB)走 resumable_upload 分片上传."""
from video_processing.oss_helpers import upload_to_oss
large_file = self._create_temp_file(100 * 1024 * 1024) # 100MB
try:
mock_bucket = MagicMock()
with (
patch.dict(
os.environ,
{
"OSS_ACCESS_KEY_ID": "test-key",
"OSS_ACCESS_KEY_SECRET": "test-secret",
"OSS_ENDPOINT": "oss-cn-hangzhou.aliyuncs.com",
"OSS_BUCKET_NAME": "test-bucket",
},
),
patch("video_processing.oss_helpers.oss2.Auth"),
patch("video_processing.oss_helpers.oss2.Bucket", return_value=mock_bucket),
patch("video_processing.oss_helpers.oss2.resumable_upload") as mock_resumable,
):
url = upload_to_oss(large_file, "test/large.mp4")
# 验证调用了分片上传
mock_resumable.assert_called_once()
# 验证没调用 put_object_from_file
mock_bucket.put_object_from_file.assert_not_called()
# 验证分片参数
call_kwargs = mock_resumable.call_args[1]
assert call_kwargs["multipart_threshold"] == 100 * 1024 * 1024
assert call_kwargs["part_size"] == 8 * 1024 * 1024
assert call_kwargs["num_threads"] == 3
# 验证返回 URL
assert url == "https://test-bucket.oss-cn-hangzhou.aliyuncs.com/test/large.mp4"
finally:
large_file.unlink()
# ── upload_to_oss 超时测试 ────────────────────────────────────────────────────
class TestUploadToOSSTimeout:
"""测试 upload_to_oss() 超时保护."""
def test_upload_timeout_returns_none(self):
"""上传超过总超时时返回 None."""
from video_processing.oss_helpers import upload_to_oss
small_file = tempfile.NamedTemporaryFile(delete=False, suffix=".mp4")
small_file.write(b"x" * 1024) # 1KB
small_file.close()
file_path = Path(small_file.name)
def slow_upload(*args, **kwargs):
"""模拟慢速上传,超过超时时间."""
time.sleep(2)
mock_bucket = MagicMock()
mock_bucket.put_object_from_file.side_effect = slow_upload
try:
with (
patch.dict(
os.environ,
{
"OSS_ACCESS_KEY_ID": "test-key",
"OSS_ACCESS_KEY_SECRET": "test-secret",
"OSS_ENDPOINT": "oss-cn-hangzhou.aliyuncs.com",
"OSS_BUCKET_NAME": "test-bucket",
},
),
patch("video_processing.oss_helpers.oss2.Auth"),
patch("video_processing.oss_helpers.oss2.Bucket", return_value=mock_bucket),
patch("video_processing.oss_helpers.OSS_UPLOAD_TOTAL_TIMEOUT", 1), # 1秒超时
):
url = upload_to_oss(file_path, "test/slow.mp4")
# 超时应返回 None
assert url is None, "上传超时应返回 None"
finally:
file_path.unlink()
def test_upload_exception_returns_none(self):
"""上传异常时返回 None."""
from video_processing.oss_helpers import upload_to_oss
small_file = tempfile.NamedTemporaryFile(delete=False, suffix=".mp4")
small_file.write(b"x" * 1024)
small_file.close()
file_path = Path(small_file.name)
mock_bucket = MagicMock()
mock_bucket.put_object_from_file.side_effect = RuntimeError("Network error")
try:
with (
patch.dict(
os.environ,
{
"OSS_ACCESS_KEY_ID": "test-key",
"OSS_ACCESS_KEY_SECRET": "test-secret",
"OSS_ENDPOINT": "oss-cn-hangzhou.aliyuncs.com",
"OSS_BUCKET_NAME": "test-bucket",
},
),
patch("video_processing.oss_helpers.oss2.Auth"),
patch("video_processing.oss_helpers.oss2.Bucket", return_value=mock_bucket),
):
url = upload_to_oss(file_path, "test/error.mp4")
assert url is None, "上传异常应返回 None"
finally:
file_path.unlink()
def test_upload_no_bucket_returns_none(self):
"""OSS 未配置时返回 None."""
from video_processing.oss_helpers import upload_to_oss
small_file = tempfile.NamedTemporaryFile(delete=False, suffix=".mp4")
small_file.write(b"x" * 1024)
small_file.close()
file_path = Path(small_file.name)
try:
with patch.dict(os.environ, {}, clear=True):
url = upload_to_oss(file_path, "test/noconfig.mp4")
assert url is None
finally:
file_path.unlink()
+7 -4
View File
@@ -15,6 +15,7 @@ from video_processing.ffmpeg_utils import build_xfade_filter_chain
# ── P0-3: build_xfade_filter_chain 安全钳制 ──────────────────────────────────
class TestBuildXfadeFilterChainSafetyClamp:
"""验证 xfade 滤镜链的安全钳制逻辑,防止 exit 234。"""
@@ -123,7 +124,9 @@ class TestBuildXfadeFilterChainSafetyClamp:
durations_found.append(float(m.group(1)))
# 第一个 xfade: td 必须 ≤ 0.3 (第二个输入 clip_durations[1]=0.3)
assert durations_found[0] <= 0.3 + 0.001, f"第一个 xfade td={durations_found[0]} 超过 clip_durations[1]=0.3"
assert durations_found[0] <= 0.3 + 0.001, (
f"第一个 xfade td={durations_found[0]} 超过 clip_durations[1]=0.3"
)
# 第二个 xfade: td 可以 = 0.5 (clip_durations[2]=5.0)
assert durations_found[1] <= 0.5 + 0.001
assert dur > 0
@@ -170,9 +173,9 @@ class TestBuildXfadeFilterChainSafetyClamp:
assert dur_val >= 0.001 # 至少 1ms
# P1 修复验证: td 不能超过第二个输入片段时长
second_input_idx = xfade_idx + 1
assert (
dur_val <= durations[second_input_idx] + 0.001
), f"td={dur_val} > clip_durations[{second_input_idx}]={durations[second_input_idx]}"
assert dur_val <= durations[second_input_idx] + 0.001, (
f"td={dur_val} > clip_durations[{second_input_idx}]={durations[second_input_idx]}"
)
xfade_idx += 1
+91 -104
View File
@@ -13,6 +13,7 @@ from unittest.mock import MagicMock, patch
import pytest
# ── oss_bucket endpoint scheme 修复 ──────────────────────────────────────────
@@ -24,19 +25,17 @@ class TestOSSBucketEndpointScheme:
from video_processing.oss_helpers import oss_bucket
mock_bucket_instance = MagicMock()
with (
patch.dict(
os.environ,
{
"OSS_ACCESS_KEY_ID": "test-key",
"OSS_ACCESS_KEY_SECRET": "test-secret",
"OSS_ENDPOINT": "oss-cn-hangzhou.aliyuncs.com",
"OSS_BUCKET_NAME": "test-bucket",
},
),
patch("video_processing.oss_helpers.oss2.Auth") as mock_auth,
patch("video_processing.oss_helpers.oss2.Bucket", return_value=mock_bucket_instance) as mock_bucket_cls,
):
with patch.dict(
os.environ,
{
"OSS_ACCESS_KEY_ID": "test-key",
"OSS_ACCESS_KEY_SECRET": "test-secret",
"OSS_ENDPOINT": "oss-cn-hangzhou.aliyuncs.com",
"OSS_BUCKET_NAME": "test-bucket",
},
), patch("video_processing.oss_helpers.oss2.Auth") as mock_auth, patch(
"video_processing.oss_helpers.oss2.Bucket", return_value=mock_bucket_instance
) as mock_bucket_cls:
# 清除缓存,确保重新创建
import video_processing.oss_helpers as oss_mod
@@ -46,7 +45,9 @@ class TestOSSBucketEndpointScheme:
# 验证 endpoint 传的是带 https:// 的
call_args = mock_bucket_cls.call_args
endpoint_arg = call_args[0][1] # 第 2 个位置参数是 endpoint
assert endpoint_arg.startswith("https://"), f"endpoint 应该带 https:// 前缀,实际为: {endpoint_arg}"
assert endpoint_arg.startswith("https://"), (
f"endpoint 应该带 https:// 前缀,实际为: {endpoint_arg}"
)
assert "oss-cn-hangzhou.aliyuncs.com" in endpoint_arg
def test_endpoint_with_https_keeps_as_is(self):
@@ -54,19 +55,17 @@ class TestOSSBucketEndpointScheme:
from video_processing.oss_helpers import oss_bucket
mock_bucket_instance = MagicMock()
with (
patch.dict(
os.environ,
{
"OSS_ACCESS_KEY_ID": "test-key",
"OSS_ACCESS_KEY_SECRET": "test-secret",
"OSS_ENDPOINT": "https://oss-cn-hangzhou.aliyuncs.com",
"OSS_BUCKET_NAME": "test-bucket",
},
),
patch("video_processing.oss_helpers.oss2.Auth"),
patch("video_processing.oss_helpers.oss2.Bucket", return_value=mock_bucket_instance) as mock_bucket_cls,
):
with patch.dict(
os.environ,
{
"OSS_ACCESS_KEY_ID": "test-key",
"OSS_ACCESS_KEY_SECRET": "test-secret",
"OSS_ENDPOINT": "https://oss-cn-hangzhou.aliyuncs.com",
"OSS_BUCKET_NAME": "test-bucket",
},
), patch("video_processing.oss_helpers.oss2.Auth"), patch(
"video_processing.oss_helpers.oss2.Bucket", return_value=mock_bucket_instance
) as mock_bucket_cls:
import video_processing.oss_helpers as oss_mod
bucket = oss_bucket()
@@ -82,19 +81,17 @@ class TestOSSBucketEndpointScheme:
from video_processing.oss_helpers import oss_bucket
mock_bucket_instance = MagicMock()
with (
patch.dict(
os.environ,
{
"OSS_ACCESS_KEY_ID": "test-key",
"OSS_ACCESS_KEY_SECRET": "test-secret",
"OSS_ENDPOINT": "http://oss-cn-hangzhou.aliyuncs.com",
"OSS_BUCKET_NAME": "test-bucket",
},
),
patch("video_processing.oss_helpers.oss2.Auth"),
patch("video_processing.oss_helpers.oss2.Bucket", return_value=mock_bucket_instance) as mock_bucket_cls,
):
with patch.dict(
os.environ,
{
"OSS_ACCESS_KEY_ID": "test-key",
"OSS_ACCESS_KEY_SECRET": "test-secret",
"OSS_ENDPOINT": "http://oss-cn-hangzhou.aliyuncs.com",
"OSS_BUCKET_NAME": "test-bucket",
},
), patch("video_processing.oss_helpers.oss2.Auth"), patch(
"video_processing.oss_helpers.oss2.Bucket", return_value=mock_bucket_instance
) as mock_bucket_cls:
import video_processing.oss_helpers as oss_mod
bucket = oss_bucket()
@@ -136,18 +133,16 @@ class TestGetSignedDownloadUrl:
mock_bucket = MagicMock()
mock_bucket.sign_url.return_value = "https://test-bucket.oss-cn-hangzhou.aliyuncs.com/generated/test.mp4?OSSAccessKeyId=xxx&Expires=xxx&Signature=xxx"
with (
patch.dict(
os.environ,
{
"OSS_ACCESS_KEY_ID": "test-key",
"OSS_ACCESS_KEY_SECRET": "test-secret",
"OSS_ENDPOINT": "oss-cn-hangzhou.aliyuncs.com",
"OSS_BUCKET_NAME": "test-bucket",
},
),
patch("video_processing.oss_helpers.oss2.Auth"),
patch("video_processing.oss_helpers.oss2.Bucket", return_value=mock_bucket),
with patch.dict(
os.environ,
{
"OSS_ACCESS_KEY_ID": "test-key",
"OSS_ACCESS_KEY_SECRET": "test-secret",
"OSS_ENDPOINT": "oss-cn-hangzhou.aliyuncs.com",
"OSS_BUCKET_NAME": "test-bucket",
},
), patch("video_processing.oss_helpers.oss2.Auth"), patch(
"video_processing.oss_helpers.oss2.Bucket", return_value=mock_bucket
):
result = get_signed_download_url("generated/test.mp4", expires_seconds=3600)
@@ -160,24 +155,22 @@ class TestGetSignedDownloadUrl:
from video_processing.oss_helpers import get_signed_download_url
mock_bucket = MagicMock()
mock_bucket.sign_url.return_value = (
"https://test-bucket.oss-cn-hangzhou.aliyuncs.com/generated/test.mp4?sign=xxx"
)
mock_bucket.sign_url.return_value = "https://test-bucket.oss-cn-hangzhou.aliyuncs.com/generated/test.mp4?sign=xxx"
with (
patch.dict(
os.environ,
{
"OSS_ACCESS_KEY_ID": "test-key",
"OSS_ACCESS_KEY_SECRET": "test-secret",
"OSS_ENDPOINT": "oss-cn-hangzhou.aliyuncs.com",
"OSS_BUCKET_NAME": "test-bucket",
},
),
patch("video_processing.oss_helpers.oss2.Auth"),
patch("video_processing.oss_helpers.oss2.Bucket", return_value=mock_bucket),
with patch.dict(
os.environ,
{
"OSS_ACCESS_KEY_ID": "test-key",
"OSS_ACCESS_KEY_SECRET": "test-secret",
"OSS_ENDPOINT": "oss-cn-hangzhou.aliyuncs.com",
"OSS_BUCKET_NAME": "test-bucket",
},
), patch("video_processing.oss_helpers.oss2.Auth"), patch(
"video_processing.oss_helpers.oss2.Bucket", return_value=mock_bucket
):
result = get_signed_download_url("https://test-bucket.oss-cn-hangzhou.aliyuncs.com/generated/test.mp4")
result = get_signed_download_url(
"https://test-bucket.oss-cn-hangzhou.aliyuncs.com/generated/test.mp4"
)
mock_bucket.sign_url.assert_called_once()
# 验证传给 sign_url 的是纯 storage key,不是完整 URL
@@ -200,18 +193,16 @@ class TestGetSignedDownloadUrl:
mock_bucket = MagicMock()
mock_bucket.sign_url.side_effect = Exception("sign failed")
with (
patch.dict(
os.environ,
{
"OSS_ACCESS_KEY_ID": "test-key",
"OSS_ACCESS_KEY_SECRET": "test-secret",
"OSS_ENDPOINT": "oss-cn-hangzhou.aliyuncs.com",
"OSS_BUCKET_NAME": "test-bucket",
},
),
patch("video_processing.oss_helpers.oss2.Auth"),
patch("video_processing.oss_helpers.oss2.Bucket", return_value=mock_bucket),
with patch.dict(
os.environ,
{
"OSS_ACCESS_KEY_ID": "test-key",
"OSS_ACCESS_KEY_SECRET": "test-secret",
"OSS_ENDPOINT": "oss-cn-hangzhou.aliyuncs.com",
"OSS_BUCKET_NAME": "test-bucket",
},
), patch("video_processing.oss_helpers.oss2.Auth"), patch(
"video_processing.oss_helpers.oss2.Bucket", return_value=mock_bucket
):
result = get_signed_download_url("generated/test.mp4")
assert result is None
@@ -232,18 +223,16 @@ class TestUploadToOSSReturnsHTTPS:
from pathlib import Path
with (
patch.dict(
os.environ,
{
"OSS_ACCESS_KEY_ID": "test-key",
"OSS_ACCESS_KEY_SECRET": "test-secret",
"OSS_ENDPOINT": "oss-cn-hangzhou.aliyuncs.com",
"OSS_BUCKET_NAME": "test-bucket",
},
),
patch("video_processing.oss_helpers.oss2.Auth"),
patch("video_processing.oss_helpers.oss2.Bucket", return_value=mock_bucket),
with patch.dict(
os.environ,
{
"OSS_ACCESS_KEY_ID": "test-key",
"OSS_ACCESS_KEY_SECRET": "test-secret",
"OSS_ENDPOINT": "oss-cn-hangzhou.aliyuncs.com",
"OSS_BUCKET_NAME": "test-bucket",
},
), patch("video_processing.oss_helpers.oss2.Auth"), patch(
"video_processing.oss_helpers.oss2.Bucket", return_value=mock_bucket
):
result = upload_to_oss(Path("/tmp/test.mp4"), "generated/test.mp4")
@@ -260,18 +249,16 @@ class TestUploadToOSSReturnsHTTPS:
from pathlib import Path
with (
patch.dict(
os.environ,
{
"OSS_ACCESS_KEY_ID": "test-key",
"OSS_ACCESS_KEY_SECRET": "test-secret",
"OSS_ENDPOINT": "https://oss-cn-hangzhou.aliyuncs.com",
"OSS_BUCKET_NAME": "test-bucket",
},
),
patch("video_processing.oss_helpers.oss2.Auth"),
patch("video_processing.oss_helpers.oss2.Bucket", return_value=mock_bucket),
with patch.dict(
os.environ,
{
"OSS_ACCESS_KEY_ID": "test-key",
"OSS_ACCESS_KEY_SECRET": "test-secret",
"OSS_ENDPOINT": "https://oss-cn-hangzhou.aliyuncs.com",
"OSS_BUCKET_NAME": "test-bucket",
},
), patch("video_processing.oss_helpers.oss2.Auth"), patch(
"video_processing.oss_helpers.oss2.Bucket", return_value=mock_bucket
):
result = upload_to_oss(Path("/tmp/test.mp4"), "generated/test.mp4")
+15 -15
View File
@@ -64,7 +64,7 @@ class TestPresetVoice:
def test_preset_voice_to_dict(self) -> None:
"""序列化。"""
voice = PresetVoice(
voice_id="longxiaochun_v3",
voice_id="longxiaochun",
name="龙小淳",
description="温柔女声",
gender="female",
@@ -73,7 +73,7 @@ class TestPresetVoice:
result = voice.to_dict()
assert result["voice_id"] == "longxiaochun_v3"
assert result["voice_id"] == "longxiaochun"
assert result["name"] == "龙小淳"
assert result["description"] == "温柔女声"
assert result["gender"] == "female"
@@ -127,14 +127,14 @@ class TestPresetVoicesConfig:
def test_cosyvoice_voice_ids(self) -> None:
"""音色 ID 应为 CosyVoice 真实可用的音色名。"""
expected_ids = {
"longxiaochun_v3",
"longxiaoxia_v3",
"longxiaochen_v3",
"longyue_v3",
"longshu_v3",
"longjing_v3",
"longbo_v3",
"longtian_v3",
"longxiaochun",
"longxiaoxia",
"longxiaochen",
"longyue",
"longshu",
"longjing",
"longbo",
"longtian",
}
actual_ids = {v.voice_id for v in PRESET_VOICES}
assert actual_ids == expected_ids
@@ -164,10 +164,10 @@ class TestPresetVoiceHelpers:
def test_get_preset_voice_by_id_found(self) -> None:
"""按 ID 查找存在的音色。"""
voice = get_preset_voice_by_id("longxiaochun_v3")
voice = get_preset_voice_by_id("longxiaochun")
assert voice is not None
assert voice.name == "龙小淳"
assert voice.voice_id == "longxiaochun_v3"
assert voice.voice_id == "longxiaochun"
def test_get_preset_voice_by_id_not_found(self) -> None:
"""按 ID 查找不存在的音色。"""
@@ -176,9 +176,9 @@ class TestPresetVoiceHelpers:
def test_is_preset_voice_true(self) -> None:
"""判断预置音色返回 True。"""
assert is_preset_voice("longxiaochun_v3") is True
assert is_preset_voice("longxiaoxia_v3") is True
assert is_preset_voice("longbo_v3") is True
assert is_preset_voice("longxiaochun") is True
assert is_preset_voice("longxiaoxia") is True
assert is_preset_voice("longbo") is True
def test_is_preset_voice_false(self) -> None:
"""判断非预置音色返回 False。"""
-424
View File
@@ -1,424 +0,0 @@
"""RenderAdapter 单元测试 — Phase 2.
测试适配层的计划加载素材下载引擎调用结果上传等逻辑
"""
from __future__ import annotations
from dataclasses import dataclass, field
from pathlib import Path
from typing import Any
from unittest.mock import MagicMock, patch
import pytest
from video_processing.render_adapter import RenderAdapter, RenderAdapterResult
# ── Fixtures ──────────────────────────────────────────────────────────────────
@dataclass
class FakeClip:
"""模拟 EditPlanClip。"""
id: str
plan_id: str = "plan_001"
clip_type: str = "main"
order: int = 0
asset_id: str = ""
text_content: str = ""
start_time: float = 0.0
duration: float = 0.0
transition_effect: str = "cut"
status: str = "ready"
config: dict[str, Any] = field(default_factory=dict)
@dataclass
class FakePlan:
"""模拟 EditPlan。"""
id: str = "plan_001"
name: str = "测试计划"
status: str = "editing"
config: dict[str, Any] = field(default_factory=dict)
def _make_clip(
clip_id: str,
clip_type: str = "main",
order: int = 0,
asset_id: str | None = None,
duration: float = 5.0,
status: str = "ready",
transition_effect: str = "cut",
config: dict[str, Any] | None = None,
) -> FakeClip:
# asset_id 为 None 时生成默认值,为空字符串时保留空串
if asset_id is None:
asset_id = f"asset_{clip_id}.mp4"
return FakeClip(
id=clip_id,
clip_type=clip_type,
order=order,
asset_id=asset_id,
duration=duration,
status=status,
transition_effect=transition_effect,
config=config or {},
)
def _make_adapter(
plan: FakePlan | None = None,
clips: list[FakeClip] | None = None,
) -> tuple[RenderAdapter, MagicMock, MagicMock]:
"""创建测试用的 RenderAdapter 及 mock repo。
Returns:
(adapter, mock_plan_repo, mock_clip_repo)
"""
mock_db = MagicMock()
adapter = RenderAdapter(mock_db)
# 替换内部 repo
mock_plan_repo = MagicMock()
mock_clip_repo = MagicMock()
adapter._plan_repo = mock_plan_repo
adapter._clip_repo = mock_clip_repo
# 设置默认返回
if plan is not None:
mock_plan_repo.get.return_value = plan
if clips is not None:
mock_clip_repo.list_by_plan.return_value = clips
return adapter, mock_plan_repo, mock_clip_repo
# ── validate_plan 测试 ───────────────────────────────────────────────────────
class TestValidatePlan:
def test_plan_not_found(self):
"""计划不存在时校验失败。"""
adapter, mock_plan_repo, _ = _make_adapter(plan=None)
mock_plan_repo.get.return_value = None
valid, errors, warnings, ready_count, total_count = adapter.validate_plan("plan_001")
assert not valid
assert len(errors) == 1
assert "不存在" in errors[0]
assert ready_count == 0
assert total_count == 0
def test_no_clips(self):
"""没有任何片段时校验失败。"""
plan = FakePlan(id="plan_001", status="editing")
adapter, _, mock_clip_repo = _make_adapter(plan=plan, clips=[])
valid, errors, warnings, ready_count, total_count = adapter.validate_plan("plan_001")
assert not valid
assert any("没有任何片段" in e for e in errors)
def test_no_ready_clips(self):
"""没有 ready 片段时校验失败。"""
plan = FakePlan(id="plan_001", status="editing")
clips = [
_make_clip("c1", status="pending"),
_make_clip("c2", status="pending"),
]
adapter, _, _ = _make_adapter(plan=plan, clips=clips)
valid, errors, warnings, ready_count, total_count = adapter.validate_plan("plan_001")
assert not valid
assert any("没有就绪" in e for e in errors)
assert ready_count == 0
assert total_count == 2
def test_ready_clip_no_asset(self):
"""ready 片段没有 asset_id 时报错。"""
plan = FakePlan(id="plan_001", status="editing")
clips = [
_make_clip("c1", asset_id=""),
]
adapter, _, _ = _make_adapter(plan=plan, clips=clips)
valid, errors, warnings, ready_count, total_count = adapter.validate_plan("plan_001")
assert not valid
assert any("没有分配素材" in e for e in errors)
def test_valid_plan(self):
"""正常计划校验通过。"""
plan = FakePlan(id="plan_001", status="editing")
clips = [
_make_clip("c1", order=0, duration=3.0),
_make_clip("c2", order=1, duration=4.0),
]
adapter, _, _ = _make_adapter(plan=plan, clips=clips)
valid, errors, warnings, ready_count, total_count = adapter.validate_plan("plan_001")
assert valid
assert len(errors) == 0
assert ready_count == 2
assert total_count == 2
def test_wrong_status(self):
"""计划状态不正确时报错。"""
plan = FakePlan(id="plan_001", status="draft")
clips = [_make_clip("c1")]
adapter, _, _ = _make_adapter(plan=plan, clips=clips)
valid, errors, _, _, _ = adapter.validate_plan("plan_001")
assert not valid
assert any("状态不正确" in e for e in errors)
def test_mixed_status_with_warnings(self):
"""混合状态时有 pending/failed 警告。"""
plan = FakePlan(id="plan_001", status="editing")
clips = [
_make_clip("c1", order=0, status="ready"),
_make_clip("c2", order=1, status="pending"),
_make_clip("c3", order=2, status="failed"),
]
adapter, _, _ = _make_adapter(plan=plan, clips=clips)
valid, errors, warnings, ready_count, total_count = adapter.validate_plan("plan_001")
assert valid
assert any("pending" in w for w in warnings)
assert any("failed" in w for w in warnings)
assert ready_count == 1
assert total_count == 3
# ── render_plan 测试 ─────────────────────────────────────────────────────────
class TestRenderPlan:
def test_plan_not_found(self):
"""计划不存在时返回失败。"""
adapter, mock_plan_repo, _ = _make_adapter(plan=None)
mock_plan_repo.get.return_value = None
result = adapter.render_plan("plan_001")
assert not result.success
assert "不存在" in result.error_message
def test_no_ready_clips(self):
"""没有 ready 片段时返回失败。"""
plan = FakePlan(id="plan_001", status="editing")
clips = [_make_clip("c1", status="pending")]
adapter, _, _ = _make_adapter(plan=plan, clips=clips)
result = adapter.render_plan("plan_001")
assert not result.success
assert "没有可渲染" in result.error_message
assert result.clip_count == 0
@patch("video_processing.render_adapter.download_asset")
def test_all_assets_download_fail(self, mock_download):
"""所有素材下载失败时返回失败。"""
mock_download.return_value = False
plan = FakePlan(id="plan_001", status="editing")
clips = [_make_clip("c1", order=0, duration=5.0)]
adapter, _, _ = _make_adapter(plan=plan, clips=clips)
result = adapter.render_plan("plan_001")
assert not result.success
assert "素材下载失败" in result.error_message
@patch("video_processing.render_adapter.upload_to_oss")
@patch("video_processing.render_adapter.UnifiedRenderService")
@patch("video_processing.render_adapter.download_asset")
def test_successful_render(self, mock_download, mock_render_cls, mock_upload, tmp_path):
"""完整渲染流程成功。"""
# 素材下载成功
def _fake_download(asset_id, local_path):
local_path.parent.mkdir(parents=True, exist_ok=True)
local_path.write_bytes(b"fake video data")
return True
mock_download.side_effect = _fake_download
# 渲染成功
mock_render = MagicMock()
mock_render.render.return_value = MagicMock(
output_path=tmp_path / "output.mp4",
duration=10.0,
file_size=102400,
width=1280,
height=720,
)
mock_render_cls.return_value = mock_render
# 上传成功
mock_upload.return_value = "https://oss.example.com/rendered/plan_001/job_001.mp4"
plan = FakePlan(id="plan_001", status="editing")
clips = [
_make_clip("c1", order=0, duration=5.0),
_make_clip("c2", order=1, duration=5.0),
]
adapter, _, _ = _make_adapter(plan=plan, clips=clips)
result = adapter.render_plan(
"plan_001",
job_id="job_001",
work_dir=tmp_path / "work",
)
assert result.success
assert result.output_url.startswith("https://")
assert result.duration == 10.0
assert result.width == 1280
assert result.height == 720
assert result.clip_count == 2
# 验证 UnifiedRenderService 被正确调用
mock_render_cls.assert_called_once()
call_kwargs = mock_render_cls.call_args
assert call_kwargs.kwargs["plan"] is plan
assert len(call_kwargs.kwargs["clips"]) == 2
assert len(call_kwargs.kwargs["asset_path_map"]) == 2
@patch("video_processing.render_adapter.download_asset")
def test_progress_callback(self, mock_download, tmp_path):
"""进度回调被正确触发。"""
def _fake_download(asset_id, local_path):
local_path.parent.mkdir(parents=True, exist_ok=True)
local_path.write_bytes(b"fake data")
return True
mock_download.side_effect = _fake_download
# 模拟渲染异常,避免走到最后
with patch("video_processing.render_adapter.UnifiedRenderService") as mock_render_cls:
mock_render = MagicMock()
mock_render.render.side_effect = RuntimeError("render error")
mock_render_cls.return_value = mock_render
plan = FakePlan(id="plan_001", status="editing")
clips = [_make_clip("c1", order=0, duration=5.0)]
adapter, _, _ = _make_adapter(plan=plan, clips=clips)
progress_values = []
def progress_cb(progress: float, stage: str) -> None:
progress_values.append((progress, stage))
result = adapter.render_plan(
"plan_001",
work_dir=tmp_path / "work",
progress_cb=progress_cb,
)
# 即使渲染失败,前期进度也应该上报了
assert len(progress_values) > 0
# 第一个进度应该是加载计划
assert progress_values[0][1] == "加载剪辑计划"
@patch("video_processing.render_adapter.download_asset")
def test_partial_asset_download(self, mock_download, tmp_path):
"""部分素材下载失败时,只使用成功的素材。"""
download_results = [True, False, True] # 3个素材中2个成功
def _fake_download(asset_id, local_path):
idx = hash(asset_id) % 3
if download_results[idx]:
local_path.parent.mkdir(parents=True, exist_ok=True)
local_path.write_bytes(b"fake data")
return True
return False
mock_download.side_effect = _fake_download
with patch("video_processing.render_adapter.UnifiedRenderService") as mock_render_cls:
mock_render = MagicMock()
mock_render.render.return_value = MagicMock(
output_path=tmp_path / "out.mp4",
duration=5.0,
file_size=1024,
width=1280,
height=720,
)
mock_render_cls.return_value = mock_render
with patch("video_processing.render_adapter.upload_to_oss", return_value="https://example.com/out.mp4"):
plan = FakePlan(id="plan_001", status="editing")
clips = [
_make_clip("c1", order=0, duration=3.0, asset_id="asset_001.mp4"),
_make_clip("c2", order=1, duration=3.0, asset_id="asset_002.mp4"),
_make_clip("c3", order=2, duration=3.0, asset_id="asset_003.mp4"),
]
adapter, _, _ = _make_adapter(plan=plan, clips=clips)
result = adapter.render_plan(
"plan_001",
work_dir=tmp_path / "work",
)
# 至少有部分素材成功,渲染应该进行
# (具体成功数量取决于 hash 结果,但至少1个成功就能渲染)
assert result.success or "素材下载失败" in result.error_message
# ── _download_assets 测试 ────────────────────────────────────────────────────
class TestDownloadAssets:
@patch("video_processing.render_adapter.download_asset")
def test_all_download_success(self, mock_download, tmp_path):
"""全部素材下载成功。"""
mock_download.return_value = True
clips = [
_make_clip("c1", order=0, asset_id="key1.mp4"),
_make_clip("c2", order=1, asset_id="key2.mp4"),
]
result = RenderAdapter._download_assets(clips, tmp_path)
assert len(result) == 2
assert "key1.mp4" in result
assert "key2.mp4" in result
assert mock_download.call_count == 2
@patch("video_processing.render_adapter.download_asset")
def test_empty_asset_id_skipped(self, mock_download, tmp_path):
"""空 asset_id 的片段被跳过。"""
clips = [
_make_clip("c1", order=0, asset_id=""),
_make_clip("c2", order=1, asset_id="key2.mp4"),
]
mock_download.return_value = True
result = RenderAdapter._download_assets(clips, tmp_path)
assert len(result) == 1
assert "key2.mp4" in result
assert mock_download.call_count == 1 # 只调用了一次下载
@patch("video_processing.render_adapter.download_asset")
def test_all_download_fail(self, mock_download, tmp_path):
"""全部下载失败返回空字典。"""
mock_download.return_value = False
clips = [
_make_clip("c1", order=0, asset_id="key1.mp4"),
]
result = RenderAdapter._download_assets(clips, tmp_path)
assert len(result) == 0
-351
View File
@@ -1,351 +0,0 @@
"""任务队列限流防护单元测试。"""
from __future__ import annotations
import os
import sys
from unittest.mock import MagicMock
import pytest
sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..", "..", "apps", "api"))
from app.core.task_enqueue import (
GLOBAL_PENDING_LIMIT,
USER_PENDING_LIMIT,
GlobalQueueFull,
UserPendingLimitExceeded,
check_queue_limits,
safe_enqueue_generation_task,
)
# ---------------------------------------------------------------------------
# Mock helpers
# ---------------------------------------------------------------------------
class MockRepository:
"""支持 pending 计数的 mock repository。
支持通过 set_pending 动态修改计数用于模拟入队后计数变化的并发场景
"""
def __init__(self, user_pending: int = 0, global_pending: int = 0):
self._user_pending = user_pending
self._global_pending = global_pending
self._send_task_called = False
self.updated_tasks = []
def count_pending_by_user(self, user_id: str) -> int:
return self._user_pending
def count_pending_total(self) -> int:
return self._global_pending
def update(self, task):
self.updated_tasks.append(task)
return task
def set_pending(self, *, user_pending: int | None = None, global_pending: int | None = None):
"""动态修改 pending 计数,模拟并发场景。"""
if user_pending is not None:
self._user_pending = user_pending
if global_pending is not None:
self._global_pending = global_pending
class MockTask:
def __init__(self, task_id: str = "task-1", status: str = "pending"):
self.id = task_id
self.status = status
self.error_message = ""
def mark_failed(self, reason: str):
self.status = "failed"
self.error_message = reason
@pytest.fixture(autouse=True)
def mock_celery(monkeypatch):
"""mock 掉 celery_app.send_task,避免真实发送。"""
mock_send = MagicMock()
monkeypatch.setattr("app.core.celery_app.celery_app.send_task", mock_send)
return mock_send
# ---------------------------------------------------------------------------
# 常量导出测试
# ---------------------------------------------------------------------------
def test_limit_constants_are_exported():
"""限流阈值常量已导出,供业务代码引用。"""
assert USER_PENDING_LIMIT == 3
assert GLOBAL_PENDING_LIMIT == 20
# ---------------------------------------------------------------------------
# check_queue_limits 单元测试(预检查用,>= 边界)
# ---------------------------------------------------------------------------
class TestCheckQueueLimits:
"""队列限流检查函数测试(预检查语义,>= 上限即拒绝)。"""
def test_normal_passes_through(self):
"""正常范围内的任务不受限制。"""
repo = MockRepository(user_pending=1, global_pending=5)
check_queue_limits("user-1", repo)
def test_user_limit_exceeded_raises(self):
"""用户 pending 超过上限抛 UserPendingLimitExceeded。"""
repo = MockRepository(user_pending=4, global_pending=5)
with pytest.raises(UserPendingLimitExceeded) as exc_info:
check_queue_limits("user-1", repo)
assert exc_info.value.user_id == "user-1"
assert exc_info.value.pending_count == 4
assert exc_info.value.limit == 3
def test_user_at_limit_also_raises(self):
"""用户 pending 刚好等于上限也拒绝(>= 边界)。"""
repo = MockRepository(user_pending=3, global_pending=5)
with pytest.raises(UserPendingLimitExceeded):
check_queue_limits("user-1", repo)
def test_user_below_limit_passes(self):
"""用户 pending 比上限少 1,通过。"""
repo = MockRepository(user_pending=2, global_pending=5)
check_queue_limits("user-1", repo)
def test_global_limit_exceeded_raises(self):
"""全局 pending 超过上限抛 GlobalQueueFull。"""
repo = MockRepository(user_pending=1, global_pending=21)
with pytest.raises(GlobalQueueFull) as exc_info:
check_queue_limits("user-1", repo)
assert exc_info.value.pending_count == 21
assert exc_info.value.limit == 20
def test_global_at_limit_also_raises(self):
"""全局 pending 刚好等于上限也拒绝(>= 边界)。"""
repo = MockRepository(user_pending=1, global_pending=20)
with pytest.raises(GlobalQueueFull):
check_queue_limits("user-1", repo)
def test_global_below_limit_passes(self):
"""全局 pending 比上限少 1,通过。"""
repo = MockRepository(user_pending=1, global_pending=19)
check_queue_limits("user-1", repo)
def test_global_takes_priority_over_user(self):
"""全局和用户都超限时,优先抛全局异常。"""
repo = MockRepository(user_pending=5, global_pending=25)
with pytest.raises(GlobalQueueFull):
check_queue_limits("user-1", repo)
def test_empty_user_id_skips_user_check(self):
"""不传 user_id 时跳过用户级检查,只做全局检查。"""
repo = MockRepository(user_pending=10, global_pending=5)
# 用户超限但不传 user_id → 全局未超限,应该通过
check_queue_limits("", repo)
# ---------------------------------------------------------------------------
# safe_enqueue_generation_task 限流集成测试(入队前用 >,包含当前任务)
# ---------------------------------------------------------------------------
class TestSafeEnqueueWithLimits:
"""安全入队函数的限流功能测试。"""
def test_normal_task_enqueues_successfully(self, mock_celery):
"""正常任务入队成功,返回 True。"""
repo = MockRepository(user_pending=0, global_pending=0)
task = MockTask("task-1")
result = safe_enqueue_generation_task(task, repo, user_id="user-1")
assert result is True
mock_celery.assert_called_once_with("worker.generate_video", args=["task-1"])
assert len(repo.updated_tasks) == 0 # 成功不需要更新状态
def test_user_limit_rejected_with_failed_status(self, mock_celery):
"""用户超限:任务标记为 failed,抛 UserPendingLimitExceeded。"""
repo = MockRepository(user_pending=5, global_pending=5)
task = MockTask("task-1")
with pytest.raises(UserPendingLimitExceeded):
safe_enqueue_generation_task(task, repo, user_id="user-1")
mock_celery.assert_not_called()
assert task.status == "failed"
assert "限流" in task.error_message
assert len(repo.updated_tasks) == 1
def test_user_at_limit_still_passes(self, mock_celery):
"""用户 pending 刚好等于上限:入队前检查用 >,包含当前任务,刚好到上限不算超。
与预检查的 >= 语义一致预检查时 pending=3 拒绝不能再加新的
safe_enqueue 被调用时任务已是 pending就是第3个
pending=3 不满足 >3所以通过
"""
repo = MockRepository(user_pending=3, global_pending=5)
task = MockTask("task-1")
result = safe_enqueue_generation_task(task, repo, user_id="user-1")
assert result is True
mock_celery.assert_called_once()
def test_user_one_over_limit_rejected(self, mock_celery):
"""用户 pending = limit + 1:超限被拒。"""
repo = MockRepository(user_pending=4, global_pending=5)
task = MockTask("task-1")
with pytest.raises(UserPendingLimitExceeded):
safe_enqueue_generation_task(task, repo, user_id="user-1")
mock_celery.assert_not_called()
def test_global_limit_rejected_with_failed_status(self, mock_celery):
"""全局超限:任务标记为 failed,抛 GlobalQueueFull。"""
repo = MockRepository(user_pending=1, global_pending=21)
task = MockTask("task-1")
with pytest.raises(GlobalQueueFull):
safe_enqueue_generation_task(task, repo, user_id="user-1")
mock_celery.assert_not_called()
assert task.status == "failed"
assert len(repo.updated_tasks) == 1
def test_global_at_limit_still_passes(self, mock_celery):
"""全局 pending 刚好等于上限:入队前检查用 >,包含当前任务,刚好到上限不算超。"""
repo = MockRepository(user_pending=1, global_pending=20)
task = MockTask("task-1")
result = safe_enqueue_generation_task(task, repo, user_id="user-1")
assert result is True
mock_celery.assert_called_once()
def test_no_user_id_skips_user_limit(self, mock_celery):
"""不传 user_id 时跳过用户级限流,只做全局检查。"""
repo = MockRepository(user_pending=10, global_pending=5)
task = MockTask("task-1")
result = safe_enqueue_generation_task(task, repo, user_id="")
assert result is True
mock_celery.assert_called_once()
def test_no_user_id_still_checks_global(self, mock_celery):
"""不传 user_id 时全局超限仍然被拦。"""
repo = MockRepository(user_pending=10, global_pending=25)
task = MockTask("task-1")
with pytest.raises(GlobalQueueFull):
safe_enqueue_generation_task(task, repo, user_id="")
mock_celery.assert_not_called()
def test_default_limits_match_constants(self, mock_celery):
"""默认配置与导出常量一致。"""
# 刚好在默认限制内(limit - 1)
repo = MockRepository(user_pending=2, global_pending=19)
task = MockTask("task-1")
result = safe_enqueue_generation_task(task, repo, user_id="user-1")
assert result is True
def test_update_failure_does_not_crash(self, mock_celery):
"""repository.update 失败也不崩溃,异常继续向上抛。"""
class BadRepo(MockRepository):
def update(self, task):
raise RuntimeError("db down")
repo = BadRepo(user_pending=5, global_pending=5)
task = MockTask("task-1")
# 仍然抛 UserPendingLimitExceeded,不会被 update 失败掩盖
with pytest.raises(UserPendingLimitExceeded):
safe_enqueue_generation_task(task, repo, user_id="user-1")
mock_celery.assert_not_called()
# 任务状态还是变了(内存里改了)
assert task.status == "failed"
# ---------------------------------------------------------------------------
# 入队后最终校验(并发竞态兜底)测试
# ---------------------------------------------------------------------------
class TestPostEnqueueFinalCheck:
"""入队后最终校验:模拟并发场景,Celery发送后计数增加被兜住。"""
def test_post_enqueue_global_overflow_rollback(self, mock_celery):
"""并发场景:入队前检查通过,但发送Celery后全局计数超限 → 回滚为failed。
模拟两个请求同时通过入队前检查都查到 global=19
都创建了任务DB里变成 21先发送Celery的那个在最终校验时被兜住
"""
repo = MockRepository(user_pending=1, global_pending=20) # 入队前:20 > 20?否
task = MockTask("task-1")
# 模拟发送Celery后,另一个并发请求也创建了任务,全局变成21
def side_effect(*args, **kwargs):
repo.set_pending(global_pending=21)
mock_celery.side_effect = side_effect
with pytest.raises(GlobalQueueFull) as exc_info:
safe_enqueue_generation_task(task, repo, user_id="user-1")
# Celery 确实发出去了(兜底不撤销 Celery,只回滚 DB 状态)
mock_celery.assert_called_once()
# 任务被标记为 failed
assert task.status == "failed"
assert "入队后" in task.error_message
assert exc_info.value.pending_count == 21
assert len(repo.updated_tasks) == 1
def test_post_enqueue_user_overflow_rollback(self, mock_celery):
"""并发场景:入队前检查通过,但发送Celery后用户计数超限 → 回滚为failed。"""
repo = MockRepository(user_pending=3, global_pending=5) # 入队前:3 > 3?否
task = MockTask("task-1")
def side_effect(*args, **kwargs):
repo.set_pending(user_pending=4)
mock_celery.side_effect = side_effect
with pytest.raises(UserPendingLimitExceeded) as exc_info:
safe_enqueue_generation_task(task, repo, user_id="user-1")
mock_celery.assert_called_once()
assert task.status == "failed"
assert "入队后" in task.error_message
assert exc_info.value.user_id == "user-1"
assert exc_info.value.pending_count == 4
def test_post_enqueue_global_priority_over_user(self, mock_celery):
"""入队后校验:全局和用户都超限时,优先抛全局异常。"""
repo = MockRepository(user_pending=3, global_pending=20)
task = MockTask("task-1")
def side_effect(*args, **kwargs):
repo.set_pending(user_pending=5, global_pending=22)
mock_celery.side_effect = side_effect
with pytest.raises(GlobalQueueFull):
safe_enqueue_generation_task(task, repo, user_id="user-1")
assert task.status == "failed"
def test_post_enqueue_no_change_still_passes(self, mock_celery):
"""入队后计数没变 → 正常通过,不回滚。"""
repo = MockRepository(user_pending=2, global_pending=10)
task = MockTask("task-1")
result = safe_enqueue_generation_task(task, repo, user_id="user-1")
assert result is True
mock_celery.assert_called_once()
assert task.status == "pending" # 状态没变
assert len(repo.updated_tasks) == 0 # 没更新 DB
def test_post_enqueue_no_user_id_skips_user_check(self, mock_celery):
"""不传 user_id 时,入队后校验也跳过用户级,只查全局。"""
repo = MockRepository(user_pending=10, global_pending=5)
task = MockTask("task-1")
def side_effect(*args, **kwargs):
repo.set_pending(user_pending=15, global_pending=5) # 用户超限但全局没超
mock_celery.side_effect = side_effect
result = safe_enqueue_generation_task(task, repo, user_id="")
assert result is True # 用户级不检查,全局没超限 → 通过
+4 -4
View File
@@ -15,7 +15,7 @@ class TestTTSJobCreate:
job = TTSJob.create(
user_id="user_001",
input_text="这是一段测试文本",
voice_id="longxiaochun_v3",
voice_id="longxiaochun",
voice_model="cosyvoice-v1",
project_id="project_001",
voice_clone_profile_id="profile_001",
@@ -26,7 +26,7 @@ class TestTTSJobCreate:
assert job.id
assert job.user_id == "user_001"
assert job.input_text == "这是一段测试文本"
assert job.voice_id == "longxiaochun_v3"
assert job.voice_id == "longxiaochun"
assert job.voice_model == "cosyvoice-v1"
assert job.project_id == "project_001"
assert job.voice_clone_profile_id == "profile_001"
@@ -291,7 +291,7 @@ class TestTTSJobToDict:
job = TTSJob.create(
user_id="user_001",
input_text="测试文本",
voice_id="longxiaochun_v3",
voice_id="longxiaochun",
voice_model="cosyvoice-v1",
project_id="project_001",
voice_clone_profile_id="profile_001",
@@ -306,7 +306,7 @@ class TestTTSJobToDict:
assert result["id"] == job.id
assert result["user_id"] == "user_001"
assert result["input_text"] == "测试文本"
assert result["voice_id"] == "longxiaochun_v3"
assert result["voice_id"] == "longxiaochun"
assert result["voice_model"] == "cosyvoice-v1"
assert result["project_id"] == "project_001"
assert result["voice_clone_profile_id"] == "profile_001"
View File
File diff suppressed because it is too large Load Diff
+247
View File
@@ -0,0 +1,247 @@
"""
视频合成服务安全校验单元测试
针对 PR #159 安全审计发现的问题进行测试
"""
import os
import tempfile
from unittest.mock import MagicMock, patch
import pytest
# 导入被测试的模块
from apps.worker.video_processing.video_compose_service import (
ALLOWED_INPUT_PREFIXES,
ALLOWED_OUTPUT_DIRS,
ALLOWED_TRANSITIONS,
Clip,
EditingMode,
EditingModeConfig,
VideoComposeService,
)
class TestOutputPathValidation:
"""P0: 输出路径穿越校验测试"""
def setup_method(self):
"""测试前设置"""
self.config = EditingModeConfig(mode=EditingMode.ONE_TAKE)
self.service = VideoComposeService(self.config)
def test_valid_output_path_in_allowed_dir(self):
"""测试合法的输出路径"""
valid_path = "/tmp/video_output/test.mp4"
result = self.service._validate_output_path(valid_path)
assert result == os.path.abspath(valid_path)
def test_valid_output_path_with_relative_components(self):
"""测试带相对路径成分但最终在允许目录内的路径"""
valid_path = "/tmp/video_output/subdir/../test.mp4"
result = self.service._validate_output_path(valid_path)
assert result == os.path.abspath(valid_path)
def test_path_traversal_attack_blocked(self):
"""测试路径穿越攻击被阻止"""
# 尝试穿越到 /etc/passwd
malicious_path = "/tmp/video_output/../../../etc/passwd"
with pytest.raises(ValueError, match="输出路径不在允许范围内"):
self.service._validate_output_path(malicious_path)
def test_path_traversal_attack_blocked_var_app(self):
"""测试针对 /var/app 的路径穿越攻击被阻止"""
malicious_path = "/var/app/rendered/../../config/../../../etc/passwd"
with pytest.raises(ValueError, match="输出路径不在允许范围内"):
self.service._validate_output_path(malicious_path)
def test_absolute_path_to_forbidden_location(self):
"""测试直接访问禁止位置"""
forbidden_path = "/etc/shadow"
with pytest.raises(ValueError, match="输出路径不在允许范围内"):
self.service._validate_output_path(forbidden_path)
def test_root_path_blocked(self):
"""测试根目录被阻止"""
root_path = "/"
with pytest.raises(ValueError, match="输出路径不在允许范围内"):
self.service._validate_output_path(root_path)
def test_absolute_path_to_tmp_not_allowed(self):
"""测试 /tmp 不在白名单中时应被阻止"""
# /tmp 不在 ALLOWED_OUTPUT_DIRS 中
tmp_path = "/tmp/test.mp4"
with pytest.raises(ValueError, match="输出路径不在允许范围内"):
self.service._validate_output_path(tmp_path)
class TestInputPathValidation:
"""P1-1: 输入路径格式校验测试"""
def setup_method(self):
"""测试前设置"""
self.config = EditingModeConfig(mode=EditingMode.ONE_TAKE)
self.service = VideoComposeService(self.config)
def test_valid_s3_path(self):
"""测试 S3 路径"""
assert self.service._validate_input_path("s3://bucket/key.mp4") is True
def test_valid_oss_path(self):
"""测试 OSS 路径"""
assert self.service._validate_input_path("oss://bucket/key.mp4") is True
def test_valid_local_path(self):
"""测试 local:// 路径"""
assert self.service._validate_input_path("local://asset/123.mp4") is True
def test_valid_var_storage_path(self):
"""测试 /var/storage/ 路径"""
assert self.service._validate_input_path("/var/storage/assets/123.mp4") is True
def test_path_traversal_in_input_rejected(self):
"""测试输入路径中的路径穿越尝试被拒绝"""
malicious_path = "s3://bucket/../../etc/passwd"
# 这会通过前缀检查,但实际使用时文件系统访问会失败
# 安全设计:只校验格式前缀
assert self.service._validate_input_path(malicious_path) is True
def test_malicious_input_path_blocked(self):
"""测试恶意输入路径被阻止"""
assert self.service._validate_input_path("/etc/passwd") is False
assert self.service._validate_input_path("file:///etc/passwd") is False
assert self.service._validate_input_path("http://evil.com/shell.sh") is False
def test_empty_path_rejected(self):
"""测试空路径被拒绝"""
assert self.service._validate_input_path("") is False
def test_random_string_rejected(self):
"""测试随机字符串被拒绝"""
assert self.service._validate_input_path("random123") is False
assert self.service._validate_input_path("abc../../../etc") is False
class TestTransitionValidation:
"""P1-2: 转场参数白名单校验测试"""
def setup_method(self):
"""测试前设置"""
self.config = EditingModeConfig(mode=EditingMode.ONE_TAKE)
self.service = VideoComposeService(self.config)
@pytest.mark.parametrize("transition", list(ALLOWED_TRANSITIONS))
def test_valid_transitions(self, transition):
"""测试所有合法的转场效果"""
result = self.service._validate_transition(transition)
assert result == transition
def test_invalid_transition_defaults_to_fade(self):
"""测试非法转场效果默认为 fade"""
result = self.service._validate_transition("random_transition")
assert result == "fade"
def test_sql_injection_in_transition_blocked(self):
"""测试 SQL 注入尝试被阻止"""
result = self.service._validate_transition("fade; DROP TABLE videos;--")
assert result == "fade"
def test_shell_injection_in_transition_blocked(self):
"""测试 Shell 注入尝试被阻止"""
result = self.service._validate_transition("fade$(whoami)")
assert result == "fade"
def test_empty_transition_handled(self):
"""测试空转场名称"""
result = self.service._validate_transition("")
assert result == "fade"
def test_none_transition_handled(self):
"""测试 None 转场名称"""
result = self.service._validate_transition(None)
assert result == "fade"
def test_get_validated_transition_returns_mapped(self):
"""测试 _get_validated_transition 返回映射后的值"""
# "fade" 应该映射为 "fade"
result = self.service._get_validated_transition("fade")
assert result == "fade"
class TestComposeSecurityIntegration:
"""安全集成测试"""
def setup_method(self):
"""测试前设置"""
self.config = EditingModeConfig(mode=EditingMode.ONE_TAKE)
self.service = VideoComposeService(self.config)
def test_compose_rejects_malicious_output_path(self):
"""测试 compose 方法拒绝恶意输出路径"""
clips = [
Clip(asset_id="s3://bucket/video1.mp4"),
Clip(asset_id="s3://bucket/video2.mp4"),
]
with pytest.raises(ValueError, match="输出路径不在允许范围内"):
self.service.compose(clips, output_path="/etc/passwd")
def test_compose_rejects_invalid_input_path(self):
"""测试 compose 方法拒绝非法输入路径"""
clips = [
Clip(asset_id="/etc/shadow"), # 非法路径
]
with pytest.raises(ValueError, match="不合法的输入路径"):
self.service.compose(clips)
def test_compose_with_valid_paths(self):
"""测试合法路径可以正常处理"""
with tempfile.TemporaryDirectory() as tmpdir:
# 创建临时视频文件
video_path = os.path.join(tmpdir, "input.mp4")
output_path = os.path.join("/tmp/video_output", "output.mp4")
# 创建空的测试文件(实际测试需要真实视频)
with open(video_path, "wb") as f:
f.write(b"fake video data")
clips = [
Clip(asset_id=f"local://{video_path}"),
]
# 验证输入校验通过
assert self.service._validate_input_path(f"local://{video_path}") is True
def test_compose_empty_clips_rejected(self):
"""测试空片段列表被拒绝"""
with pytest.raises(ValueError, match="clips 不能为空"):
self.service.compose([])
class TestWhiteListConstants:
"""白名单常量测试"""
def test_allowed_output_dirs_not_empty(self):
"""测试输出目录白名单不为空"""
assert len(ALLOWED_OUTPUT_DIRS) > 0
assert "/tmp/video_output" in ALLOWED_OUTPUT_DIRS
assert "/var/app/rendered" in ALLOWED_OUTPUT_DIRS
def test_allowed_input_prefixes_not_empty(self):
"""测试输入路径前缀白名单不为空"""
assert len(ALLOWED_INPUT_PREFIXES) > 0
assert "s3://" in ALLOWED_INPUT_PREFIXES
assert "oss://" in ALLOWED_INPUT_PREFIXES
assert "local://" in ALLOWED_INPUT_PREFIXES
assert "/var/storage/" in ALLOWED_INPUT_PREFIXES
def test_allowed_transitions_not_empty(self):
"""测试转场效果白名单不为空"""
assert len(ALLOWED_TRANSITIONS) > 0
assert "fade" in ALLOWED_TRANSITIONS
assert "dissolve" in ALLOWED_TRANSITIONS
assert "slideleft" in ALLOWED_TRANSITIONS
if __name__ == "__main__":
pytest.main([__file__, "-v"])
View File