Compare commits

..

1 Commits

Author SHA1 Message Date
灵应 2cde32794f fix(ci): check-push-paths 恢复脚本下载——commit e052c054 移除了 checkout 导致 runner 上无脚本文件
CI/CD Pipeline / Check push changed paths (pull_request) Has been skipped
CI/CD Pipeline / Dedup Check - skip PR tests when covered by push pipeline (pull_request) Successful in 1s
CI/CD Pipeline / Check if frontend-only change (pull_request) Successful in 1s
CI/CD Pipeline / Frontend Lint (pull_request) Has been skipped
CI/CD Pipeline / Frontend Unit Tests (pull_request) Has been skipped
CI/CD Pipeline / PR Build API Image (pull_request) Successful in 45s
CI/CD Pipeline / PR Build Web Image (pull_request) Has been skipped
CI/CD Pipeline / PR Build Worker Image (pull_request) Successful in 25s
CI/CD Pipeline / Build Staging API Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Web Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Worker Image (pull_request) Has been skipped
PR Automation / Auto Merge on CI Green + Approved (pull_request) Successful in 4m51s
CI/CD Pipeline / Retag skipped Staging API Image (pull_request) Has been skipped
CI/CD Pipeline / Retag skipped Staging Web Image (pull_request) Has been skipped
CI/CD Pipeline / Retag skipped Staging Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Deploy Staging (Watchtower auto-deploy) (pull_request) Has been skipped
CI/CD Pipeline / Staging E2E Tests (pull_request) Has been skipped
Preview Deploy / Deploy Preview Environment (pull_request) Successful in 5m18s
CI/CD Pipeline / Staging API Integration Tests (pull_request) Has been skipped
CI/CD Pipeline / ACR Image Cleanup (pull_request) Has been skipped
PR Automation / Auto Approve on CI Green (pull_request) Successful in 5m53s
CI/CD Pipeline / Validate - Style (pull_request) Successful in 9m5s
CI/CD Pipeline / Validate - Python (mypy + alembic) (pull_request) Successful in 9m33s
AI Code Review / AI Code Review (pull_request) Override - CI config change only, no code review needed
ACR Cleanup / ACR Image Cleanup (pull_request_target) Successful in 48s
Preview Cleanup / Cleanup Preview Environment (pull_request) Successful in 1m51s
CI/CD Pipeline / Validate - Security (pull_request) Successful in 23m33s
CI/CD Pipeline / Integration Tests (pull_request) Successful in 28m8s
CI/CD Pipeline / Unit Tests (pull_request) Successful in 37m45s
CI/CD Pipeline / Build Production API Image (pull_request) Has been skipped
CI/CD Pipeline / Build Production Web Image (pull_request) Has been skipped
CI/CD Pipeline / Build Production Worker Image (pull_request) Has been skipped
CI/CD Pipeline / CI Gate (pull_request) Successful in 1s
CI/CD Pipeline / Deploy Production (pull_request) Has been skipped
CI/CD Pipeline / Canary Release to Production (pull_request) Has been skipped
CI/CD Pipeline / Production Browser E2E (pull_request) Has been skipped
2026-08-31 19:44:38 +08:00
540 changed files with 6216 additions and 52727 deletions
-1
View File
@@ -1 +0,0 @@
CI re-trigger after runner add-host/DNS fix. This file is harmless and not referenced.
-105
View File
@@ -1,105 +0,0 @@
name: CI Base Image Build
on:
push:
branches:
- develop
- main
paths:
- 'requirements-base.txt'
- 'requirements-dev.txt'
- 'infra/docker/ci.Dockerfile'
workflow_dispatch:
inputs:
reason:
description: "触发原因"
required: false
default: "手动触发 - ci-base 镜像重建"
concurrency:
group: ci-base-image-build
cancel-in-progress: false
jobs:
build-ci-base:
name: Build CI Base Image
runs-on: runtime-builder
timeout-minutes: 60
steps:
- name: Checkout code
shell: sh
env:
GITHUB_TOKEN: ${{ github.token }}
run: |
curl -sH "Authorization: token $GITHUB_TOKEN" \
"${GITHUB_API_URL}/repos/${GITHUB_REPOSITORY}/raw/scripts/ci/step_checkout.sh?ref=${GITHUB_SHA}" \
| bash
- name: Docker login to Gitea Registry
shell: sh
env:
GITEA_REGISTRY_USER: xiaoxia
GITEA_REGISTRY_TOKEN: ${{ secrets.REGISTRY_TOKEN }}
run: |
set -eu
for i in 1 2 3; do
echo "=== Docker login 尝试 $i/3 ==="
if docker login git.xiaoxiajianji.com -u "${GITEA_REGISTRY_USER}" -p "${GITEA_REGISTRY_TOKEN}"; then
echo "✅ Docker login successful"
break
fi
echo "❌ Docker login 失败(尝试 $i/3),5s 后重试..."
sleep 5
done
- name: Build and push CI base image
shell: sh
run: |
set -eu
IMAGE="git.xiaoxiajianji.com/xiaoxia/xiaoxia-saas/ci-base"
VERSION_TAG="deps-$(date +%Y%m%d-%H%M)-${GITHUB_SHA::8}"
echo "=== Building CI base image (tags: latest, ${VERSION_TAG}) ==="
docker build --progress=plain \
-f infra/docker/ci.Dockerfile \
-t "${IMAGE}:latest" \
-t "${IMAGE}:${VERSION_TAG}" \
.
echo "✅ Image built successfully"
echo "=== Pushing ${VERSION_TAG} ==="
docker push "${IMAGE}:${VERSION_TAG}"
echo "=== Pushing latest ==="
docker push "${IMAGE}:latest"
echo "✅ Pushed to Gitea Registry"
- name: Verify image
shell: sh
run: |
set -eu
IMAGE="git.xiaoxiajianji.com/xiaoxia/xiaoxia-saas/ci-base:latest"
echo "=== Verifying pinned deps in fresh image ==="
docker run --rm "${IMAGE}" /opt/xiaoxia-ci-venv/bin/python -c \
"import httpcore, h2, numpy, httpx; print('VERSIONS:', httpcore.__version__, h2.__version__, numpy.__version__, httpx.__version__)"
- name: Notify result
if: always()
continue-on-error: true
shell: sh
env:
CI_NOTIFY_WEBHOOK: ${{ secrets.CI_NOTIFY_WEBHOOK }}
run: |
set +e
if [ "${{ job.status }}" = "success" ]; then
NOTIFY_MODE=success JOB_NAME="CI Base Image Build" python3 scripts/ci_notify.py
else
NOTIFY_MODE=failure JOB_NAME="CI Base Image Build" python3 scripts/ci_notify.py
fi
- name: Cleanup
if: always()
shell: sh
run: |
IMAGE="git.xiaoxiajianji.com/xiaoxia/xiaoxia-saas/ci-base"
docker rmi "${IMAGE}:latest" 2>/dev/null || true
echo "Cleanup done"
-51
View File
@@ -1,51 +0,0 @@
name: CI Canary Check
on:
schedule:
- cron: '*/30 * * * *'
workflow_dispatch:
jobs:
canary:
runs-on: ci-l2
timeout-minutes: 10
steps:
- name: Canary (runner -> docker -> network -> gitea)
run: |
set -e
echo "== runner/container basic =="
date; hostname; whoami
echo "== gitea api reachability =="
code=$(curl -s -o /tmp/v.json -w '%{http_code}' -m 15 "$GITHUB_API_URL/version")
echo "gitea api http_code=$code"
[ "$code" = "200" ] || { echo "::error::Gitea API unreachable, http_code=$code"; exit 1; }
cat /tmp/v.json; echo
echo "== external egress =="
ext=$(curl -s -o /dev/null -w '%{http_code}' -m 15 https://www.baidu.com || echo 000)
echo "external http_code=$ext"
echo "== gitea domain resolves NOT to loopback =="
set -o pipefail
ip=$(getent hosts git.xiaoxiajianji.com | awk '{print $1}' | head -1)
echo "git.xiaoxiajianji.com -> $ip"
if [ -z "$ip" ]; then
echo "::error::DNS resolution failed, git.xiaoxiajianji.com unresolvable"; exit 1
fi
if [ "$ip" = "127.0.0.1" ] || [ "$ip" = "::1" ]; then
echo "::error::Gitea domain resolves to loopback inside job container (hosts/DNS leak)"; exit 1
fi
echo "CANARY OK"
- name: Notify failure
if: failure()
env:
CI_NOTIFY_WEBHOOK: ${{ secrets.CI_NOTIFY_WEBHOOK }}
run: |
set +e
if [ -n "$CI_NOTIFY_WEBHOOK" ]; then
MSG="🚨 CI 金丝雀失败:runner->docker->网络->Gitea 链路异常,时间 $(date '+%Y-%m-%d %H:%M:%S'),请立即检查构建服务器"
python3 - "$CI_NOTIFY_WEBHOOK" "$MSG" <<'PY'
import json,sys,urllib.request
hook,msg=sys.argv[1],sys.argv[2]
data=json.dumps({"msg_type":"text","content":{"text":msg}}).encode()
urllib.request.urlopen(urllib.request.Request(hook,data=data,headers={"Content-Type":"application/json"}),timeout=10)
PY
fi
exit 0
+23 -84
View File
@@ -20,7 +20,6 @@ on:
default: "手动触发 - CI漏触发补跑" default: "手动触发 - CI漏触发补跑"
permissions: permissions:
contents: read contents: read
pull-requests: read
concurrency: concurrency:
group: ci-pipeline-${{ gitea.ref }} group: ci-pipeline-${{ gitea.ref }}
cancel-in-progress: true cancel-in-progress: true
@@ -89,22 +88,9 @@ jobs:
GITHUB_TOKEN: ${{ github.token }} GITHUB_TOKEN: ${{ github.token }}
run: | run: |
set -eu set -eu
# 优先用 git diff 判断 PR 改动范围(比 API 稳定)
PR_NUMBER=$(echo "$GITHUB_REF" | sed 's|refs/pull/||; s|/.*||') PR_NUMBER=$(echo "$GITHUB_REF" | sed 's|refs/pull/||; s|/.*||')
if command -v git >/dev/null 2>&1 && [ -d .git ]; then API_URL="${GITHUB_API_URL}/repos/${GITHUB_REPOSITORY}/pulls/${PR_NUMBER}/files?limit=300"
FILES=$(git diff --name-only origin/develop...HEAD 2>/dev/null || true) FILES=$(curl -s -H "Authorization: token ${GITHUB_TOKEN}" "$API_URL" | python3 -c "import sys,json; [print(f['filename']) for f in json.load(sys.stdin)]")
fi
if [ -z "${FILES:-}" ]; then
# fallback 到 API
API_URL="${GITHUB_API_URL}/repos/${GITHUB_REPOSITORY}/pulls/${PR_NUMBER}/files?limit=300"
FILES=$(curl -sf -H "Authorization: token ${GITHUB_TOKEN}" "$API_URL" | python3 -c "import sys,json; [print(f['filename']) for f in json.load(sys.stdin)]" 2>/dev/null || true)
fi
if [ -z "${FILES:-}" ]; then
echo "⚠️ 无法获取变更文件列表,保守运行完整 CI"
echo "skip_backend=false" >> $GITHUB_OUTPUT
echo "skip_frontend=false" >> $GITHUB_OUTPUT
exit 0
fi
FRONTEND_COUNT=$(echo "$FILES" | grep -c '^apps/web/' || true) FRONTEND_COUNT=$(echo "$FILES" | grep -c '^apps/web/' || true)
BACKEND_COUNT=$(echo "$FILES" | grep -cv '^apps/web/' || true) BACKEND_COUNT=$(echo "$FILES" | grep -cv '^apps/web/' || true)
TOTAL=$(echo "$FILES" | grep -cv '^$' || true) TOTAL=$(echo "$FILES" | grep -cv '^$' || true)
@@ -164,7 +150,6 @@ jobs:
run: bash scripts/ci/step_timer_start.sh run: bash scripts/ci/step_timer_start.sh
- name: Cache pip dependencies - name: Cache pip dependencies
uses: actions/cache@v4 uses: actions/cache@v4
continue-on-error: true
with: with:
path: /root/.cache/pip path: /root/.cache/pip
key: ${{ runner.os }}-pip-style-${{ hashFiles('requirements*.txt') }} key: ${{ runner.os }}-pip-style-${{ hashFiles('requirements*.txt') }}
@@ -196,7 +181,7 @@ jobs:
- name: Run style checks - name: Run style checks
shell: bash shell: bash
run: bash scripts/ci/validate_style.sh run: bash scripts/ci/validate_style.sh
- name: Auto-fix formatting (black + isort + ruff) - name: Auto-fix formatting (black + isort)
if: failure() if: failure()
shell: sh shell: sh
env: env:
@@ -267,7 +252,6 @@ jobs:
run: bash scripts/ci/step_timer_start.sh run: bash scripts/ci/step_timer_start.sh
- name: Cache pip dependencies - name: Cache pip dependencies
uses: actions/cache@v4 uses: actions/cache@v4
continue-on-error: true
with: with:
path: /root/.cache/pip path: /root/.cache/pip
key: ${{ runner.os }}-pip-security-${{ hashFiles('requirements*.txt') }} key: ${{ runner.os }}-pip-security-${{ hashFiles('requirements*.txt') }}
@@ -297,7 +281,6 @@ jobs:
sleep 5 sleep 5
done done
- name: Run security checks - name: Run security checks
continue-on-error: true # Security scan is advisory; runner failure must not block deploy
shell: bash shell: bash
env: env:
GITHUB_TOKEN: ${{ github.token }} GITHUB_TOKEN: ${{ github.token }}
@@ -348,7 +331,7 @@ jobs:
PIP_NO_CACHE_DIR: '' PIP_NO_CACHE_DIR: ''
DATABASE_URL: postgresql+psycopg://postgres:postgres@host.docker.internal:5432/xiaoxia_saas DATABASE_URL: postgresql+psycopg://postgres:postgres@host.docker.internal:5432/xiaoxia_saas
USE_IN_MEMORY_DB: 'false' USE_IN_MEMORY_DB: 'false'
CI_USE_SHARED_PG: 'false' CI_USE_SHARED_PG: 'true'
permissions: permissions:
contents: read contents: read
steps: steps:
@@ -369,7 +352,6 @@ jobs:
run: bash scripts/ci/step_timer_start.sh run: bash scripts/ci/step_timer_start.sh
- name: Cache pip dependencies - name: Cache pip dependencies
uses: actions/cache@v4 uses: actions/cache@v4
continue-on-error: true
with: with:
path: /root/.cache/pip path: /root/.cache/pip
key: ${{ runner.os }}-pip-python-${{ hashFiles('requirements*.txt') }} key: ${{ runner.os }}-pip-python-${{ hashFiles('requirements*.txt') }}
@@ -474,7 +456,6 @@ jobs:
run: bash scripts/ci/step_install_ffmpeg.sh run: bash scripts/ci/step_install_ffmpeg.sh
- name: Cache pip dependencies - name: Cache pip dependencies
uses: actions/cache@v4 uses: actions/cache@v4
continue-on-error: true
with: with:
path: /root/.cache/pip path: /root/.cache/pip
key: ${{ runner.os }}-pip-unittests-${{ hashFiles('requirements*.txt') }} key: ${{ runner.os }}-pip-unittests-${{ hashFiles('requirements*.txt') }}
@@ -532,7 +513,7 @@ jobs:
env: env:
DATABASE_URL: postgresql+psycopg://postgres:postgres@host.docker.internal:5432/xiaoxia_saas DATABASE_URL: postgresql+psycopg://postgres:postgres@host.docker.internal:5432/xiaoxia_saas
USE_IN_MEMORY_DB: 'false' USE_IN_MEMORY_DB: 'false'
CI_USE_SHARED_PG: 'false' CI_USE_SHARED_PG: 'true'
OSS_ACCESS_KEY_ID: placeholder OSS_ACCESS_KEY_ID: placeholder
OSS_ACCESS_KEY_SECRET: placeholder OSS_ACCESS_KEY_SECRET: placeholder
OSS_BUCKET_NAME: xiaoxia-autocut OSS_BUCKET_NAME: xiaoxia-autocut
@@ -683,7 +664,6 @@ jobs:
run: bash scripts/ci/step_timer_start.sh run: bash scripts/ci/step_timer_start.sh
- name: Cache npm dependencies - name: Cache npm dependencies
uses: actions/cache@v4 uses: actions/cache@v4
continue-on-error: true
with: with:
path: /root/.npm path: /root/.npm
key: ${{ runner.os }}-npm-vitest-${{ hashFiles('apps/web/package-lock.json') }} key: ${{ runner.os }}-npm-vitest-${{ hashFiles('apps/web/package-lock.json') }}
@@ -827,6 +807,9 @@ jobs:
CACHE_REF="${REGISTRY}/${{ matrix.cache_name }}:develop" CACHE_REF="${REGISTRY}/${{ matrix.cache_name }}:develop"
EXTRA_BUILD_ARGS="APP_VERSION=\"${GITHUB_SHA}\"" EXTRA_BUILD_ARGS="APP_VERSION=\"${GITHUB_SHA}\""
if [ "${{ matrix.service }}" = "web" ]; then
EXTRA_BUILD_ARGS="$EXTRA_BUILD_ARGS NGINX_CONF=infra/docker/nginx-staging.conf"
fi
# Worker 与 API/Web 统一走持久 builderci-builder-persist),共享宿主机层缓存 # Worker 与 API/Web 统一走持久 builderci-builder-persist),共享宿主机层缓存
NO_CACHE_FLAG="" NO_CACHE_FLAG=""
@@ -1023,6 +1006,9 @@ jobs:
CACHE_REF="${REGISTRY}/${{ matrix.cache_name }}:${GITHUB_REF_NAME}" CACHE_REF="${REGISTRY}/${{ matrix.cache_name }}:${GITHUB_REF_NAME}"
EXTRA_BUILD_ARGS="APP_VERSION=\"${GITHUB_SHA}\"" EXTRA_BUILD_ARGS="APP_VERSION=\"${GITHUB_SHA}\""
if [ "${{ matrix.service }}" = "web" ]; then
EXTRA_BUILD_ARGS="$EXTRA_BUILD_ARGS NGINX_CONF=infra/docker/nginx-staging.conf"
fi
NO_CACHE_FLAG="" NO_CACHE_FLAG=""
for i in 1 2 3; do for i in 1 2 3; do
@@ -1178,33 +1164,6 @@ jobs:
run: | run: |
set +e set +e
NOTIFY_MODE=start JOB_NAME="Deploy Staging" python3 scripts/ci_notify.py NOTIFY_MODE=start JOB_NAME="Deploy Staging" python3 scripts/ci_notify.py
- name: Render .env from template
shell: sh
env:
STAGING_DATABASE_URL: ${{ secrets.STAGING_DATABASE_URL }}
STAGING_REDIS_URL: ${{ secrets.STAGING_REDIS_URL }}
STAGING_CELERY_BROKER_URL: ${{ secrets.STAGING_CELERY_BROKER_URL }}
STAGING_CELERY_RESULT_BACKEND: ${{ secrets.STAGING_CELERY_RESULT_BACKEND }}
STAGING_JWT_SECRET_KEY: ${{ secrets.STAGING_JWT_SECRET_KEY }}
STAGING_MINIO_ENDPOINT: ${{ secrets.STAGING_MINIO_ENDPOINT }}
STAGING_MINIO_ACCESS_KEY: ${{ secrets.STAGING_MINIO_ACCESS_KEY }}
STAGING_MINIO_SECRET_KEY: ${{ secrets.STAGING_MINIO_SECRET_KEY }}
STAGING_MINIO_BUCKET: ${{ secrets.STAGING_MINIO_BUCKET }}
OSS_ACCESS_KEY_ID: ${{ secrets.OSS_ACCESS_KEY_ID }}
OSS_ACCESS_KEY_SECRET: ${{ secrets.OSS_ACCESS_KEY_SECRET }}
COSYVOICE_API_KEY: ${{ secrets.COSYVOICE_API_KEY }}
DASHSCOPE_API_KEY: ${{ secrets.DASHSCOPE_API_KEY }}
MEDIAKIT_API_KEY: ${{ secrets.MEDIAKIT_API_KEY }}
WECHAT_APP_ID: ${{ secrets.WECHAT_APP_ID }}
WECHAT_APP_SECRET: ${{ secrets.WECHAT_APP_SECRET }}
run: |
set -eu
echo "Rendering .env from template + secrets..."
bash scripts/render_env.sh staging
echo "✅ .env rendered (file contains secrets, not printed to log)"
# 验证文件存在且非空
test -s .env.rendered
echo "✅ .env.rendered validated ($(wc -l < .env.rendered) lines)"
- name: Docker login to Registry - name: Docker login to Registry
shell: sh shell: sh
env: env:
@@ -1277,31 +1236,9 @@ jobs:
ssh -p "$staging_port" -i "$key_path" -o StrictHostKeyChecking=no "${staging_user}@${staging_host}" "echo SSH_CONNECTION_OK && hostname" ssh -p "$staging_port" -i "$key_path" -o StrictHostKeyChecking=no "${staging_user}@${staging_host}" "echo SSH_CONNECTION_OK && hostname"
echo "SSH connection verified" echo "SSH connection verified"
# 配置 Diff 检查:下载服务器当前 .env,对比渲染结果,检测漂移
echo "Running config diff check..."
scp -P "$staging_port" -i "$key_path" -o StrictHostKeyChecking=no \
"${staging_user}@${staging_host}:/var/lib/xiaoxia-saas-staging/.env" .env.current 2>/dev/null \
|| touch .env.current # 首次部署时文件不存在,创建空文件
bash scripts/config_diff_check.sh .env.rendered .env.current
rm -f .env.current
echo "Config diff check done"
# 上传渲染后的 .env 到服务器(替代服务器上旧的 .env)
echo "Uploading rendered .env to staging server..."
# 备份旧 .env
ssh -p "$staging_port" -i "$key_path" -o StrictHostKeyChecking=no "${staging_user}@${staging_host}" \
"cp -f /var/lib/xiaoxia-saas-staging/.env /var/lib/xiaoxia-saas-staging/.env.bak.\$(date +%Y%m%d%H%M%S) 2>/dev/null || true"
# 上传新 .env
scp -P "$staging_port" -i "$key_path" -o StrictHostKeyChecking=no .env.rendered \
"${staging_user}@${staging_host}:/var/lib/xiaoxia-saas-staging/.env"
echo "✅ .env uploaded to staging server"
# 通过环境变量传递凭证,避免命令行引号转义问题 # 通过环境变量传递凭证,避免命令行引号转义问题
cat scripts/ci_staging_deploy.sh | ssh -p "$staging_port" -i "$key_path" -o StrictHostKeyChecking=no "${staging_user}@${staging_host}" "IMAGE_TAG=${GITHUB_SHA} ACR_USERNAME=${ACR_USERNAME} ACR_PASSWORD=${ACR_PASSWORD} sh" cat scripts/ci_staging_deploy.sh | ssh -p "$staging_port" -i "$key_path" -o StrictHostKeyChecking=no "${staging_user}@${staging_host}" "IMAGE_TAG=${GITHUB_SHA} ACR_USERNAME=${ACR_USERNAME} ACR_PASSWORD=${ACR_PASSWORD} sh"
# 清理 CI runner 上的渲染文件
rm -f .env.rendered
- name: Staging health check + auto rollback - name: Staging health check + auto rollback
if: success() if: success()
shell: sh shell: sh
@@ -1472,7 +1409,7 @@ jobs:
- unit-tests - unit-tests
- frontend-lint - frontend-lint
- frontend-unit-test - frontend-unit-test
if: github.event_name == 'push' && github.ref_name == 'main' && !failure() && !cancelled() if: startsWith(github.ref, 'refs/tags/v') || (github.event_name == 'push' && github.ref_name == 'main')
strategy: strategy:
fail-fast: false fail-fast: false
matrix: matrix:
@@ -1556,11 +1493,18 @@ jobs:
set -eu set -eu
REGISTRY="xiaoxia-registry.cn-hangzhou.cr.aliyuncs.com/xiaoxiakeji" REGISTRY="xiaoxia-registry.cn-hangzhou.cr.aliyuncs.com/xiaoxiakeji"
# 根据ref类型设置镜像标签:tag用版本号,分支用分支名+sha # 根据ref类型设置镜像标签:tag用版本号,分支用分支名+sha
TAG_NAME="${GITHUB_SHA}" if [[ "$GITHUB_REF" == refs/tags/* ]]; then
TAG_NAME="${GITHUB_REF_NAME}"
else
TAG_NAME="${GITHUB_REF_NAME}-${GITHUB_SHA::8}"
fi
IMAGE_TAG="${REGISTRY}/${{ matrix.image_name }}:${TAG_NAME}" IMAGE_TAG="${REGISTRY}/${{ matrix.image_name }}:${TAG_NAME}"
CACHE_REF="${REGISTRY}/${{ matrix.cache_name }}:main" CACHE_REF="${REGISTRY}/${{ matrix.cache_name }}:main"
EXTRA_BUILD_ARGS="APP_VERSION=\"${TAG_NAME}\"" EXTRA_BUILD_ARGS="APP_VERSION=\"${TAG_NAME}\""
if [ "${{ matrix.service }}" = "web" ]; then
EXTRA_BUILD_ARGS="$EXTRA_BUILD_ARGS NGINX_CONF=infra/docker/nginx-production.conf"
fi
# Docker build 带重试:失败自动重试2次,第2次重试加--no-cache # Docker build 带重试:失败自动重试2次,第2次重试加--no-cache
NO_CACHE_FLAG="" NO_CACHE_FLAG=""
@@ -1615,7 +1559,7 @@ jobs:
concurrency: concurrency:
group: deploy-production-${{ gitea.ref }} group: deploy-production-${{ gitea.ref }}
cancel-in-progress: false cancel-in-progress: false
if: github.event_name == 'push' && github.ref_name == 'main' if: startsWith(github.ref, 'refs/tags/v')
needs: needs:
- build-production - build-production
steps: steps:
@@ -1688,7 +1632,7 @@ jobs:
echo "SSH connection verified" echo "SSH connection verified"
# 通过环境变量传递凭证,避免命令行引号转义问题 # 通过环境变量传递凭证,避免命令行引号转义问题
cat scripts/ci_production_deploy.sh | ssh -p "$production_port" -i "$key_path" -o StrictHostKeyChecking=no "${production_user}@${production_host}" "IMAGE_TAG=${GITHUB_SHA} ACR_USERNAME=${ACR_USERNAME} ACR_PASSWORD=${ACR_PASSWORD} sh" cat scripts/ci_production_deploy.sh | ssh -p "$production_port" -i "$key_path" -o StrictHostKeyChecking=no "${production_user}@${production_host}" "IMAGE_TAG=${GITHUB_REF_NAME} ACR_USERNAME=${ACR_USERNAME} ACR_PASSWORD=${ACR_PASSWORD} sh"
- name: Production health check + auto rollback - name: Production health check + auto rollback
if: success() if: success()
@@ -1747,7 +1691,7 @@ jobs:
name: Production Browser E2E name: Production Browser E2E
runs-on: runtime-builder runs-on: runtime-builder
timeout-minutes: 15 timeout-minutes: 15
# if: removed - runs after deploy-production succeeds if: startsWith(github.ref, 'refs/tags/v')
needs: deploy-production needs: deploy-production
steps: steps:
- name: Checkout code - name: Checkout code
@@ -2067,11 +2011,6 @@ jobs:
echo " ⏳ $name: pending(审查中,暂不阻塞)" echo " ⏳ $name: pending(审查中,暂不阻塞)"
continue continue
fi fi
# Security scan cancelled/failed时不阻塞部署(runner故障不应卡住流水线)
if [ "$name" = "validate-security" ] && { [ "$result" = "cancelled" ] || [ "$result" = "failure" ]; }; then
echo " ⚠️ $name: $result(安全扫描为非阻塞项,不卡住部署)"
continue
fi
check_job "$name" "$result" check_job "$name" "$result"
done done
-60
View File
@@ -1,60 +0,0 @@
name: "Debug: Web container v2 (mount conflict)"
on:
push:
branches: [debug/web-crash-v2]
workflow_dispatch:
jobs:
web-diag:
runs-on: runtime-builder
timeout-minutes: 10
steps:
- name: Setup SSH and diagnose
shell: bash
env:
STAGING_SSH_KEY: ${{ secrets.PREVIEW_SSH_KEY }}
run: |
set -x
which ssh || (apt-get update -qq && apt-get install -y -qq openssh-client)
mkdir -p ~/.ssh && chmod 700 ~/.ssh
printf "%s" "$STAGING_SSH_KEY" > ~/.ssh/id_rsa
chmod 600 ~/.ssh/id_rsa
H=47.98.113.167; P=22222
ssh-keyscan -p $P -H $H >> ~/.ssh/known_hosts 2>/dev/null
ssh -p $P -i ~/.ssh/id_rsa -o StrictHostKeyChecking=no root@$H 'bash -s' <<'REMOTE'
set -x
echo "=== Current staging containers ==="
docker ps -a --filter name=xiaoxia-*-staging --format "table {{.Names}}\t{{.Status}}\t{{.Image}}"
echo ""
echo "=== Web container logs (current/current-rolledback) ==="
docker logs xiaoxia-web-staging 2>&1 | tail -40
echo ""
echo "=== Web inspect: env & mounts ==="
docker inspect xiaoxia-web-staging --format 'Entrypoint: {{.Config.Entrypoint}} Cmd: {{.Config.Cmd}}'
docker inspect xiaoxia-web-staging --format '{{range .Config.Env}}{{.}}{{"\n"}}{{end}}' | grep -E "APP_ENV|VERSION"
echo "Mounts:"
docker inspect xiaoxia-web-staging --format '{{range .Mounts}}{{.Type}} {{.Source}} -> {{.Destination}} (rw={{.RW}}){{"\n"}}{{end}}'
echo ""
echo "=== Reproduce: rm on read-only bind mount ==="
docker run --rm --name nginx-ro-test \
-v /var/lib/xiaoxia-saas-staging/nginx-staging.conf:/etc/nginx/conf.d/default.conf:ro \
git.xiaoxiajianji.com/xiaoxia/xiaoxia-saas/xiaoxia-saas-web:387514c \
sh -c '
set -x
echo "Before:"
ls -la /etc/nginx/conf.d/
echo "Try rm (as entrypoint does):"
rm -f /etc/nginx/conf.d/default.conf
echo "rm exitcode=$?"
echo "After rm:"
ls -la /etc/nginx/conf.d/
echo "Test ln:"
ln -s /etc/nginx/nginx-staging.conf /etc/nginx/conf.d/default.conf
echo "ln exitcode=$?"
ls -la /etc/nginx/conf.d/
echo "nginx -t:"
nginx -t 2>&1
' 2>&1
echo ""
echo "=== Also test with NEW fixed image (9c0d4b1 if present) ==="
docker images | grep xiaoxia-saas-web | head -5
REMOTE
-28
View File
@@ -1,28 +0,0 @@
name: "E2E lipsync verify v3"
on:
push:
branches: [debug/e2e-lipsync]
workflow_dispatch:
jobs:
e2e:
runs-on: runtime-builder
timeout-minutes: 10
steps:
- name: Setup SSH
shell: bash
env:
STAGING_SSH_KEY: ${{ secrets.PREVIEW_SSH_KEY }}
run: |
set -eux
which ssh || (apt-get update -qq && apt-get install -y -qq openssh-client)
mkdir -p ~/.ssh && chmod 700 ~/.ssh
printf "%s" "$STAGING_SSH_KEY" > ~/.ssh/id_rsa
chmod 600 ~/.ssh/id_rsa
ssh-keyscan -p 22222 -H 47.98.113.167 >> ~/.ssh/known_hosts 2>/dev/null
- name: Run
shell: bash
run: |
set -x
echo 'IyEvYmluL2Jhc2gKc2V0IC14CmVjaG8gIj09PSAxLiBSZWNlbnQgbGlwc3luYyBqb2JzIChjb3JyZWN0IG1vZGVsIHBhdGgpID09PSIKZG9ja2VyIGV4ZWMgeGlhb3hpYS1hcGktc3RhZ2luZyBzaCAtYyAiY2QgL2FwcCAmJiBweXRob24gLWMgXCIKZnJvbSBhcHAuZGIgaW1wb3J0IFNlc3Npb25Mb2NhbApmcm9tIHBhY2thZ2VzLmFkYXB0ZXJzLnNxbGFsY2hlbXlfaW1wbC5tb2RlbHMgaW1wb3J0IExpcHN5bmNKb2JNb2RlbCBhcyBMSgpkYiA9IFNlc3Npb25Mb2NhbCgpCnRvdGFsID0gZGIucXVlcnkoTEopLmNvdW50KCkKcHJpbnQoZidUb3RhbCBsaXBzeW5jIGpvYnM6IHt0b3RhbH0nKQpqb2JzID0gZGIucXVlcnkoTEopLm9yZGVyX2J5KExKLmNyZWF0ZWRfYXQuZGVzYygpKS5saW1pdCgxMCkuYWxsKCkKZm9yIGogaW4gam9iczoKICAgIGVyciA9IChqLmVycm9yX21lc3NhZ2Ugb3IgJycpWzoxMDBdCiAgICBwcmludChmJyAgaWQ9e2ouaWR9IHVzZXI9e2oudXNlcl9pZH0gc3RhdHVzPXtqLnN0YXR1c30gZXJyX2NvZGU9e2ouZXJyb3JfY29kZX0gZXJyPXtlcnJ9JykKZGIuY2xvc2UoKQpcIiIgMj4mMSB8IHRhaWwgLTMwCmVjaG8gIiIKZWNobyAiPT09IDIuIFRlc3QgdXNlcnMgPT09Igpkb2NrZXIgZXhlYyB4aWFveGlhLWFwaS1zdGFnaW5nIHNoIC1jICJjZCAvYXBwICYmIHB5dGhvbiAtYyBcIgpmcm9tIGFwcC5kYiBpbXBvcnQgU2Vzc2lvbkxvY2FsCmZyb20gcGFja2FnZXMuYWRhcHRlcnMuc3FsYWxjaGVteV9pbXBsLm1vZGVscyBpbXBvcnQgVXNlck1vZGVsIGFzIFUKZGIgPSBTZXNzaW9uTG9jYWwoKQp1c2VycyA9IGRiLnF1ZXJ5KFUpLm9yZGVyX2J5KFUuY3JlYXRlZF9hdC5kZXNjKCkpLmxpbWl0KDUpLmFsbCgpCmZvciB1IGluIHVzZXJzOgogICAgcHJpbnQoZicgIGlkPXt1LmlkfSBlbWFpbD17dS5lbWFpbH0gcGhvbmU9e2dldGF0dHIodSxcInBob25lX251bWJlclwiLE5vbmUpfSBhY3RpdmU9e2dldGF0dHIodSxcImlzX2FjdGl2ZVwiLE5vbmUpfScpCmRiLmNsb3NlKCkKXCIiIDI+JjEgfCB0YWlsIC0xNQplY2hvICIiCmVjaG8gIj09PSAzLiBBUEkgd29ya2VyIGxvZ3M6IGFueSBsaXBzeW5jIHRhc2sgZXhlY3V0aW9uIGhpc3Rvcnk/ID09PSIKZG9ja2VyIGxvZ3MgeGlhb3hpYS13b3JrZXItc3RhZ2luZyAyPiYxIHwgZ3JlcCAtaUUgImxpcHN5bmNfdHRzfHN5bnRoZXNpemVfYW5kX3N1Ym1pdHxBc3luY0Rpc3BhdGNoRmFpbGVkfHR0c19wcm9jZXNzaW5nfHR0c19wcm94eSIgfCB0YWlsIC0yMAplY2hvICIiCmVjaG8gIj09PSA0LiBQdWJsaWMgR0VUIC9oZWFsdGggd2l0aCBHRVQgbWV0aG9kID09PSIKY3VybCAtc2YgLS1tYXgtdGltZSAxMCBodHRwczovL3N0YWdpbmctYXBpLnhpYW94aWFqaWFuamkuY29tL2hlYWx0aCAyPiYxIHwgaGVhZCAtNQplY2hvICIiCmVjaG8gIj09PSA1LiBQdWJsaWMgd2ViIHJvb3QgPT09IgpjdXJsIC1zZiAtLW1heC10aW1lIDEwIGh0dHBzOi8vc3RhZ2luZy54aWFveGlhamlhbmppLmNvbS8gMj4mMSB8IGhlYWQgLTUK' | base64 -d > /tmp/e2e.sh
chmod +x /tmp/e2e.sh
ssh -p 22222 -i ~/.ssh/id_rsa -o StrictHostKeyChecking=no root@47.98.113.167 'bash -s' < /tmp/e2e.sh
@@ -1,59 +0,0 @@
name: Playwright Base Image Build
on:
workflow_dispatch:
inputs:
reason:
description: "触发原因"
required: false
default: "构建 playwright 基础镜像"
jobs:
build-playwright:
name: Build Playwright Base Image
runs-on: runtime-builder
timeout-minutes: 30
steps:
- name: Docker login to Gitea Registry
shell: sh
env:
GITEA_REGISTRY_USER: xiaoxia
GITEA_REGISTRY_TOKEN: ${{ secrets.REGISTRY_TOKEN }}
run: |
set -eu
for i in 1 2 3; do
echo "=== Docker login attempt $i/3 ==="
if printf '%s' "${GITEA_REGISTRY_TOKEN}" | docker login git.xiaoxiajianji.com -u "${GITEA_REGISTRY_USER}" --password-stdin; then
echo "Docker login successful"
break
fi
echo "Docker login failed (attempt $i/3), retrying in 5s..."
sleep 5
[ $i -eq 3 ] && exit 1
done
- name: Pull, retag and push Playwright image
shell: sh
run: |
set -eu
OFFICIAL_IMAGE="mcr.microsoft.com/playwright:v1.45.0-jammy"
GITEA_IMAGE="git.xiaoxiajianji.com/xiaoxia/base/playwright:v1.45.0-jammy"
echo "=== Pulling official Playwright image ==="
docker pull "${OFFICIAL_IMAGE}"
echo "=== Tagging ==="
docker tag "${OFFICIAL_IMAGE}" "${GITEA_IMAGE}"
echo "=== Pushing to Gitea Registry ==="
docker push "${GITEA_IMAGE}"
echo "Done: ${GITEA_IMAGE}"
- name: Cleanup
if: always()
shell: sh
run: |
docker rmi "mcr.microsoft.com/playwright:v1.45.0-jammy" 2>/dev/null || true
docker rmi "git.xiaoxiajianji.com/xiaoxia/base/playwright:v1.45.0-jammy" 2>/dev/null || true
echo "Cleanup done"
+1 -1
View File
@@ -3,7 +3,7 @@ name: PR Auto Scan
# 作为短作业模式的兜底,防止事件驱动遗漏 # 作为短作业模式的兜底,防止事件驱动遗漏
on: on:
schedule: schedule:
# - cron: "*/15 * * * *" # DISABLED: temporarily to stop failure spam (2026-09-02) # 每10分钟扫描一次(脚本自带240s墙钟上限,降频减负) - cron: "*/10 * * * *" # 每10分钟扫描一次(脚本自带240s墙钟上限,降频减负)
workflow_dispatch: workflow_dispatch:
permissions: permissions:
+2 -3
View File
@@ -18,7 +18,7 @@ jobs:
name: Auto Approve on CI Green name: Auto Approve on CI Green
runs-on: ci-check runs-on: ci-check
if: github.event_name == 'pull_request' && !github.event.pull_request.draft if: github.event_name == 'pull_request' && !github.event.pull_request.draft
timeout-minutes: 10 # 等待CI全绿+审批,需要充足时间 timeout-minutes: 3 # 等待模式:等CI全绿后自动合并,不遗漏任何PR
steps: steps:
- name: Checkout code - name: Checkout code
shell: sh shell: sh
@@ -61,8 +61,7 @@ jobs:
name: Auto Merge on CI Green + Approved name: Auto Merge on CI Green + Approved
runs-on: ci-check runs-on: ci-check
if: github.event_name == 'pull_request' && !github.event.pull_request.draft && github.event.pull_request.base.ref == 'develop' if: github.event_name == 'pull_request' && !github.event.pull_request.draft && github.event.pull_request.base.ref == 'develop'
needs: [auto-approve] # 修复竞态:必须等审批完成后再尝试合并 timeout-minutes: 3 # 短作业模式:检查一次,不满足就退出,由pr-auto-scan每5分钟定时兜底
timeout-minutes: 15 # 等待审批+CI就绪+合并,需要充足时间
steps: steps:
- name: Checkout code - name: Checkout code
shell: sh shell: sh
-6
View File
@@ -24,11 +24,6 @@ ruff_cache/
.env.production .env.production
.env.staging .env.staging
!.env.example !.env.example
# 配置模板不受忽略规则限制
!deploy/configs/.env.staging
!deploy/configs/.env.production
# 渲染后的 env 文件包含真实密钥,绝不能提交
.env.rendered
# OS / editor # OS / editor
.DS_Store .DS_Store
@@ -59,4 +54,3 @@ frontend-v21-ui-prototype-final.html
!.vscode/settings.json !.vscode/settings.json
.vscode/extensions.json .vscode/extensions.json
.coverage .coverage
.env.current
@@ -1,57 +0,0 @@
"""migrate template_segments data to template_clip_configs
Revision ID: 060_migrate_segments
Revises: 059_duplicate_rate
Create Date: 2026-08-31
"""
import sqlalchemy as sa
from alembic import op
revision = "060_migrate_segments"
down_revision = "059_duplicate_rate"
branch_labels = None
depends_on = None
def upgrade() -> None:
dialect = op.get_bind().dialect.name
if dialect == "postgresql":
config_expr = (
"CASE WHEN s.material_type IS NOT NULL AND s.material_type != '' "
"THEN json_build_object('material_type', s.material_type)::jsonb "
"ELSE '{}'::jsonb END"
)
empty_json = "'{}'::jsonb"
else:
config_expr = (
"CASE WHEN s.material_type IS NOT NULL AND s.material_type != '' "
"THEN JSON_OBJECT('material_type', s.material_type) "
"ELSE '{}' END"
)
empty_json = "'{}'"
sql_str = (
"INSERT INTO template_clip_configs "
'(id, template_id, clip_type, "order", min_duration, max_duration, '
"text_template, material_requirements, transition_effect, config, "
"created_at, updated_at) "
"SELECT "
"s.id, s.template_id, 'main', s.segment_order, "
"s.duration_min, s.duration_max, "
"'', " + empty_json + ", "
"'cut', " + config_expr + ", "
"s.created_at, s.updated_at "
"FROM template_segments s "
"WHERE NOT EXISTS ("
" SELECT 1 FROM template_clip_configs c "
" WHERE c.template_id = s.template_id"
")"
)
op.execute(sa.text(sql_str))
def downgrade() -> None:
pass
@@ -1,26 +0,0 @@
"""add sort_order to template_categories
Revision ID: 061_sort_order
Revises: 060_migrate_segments
Create Date: 2026-09-02
"""
import sqlalchemy as sa
from alembic import op
revision = "061_sort_order"
down_revision = "060_migrate_segments"
branch_labels = None
depends_on = None
def upgrade() -> None:
op.add_column(
"template_categories",
sa.Column("sort_order", sa.Integer, nullable=False, server_default="0"),
)
def downgrade() -> None:
op.drop_column("template_categories", "sort_order")
@@ -1,28 +0,0 @@
"""re-add edit_plan_id to generation_tasks (align staging with production)
Revision ID: 062_edit_plan_id
Revises: 061_sort_order
Create Date: 2026-09-02
"""
import sqlalchemy as sa
from alembic import op
revision = "062_edit_plan_id"
down_revision = "061_sort_order"
branch_labels = None
depends_on = None
def upgrade() -> None:
op.add_column(
"generation_tasks",
sa.Column("edit_plan_id", sa.String(36), nullable=True),
)
op.create_index("ix_generation_tasks_edit_plan_id_2", "generation_tasks", ["edit_plan_id"])
def downgrade() -> None:
op.drop_index("ix_generation_tasks_edit_plan_id_2", table_name="generation_tasks")
op.drop_column("generation_tasks", "edit_plan_id")
@@ -1,46 +0,0 @@
"""add video_fingerprint_chunks table for per-chunk fingerprint storage
Revision ID: 063_fingerprint_chunks
Revises: 062_edit_plan_id
Create Date: 2026-09-03
"""
import sqlalchemy as sa
from alembic import op
revision = "063_fingerprint_chunks"
down_revision = "062_edit_plan_id"
branch_labels = None
depends_on = None
def upgrade() -> None:
op.create_table(
"video_fingerprint_chunks",
sa.Column("id", sa.String(36), primary_key=True),
sa.Column("video_id", sa.String(36), nullable=False),
sa.Column("project_id", sa.String(36), nullable=False),
sa.Column("user_id", sa.String(36), nullable=False, server_default=""),
sa.Column("start_time_ms", sa.Integer, nullable=False),
sa.Column("end_time_ms", sa.Integer, nullable=False),
sa.Column("phash_binary", sa.String(16), nullable=False),
sa.Column("color_histogram", sa.JSON, nullable=False),
sa.Column("frame_count", sa.Integer, nullable=False, server_default="1"),
sa.Column(
"created_at",
sa.DateTime,
nullable=False,
server_default=sa.func.now(),
),
)
op.create_index("ix_vfc_video_id", "video_fingerprint_chunks", ["video_id"])
op.create_index("ix_vfc_project_id", "video_fingerprint_chunks", ["project_id"])
op.create_index("ix_vfc_user_id", "video_fingerprint_chunks", ["user_id"])
def downgrade() -> None:
op.drop_index("ix_vfc_user_id", table_name="video_fingerprint_chunks")
op.drop_index("ix_vfc_project_id", table_name="video_fingerprint_chunks")
op.drop_index("ix_vfc_video_id", table_name="video_fingerprint_chunks")
op.drop_table("video_fingerprint_chunks")
@@ -1,25 +0,0 @@
"""add match_count and visual_similarity to generated_videos
Revision ID: 064_match_count_visual_sim
Revises: 063_fingerprint_chunks
Create Date: 2026-09-03
"""
import sqlalchemy as sa
from alembic import op
revision = "064_match_count_visual_sim"
down_revision = "063_fingerprint_chunks"
branch_labels = None
depends_on = None
def upgrade() -> None:
op.add_column("generated_videos", sa.Column("match_count", sa.Integer(), nullable=True, server_default="0"))
op.add_column("generated_videos", sa.Column("visual_similarity", sa.Float(), nullable=True, server_default="0.0"))
def downgrade() -> None:
op.drop_column("generated_videos", "visual_similarity")
op.drop_column("generated_videos", "match_count")
@@ -1,25 +0,0 @@
"""add visual_similarity and match_count to duplication_records
Revision ID: 065_dup_record_sim_match
Revises: 064_match_count_visual_sim
Create Date: 2026-09-04
"""
import sqlalchemy as sa
from alembic import op
revision = "065_dup_record_sim_match"
down_revision = "064_match_count_visual_sim"
branch_labels = None
depends_on = None
def upgrade() -> None:
op.add_column("duplication_records", sa.Column("visual_similarity", sa.Float(), nullable=True))
op.add_column("duplication_records", sa.Column("match_count", sa.Integer(), nullable=True))
def downgrade() -> None:
op.drop_column("duplication_records", "match_count")
op.drop_column("duplication_records", "visual_similarity")
@@ -1,34 +0,0 @@
"""add client_upload_id to assets and asset_id to ingest_jobs
Issue #1714:上传 complete 幂等 + worker 转码回写关联。
- assets.client_upload_id:客户端幂等 tokencomplete 去重)
- ingest_jobs.asset_idcomplete 阶段创建的占位 asset idworker 回写关联,
防止 HEVC 转码改写 storage_key 后找不到占位而兜底新建 READY 记录)
Revision ID: 066_upload_idempotency
Revises: 065_dup_record_sim_match
Create Date: 2026-09-05
"""
import sqlalchemy as sa
from alembic import op
revision = "066_upload_idempotency"
down_revision = "065_dup_record_sim_match"
branch_labels = None
depends_on = None
def upgrade() -> None:
op.add_column("assets", sa.Column("client_upload_id", sa.String(64), nullable=True))
op.create_index("ix_assets_client_upload_id", "assets", ["client_upload_id"])
op.add_column("ingest_jobs", sa.Column("asset_id", sa.String(36), nullable=False, server_default=""))
op.create_index("ix_ingest_jobs_asset_id", "ingest_jobs", ["asset_id"])
def downgrade() -> None:
op.drop_index("ix_ingest_jobs_asset_id", table_name="ingest_jobs")
op.drop_column("ingest_jobs", "asset_id")
op.drop_index("ix_assets_client_upload_id", table_name="assets")
op.drop_column("assets", "client_upload_id")
@@ -1,35 +0,0 @@
"""add celery_task_id to generation_tasks and ingest_jobs
Issue #1714:孤儿恢复/超时清理撤销队列消息。
- generation_tasks.celery_task_id:入队时记录的 Celery 消息 ID,清理时 revoke
- ingest_jobs.celery_task_id:同上(素材转码任务)
Revision ID: 067_celery_task_id
Revises: 066_upload_idempotency
Create Date: 2026-09-05
"""
import sqlalchemy as sa
from alembic import op
revision = "067_celery_task_id"
down_revision = "066_upload_idempotency"
branch_labels = None
depends_on = None
def upgrade() -> None:
op.add_column(
"generation_tasks",
sa.Column("celery_task_id", sa.String(64), nullable=False, server_default=""),
)
op.add_column(
"ingest_jobs",
sa.Column("celery_task_id", sa.String(64), nullable=False, server_default=""),
)
def downgrade() -> None:
op.drop_column("ingest_jobs", "celery_task_id")
op.drop_column("generation_tasks", "celery_task_id")
@@ -1,26 +0,0 @@
"""add profile_completed to users
Issue #1718:微信新用户首次登录需设置昵称(PATCH /auth/me)。
- users.profile_completed:资料是否已完善;存量行默认 True(不触发引导),
微信新建用户在应用层置 False。
"""
import sqlalchemy as sa
from alembic import op
revision = "068_user_profile_completed"
down_revision = "067_celery_task_id"
branch_labels = None
depends_on = None
def upgrade() -> None:
op.add_column(
"users",
sa.Column("profile_completed", sa.Boolean(), nullable=False, server_default=sa.text("true")),
)
def downgrade() -> None:
op.drop_column("users", "profile_completed")
@@ -1,72 +0,0 @@
"""Projects is_default + partial unique index for idempotent default project (Issue #1775)
Revision ID: 069_project_is_default
Revises: 068_user_profile_completed
Create Date: 2026-09-08
背景:
小程序端 getOrCreateDefaultProject 在重试/并发/前端重复调用下,
仅靠应用层"先查再插"不保证幂等,会给同一用户重复创建默认项目。
改动:
1. projects 表新增 is_default 布尔列(默认 false
2. 部分唯一索引 uq_projects_owner_default(owner_user_id) WHERE is_default = true
—— 保证每个用户至多一个默认项目
3. 存量数据回填:把名为"默认项目"的存量项目按创建时间最早者标记为 is_default=true
(只标记不删除;存量重复项目的清理另行确认后单独执行)
注意:部分唯一索引依赖 PostgreSQL,不支持 downgrade 到其他方言。
"""
import sqlalchemy as sa
from alembic import op
revision = "069_project_is_default"
down_revision = "068_user_profile_completed"
branch_labels = None
depends_on = None
def upgrade() -> None:
# 1. 新增 is_default 列
op.add_column(
"projects",
sa.Column(
"is_default",
sa.Boolean(),
nullable=False,
server_default=sa.text("false"),
),
)
# 2. 存量回填:每个拥有"默认项目"的用户,只把最早创建的那一个标记为默认。
# 用 ROW_NUMBER() 取每组第一条;非"默认项目"命名的项目不标记(保守,不动用户自建项目)。
op.execute("""
UPDATE projects p
SET is_default = true
WHERE p.id IN (
SELECT id FROM (
SELECT id,
ROW_NUMBER() OVER (
PARTITION BY owner_user_id
ORDER BY created_at ASC, id ASC
) AS rn
FROM projects
WHERE name = '默认项目'
) t
WHERE t.rn = 1
)
""")
# 3. 部分唯一索引:每用户至多一个默认项目(只约束 is_default = true 的行)
op.execute("""
CREATE UNIQUE INDEX uq_projects_owner_default
ON projects (owner_user_id)
WHERE is_default = true
""")
def downgrade() -> None:
op.execute("DROP INDEX IF EXISTS uq_projects_owner_default")
op.drop_column("projects", "is_default")
-48
View File
@@ -1,48 +0,0 @@
"""Add scripts table for oral broadcast script library (Issue #1795)
Revision ID: 070_add_scripts
Revises: 069_project_is_default
Create Date: 2026-09-08
新建 scripts 表,支持口播文案 CRUD + 分段存储。
"""
import sqlalchemy as sa
from alembic import op
revision = "070_add_scripts"
down_revision = "069_project_is_default"
branch_labels = None
depends_on = None
def upgrade() -> None:
op.create_table(
"scripts",
sa.Column("id", sa.String(36), nullable=False),
sa.Column("user_id", sa.String(36), nullable=False),
sa.Column("title", sa.String(255), nullable=False),
sa.Column("content", sa.Text(), nullable=False, server_default=""),
sa.Column("segments", sa.JSON(), nullable=False, server_default="[]"),
sa.Column("tags", sa.JSON(), nullable=False, server_default="[]"),
sa.Column(
"created_at",
sa.DateTime(timezone=True),
nullable=False,
server_default=sa.func.now(),
),
sa.Column(
"updated_at",
sa.DateTime(timezone=True),
nullable=False,
server_default=sa.func.now(),
),
sa.PrimaryKeyConstraint("id"),
)
op.create_index("ix_scripts_user_id", "scripts", ["user_id"])
def downgrade() -> None:
op.drop_index("ix_scripts_user_id", table_name="scripts")
op.drop_table("scripts")
@@ -1,47 +0,0 @@
"""add lipsync jobs table
Revision ID: 071_add_lipsync_jobs
Revises: 070_add_scripts
Create Date: 2026-09-08
"""
import sqlalchemy as sa
from alembic import op
revision = "071_add_lipsync_jobs"
down_revision = "070_add_scripts"
branch_labels = None
depends_on = None
def upgrade() -> None:
op.create_table(
"lipsync_jobs",
sa.Column("id", sa.String(36), primary_key=True),
sa.Column("user_id", sa.String(36), nullable=False, index=True),
sa.Column("project_id", sa.String(36), nullable=False, server_default=""),
sa.Column("video_url", sa.Text(), nullable=False),
sa.Column("audio_url", sa.Text(), nullable=False),
sa.Column("enable_video_loop", sa.Boolean(), nullable=False, server_default=sa.text("false")),
sa.Column("mediakit_task_id", sa.String(200), nullable=False, server_default="", index=True),
sa.Column("status", sa.String(20), nullable=False, server_default="pending", index=True),
sa.Column("output_video_url", sa.Text(), nullable=False, server_default=""),
sa.Column("output_duration", sa.Float(), nullable=False, server_default=sa.text("0.0")),
sa.Column("error_message", sa.Text(), nullable=False, server_default=""),
sa.Column("error_code", sa.String(100), nullable=False, server_default=""),
sa.Column("submitted_at", sa.DateTime(), nullable=True),
sa.Column("completed_at", sa.DateTime(), nullable=True),
sa.Column("created_at", sa.DateTime(), nullable=False, server_default=sa.func.now()),
sa.Column("updated_at", sa.DateTime(), nullable=False, server_default=sa.func.now()),
)
# 复合索引:用户 + 状态(列表查询常用)
op.create_index("ix_lipsync_jobs_user_status", "lipsync_jobs", ["user_id", "status"])
# 项目 + 用户(项目维度查询)
op.create_index("ix_lipsync_jobs_project_user", "lipsync_jobs", ["project_id", "user_id"])
def downgrade() -> None:
op.drop_index("ix_lipsync_jobs_project_user", table_name="lipsync_jobs")
op.drop_index("ix_lipsync_jobs_user_status", table_name="lipsync_jobs")
op.drop_table("lipsync_jobs")
@@ -1,48 +0,0 @@
"""add ai avatar render jobs table
Revision ID: 072_add_ai_avatar_render
Revises: 071_add_lipsync_jobs
Create Date: 2026-09-09
"""
import sqlalchemy as sa
from alembic import op
revision = "072_add_ai_avatar_render"
down_revision = "071_add_lipsync_jobs"
branch_labels = None
depends_on = None
def upgrade() -> None:
op.create_table(
"ai_avatar_render_jobs",
sa.Column("id", sa.String(36), primary_key=True),
sa.Column("user_id", sa.String(36), nullable=False, index=True),
sa.Column("project_id", sa.String(36), nullable=False, server_default=""),
sa.Column("lipsync_job_id", sa.String(36), nullable=False),
sa.Column("script_id", sa.String(36), nullable=False),
sa.Column("b_roll_segments", sa.JSON(), nullable=False, server_default="[]"),
sa.Column("title_config", sa.JSON(), nullable=False, server_default="{}"),
sa.Column("cover_config", sa.JSON(), nullable=False, server_default="{}"),
sa.Column("status", sa.String(20), nullable=False, server_default="pending", index=True),
sa.Column("progress", sa.Integer(), nullable=False, server_default=sa.text("0")),
sa.Column("output_video_url", sa.Text(), nullable=False, server_default=""),
sa.Column("output_cover_url", sa.Text(), nullable=False, server_default=""),
sa.Column("output_duration", sa.Float(), nullable=False, server_default=sa.text("0.0")),
sa.Column("error_message", sa.Text(), nullable=False, server_default=""),
sa.Column("submitted_at", sa.DateTime(), nullable=True),
sa.Column("started_at", sa.DateTime(), nullable=True),
sa.Column("completed_at", sa.DateTime(), nullable=True),
sa.Column("created_at", sa.DateTime(), nullable=False, server_default=sa.func.now()),
sa.Column("updated_at", sa.DateTime(), nullable=False, server_default=sa.func.now()),
)
op.create_index("ix_ai_avatar_render_user_status", "ai_avatar_render_jobs", ["user_id", "status"])
op.create_index("ix_ai_avatar_render_project_user", "ai_avatar_render_jobs", ["project_id", "user_id"])
def downgrade() -> None:
op.drop_index("ix_ai_avatar_render_project_user", table_name="ai_avatar_render_jobs")
op.drop_index("ix_ai_avatar_render_user_status", table_name="ai_avatar_render_jobs")
op.drop_table("ai_avatar_render_jobs")
@@ -1,45 +0,0 @@
"""lipsync_jobs 增加 TTS 直生字段(voice_id/script_text/speed/emotion
Revision ID: 073_add_lipsync_tts_fields
Revises: 072_add_ai_avatar_render
Create Date: 2026-09-09
"""
import sqlalchemy as sa
from alembic import op
revision = "073_add_lipsync_tts_fields"
down_revision = "072_add_ai_avatar_render"
branch_labels = None
depends_on = None
def upgrade() -> None:
# 对口型支持「传音色 + 文案直接生成」:后端内部先 TTS 合成音频再提交对口型
op.add_column(
"lipsync_jobs",
sa.Column("voice_id", sa.String(200), nullable=False, server_default=""),
)
op.add_column(
"lipsync_jobs",
sa.Column("script_text", sa.Text(), nullable=False, server_default=""),
)
op.add_column(
"lipsync_jobs",
sa.Column("speed", sa.Float(), nullable=False, server_default=sa.text("1.0")),
)
op.add_column(
"lipsync_jobs",
sa.Column("emotion", sa.String(20), nullable=False, server_default=""),
)
# audio_url 改为可空:直生模式下音频由后端 TTS 合成后回填
op.alter_column("lipsync_jobs", "audio_url", existing_type=sa.Text(), nullable=True)
def downgrade() -> None:
op.alter_column("lipsync_jobs", "audio_url", existing_type=sa.Text(), nullable=False)
op.drop_column("lipsync_jobs", "emotion")
op.drop_column("lipsync_jobs", "speed")
op.drop_column("lipsync_jobs", "script_text")
op.drop_column("lipsync_jobs", "voice_id")
@@ -1,36 +0,0 @@
"""ai_avatar_render_jobs.script_id 放宽为可空串(手动文案直生场景不关联文案库)
Revision ID: 074_render_script_id_optional
Revises: 073_add_lipsync_tts_fields
Create Date: 2026-09-09
"""
import sqlalchemy as sa
from alembic import op
revision = "074_render_script_id_optional"
down_revision = "073_add_lipsync_tts_fields"
branch_labels = None
depends_on = None
def upgrade() -> None:
# 列保持 NOT NULL(空串占位),仅应用层允许不传;这里显式补 server_default 防止历史约束歧义
with op.batch_alter_table("ai_avatar_render_jobs") as batch:
batch.alter_column(
"script_id",
existing_type=sa.String(length=36),
nullable=False,
server_default="",
)
def downgrade() -> None:
with op.batch_alter_table("ai_avatar_render_jobs") as batch:
batch.alter_column(
"script_id",
existing_type=sa.String(length=36),
nullable=False,
server_default=None,
)
View File
-24
View File
@@ -1,5 +1,4 @@
from app.api.routes.ai import router as ai_router from app.api.routes.ai import router as ai_router
from app.api.routes.ai_avatar_render import router as ai_avatar_render_router
from app.api.routes.asset_diagnosis import router as asset_diagnosis_router from app.api.routes.asset_diagnosis import router as asset_diagnosis_router
from app.api.routes.asset_libraries import router as asset_libraries_router from app.api.routes.asset_libraries import router as asset_libraries_router
from app.api.routes.assets import router as assets_router from app.api.routes.assets import router as assets_router
@@ -12,13 +11,10 @@ from app.api.routes.feature_flags import router as feature_flags_router
from app.api.routes.generation_cover import router as generation_cover_router from app.api.routes.generation_cover import router as generation_cover_router
from app.api.routes.generation_preview import router as generation_preview_router from app.api.routes.generation_preview import router as generation_preview_router
from app.api.routes.generation_tasks import router as generation_tasks_router from app.api.routes.generation_tasks import router as generation_tasks_router
from app.api.routes.generation_variant_plans import router as generation_variant_plans_router
from app.api.routes.health import router as health_check_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.ingest_jobs import router as ingest_jobs_router
from app.api.routes.internal_render import router as internal_render_router from app.api.routes.internal_render import router as internal_render_router
from app.api.routes.lipsync import router as lipsync_router
from app.api.routes.projects import router as projects_router from app.api.routes.projects import router as projects_router
from app.api.routes.scripts import router as scripts_router
from app.api.routes.share import router as share_router from app.api.routes.share import router as share_router
from app.api.routes.subscription import router as subscription_router from app.api.routes.subscription import router as subscription_router
from app.api.routes.tags import router as tags_router from app.api.routes.tags import router as tags_router
@@ -41,11 +37,6 @@ api_router.include_router(
auth_router, auth_router,
tags=["Auth"], tags=["Auth"],
) )
api_router.include_router(
lipsync_router,
prefix="/lipsync",
tags=["Lipsync"],
)
api_router.include_router( api_router.include_router(
projects_router, projects_router,
prefix="/projects", prefix="/projects",
@@ -108,11 +99,6 @@ api_router.include_router(
prefix="/generation", prefix="/generation",
tags=["Generation"], tags=["Generation"],
) )
api_router.include_router(
generation_variant_plans_router,
prefix="/generation",
tags=["Generation"],
)
api_router.include_router( api_router.include_router(
generation_cover_router, generation_cover_router,
prefix="/generation", prefix="/generation",
@@ -179,13 +165,3 @@ api_router.include_router(
internal_render_router, internal_render_router,
tags=["Internal"], tags=["Internal"],
) )
api_router.include_router(
scripts_router,
prefix="/scripts",
tags=["ScriptLibrary"],
)
api_router.include_router(
ai_avatar_render_router,
prefix="/ai-avatar/render",
tags=["AI Avatar Render"],
)
-218
View File
@@ -1,218 +0,0 @@
"""AI数字人渲染合成 API 路由 — #1798.
接口:
POST /api/v1/ai-avatar/render 提交渲染任务
GET /api/v1/ai-avatar/render/jobs 任务列表
GET /api/v1/ai-avatar/render/{job_id} 任务详情
POST /api/v1/ai-avatar/render/{job_id}/cancel 取消任务
POST /api/v1/ai-avatar/render/{job_id}/retry 重试失败任务
"""
from __future__ import annotations
import logging
from app.auth import AuthenticatedUser, get_current_user
from app.dependencies import get_db_session
from app.schemas.ai_avatar_render import (
AiAvatarRenderJobResponse,
CreateAiAvatarRenderRequest,
SmartCoverRequest,
SmartCoverResponse,
)
from app.services.ai_avatar_cover_service import generate_smart_cover
from app.services.ai_avatar_render_service import (
AiAvatarRenderError,
AiAvatarRenderService,
)
from fastapi import APIRouter, Depends, HTTPException, Query
from sqlalchemy.orm import Session
logger = logging.getLogger(__name__)
router = APIRouter()
def _get_service(db: Session = Depends(get_db_session)) -> AiAvatarRenderService:
return AiAvatarRenderService(db)
# ── POST / — 提交渲染任务 ────────────────────────────────────────────────
@router.post("", response_model=AiAvatarRenderJobResponse, status_code=201)
def create_render_job(
body: CreateAiAvatarRenderRequest,
current_user: AuthenticatedUser = Depends(get_current_user),
svc: AiAvatarRenderService = Depends(_get_service),
):
"""提交 AI 数字人渲染任务.
将对口型视频 + B-roll 素材 + 标题叠加 + 封面提取合成最终输出视频。
"""
try:
job = svc.create_render_job(
user_id=current_user.user.id,
lipsync_job_id=body.lipsync_job_id,
script_id=body.script_id,
b_roll_segments=[s.model_dump() for s in body.b_roll_segments],
title_config=body.title_config,
cover_config=body.cover_config,
project_id=body.project_id,
)
except AiAvatarRenderError as exc:
status_map = {
"LipsyncJobNotFound": 404,
"LipsyncJobNotCompleted": 400,
"LipsyncJobNoOutput": 400,
"ScriptNotFound": 404,
}
raise HTTPException(
status_code=status_map.get(exc.code, 400),
detail={"code": exc.code, "message": str(exc)},
) from exc
# 异步触发渲染
try:
from app.tasks.ai_avatar_render import execute_ai_avatar_render
execute_ai_avatar_render.delay(job.id)
except Exception:
logger.warning("Celery 任务提交失败,渲染任务已创建但未触发执行: %s", job.id)
return job
# ── GET /jobs — 任务列表 ─────────────────────────────────────────────────
@router.get("/jobs", response_model=dict)
def list_render_jobs(
project_id: str = Query("", description="项目 ID 过滤"),
status: str = Query("", description="状态过滤"),
offset: int = Query(0, ge=0),
limit: int = Query(20, ge=1, le=100),
current_user: AuthenticatedUser = Depends(get_current_user),
svc: AiAvatarRenderService = Depends(_get_service),
):
"""获取 AI 数字人渲染任务列表."""
items, total = svc.list_render_jobs(
user_id=current_user.user.id,
project_id=project_id,
status=status,
offset=offset,
limit=limit,
)
return {
"items": [AiAvatarRenderJobResponse.model_validate(j) for j in items],
"total": total,
"offset": offset,
"limit": limit,
}
# ── GET /{job_id} — 任务详情 ─────────────────────────────────────────────
@router.get("/{job_id}", response_model=AiAvatarRenderJobResponse)
def get_render_job(
job_id: str,
current_user: AuthenticatedUser = Depends(get_current_user),
svc: AiAvatarRenderService = Depends(_get_service),
):
"""获取渲染任务详情."""
job = svc.get_render_job(job_id, current_user.user.id)
if job is None:
raise HTTPException(status_code=404, detail="渲染任务不存在")
return job
# ── POST /{job_id}/cancel — 取消任务 ─────────────────────────────────────
@router.post("/{job_id}/cancel", response_model=AiAvatarRenderJobResponse)
def cancel_render_job(
job_id: str,
current_user: AuthenticatedUser = Depends(get_current_user),
svc: AiAvatarRenderService = Depends(_get_service),
):
"""取消渲染任务(仅 pending 状态可取消)."""
job = svc.cancel_render_job(job_id, current_user.user.id)
if job is None:
raise HTTPException(status_code=404, detail="渲染任务不存在")
if job.status != "cancelled":
raise HTTPException(
status_code=400,
detail=f"任务状态 {job.status} 不可取消,仅 pending 可取消",
)
return job
# ── POST /{job_id}/retry — 重试失败任务 ──────────────────────────────────
@router.post("/{job_id}/retry", response_model=AiAvatarRenderJobResponse)
def retry_render_job(
job_id: str,
current_user: AuthenticatedUser = Depends(get_current_user),
svc: AiAvatarRenderService = Depends(_get_service),
):
"""重试失败的渲染任务."""
job = svc.retry_render_job(job_id, current_user.user.id)
if job is None:
raise HTTPException(status_code=404, detail="渲染任务不存在")
if job.status != "pending":
raise HTTPException(
status_code=400,
detail=f"仅 failed 状态的任务可重试,当前状态: {job.status}",
)
# 重新触发渲染
try:
from app.tasks.ai_avatar_render import execute_ai_avatar_render
execute_ai_avatar_render.delay(job.id)
except Exception:
logger.warning("Celery 任务提交失败,重试任务已重置但未触发执行: %s", job.id)
return job
# ── POST /smart-cover — 智能获取封面(MediaKit 抽帧 + 评分选帧)────────
@router.post("/smart-cover", response_model=SmartCoverResponse)
def generate_avatar_smart_cover(
body: SmartCoverRequest,
current_user: AuthenticatedUser = Depends(get_current_user),
) -> SmartCoverResponse:
"""智能获取数字人视频封面.
复用智能剪辑的 MediaKit 抽帧 + 质量评分选最佳帧逻辑(非 FFmpeg 简单截帧),
并将选中帧转存到自家 OSS,返回非临时的封面公网 URL。
前端「智能获取封面」按钮可直接调用本接口;不依赖渲染任务完成。
"""
video_url = (body.video_url or "").strip()
if not video_url.startswith(("http://", "https://")):
raise HTTPException(status_code=400, detail="video_url 必须是合法的 HTTP/HTTPS URL")
try:
cover_url = generate_smart_cover(video_url, max_frames=body.max_frames)
except Exception as exc:
logger.error(
"智能封面生成异常: user=%s video_url=%s err=%s",
current_user.user.id, video_url[:80], exc,
exc_info=True,
)
cover_url = ""
if not cover_url:
return SmartCoverResponse(
cover_url="",
status="fallback_failed",
message="智能抽帧失败(MediaKit 不可用或抽帧异常),请稍后重试",
)
logger.info("智能封面生成成功: user=%s cover_url=%s", current_user.user.id, cover_url[:120])
return SmartCoverResponse(cover_url=cover_url, status="completed")
+24 -5
View File
@@ -20,7 +20,7 @@ from packages.application import (
GetProjectUseCase, GetProjectUseCase,
ListAssetLibrariesUseCase, ListAssetLibrariesUseCase,
) )
from packages.domain import AssetLibraryKind from packages.domain import AssetLibrary, AssetLibraryKind
from ._helpers import check_project_access from ._helpers import check_project_access
@@ -120,11 +120,30 @@ def ensure_default_library(
kind = AssetLibraryKind(request.kind) kind = AssetLibraryKind(request.kind)
# Issue #1775: 幂等获取/创建——依赖唯一约束 uq_asset_libraries_project_kind # 查找该项目下同 kind 的素材库,返回第一个
# 并发创建冲突时回滚重查返回已有记录,不再依赖应用层"先查后插",也不会 500。 existing = asset_library_repository.find_by_project(request.project_id)
for lib in existing:
if lib.kind == kind:
return _to_asset_library_response(lib)
# 不存在 → 自动创建
import uuid
from datetime import datetime, timezone
now = datetime.now(timezone.utc)
default_name = _DEFAULT_LIBRARY_NAMES.get(request.kind, f"{request.kind}素材库") default_name = _DEFAULT_LIBRARY_NAMES.get(request.kind, f"{request.kind}素材库")
library = asset_library_repository.get_or_create_default_library(request.project_id, kind, name=default_name) library = AssetLibrary(
return _to_asset_library_response(library) id=str(uuid.uuid4()),
project_id=request.project_id,
name=default_name,
kind=kind,
asset_count=0,
total_size=0,
created_at=now,
updated_at=now,
)
created = asset_library_repository.create(library)
return _to_asset_library_response(created)
@router.delete("/{library_id}", status_code=status.HTTP_204_NO_CONTENT, response_class=Response) @router.delete("/{library_id}", status_code=status.HTTP_204_NO_CONTENT, response_class=Response)
+35 -49
View File
@@ -290,7 +290,7 @@ def list_assets(
else: else:
total = asset_repository.count_by_project_ids(project_ids, status=status_list) total = asset_repository.count_by_project_ids(project_ids, status=status_list)
# 跨项目分页:逐项目累积直到凑够一页 # 跨项目分页:逐项目累积直到凑够一页
paged_items = [] paged_items: list = []
offset = skip offset = skip
remaining = limit remaining = limit
for pid in project_ids: for pid in project_ids:
@@ -579,45 +579,39 @@ def smart_match_assets(
filtered_assets = asset_repository.find_by_library(request.library_id, status=["ready"], limit=10000) filtered_assets = asset_repository.find_by_library(request.library_id, status=["ready"], limit=10000)
total_candidates = len(filtered_assets) total_candidates = len(filtered_assets)
# ── 过滤前置:余量 + 高频使用,过滤在评分/截取 limit 之前完成 ────────── # 调用统一智能选素材算法(kind 已在 DB 层过滤,无需重复过滤)
# 旧实现先 smart_select_assets(limit=N) 再对这 N 条做过滤,过滤后不回补, results = smart_select_assets(
# 当排名靠前的素材恰好都被排除时返回空 items(前端回退全选,smart-match 名存实亡)。 filtered_assets,
# 现在先过滤全量候选,每级过滤后为空/不足则回退上一级,最后才评分截取。 limit=request.limit,
kind=None,
)
# 1) 余量过滤:usable=False(零重复可切区间耗尽且历史区间均达复用上限)的素材排除 # 结果层过滤:usable=false(零重复可切区间耗尽且历史区间均达复用上限)的素材
usable_assets = [] # 不返回给前端;不动 smart_select_assets 评分逻辑本身
exhausted_assets = [] filtered_results = []
for a in filtered_assets: for r in results:
try: try:
avail = compute_asset_availability(a) avail = compute_asset_availability(r.asset)
except Exception: except Exception:
logger.warning( logger.warning(
"smart-match 余量计算失败,按可用处理: asset_id=%s", "smart-match 余量计算失败,按可用处理: asset_id=%s",
getattr(a, "id", "?"), getattr(r.asset, "id", "?"),
exc_info=True, exc_info=True,
) )
avail = None avail = None
if avail is not None and not avail["usable"]: if avail is not None and not avail["usable"]:
exhausted_assets.append(a) logger.info(
else: "smart-match 排除已用尽素材: asset_id=%s name=%s",
usable_assets.append(a) getattr(r.asset, "id", "?"),
getattr(r.asset, "name", ""),
)
continue
filtered_results.append(r)
if exhausted_assets: # 高频使用排除:同一素材在最近 5 个视频中出现超过 3 次则排除
logger.info(
"smart-match 余量过滤: 候选 %d,可切区间耗尽 %d",
len(filtered_assets), len(exhausted_assets),
)
# 回退策略:余量过滤后为空(全部耗尽)时,保留全部候选,不返回空结果。
# 宁可让用户在已耗尽素材上复用,也比 smart-match 空结果回退全选更可控
# (全选同样会选到这些素材,且不经过评分排序)。
pool = usable_assets if usable_assets else filtered_assets
# 2) 高频使用排除:同一素材在最近 5 个视频中出现超过 3 次则排除
MAX_RECENT_USE_COUNT = 3 MAX_RECENT_USE_COUNT = 3
high_freq_assets = set() if filtered_results:
if pool: asset_ids = [getattr(r.asset, "id", "") for r in filtered_results if getattr(r.asset, "id", "")]
asset_ids = [getattr(a, "id", "") for a in pool if getattr(a, "id", "")]
if asset_ids: if asset_ids:
try: try:
use_counts = get_asset_recent_use_counts( use_counts = get_asset_recent_use_counts(
@@ -625,35 +619,27 @@ def smart_match_assets(
asset_ids=asset_ids, asset_ids=asset_ids,
recent_video_count=5, recent_video_count=5,
) )
for a in pool: high_use_excluded = set()
aid = getattr(a, "id", "") for r in filtered_results:
aid = getattr(r.asset, "id", "")
count = use_counts.get(aid, 0) count = use_counts.get(aid, 0)
if count > MAX_RECENT_USE_COUNT: if count > MAX_RECENT_USE_COUNT:
high_freq_assets.add(aid)
logger.info( logger.info(
"smart-match 排除高频使用素材: asset_id=%s use_count=%d limit=%d", "smart-match 排除高频使用素材: asset_id=%s use_count=%d limit=%d",
aid, count, MAX_RECENT_USE_COUNT, aid, count, MAX_RECENT_USE_COUNT,
) )
# 回退策略:排除后剩余素材不足(为空或不够 limit)时, high_use_excluded.add(id(r))
# 不再全部排除,保留全部可用素材
if high_freq_assets:
remaining_count = len(pool) - len(high_freq_assets)
enough = request.limit is None or remaining_count >= request.limit
if remaining_count > 0 and enough:
pool = [a for a in pool if getattr(a, "id", "") not in high_freq_assets]
else: else:
logger.info( pass
"smart-match 高频排除后素材不足(%d<%s),保留全部 %d", # 如果排除后不够 limit,放宽到不限制
remaining_count, remaining = [r for r in filtered_results if id(r) not in high_use_excluded]
request.limit if request.limit is not None else "不限", if len(remaining) >= request.limit:
len(pool), filtered_results = remaining
) else:
logger.info("smart-match 高频排除后素材不足(%d<%d),保留全部", len(remaining), request.limit)
except Exception: except Exception:
logger.warning("smart-match 高频使用查询失败,跳过排除", exc_info=True) logger.warning("smart-match 高频使用查询失败,跳过排除", exc_info=True)
# 3) 调用统一智能选素材算法(kind 已在 DB 层过滤,无需重复过滤)
results = smart_select_assets(pool, limit=request.limit, kind=None)
# 扁平结构:SmartMatchItem 继承 AssetResponse,素材字段直接在条目顶层, # 扁平结构:SmartMatchItem 继承 AssetResponse,素材字段直接在条目顶层,
# 前端无需解析 item.asset 包装层,item.id / item.usable / 余量字段直接可读 # 前端无需解析 item.asset 包装层,item.id / item.usable / 余量字段直接可读
items = [ items = [
@@ -662,7 +648,7 @@ def smart_match_assets(
score=r.score, score=r.score,
breakdown=r.breakdown, breakdown=r.breakdown,
) )
for r in results for r in filtered_results
] ]
return SmartMatchResponse(items=items, total_candidates=total_candidates) return SmartMatchResponse(items=items, total_candidates=total_candidates)
+3 -187
View File
@@ -1,5 +1,4 @@
""" """
from __future__ import annotations
Canonical authentication API routes. Canonical authentication API routes.
The route layer is intentionally thin: repository construction lives in The route layer is intentionally thin: repository construction lives in
@@ -14,9 +13,9 @@ import jwt
from app.auth import AuthenticatedUser, blacklist_token, get_current_user from app.auth import AuthenticatedUser, blacklist_token, get_current_user
from app.config import settings from app.config import settings
from app.dependencies import get_auth_email_service, get_auth_session_store, get_user_repository from app.dependencies import get_auth_email_service, get_auth_session_store, get_user_repository
from fastapi import APIRouter, Depends, Header, HTTPException, Request, status from fastapi import APIRouter, Depends, Header, HTTPException, status
from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer
from pydantic import BaseModel, EmailStr, field_validator from pydantic import BaseModel, EmailStr
from packages.adapters.redis import NoopSessionStore from packages.adapters.redis import NoopSessionStore
from packages.adapters.smtp import NoopEmailService from packages.adapters.smtp import NoopEmailService
@@ -85,23 +84,6 @@ class CurrentUserResponse(BaseModel):
phone: str = "" phone: str = ""
phone_verified: bool = False phone_verified: bool = False
binding_complete: bool = False binding_complete: bool = False
wechat_bound: bool = False
profile_completed: bool = True
class UserProfileResponse(BaseModel):
"""用户资料负载(PATCH /me、绑定/解绑接口复用;字段与 GET /auth/me 一致,前端 normalizeUser 直接消费)"""
user_id: str
email: str
username: str
display_name: str
email_verified: bool
phone: str = ""
phone_verified: bool = False
binding_complete: bool = False
wechat_bound: bool = False
profile_completed: bool = True
class PasswordResetRequestModel(BaseModel): class PasswordResetRequestModel(BaseModel):
@@ -290,52 +272,9 @@ async def get_current_user_info(
phone=user.phone or "", phone=user.phone or "",
phone_verified=user.phone_verified, phone_verified=user.phone_verified,
binding_complete=binding_complete, binding_complete=binding_complete,
wechat_bound=bool(user.wechat_openid),
profile_completed=user.profile_completed,
) )
class UpdateProfileRequest(BaseModel):
"""更新个人资料请求(当前仅支持昵称)"""
display_name: str
@field_validator("display_name")
@classmethod
def _validate_display_name(cls, v: str) -> str:
name = (v or "").strip()
if not name:
raise ValueError("昵称不能为空白")
if len(name) > 20:
raise ValueError("昵称长度需在 1-20 个字符之间")
return name
class UpdateProfileResponse(BaseModel):
"""更新资料响应:前端 normalizeUser(response.user) 直接消费"""
user: UserProfileResponse
@router.patch("/me", response_model=UpdateProfileResponse)
async def update_current_user_profile(
request: UpdateProfileRequest,
current_user: AuthenticatedUser = Depends(get_current_user),
user_repository: UserRepository = Depends(get_user_repository),
) -> UpdateProfileResponse:
"""更新当前登录用户昵称(微信新用户首次设置昵称后置 profile_completed=True)。"""
user = current_user.user
user.display_name = request.display_name # 已 stripvalidator
if not user.profile_completed:
user.profile_completed = True
user_repository.save(user)
logger.info("[资料更新] 用户 %s 更新昵称,profile_completed=%s", user.id, user.profile_completed)
# 重新读取,确保返回的是持久化后的最新状态
fresh = user_repository.find_by_id(user.id) or user
return UpdateProfileResponse(user=_user_profile(fresh))
class _NoopSessionStore(NoopSessionStore): class _NoopSessionStore(NoopSessionStore):
pass pass
@@ -487,7 +426,6 @@ async def get_wechat_auth_url() -> WechatAuthUrlResponse:
@router.post("/wechat/callback", response_model=WechatLoginResponse) @router.post("/wechat/callback", response_model=WechatLoginResponse)
async def wechat_callback( async def wechat_callback(
request: WechatCallbackRequest, request: WechatCallbackRequest,
http_request: Request,
user_repository: UserRepository = Depends(get_user_repository), user_repository: UserRepository = Depends(get_user_repository),
) -> WechatLoginResponse: ) -> WechatLoginResponse:
"""微信登录回调处理""" """微信登录回调处理"""
@@ -495,30 +433,11 @@ async def wechat_callback(
from packages.application.auth.wechat_sync_use_case import WechatSyncRequest as SyncRequest from packages.application.auth.wechat_sync_use_case import WechatSyncRequest as SyncRequest
from packages.application.auth.wechat_sync_use_case import WechatSyncUseCase from packages.application.auth.wechat_sync_use_case import WechatSyncUseCase
# 回调可观测性:记录 UA(区分微信内置浏览器 MicroMessenger)与 state
# 便于排查"停留 open.weixin.qq.com / 回调失败"类问题(#1718
user_agent = http_request.headers.get("User-Agent", "")
is_wechat_browser = "MicroMessenger" in user_agent
logger.info(
"[微信回调] 收到回调: state=%s code_len=%d UA=%r 微信内置浏览器=%s",
(request.state or "")[:8],
len(request.code or ""),
user_agent[:200],
is_wechat_browser,
)
# 1. 用 code 换微信用户信息 # 1. 用 code 换微信用户信息
oauth_service = get_wechat_oauth_service() oauth_service = get_wechat_oauth_service()
wechat_user, err = oauth_service.handle_callback(request.code, request.state) wechat_user, err = oauth_service.handle_callback(request.code, request.state)
if err: if err:
# state 校验失败 / 微信 errcode 等错误原文已在 service 内 log,这里带上 UA 上下文
logger.warning("[微信回调] 处理失败: err=%s 微信内置浏览器=%s", err, is_wechat_browser)
raise HTTPException(status_code=400, detail=err) raise HTTPException(status_code=400, detail=err)
logger.info(
"[微信回调] state 校验通过,微信用户信息获取成功: openid=%s unionid=%s",
wechat_user.openid[:8] if wechat_user.openid else "",
bool(wechat_user.unionid),
)
# 2. 同步登录/注册(复用 wechat-sync 逻辑) # 2. 同步登录/注册(复用 wechat-sync 逻辑)
use_case = WechatSyncUseCase(user_repository=user_repository) use_case = WechatSyncUseCase(user_repository=user_repository)
@@ -537,7 +456,7 @@ async def wechat_callback(
user = user_repository.find_by_id(response.user_id) user = user_repository.find_by_id(response.user_id)
binding_complete = False binding_complete = False
if user: if user:
binding_complete = bool( binding_complete = (
user.phone_verified and user.email_verified and user.email and "@wechat.local" not in user.email user.phone_verified and user.email_verified and user.email and "@wechat.local" not in user.email
) )
@@ -553,109 +472,6 @@ async def wechat_callback(
) )
# ==================== 微信账号绑定/解绑(已登录用户) ====================
class WechatBindUrlResponse(BaseModel):
auth_url: str
state: str
class WechatBindCompleteRequest(BaseModel):
code: str
state: str = ""
class WechatBindCompleteResponse(BaseModel):
success: bool
user: UserProfileResponse
class WechatUnbindResponse(BaseModel):
success: bool
def _user_profile(user) -> UserProfileResponse:
binding_complete = bool(
user.phone_verified and user.email_verified and user.email and "@wechat.local" not in user.email
)
return UserProfileResponse(
user_id=user.id,
email=user.email,
username=user.username,
display_name=user.display_name,
email_verified=user.email_verified,
phone=user.phone or "",
phone_verified=user.phone_verified,
binding_complete=binding_complete,
wechat_bound=bool(user.wechat_openid),
profile_completed=user.profile_completed,
)
@router.get("/wechat/bind/url", response_model=WechatBindUrlResponse)
async def get_wechat_bind_url(
current_user: AuthenticatedUser = Depends(get_current_user),
) -> WechatBindUrlResponse:
"""获取微信绑定授权链接(已登录用户场景)。state 经 Redis 存储做 CSRF 校验。"""
from packages.application.auth.wechat_oauth_service import get_wechat_oauth_service
oauth_service = get_wechat_oauth_service()
auth_url, state = oauth_service.generate_auth_url()
logger.info("[微信绑定] 用户 %s 请求绑定授权链接", current_user.user.id)
return WechatBindUrlResponse(auth_url=auth_url, state=state)
@router.post("/wechat/bind", response_model=WechatBindCompleteResponse)
async def wechat_bind(
request: WechatBindCompleteRequest,
current_user: AuthenticatedUser = Depends(get_current_user),
user_repository: UserRepository = Depends(get_user_repository),
) -> WechatBindCompleteResponse:
"""微信绑定完成:扫码回调后用 code 换 openid,绑定到当前登录账号(不创建新用户)。"""
from packages.application.auth.wechat_bind_use_case import WechatBindRequest, WechatBindUseCase
from packages.application.auth.wechat_oauth_service import get_wechat_oauth_service
oauth_service = get_wechat_oauth_service()
wechat_user, err = oauth_service.handle_callback(request.code, request.state)
if err:
logger.warning("[微信绑定] 用户 %s 换取微信信息失败: %s", current_user.user.id, err)
raise HTTPException(status_code=400, detail=err)
use_case = WechatBindUseCase(user_repository=user_repository)
result, error, http_status = use_case.bind(
WechatBindRequest(
user_id=current_user.user.id,
openid=wechat_user.openid,
unionid=wechat_user.unionid or "",
)
)
if error:
logger.warning("[微信绑定] 用户 %s 绑定失败: %s", current_user.user.id, error)
raise HTTPException(status_code=http_status, detail=error)
logger.info("[微信绑定] 用户 %s 绑定成功 openid=%s", current_user.user.id, wechat_user.openid[:8])
return WechatBindCompleteResponse(success=True, user=_user_profile(result.user))
@router.delete("/wechat/bind", response_model=WechatUnbindResponse)
async def wechat_unbind(
current_user: AuthenticatedUser = Depends(get_current_user),
user_repository: UserRepository = Depends(get_user_repository),
) -> WechatUnbindResponse:
"""解绑微信:需账号仍有其他登录方式(密码/手机/真实邮箱),否则拒绝。"""
from packages.application.auth.wechat_bind_use_case import WechatUnbindUseCase
use_case = WechatUnbindUseCase(user_repository=user_repository)
result, error, http_status = use_case.unbind(current_user.user.id)
if error:
logger.warning("[微信解绑] 用户 %s 解绑失败: %s", current_user.user.id, error)
raise HTTPException(status_code=http_status, detail=error)
logger.info("[微信解绑] 用户 %s 解绑成功", current_user.user.id)
return WechatUnbindResponse(success=True)
# ==================== 验证码 & 绑定 ==================== # ==================== 验证码 & 绑定 ====================
+3 -5
View File
@@ -14,7 +14,6 @@ from typing import Any
from uuid import uuid4 from uuid import uuid4
from app.api.routes._helpers import require_project_and_library from app.api.routes._helpers import require_project_and_library
from app.api.routes.upload import _persist_celery_task_id
from app.auth import AuthenticatedUser, get_current_user from app.auth import AuthenticatedUser, get_current_user
from app.core.celery_app import celery_app from app.core.celery_app import celery_app
from app.core.storage import OSSStorageService, get_storage_service from app.core.storage import OSSStorageService, get_storage_service
@@ -177,8 +176,8 @@ def _cleanup_expired_uploads() -> int:
meta_file.unlink() meta_file.unlink()
cleaned += 1 cleaned += 1
logger.info(f"Cleaned up expired upload: {upload_id}") logger.info(f"Cleaned up expired upload: {upload_id}")
except Exception: except Exception as e:
logger.exception("Failed to cleanup upload metadata: %s", meta_file) logger.warning(f"Failed to cleanup upload metadata {meta_file}: {e}")
return cleaned return cleaned
@@ -382,8 +381,7 @@ async def complete_chunked_upload(
file_hash=request.file_hash, file_hash=request.file_hash,
) )
) )
celery_result = celery_app.send_task("worker.ingest_asset", args=[job.id]) celery_app.send_task("worker.ingest_asset", args=[job.id])
_persist_celery_task_id(ingest_job_repository, job, getattr(celery_result, "id", ""))
# Update metadata status # Update metadata status
meta["status"] = "completed" meta["status"] = "completed"
-9
View File
@@ -7,7 +7,6 @@ from typing import Any
from uuid import uuid4 from uuid import uuid4
from app.auth import AuthenticatedUser, get_current_user 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.storage import OSSStorageService, get_storage_service
from app.dependencies import get_duplication_repository from app.dependencies import get_duplication_repository
from app.schemas.duplication import ( from app.schemas.duplication import (
@@ -77,8 +76,6 @@ def _to_record_response(record: DuplicationRecord) -> DuplicationRecordResponse:
status=record.status, status=record.status,
duplicate_rate=record.duplicate_rate, duplicate_rate=record.duplicate_rate,
duplicate_count=record.duplicate_count, duplicate_count=record.duplicate_count,
visual_similarity=getattr(record, "visual_similarity", None),
match_count=getattr(record, "match_count", None),
created_at=record.created_at.isoformat(), created_at=record.created_at.isoformat(),
updated_at=record.updated_at.isoformat(), updated_at=record.updated_at.isoformat(),
) )
@@ -93,8 +90,6 @@ def _to_detail_response(record: DuplicationRecord) -> DuplicationDetailResponse:
status=record.status, status=record.status,
duplicate_rate=record.duplicate_rate, duplicate_rate=record.duplicate_rate,
duplicate_count=record.duplicate_count, duplicate_count=record.duplicate_count,
visual_similarity=getattr(record, "visual_similarity", None),
match_count=getattr(record, "match_count", None),
created_at=record.created_at.isoformat(), created_at=record.created_at.isoformat(),
updated_at=record.updated_at.isoformat(), updated_at=record.updated_at.isoformat(),
segments=[ segments=[
@@ -197,8 +192,6 @@ async def upload_for_duplication(
authenticated_user.user.id, authenticated_user.user.id,
) )
celery_app.send_task("worker.process_duplication_check", args=[record.id])
return DuplicationUploadResponse( return DuplicationUploadResponse(
id=record.id, id=record.id,
status=record.status, status=record.status,
@@ -303,8 +296,6 @@ def retry_duplication(
detail=f"查重记录 {record_id} 不存在", detail=f"查重记录 {record_id} 不存在",
) )
celery_app.send_task("worker.process_duplication_check", args=[updated.id])
return DuplicationUploadResponse( return DuplicationUploadResponse(
id=updated.id, id=updated.id,
status=updated.status, status=updated.status,
+7 -87
View File
@@ -75,86 +75,6 @@ class GenerateCoverResponse(BaseModel):
# ── Route ──────────────────────────────────────────────────────────────── # ── Route ────────────────────────────────────────────────────────────────
def _select_best_frame_from_snapshots(
snapshots: list[dict], plan_id: str
) -> str:
"""从 MediaKit 抽帧结果中,通过质量评分选出最佳帧。
降级策略:cv2 不可用或评分失败时,返回第一帧。
Args:
snapshots: MediaKit 返回的帧列表 [{"image_url": str, ...}, ...]
plan_id: 计划 ID(日志用)
Returns:
最佳帧的 image_url,或空字符串
"""
if not snapshots:
return ""
if len(snapshots) == 1:
return snapshots[0].get("image_url") or snapshots[0].get("url") or ""
try:
import tempfile
import httpx
from packages.shared.cover_frame_scorer import score_frames
scored_candidates = []
for snap in snapshots:
url = snap.get("image_url") or snap.get("url") or ""
if not url:
continue
# 下载帧到临时文件进行评分
try:
resp = httpx.get(url, timeout=15, follow_redirects=True)
resp.raise_for_status()
with tempfile.NamedTemporaryFile(suffix=".jpg", delete=False) as tmp:
tmp.write(resp.content)
tmp_path = tmp.name
scored_candidates.append({"image_path": tmp_path, "url": url})
except Exception:
# 下载失败的帧跳过,给默认低分
scored_candidates.append({"image_path": None, "url": url, "score": 0.0})
if not scored_candidates:
return snapshots[0].get("image_url") or snapshots[0].get("url") or ""
scored = score_frames(scored_candidates)
best = scored[0] if scored else None
best_url = best.get("url", "") if best else ""
best_score = best.get("score", 0.0) if best else 0.0
logger.info(
"[封面生成] 帧质量评分完成: plan_id=%s candidates=%d best_score=%.1f",
plan_id,
len(scored_candidates),
best_score,
)
# 清理临时文件
for c in scored_candidates:
path = c.get("image_path")
if path:
try:
from pathlib import Path
Path(path).unlink(missing_ok=True)
except Exception:
pass
return best_url
except Exception:
logger.warning(
"[封面生成] 帧质量评分失败,使用第一帧: plan_id=%s",
plan_id,
exc_info=True,
)
return snapshots[0].get("image_url") or snapshots[0].get("url") or ""
def _persist_cover_frame( def _persist_cover_frame(
frame_url: str, frame_url: str,
plan_id: str, plan_id: str,
@@ -593,7 +513,7 @@ def generate_cover(
if generation_task_id: if generation_task_id:
try: try:
task = gen_task_repo.get(generation_task_id) task = gen_task_repo.get(generation_task_id)
if task and getattr(task, "cover_url", ""): # type: ignore[arg-type] if task and getattr(task, "cover_url", ""):
cover_url_from_task = task.cover_url cover_url_from_task = task.cover_url
logger.info( logger.info(
"[封面生成] 统一管道封面(步骤A-direct): plan_id=%s task_id=%s url=%s", "[封面生成] 统一管道封面(步骤A-direct): plan_id=%s task_id=%s url=%s",
@@ -618,7 +538,7 @@ def generate_cover(
gv_task_id = getattr(gv, "generation_task_id", "") or "" gv_task_id = getattr(gv, "generation_task_id", "") or ""
if gv_task_id: if gv_task_id:
task_a2 = gen_task_repo.get(gv_task_id) task_a2 = gen_task_repo.get(gv_task_id)
if task_a2 and getattr(task_a2, "cover_url", ""): # type: ignore[arg-type] if task_a2 and getattr(task_a2, "cover_url", ""):
cover_url_from_task = task_a2.cover_url cover_url_from_task = task_a2.cover_url
logger.info( logger.info(
"[封面生成] 封面(步骤A2-video-task): plan_id=%s video_id=%s url=%s", "[封面生成] 封面(步骤A2-video-task): plan_id=%s video_id=%s url=%s",
@@ -732,13 +652,13 @@ def generate_cover(
snapshots = mk_client.extract_frames( snapshots = mk_client.extract_frames(
video_url=primary_video_url, video_url=primary_video_url,
strategy="SpecifiedFrames", strategy="SpecifiedFrames",
max_frames=5, # 抽 5 帧,通过质量评分选最佳 max_frames=1,
poll_interval=2.0, poll_interval=2.0,
max_poll_attempts=5, max_poll_attempts=5,
max_retries=0, max_retries=0,
) )
if snapshots: if snapshots:
raw = _select_best_frame_from_snapshots(snapshots, plan_id) raw = snapshots[0].get("image_url") or snapshots[0].get("url") or ""
if raw: if raw:
cover_url_from_task = _persist_cover_frame(raw, plan_id) cover_url_from_task = _persist_cover_frame(raw, plan_id)
logger.info( logger.info(
@@ -795,13 +715,13 @@ def generate_cover(
snapshots = mk_client.extract_frames( snapshots = mk_client.extract_frames(
video_url=src_url, video_url=src_url,
strategy="SpecifiedFrames", strategy="SpecifiedFrames",
max_frames=5, # 抽 5 帧,通过质量评分选最佳 max_frames=1,
poll_interval=2.0, poll_interval=2.0,
max_poll_attempts=5, max_poll_attempts=5,
max_retries=0, max_retries=0,
) )
if snapshots: if snapshots:
raw = _select_best_frame_from_snapshots(snapshots, plan_id) raw = snapshots[0].get("image_url") or snapshots[0].get("url") or ""
if raw: if raw:
cover_url_from_task = _persist_cover_frame( cover_url_from_task = _persist_cover_frame(
raw, raw,
@@ -827,7 +747,7 @@ def generate_cover(
if cover_url_from_task: if cover_url_from_task:
# 标题已在预览视频渲染时烧录(ASS字幕),封面帧自然包含标题 # 标题已在预览视频渲染时烧录(ASS字幕),封面帧自然包含标题
cover_data: dict[str, object] = { # type: ignore[no-redef] cover_data = {
"type": "ai_frame", "type": "ai_frame",
"image_url": cover_url_from_task, "image_url": cover_url_from_task,
"frame_time": 0.0, "frame_time": 0.0,
+118 -316
View File
@@ -14,7 +14,6 @@ from app.core.task_enqueue import (
USER_PENDING_LIMIT, USER_PENDING_LIMIT,
GlobalQueueFull, GlobalQueueFull,
UserPendingLimitExceeded, UserPendingLimitExceeded,
build_rate_limit_detail,
safe_enqueue_generation_task, safe_enqueue_generation_task,
) )
from app.dependencies import ( from app.dependencies import (
@@ -24,7 +23,6 @@ from app.dependencies import (
get_generation_task_repository, get_generation_task_repository,
) )
from app.schemas.generation_task import ( from app.schemas.generation_task import (
BatchPreviewGenerationTaskResponse,
CreatePreviewGenerationTaskRequest, CreatePreviewGenerationTaskRequest,
PreviewGenerationTaskResponse, PreviewGenerationTaskResponse,
) )
@@ -100,7 +98,7 @@ def _resolve_strategy_id_from_template(template_id: str, db: Session, user_id: s
try: try:
new_repo = SQLAlchemyEditTemplateRepository(db) new_repo = SQLAlchemyEditTemplateRepository(db)
new_template = new_repo.get(template_id) new_template = new_repo.get(template_id)
if new_template and getattr(new_template, "editing_mode", ""): # type: ignore[arg-type] if new_template and getattr(new_template, "editing_mode", ""):
mode = new_template.editing_mode.strip() mode = new_template.editing_mode.strip()
if mode: if mode:
logger.info( logger.info(
@@ -195,19 +193,11 @@ def _to_preview_response(task, generated_videos: list | None = None) -> PreviewG
if started_at and completed_at: if started_at and completed_at:
generate_duration = (completed_at - started_at).total_seconds() generate_duration = (completed_at - started_at).total_seconds()
title_cfg = getattr(task, "title_config", None)
title_cfg = title_cfg if isinstance(title_cfg, dict) else {}
extra_meta = getattr(task, "extra_meta", None)
extra_meta = extra_meta if isinstance(extra_meta, dict) else {}
voice_library_id = getattr(task, "voice_library_id", "") or ""
if not isinstance(voice_library_id, str):
voice_library_id = str(voice_library_id) if voice_library_id else ""
return PreviewGenerationTaskResponse( return PreviewGenerationTaskResponse(
task_id=task.id, task_id=task.id,
status=task.status.value if hasattr(task.status, "value") else str(task.status), status=task.status.value if hasattr(task.status, "value") else str(task.status),
progress=float(task.progress or 0.0), progress=float(task.progress or 0.0),
is_preview=bool(getattr(task, "is_preview", True)), is_preview=bool(getattr(task, "is_preview", True)),
variant_index=int(extra_meta.get("variant_index", 0) or 0),
resolution=getattr(task, "resolution", "") or "", resolution=getattr(task, "resolution", "") or "",
video_url=video_url, video_url=video_url,
duration=duration, duration=duration,
@@ -216,8 +206,6 @@ def _to_preview_response(task, generated_videos: list | None = None) -> PreviewG
transition_count=transition_count, transition_count=transition_count,
material_usage=material_usage, material_usage=material_usage,
error_message=task.error_message or "", error_message=task.error_message or "",
title_text=str(title_cfg.get("text", "") or ""),
voice_library_id=voice_library_id,
created_at=task.created_at, created_at=task.created_at,
started_at=started_at, started_at=started_at,
finished_at=completed_at, finished_at=completed_at,
@@ -225,100 +213,50 @@ def _to_preview_response(task, generated_videos: list | None = None) -> PreviewG
) )
def _resolve_preview_edit_plan_id( @router.post("/preview", response_model=PreviewGenerationTaskResponse, status_code=201)
*,
request: CreatePreviewGenerationTaskRequest,
task,
db: Session,
user_id: str,
) -> str:
"""确定任务关联的编辑计划ID:优先前端传入,否则按 template_id+user 兜底查找。"""
if task.source_edit_plan_id:
return task.source_edit_plan_id
if not request.template_id:
return ""
try:
from packages.adapters.sqlalchemy_impl.edit_plan_repository import (
SQLAlchemyEditPlanRepository,
)
_plan_repo = SQLAlchemyEditPlanRepository(db)
_plans = _plan_repo.list_by_template(request.template_id, limit=20)
for _p in _plans:
if (_p.created_by_user_id or "") == user_id:
logger.info(
"[预览生成] 自动关联编辑计划: task_id=%s plan_id=%s",
task.id,
_p.id,
)
return _p.id
except Exception:
logger.warning(
"[预览生成] 查找关联编辑计划失败(不影响主流程): task_id=%s",
task.id,
exc_info=True,
)
return ""
def _variant_value(values: list[str], index: int, fallback: str = "") -> str:
"""从变体数组中取值:长度1=共用,长度>N=按索引,空数组=回退 fallback。"""
if not values:
return fallback
if len(values) == 1:
return values[0]
return values[index] if index < len(values) else fallback
@router.post("/preview", response_model=BatchPreviewGenerationTaskResponse, status_code=201)
def create_preview_generation_task( def create_preview_generation_task(
request: CreatePreviewGenerationTaskRequest, request: CreatePreviewGenerationTaskRequest,
authenticated_user: AuthenticatedUser = Depends(get_current_user), authenticated_user: AuthenticatedUser = Depends(get_current_user),
generation_task_repository=Depends(get_generation_task_repository), generation_task_repository=Depends(get_generation_task_repository),
db: Session = Depends(get_db_session), db: Session = Depends(get_db_session),
asset_repo=Depends(get_asset_repository), asset_repo=Depends(get_asset_repository),
) -> BatchPreviewGenerationTaskResponse: ) -> PreviewGenerationTaskResponse:
"""创建预览生成任务(支持批量) """创建预览生成任务。
preview_count=1 时行为与旧版完全一致(创建 1 个任务); 预览渲染品质与正式生成一致(1080p, CRF 23, medium preset),确认生成时可直接复用预览产物。
preview_count=N 时一次创建 N 个独立变体任务:
- 每个变体克隆独立编辑计划(独立 clips、独立随机素材起点),N 个预览内容互不相同 Args:
- 每个变体拥有独立 task_id / 状态 / 预览视频 URL,前端按 task_id 分别轮询 request: 预览任务创建请求(template_id + asset_ids 等)
- 标题样式(font/color/position 等)全局共用;标题文字/配音/封面可按变体独立
titles[] / voice_library_ids[] / cover_urls[],长度1=共用,长度N=独立)
Returns: Returns:
201 + 变体任务数组 {items: [...], total: N} 201 + 预览任务详情
""" """
user_id = authenticated_user.user.id user_id = authenticated_user.user.id
count = max(1, request.preview_count)
logger.info( logger.info(
"[预览生成] 接收请求: user_id=%s, template_id=%s, asset_count=%d, preview_count=%d", "[预览生成] 接收请求: user_id=%s, template_id=%s, asset_count=%d, preview_count=%d",
user_id, user_id,
request.template_id, request.template_id,
len(request.asset_ids), len(request.asset_ids),
count, request.preview_count,
) )
# 预检查队列限流(按变体总数计) # 预检查队列限流
try: try:
user_pending = generation_task_repository.count_pending_by_user(user_id) user_pending = generation_task_repository.count_pending_by_user(user_id)
global_pending = generation_task_repository.count_pending_total() global_pending = generation_task_repository.count_pending_total()
if user_pending + count > USER_PENDING_LIMIT: if user_pending + 1 > USER_PENDING_LIMIT:
raise UserPendingLimitExceeded( raise UserPendingLimitExceeded(user_id=user_id, pending_count=user_pending + 1, limit=USER_PENDING_LIMIT)
user_id=user_id, pending_count=user_pending + count, limit=USER_PENDING_LIMIT if global_pending + 1 > GLOBAL_PENDING_LIMIT:
) raise GlobalQueueFull(pending_count=global_pending + 1, limit=GLOBAL_PENDING_LIMIT)
if global_pending + count > GLOBAL_PENDING_LIMIT:
raise GlobalQueueFull(pending_count=global_pending + count, limit=GLOBAL_PENDING_LIMIT)
except UserPendingLimitExceeded as e: except UserPendingLimitExceeded as e:
raise HTTPException( raise HTTPException(
status_code=429, status_code=429,
detail=build_rate_limit_detail(e, generation_task_repository, scope="user"), detail=f"您的待处理任务过多(当前 {e.pending_count - 1}/{e.limit}),请等待后再提交",
) from e ) from e
except GlobalQueueFull as e: except GlobalQueueFull as e:
raise HTTPException( raise HTTPException(
status_code=503, status_code=503,
detail=build_rate_limit_detail(e, generation_task_repository, scope="global"), detail="系统繁忙,请稍后再试",
) from e ) from e
# 确定视频比例:优先前端传入,否则从模板 mode 推断 # 确定视频比例:优先前端传入,否则从模板 mode 推断
@@ -335,11 +273,14 @@ def create_preview_generation_task(
w, h = int(parts[0]), int(parts[1]) w, h = int(parts[0]), int(parts[1])
base = 1920 base = 1920
if w < h: if w < h:
# 竖屏
output_width = round(base * w / h) output_width = round(base * w / h)
output_height = base output_height = base
else: else:
# 横屏
output_width = base output_width = base
output_height = round(base * h / w) output_height = round(base * h / w)
# 对齐到偶数
output_width = output_width - output_width % 2 output_width = output_width - output_width % 2
output_height = output_height - output_height % 2 output_height = output_height - output_height % 2
except (ValueError, ZeroDivisionError): except (ValueError, ZeroDivisionError):
@@ -348,70 +289,42 @@ def create_preview_generation_task(
logger.info( logger.info(
"[预览生成] 分辨率: video_ratio=%s%s (%dx%d)", "[预览生成] 分辨率: video_ratio=%s%s (%dx%d)",
video_ratio, video_ratio, resolution, output_width, output_height,
resolution,
output_width,
output_height,
) )
# 从模板读取 editing_mode / mode 作为 strategy_id(渲染 pipeline 的 mode 参数)
strategy_id = _resolve_strategy_id_from_template(request.template_id, db, user_id) strategy_id = _resolve_strategy_id_from_template(request.template_id, db, user_id)
base_title_config = request.title_config or {}
title_config = request.title_config or {}
use_case = CreateGenerationTaskUseCase(generation_task_repository) use_case = CreateGenerationTaskUseCase(generation_task_repository)
# ── 预创建第一个任务,仅用于解析源编辑计划(不落库为最终任务)──
# 先创建一个临时任务拿到 task 对象上下文,实际 N 个任务在循环中统一创建;
# 为保持与旧版一致的源 plan 解析逻辑,先创建任务0、解析源 plan,
# 再预克隆 N 个变体 plan,最后重建任务关联。
# 简化实现:直接创建全部任务,plan 关联在创建后、入队前完成。
created_tasks: list = []
variant_plan_ids: list[str] = [] # 每个变体最终关联的 plan_id(按变体顺序)
try: try:
for variant_index in range(count): task = use_case.execute(
# 变体独立标题文字:titles[] 覆盖 title_config.text CreateGenerationTaskCommand(
variant_title_text = _variant_value(request.titles, variant_index, "") project_id="",
variant_title_config = dict(base_title_config) asset_library_id="",
if variant_title_text.strip(): strategy_id=strategy_id,
variant_title_config["text"] = variant_title_text.strip() voice_library_id=request.voice_library_id,
template_id=request.template_id,
# 变体独立配音 asset_ids=list(request.asset_ids),
variant_voice_library_id = _variant_value( title_ids=list(request.title_ids),
request.voice_library_ids, variant_index, request.voice_library_id voice_ids=list(request.voice_ids),
created_by_user_id=user_id,
source_edit_plan_id=request.source_edit_plan_id,
asset_select_mode="",
batch_id="",
video_title=request.video_title,
resolution=resolution,
bgm_config=request.bgm_config or {},
auto_retry_enabled=False,
auto_retry_max=0,
is_preview=True,
title_config=title_config,
output_width=output_width,
output_height=output_height,
) )
)
task = use_case.execute(
CreateGenerationTaskCommand(
project_id="",
asset_library_id="",
strategy_id=strategy_id,
voice_library_id=variant_voice_library_id,
template_id=request.template_id,
asset_ids=list(request.asset_ids),
title_ids=list(request.title_ids),
created_by_user_id=user_id,
source_edit_plan_id=request.source_edit_plan_id,
asset_select_mode="",
batch_id="",
video_title=request.video_title,
resolution=resolution,
bgm_config=request.bgm_config or {},
auto_retry_enabled=False,
auto_retry_max=0,
is_preview=True,
title_config=variant_title_config,
output_width=output_width,
output_height=output_height,
)
)
task.extra_meta["variant_index"] = variant_index
# 解析源编辑计划(前端传入或按模板兜底查找)
source_plan_id = _resolve_preview_edit_plan_id(request=request, task=task, db=db, user_id=user_id)
task.source_edit_plan_id = source_plan_id
generation_task_repository.update(task)
created_tasks.append(task)
except ValueError as e: except ValueError as e:
logger.warning("[预览生成] 创建失败: %s", e) logger.warning("[预览生成] 创建失败: %s", e)
raise HTTPException(status_code=400, detail=str(e)) from e raise HTTPException(status_code=400, detail=str(e)) from e
@@ -419,204 +332,93 @@ def create_preview_generation_task(
logger.error("[预览生成] 创建失败: %s", e, exc_info=True) logger.error("[预览生成] 创建失败: %s", e, exc_info=True)
raise HTTPException(status_code=500, detail="创建预览生成任务失败,请稍后再试") from e raise HTTPException(status_code=500, detail="创建预览生成任务失败,请稍后再试") from e
# ── 独立变体 plan(#1743)── # 关联编辑计划:如果前端未传 source_edit_plan_id,通过 template_id + user_id 查找
# count=1:克隆源 plan(预览不污染源 plan,仅起点重算),行为与旧版一致; if not task.source_edit_plan_id and request.template_id:
# count>1:变体 0 保留源 plan,变体 1..N-1 用 reselect_plan_for_variant 完整
# 重跑单视频选片(素材洗牌+镜头洗牌+起点随机+跨变体避让+批次 20% 重叠重选),
# 所见即所得——预览变体差异即正式成片差异。
source_plan_id = created_tasks[0].source_edit_plan_id if created_tasks else ""
# #1749:各变体配音解析(严格守卫已在 schema;此处取每变体 voice 查时长)+ 时长分配
def _preview_voice_durations() -> list[float]:
try: try:
from packages.domain.variant_voice_resolver import resolve_variant_voice_ids from packages.adapters.sqlalchemy_impl.edit_plan_repository import (
SQLAlchemyEditPlanRepository,
voices = resolve_variant_voice_ids(
count=count,
voice_library_id=request.voice_library_id,
voice_library_ids=request.voice_library_ids or None,
) )
_plan_repo = SQLAlchemyEditPlanRepository(db)
_plans = _plan_repo.list_by_template(request.template_id, limit=20)
for _p in _plans:
if (_p.created_by_user_id or "") == user_id:
task.source_edit_plan_id = _p.id
generation_task_repository.update(task)
logger.info(
"[预览生成] 自动关联编辑计划: task_id=%s plan_id=%s",
task.id,
_p.id,
)
break
except Exception: except Exception:
logger.warning("[预览生成] 配音解析失败(按无配音处理)", exc_info=True) logger.warning(
return [0.0] * count "[预览生成] 查找关联编辑计划失败(不影响主流程): task_id=%s",
try: task.id,
from app.api.routes.generation_tasks import _query_voice_durations exc_info=True,
)
return _query_voice_durations(db, voices) # 每条预览都关联独立克隆 plan:多预览前端为 N 次并发调用,若共用同一 plan
except Exception: # 则 N 条预览片段完全相同;克隆时片段起点按持久化历史区间重算(含受控复用),
return [0.0] * count # 保证各预览版本内容不同
if task.source_edit_plan_id:
voice_durations = _preview_voice_durations()
if source_plan_id and count == 1:
# 单预览:克隆一份(原逻辑)+ 配音时长分配
try: try:
from app.services.edit_plan_service import EditPlanService from app.services.edit_plan_service import EditPlanService
_plan_svc = EditPlanService(db) _plan_svc = EditPlanService(db)
variant_plan = _plan_svc.clone_plan_for_variant( _preview_plan = _plan_svc.clone_plan_for_variant(
source_plan_id, task.source_edit_plan_id,
created_by_user_id=user_id, created_by_user_id=user_id,
name_suffix="预览变体", name_suffix="预览变体",
) )
if voice_durations and voice_durations[0] > 0: task.source_edit_plan_id = _preview_plan.id
try:
_plan_svc.apply_voice_duration_to_plan(variant_plan.id, voice_durations[0])
except Exception:
logger.exception("[预览生成] 变体0 配音分配失败(不阻断): plan=%s", variant_plan.id)
variant_plan_ids.append(variant_plan.id)
except Exception as e:
logger.error("[预览生成] 克隆预览 plan 异常: %s", e, exc_info=True)
for t in created_tasks:
_mark_task_failed(generation_task_repository, t, "预览计划创建失败")
raise HTTPException(
status_code=500,
detail="创建预览任务失败:无法生成独立剪辑计划,请重试",
) from e
elif source_plan_id and count > 1:
try:
from app.services.edit_plan_service import EditPlanService
_plan_svc = EditPlanService(db)
# #1749:变体 0 也 clone(不污染源 plan+ 配音分配;变体 1..N-1 独立选片
_plan0 = _plan_svc.clone_plan_for_variant(
source_plan_id,
created_by_user_id=user_id,
name_suffix="预览变体1",
)
if voice_durations and voice_durations[0] > 0:
try:
_plan_svc.apply_voice_duration_to_plan(_plan0.id, voice_durations[0])
except Exception:
logger.exception("[预览生成] 变体0 配音分配失败(不阻断): plan=%s", _plan0.id)
variant_plan_ids.append(_plan0.id)
batch_asset_pool = list(dict.fromkeys(request.asset_ids or []))
for variant_index in range(1, count):
last_err: Exception | None = None
variant_plan = None
for _attempt in range(2): # 1 次重试,抗 DB 瞬时抖动
try:
variant_plan = _plan_svc.reselect_plan_for_variant(
source_plan_id,
batch_asset_pool,
created_by_user_id=user_id,
name_suffix=f"预览变体{variant_index + 1}",
voice_duration=(
voice_durations[variant_index] if variant_index < len(voice_durations) else 0.0
),
)
break
except ValueError as ve:
logger.warning("[预览生成] 变体独立选片失败(素材不足): %s", ve)
for t in created_tasks:
_mark_task_failed(generation_task_repository, t, "预览变体选片失败")
raise HTTPException(
status_code=400,
detail=f"批量预览第 {variant_index + 1} 个视频无法独立选片:{ve}"
"请增加素材库中的视频素材后重试。",
) from ve
except Exception as reselection_err: # noqa: PERF203
last_err = reselection_err
logger.warning(
"[预览生成] 变体独立选片失败(尝试%d/2): variant=%d error=%s",
_attempt + 1,
variant_index,
reselection_err,
exc_info=True,
)
if variant_plan is None:
logger.error(
"[预览生成] 变体独立选片重试仍失败: variant=%d source=%s",
variant_index,
source_plan_id,
exc_info=last_err,
)
for t in created_tasks:
_mark_task_failed(generation_task_repository, t, "预览变体计划创建失败")
raise HTTPException(
status_code=500,
detail="创建预览任务失败:无法生成独立剪辑计划,请重试",
) from last_err
variant_plan_ids.append(variant_plan.id)
except HTTPException:
raise
except Exception as e:
logger.error("[预览生成] 变体 plan 生成异常: %s", e, exc_info=True)
for t in created_tasks:
_mark_task_failed(generation_task_repository, t, "预览变体计划创建失败")
raise HTTPException(
status_code=500,
detail="创建预览任务失败:无法生成独立剪辑计划,请重试",
) from e
# 关联变体 plan 并回写标题配置
for variant_index, task in enumerate(created_tasks):
if variant_plan_ids:
task.source_edit_plan_id = variant_plan_ids[variant_index]
generation_task_repository.update(task) generation_task_repository.update(task)
# 回写变体标题到 plan configworker 渲染时从 plan 读取 title 配置) logger.info(
if task.source_edit_plan_id and (task.title_config or {}).get("text", "").strip(): "[预览生成] 预览关联独立克隆 plan: task_id=%s clone_plan_id=%s",
try: task.id,
from app.api.routes.generation_tasks import _writeback_edit_plan_config _preview_plan.id,
_writeback_edit_plan_config(
plan_id=task.source_edit_plan_id,
task_id=task.id,
title_config=task.title_config,
db=db,
)
except Exception:
logger.warning(
"[预览生成] 回写标题配置失败(不影响主流程): task_id=%s",
task.id,
exc_info=True,
)
# ── 入队 ──
responses: list[PreviewGenerationTaskResponse] = []
rate_limit_exc: Exception | None = None # 记录首个限流异常,全部失败时返回结构化提示
for variant_index, task in enumerate(created_tasks):
try:
enqueued = safe_enqueue_generation_task(
task,
generation_task_repository,
user_id=user_id,
log_prefix=f"[预览生成][变体{variant_index + 1}]",
log_task_status=True,
) )
if not enqueued: except Exception as clone_err:
logger.warning("[预览生成] 任务入队失败: task_id=%s", task.id) # 不退回共用原 plan(否则多条预览内容相同,违反去重诉求):
_mark_task_failed(generation_task_repository, task, "任务入队失败") # 标记任务失败并中断,前端可重新发起预览
except UserPendingLimitExceeded as e: logger.error(
_mark_task_failed(generation_task_repository, task, "待处理任务超限") "[预览生成] 克隆预览变体 plan 失败,任务标记失败: task_id=%s error=%s",
rate_limit_exc = rate_limit_exc or e task.id,
except GlobalQueueFull as e: clone_err,
_mark_task_failed(generation_task_repository, task, "系统队列已满") exc_info=True,
rate_limit_exc = rate_limit_exc or e )
except Exception: _mark_task_failed(generation_task_repository, task, "预览变体计划创建失败")
logger.exception("[预览生成] 入队异常: task_id=%s", task.id)
_mark_task_failed(generation_task_repository, task, "任务入队异常")
# enqueue 会原地更新 task 状态/进度,直接用 task 构造响应
responses.append(_to_preview_response(task))
# 队列满/限流时若全部失败,返回结构化错误码(前端区分"排队"与"创建失败"
if all(r.status == "failed" for r in responses) and rate_limit_exc is not None:
if isinstance(rate_limit_exc, UserPendingLimitExceeded):
raise HTTPException( raise HTTPException(
status_code=429, status_code=500,
detail=build_rate_limit_detail(rate_limit_exc, generation_task_repository, scope="user"), detail="创建预览任务失败:无法生成独立剪辑计划,请重试",
) ) from clone_err
# 入队执行;若入队失败则标记任务为 failed 避免僵尸数据
try:
if not safe_enqueue_generation_task(
task,
generation_task_repository,
user_id=user_id,
log_prefix="[预览生成]",
log_task_status=True,
):
logger.warning("[预览生成] 任务入队失败: task_id=%s", task.id)
_mark_task_failed(generation_task_repository, task, "任务入队失败")
raise HTTPException(status_code=500, detail="任务入队失败,请稍后重试")
except UserPendingLimitExceeded as e:
_mark_task_failed(generation_task_repository, task, "待处理任务超限")
raise HTTPException(
status_code=429,
detail=f"您的待处理任务过多(当前 {e.pending_count - 1}/{e.limit}),请等待后再提交",
) from None
except GlobalQueueFull:
_mark_task_failed(generation_task_repository, task, "系统队列已满")
raise HTTPException( raise HTTPException(
status_code=503, status_code=503,
detail=build_rate_limit_detail(rate_limit_exc, generation_task_repository, scope="global"), detail="系统繁忙,请稍后再试",
) ) from None
logger.info( return _to_preview_response(task)
"[预览生成] 创建完成: %d 个变体任务, task_ids=%s",
len(responses),
[r.task_id for r in responses],
)
return BatchPreviewGenerationTaskResponse(items=responses, total=len(responses))
@router.get("/preview/{task_id}", response_model=PreviewGenerationTaskResponse) @router.get("/preview/{task_id}", response_model=PreviewGenerationTaskResponse)
+69 -287
View File
@@ -10,7 +10,6 @@ from app.core.task_enqueue import (
USER_PENDING_LIMIT, USER_PENDING_LIMIT,
GlobalQueueFull, GlobalQueueFull,
UserPendingLimitExceeded, UserPendingLimitExceeded,
build_rate_limit_detail,
safe_enqueue_generation_task, safe_enqueue_generation_task,
) )
from app.dependencies import ( from app.dependencies import (
@@ -48,39 +47,6 @@ logger = logging.getLogger(__name__)
router = APIRouter() router = APIRouter()
def _variant_value(values: list[str], index: int, fallback: str = "") -> str:
"""从变体数组中取值:长度1=共用,长度>N=按索引,空数组=回退 fallback。"""
if not values:
return fallback
if len(values) == 1:
return values[0]
return values[index] if index < len(values) else fallback
def _query_voice_durations(db: Session, voice_ids: list[str]) -> list[float]:
"""批量查询配音素材时长(秒),#1749 配音时长分配用。
逐项 try/float 硬化:MagicMock/异常/缺失 → 0.0(无配音不分配,不阻断)。
"""
ids = [v for v in dict.fromkeys(voice_ids or []) if v]
if not ids:
return []
try:
from packages.adapters.sqlalchemy_impl.models import AssetModel
rows = db.query(AssetModel.id, AssetModel.duration).filter(AssetModel.id.in_(ids)).all()
dur_map: dict[str, float] = {}
for row in rows:
try:
dur_map[row[0]] = float(row[1] or 0.0)
except (TypeError, ValueError):
dur_map[row[0]] = 0.0
return [dur_map.get(v, 0.0) for v in ids]
except Exception:
logger.warning("[生成任务] 配音时长查询失败(按无配音处理,不阻断)", exc_info=True)
return [0.0 for _ in ids]
def _to_generation_task_response(task) -> GenerationTaskResponse: def _to_generation_task_response(task) -> GenerationTaskResponse:
return GenerationTaskResponse( return GenerationTaskResponse(
id=task.id, id=task.id,
@@ -126,9 +92,6 @@ def _to_generated_video_response(item, download_url: str | None = None) -> Gener
height=item.height, height=item.height,
fps=item.fps, fps=item.fps,
download_url=download_url, download_url=download_url,
duplicate_rate=getattr(item, "duplicate_rate", None),
visual_similarity=getattr(item, "visual_similarity", None),
match_count=getattr(item, "match_count", None),
) )
@@ -147,7 +110,6 @@ def _select_assets_from_library(
assets: list, assets: list,
mode: str, mode: str,
count: int, count: int,
rng=None,
) -> list[str]: ) -> list[str]:
"""根据选取模式从素材库中选取 ready 状态的视频素材 ID。 """根据选取模式从素材库中选取 ready 状态的视频素材 ID。
@@ -155,8 +117,6 @@ def _select_assets_from_library(
assets: 素材库中所有素材(Asset 实体列表) assets: 素材库中所有素材(Asset 实体列表)
mode: 选取模式 — all=全部, smart=智能匹配(多维度评分+多样性) mode: 选取模式 — all=全部, smart=智能匹配(多维度评分+多样性)
count: 选取数量,0 表示全部(仅 smart 模式有效) count: 选取数量,0 表示全部(仅 smart 模式有效)
rng: 可选随机源(smart 模式排序噪声用),生产环境不传则内部随机;
测试可注入固定种子或零噪声随机源获得确定性结果。
Returns: Returns:
选中的素材 ID 列表 选中的素材 ID 列表
@@ -169,15 +129,15 @@ def _select_assets_from_library(
if mode == "smart": if mode == "smart":
# 智能匹配:统一使用 packages/domain/smart_match.py 的多维评分+多样性选取 # 智能匹配:统一使用 packages/domain/smart_match.py 的多维评分+多样性选取
# 评分维度:质量分(40%) + 时长适配(30%) + 新鲜度(20%) + 未使用加分(10%) # 评分维度:质量分(40%) + 时长适配(30%) + 新鲜度(20%) + 未使用加分(10%)
# 排序注入随机噪声(#1743):同分素材每次选出不同组合,从素材组合层面降重
limit = count if count > 0 else None limit = count if count > 0 else None
results = smart_select_assets(ready_video_assets, limit=limit, kind="video", rng=rng) results = smart_select_assets(ready_video_assets, limit=limit, kind="video")
return [r.asset.id for r in results] return [r.asset.id for r in results]
# 默认 all 模式:返回全部 ready 视频素材 # 默认 all 模式:返回全部 ready 视频素材
return [a.id for a in ready_video_assets] return [a.id for a in ready_video_assets]
def _writeback_edit_plan_config( def _writeback_edit_plan_config(
plan_id: str, plan_id: str,
task_id: str, task_id: str,
@@ -202,7 +162,7 @@ def _writeback_edit_plan_config(
current_config = plan_model.config if isinstance(plan_model.config, dict) else {} current_config = plan_model.config if isinstance(plan_model.config, dict) else {}
merged = dict(current_config) merged = dict(current_config)
merged["generation_task_id"] = task_id merged["generation_task_id"] = task_id
# 检查标题是否发生变化,如果变化则清除 cover 字段强制重新生成封面 # 检查标题是否发生变化,如果变化则清除 cover 字段强制重新生成封面
if title_config: if title_config:
old_title_config = merged.get("title_config", {}) or {} old_title_config = merged.get("title_config", {}) or {}
@@ -214,12 +174,10 @@ def _writeback_edit_plan_config(
del merged["cover"] del merged["cover"]
logger.info( logger.info(
"[生成任务] 标题变化,清除旧封面: plan_id=%s old_title=%s new_title=%s", "[生成任务] 标题变化,清除旧封面: plan_id=%s old_title=%s new_title=%s",
plan_id, plan_id, old_title_text, new_title_text,
old_title_text,
new_title_text,
) )
merged["title_config"] = title_config merged["title_config"] = title_config
plan_model.config = merged plan_model.config = merged
db.commit() db.commit()
logger.info( logger.info(
@@ -427,7 +385,7 @@ def create_generation_task(
use_case = CreateGenerationTaskUseCase(generation_task_repository) use_case = CreateGenerationTaskUseCase(generation_task_repository)
count = request.count count = request.count
created_tasks: list = [] created_tasks = []
failed_tasks = [] failed_tasks = []
user_id = authenticated_user.user.id user_id = authenticated_user.user.id
# 同批次任务共享 batch_id,用于视频查重时批次内比对 # 同批次任务共享 batch_id,用于视频查重时批次内比对
@@ -446,12 +404,12 @@ def create_generation_task(
except UserPendingLimitExceeded as e: except UserPendingLimitExceeded as e:
raise HTTPException( raise HTTPException(
status_code=429, status_code=429,
detail=build_rate_limit_detail(e, generation_task_repository, scope="user"), detail=f"您的待处理任务过多(当前 {e.pending_count - count}/{e.limit},本次提交 {count} 个),请等待完成后再提交",
) from e ) from e
except GlobalQueueFull as e: except GlobalQueueFull as e:
raise HTTPException( raise HTTPException(
status_code=503, status_code=503,
detail=build_rate_limit_detail(e, generation_task_repository, scope="global"), detail="系统繁忙,请稍后再试",
) from e ) from e
# 画中画已下线:strategy_id 中的 pip/voice_pip 统一映射为 one_take # 画中画已下线:strategy_id 中的 pip/voice_pip 统一映射为 one_take
@@ -460,218 +418,66 @@ def create_generation_task(
logger.info("画中画已下线,strategy_id %s → one_take", effective_strategy_id) logger.info("画中画已下线,strategy_id %s → one_take", effective_strategy_id)
effective_strategy_id = "one_take" effective_strategy_id = "one_take"
# 批量生成(count>1):每个变体必须走与单视频完全相同的独立选片流程(#1743/#1749)。 # 批量生成时每个任务关联独立克隆 plan(片段起点重算),
# - 变体 0clone 源 plan(不污染源 plan),变体 1..N-1 用 reselect_plan_for_variant # 禁止 N 条任务共用同一 source_edit_plan_id 导致片段一模一样。
# 完整重跑选片(素材级去重:fresh 优先 → 受控复用 overlap≤20% → 短素材禁复用); # 在创建任何任务【之前】预克隆全部变体:克隆失败直接中断(此时无脏数据),
# - #1749:前端可回传 variant-plans 接口预生成的 plan_idvariant_plan_ids),直接复用; # 绝不静默退回共用源 plan(否则批量视频内容重复,违反去重诉求)。
# 回传 plan 仍按各变体配音幂等重分配段长(防 variant-plans 阶段未带配音/占位时长);
# - 配音时长:独立配音各自时长、统一配音同值,逐变体 apply_voice_duration_to_plan
# 成片总时长=配音时长(素材短→末帧冻结,禁慢放/禁截配音);
# - count>1 但没有源 plan 时,不允许 N 个任务兜底共用同一 plan,直接 4xx 中断。
# 在创建任何任务【之前】预生成/校验全部变体 plan:失败直接中断(此时无脏数据)。
variant_plan_ids: list[str] = [] variant_plan_ids: list[str] = []
if count > 1: if count > 1 and request.source_edit_plan_id:
from app.services.edit_plan_service import EditPlanService from app.services.edit_plan_service import EditPlanService
from packages.domain.variant_voice_resolver import VariantVoiceError, resolve_variant_voice_ids
_plan_svc = EditPlanService(db) _plan_svc = EditPlanService(db)
for task_index in range(1, count):
# 解析每变体配音(严格守卫:独立配音长度/缺值 → 400,禁静默 fallback variant = None
try: last_err: Exception | None = None
variant_voices = resolve_variant_voice_ids( for _attempt in range(2): # 1 次重试,抗 DB 瞬时抖动
count=count, try:
voice_library_id=request.voice_library_id, variant = _plan_svc.clone_plan_for_variant(
voice_library_ids=request.voice_library_ids or None, request.source_edit_plan_id,
) created_by_user_id=user_id,
except VariantVoiceError as ve: name_suffix=f"批量{task_index + 1}",
raise HTTPException(status_code=400, detail=str(ve)) from ve
# 各变体配音时长(查询硬化:异常 → 0.0 不阻断)
voice_durations = _query_voice_durations(db, variant_voices)
# 解析批量源 plan:优先前端传入;否则按 template_id + user 查最新(与单任务兜底同源)
batch_source_plan_id = request.source_edit_plan_id
if not batch_source_plan_id and request.template_id:
try:
from packages.adapters.sqlalchemy_impl.models import EditPlanModel
_latest = (
db.query(EditPlanModel)
.filter(
EditPlanModel.template_id == request.template_id,
EditPlanModel.created_by_user_id == user_id,
) )
.order_by(EditPlanModel.created_at.desc()) break
.first() except Exception as clone_err: # noqa: PERF203
last_err = clone_err
logger.warning(
"[生成任务] 克隆变体 plan 失败(尝试%d/2): source=%s error=%s",
_attempt + 1,
request.source_edit_plan_id,
clone_err,
exc_info=True,
)
if variant is None:
logger.error(
"[生成任务] 克隆变体 plan 重试仍失败,中断批量创建: source=%s",
request.source_edit_plan_id,
exc_info=last_err,
) )
if _latest:
batch_source_plan_id = _latest.id
except Exception:
logger.warning("[生成任务] 批量源 plan 解析失败", exc_info=True)
if not batch_source_plan_id and not request.variant_plan_ids:
# 无任何可用源 plan:批量变体无从选片,明确报错,严禁静默共用/同源
logger.error("[生成任务] 批量 count=%d 但无可编辑计划(无 source_edit_plan_id/template plan", count)
raise HTTPException(
status_code=400,
detail="批量生成需要先完成预览生成(缺少剪辑计划)。请先生成预览后再批量创建。",
)
# 批次素材池:请求显式素材 + 库自动匹配素材(resolved_asset_ids
batch_asset_pool = list(dict.fromkeys(resolved_asset_ids or []))
if request.variant_plan_ids:
# ① 前端回传 variant-plans 预生成结果:直接复用(轻量选片接口已建好 plan)
if len(request.variant_plan_ids) != count:
raise HTTPException( raise HTTPException(
status_code=400, status_code=500,
detail=f"variant_plan_ids 数量({len(request.variant_plan_ids)})与视频数量({count})不一致", detail="创建批量任务失败:无法生成独立剪辑计划,请重试",
) ) from last_err
# 校验归属权 variant_plan_ids.append(variant.id)
from packages.adapters.sqlalchemy_impl.models import EditPlanModel
for _pid in request.variant_plan_ids:
_pm = db.query(EditPlanModel).filter(EditPlanModel.id == _pid).first()
if _pm is None:
raise HTTPException(status_code=400, detail=f"剪辑计划不存在: {_pid}")
if _pm.created_by_user_id and _pm.created_by_user_id != user_id:
raise HTTPException(status_code=403, detail=f"无权使用剪辑计划: {_pid}")
variant_plan_ids = list(request.variant_plan_ids)
else:
# ② 服务端选片:变体 0 clone 源 plan(不污染源 plan
try:
_plan0 = _plan_svc.clone_plan_for_variant(
batch_source_plan_id,
created_by_user_id=user_id,
name_suffix="批量1",
)
except Exception as clone_err:
logger.error("[生成任务] 变体0 clone 失败: %s", clone_err, exc_info=True)
raise HTTPException(
status_code=500, detail="创建批量任务失败:无法生成独立剪辑计划,请重试"
) from clone_err
variant_plan_ids.append(_plan0.id)
# 变体 1..N-1 独立选片
for task_index in range(1, count):
variant = None
last_err: Exception | None = None
for _attempt in range(2): # 1 次重试,抗 DB 瞬时抖动
try:
variant = _plan_svc.reselect_plan_for_variant(
batch_source_plan_id,
batch_asset_pool,
created_by_user_id=user_id,
name_suffix=f"批量{task_index + 1}",
voice_duration=voice_durations[task_index] if task_index < len(voice_durations) else 0.0,
)
break
except ValueError as ve:
# 素材不足等可预期错误:不重试,直接中断并给出明确提示
logger.warning("[生成任务] 变体独立选片失败(素材不足): %s", ve)
raise HTTPException(
status_code=400,
detail=f"批量生成第 {task_index + 1} 个视频无法独立选片:{ve}"
"请增加素材库中的视频素材后重试。",
) from ve
except Exception as reselection_err: # noqa: PERF203
last_err = reselection_err
logger.warning(
"[生成任务] 变体独立选片失败(尝试%d/2): source=%s error=%s",
_attempt + 1,
batch_source_plan_id,
reselection_err,
exc_info=True,
)
if variant is None:
logger.error(
"[生成任务] 变体独立选片重试仍失败,中断批量创建: source=%s",
batch_source_plan_id,
exc_info=last_err,
)
raise HTTPException(
status_code=500,
detail="创建批量任务失败:无法生成独立剪辑计划,请重试",
) from last_err
variant_plan_ids.append(variant.id)
# ③ 配音时长分配(回传 plan / clone 变体0 均需幂等分配;reselect 已在选片时分配)
for _vi, _pid in enumerate(variant_plan_ids):
_vd = voice_durations[_vi] if _vi < len(voice_durations) else 0.0
if _vd > 0:
try:
_plan_svc.apply_voice_duration_to_plan(_pid, _vd)
except Exception:
logger.exception("[生成任务] 变体%d 配音时长分配失败(不阻断): plan=%s", _vi, _pid)
# N=1 正式生成:渲染侧全局慢放兜底已删除(#1749),enqueue 前也必须按配音分配段长
if count == 1 and not request.is_preview:
from packages.domain.variant_voice_resolver import VariantVoiceError, resolve_variant_voice_ids
try:
_voices = resolve_variant_voice_ids(
count=1,
voice_library_id=request.voice_library_id,
voice_library_ids=request.voice_library_ids or None,
)
_single_vd: list[float] = _query_voice_durations(db, _voices)
_single_dur = _single_vd[0] if _single_vd else 0.0
_single_plan = request.source_edit_plan_id
if not _single_plan and request.template_id:
try:
from packages.adapters.sqlalchemy_impl.models import EditPlanModel
_latest = (
db.query(EditPlanModel)
.filter(
EditPlanModel.template_id == request.template_id,
EditPlanModel.created_by_user_id == user_id,
)
.order_by(EditPlanModel.created_at.desc())
.first()
)
if _latest:
_single_plan = _latest.id
except Exception:
logger.warning("[生成任务] 单任务源 plan 解析失败", exc_info=True)
if _single_dur > 0 and _single_plan:
from app.services.edit_plan_service import EditPlanService
try:
EditPlanService(db).apply_voice_duration_to_plan(_single_plan, _single_dur)
except Exception:
logger.exception("[生成任务] N=1 配音时长分配失败(不阻断): plan=%s", _single_plan)
except VariantVoiceError as ve:
raise HTTPException(status_code=400, detail=str(ve)) from ve
except Exception:
logger.exception("[生成任务] N=1 配音分配兜底异常(不阻断)")
try: try:
for task_index in range(count): for task_index in range(count):
# #1749count>1 时每个变体(含变体0)都关联各自独立 planclone/reselect/variant-plans # 第 1 条复用源 plan(保留用户编辑结果);其余使用预克隆的独立变体 plan。
if count > 1 and variant_plan_ids: # 无源 plansource_edit_plan_id 为空)时无可克隆对象,variant_plan_ids
effective_plan_id = variant_plan_ids[task_index] # 为空列表:各任务走自身随机选片流程,不做索引访问(防 IndexError)
else: effective_plan_id = request.source_edit_plan_id
effective_plan_id = request.source_edit_plan_id if task_index > 0 and variant_plan_ids:
effective_plan_id = variant_plan_ids[task_index - 1]
# 变体级独立配置:titles[]/voice_library_ids[]/cover_urls[]
# 长度1=所有变体共用,长度=count=每个变体独立,空数组=回退单值字段
variant_title_text = _variant_value(request.titles, task_index, "")
variant_title_config = dict(request.title_config or {})
if variant_title_text.strip():
variant_title_config["text"] = variant_title_text.strip()
variant_voice_library_id = _variant_value(request.voice_library_ids, task_index, request.voice_library_id)
variant_cover_url = _variant_value(request.cover_urls, task_index, request.cover_url)
task = use_case.execute( task = use_case.execute(
CreateGenerationTaskCommand( CreateGenerationTaskCommand(
project_id=project_id, project_id=project_id,
asset_library_id=asset_library_id, asset_library_id=asset_library_id,
strategy_id=effective_strategy_id, strategy_id=effective_strategy_id,
voice_library_id=variant_voice_library_id, voice_library_id=request.voice_library_id,
template_id=request.template_id, template_id=request.template_id,
asset_ids=resolved_asset_ids, asset_ids=resolved_asset_ids,
title_ids=request.title_ids, title_ids=request.title_ids,
voice_ids=request.voice_ids,
created_by_user_id=user_id, created_by_user_id=user_id,
source_edit_plan_id=effective_plan_id, source_edit_plan_id=effective_plan_id,
asset_select_mode=request.asset_select_mode, asset_select_mode=request.asset_select_mode,
@@ -685,28 +491,14 @@ def create_generation_task(
source_task_id=request.source_task_id, source_task_id=request.source_task_id,
output_width=request.output_width, output_width=request.output_width,
output_height=request.output_height, output_height=request.output_height,
cover_url=variant_cover_url, cover_url=request.cover_url,
title_config=variant_title_config, title_config=request.title_config or {},
) )
) )
# 变体序号写入 extra_meta(响应/排查时可辨识)
task.extra_meta["variant_index"] = task_index
try: try:
# 兜底关联编辑计划:前端未传 source_edit_plan_id 时, # 兜底关联编辑计划:前端未传 source_edit_plan_id 时,
# 通过 template_id + user_id 在 DB 层直接查找最新的 plan。 # 通过 template_id + user_id 在 DB 层直接查找最新的 plan。
# 必须在 enqueue 之前执行,避免 worker 读取时 source_edit_plan_id 为空(竞态条件) # 必须在 enqueue 之前执行,避免 worker 读取时 source_edit_plan_id 为空(竞态条件)
# #1743:批量(count>1)场景严禁兜底共用——变体 plan 已在上方预生成,
# 走到这里还缺 plan 说明预生成漏配,直接报错中断,不允许 N 任务关联同一 plan。
if not task.source_edit_plan_id and count > 1:
logger.error(
"[生成任务] 批量任务缺少独立 plan(禁止共用兜底): task_index=%d task_id=%s",
task_index,
task.id,
)
raise HTTPException(
status_code=500,
detail="创建批量任务失败:变体剪辑计划缺失,请重新预览后再批量生成。",
)
if not task.source_edit_plan_id and request.template_id: if not task.source_edit_plan_id and request.template_id:
try: try:
from packages.adapters.sqlalchemy_impl.models import EditPlanModel from packages.adapters.sqlalchemy_impl.models import EditPlanModel
@@ -737,13 +529,13 @@ def create_generation_task(
# 回写 plan.config:必须在 enqueue 之前执行, # 回写 plan.config:必须在 enqueue 之前执行,
# 确保 worker 读取 plan 时 config 中已包含 generation_task_id。 # 确保 worker 读取 plan 时 config 中已包含 generation_task_id。
# 批量场景下每个变体关联独立 plan,需各自回写自己的变体标题配置 # 只在首个任务时回写一次,避免批量生成时循环覆盖
_effective_plan_id = task.source_edit_plan_id _effective_plan_id = task.source_edit_plan_id
if _effective_plan_id: if _effective_plan_id and len(created_tasks) == 0:
_writeback_edit_plan_config( _writeback_edit_plan_config(
plan_id=_effective_plan_id, plan_id=_effective_plan_id,
task_id=task.id, task_id=task.id,
title_config=variant_title_config, title_config=request.title_config,
db=db, db=db,
) )
@@ -763,7 +555,7 @@ def create_generation_task(
if not created_tasks: if not created_tasks:
raise HTTPException( raise HTTPException(
status_code=429, status_code=429,
detail=build_rate_limit_detail(_e, generation_task_repository, scope="user"), detail="您的待处理任务过多,请等待完成后再提交",
) from _e ) from _e
break break
except GlobalQueueFull as _e: except GlobalQueueFull as _e:
@@ -771,7 +563,7 @@ def create_generation_task(
if not created_tasks: if not created_tasks:
raise HTTPException( raise HTTPException(
status_code=503, status_code=503,
detail=build_rate_limit_detail(_e, generation_task_repository, scope="global"), detail="系统繁忙,请稍后再试",
) from _e ) from _e
break break
except HTTPException: except HTTPException:
@@ -873,6 +665,7 @@ def confirm_generation(
template_id=source_task.template_id, template_id=source_task.template_id,
asset_ids=source_task.asset_ids, asset_ids=source_task.asset_ids,
title_ids=source_task.title_ids, title_ids=source_task.title_ids,
voice_ids=source_task.voice_ids,
created_by_user_id=authenticated_user.user.id, created_by_user_id=authenticated_user.user.id,
source_edit_plan_id=source_task.source_edit_plan_id or "", source_edit_plan_id=source_task.source_edit_plan_id or "",
asset_select_mode=source_task.asset_select_mode, asset_select_mode=source_task.asset_select_mode,
@@ -896,15 +689,15 @@ def confirm_generation(
log_task_status=True, log_task_status=True,
): ):
logger.warning("[确认生成] 入队失败: task_id=%s", new_task.id) logger.warning("[确认生成] 入队失败: task_id=%s", new_task.id)
except UserPendingLimitExceeded as _e: except UserPendingLimitExceeded:
raise HTTPException( raise HTTPException(
status_code=429, status_code=429,
detail=build_rate_limit_detail(_e, generation_task_repository, scope="user"), detail="您的待处理任务过多,请等待完成后再提交",
) from None ) from None
except GlobalQueueFull as _e: except GlobalQueueFull:
raise HTTPException( raise HTTPException(
status_code=503, status_code=503,
detail=build_rate_limit_detail(_e, generation_task_repository, scope="global"), detail="系统繁忙,请稍后再试",
) from None ) from None
return BatchGenerationTaskResponse( return BatchGenerationTaskResponse(
@@ -986,24 +779,12 @@ def retry_generation_task(
if user_pending >= USER_PENDING_LIMIT: if user_pending >= USER_PENDING_LIMIT:
raise HTTPException( raise HTTPException(
status_code=429, status_code=429,
detail=build_rate_limit_detail( detail=f"您的待处理任务过多(当前 {user_pending}/{USER_PENDING_LIMIT}),请等待完成后再提交",
UserPendingLimitExceeded(
user_id=user_id,
pending_count=user_pending,
limit=USER_PENDING_LIMIT,
),
generation_task_repository,
scope="user",
),
) )
if global_pending >= GLOBAL_PENDING_LIMIT: if global_pending >= GLOBAL_PENDING_LIMIT:
raise HTTPException( raise HTTPException(
status_code=503, status_code=503,
detail=build_rate_limit_detail( detail="系统繁忙,请稍后再试",
GlobalQueueFull(pending_count=global_pending, limit=GLOBAL_PENDING_LIMIT),
generation_task_repository,
scope="global",
),
) )
use_case = CreateGenerationTaskUseCase(generation_task_repository) use_case = CreateGenerationTaskUseCase(generation_task_repository)
@@ -1016,6 +797,7 @@ def retry_generation_task(
template_id=task.template_id, template_id=task.template_id,
asset_ids=task.asset_ids, asset_ids=task.asset_ids,
title_ids=task.title_ids, title_ids=task.title_ids,
voice_ids=task.voice_ids,
created_by_user_id=user_id, created_by_user_id=user_id,
source_edit_plan_id=task.source_edit_plan_id or "", source_edit_plan_id=task.source_edit_plan_id or "",
asset_select_mode=getattr(task, "asset_select_mode", ""), asset_select_mode=getattr(task, "asset_select_mode", ""),
@@ -1037,15 +819,15 @@ def retry_generation_task(
log_task_status=True, log_task_status=True,
): ):
logger.warning("[生成任务] 重试入队失败: task_id=%s", retried.id) logger.warning("[生成任务] 重试入队失败: task_id=%s", retried.id)
except UserPendingLimitExceeded as _e: except UserPendingLimitExceeded:
raise HTTPException( raise HTTPException(
status_code=429, status_code=429,
detail=build_rate_limit_detail(_e, generation_task_repository, scope="user"), detail="您的待处理任务过多,请等待完成后再提交",
) from None ) from None
except GlobalQueueFull as _e: except GlobalQueueFull:
raise HTTPException( raise HTTPException(
status_code=503, status_code=503,
detail=build_rate_limit_detail(_e, generation_task_repository, scope="global"), detail="系统繁忙,请稍后再试",
) from None ) from None
return _to_generation_task_response(retried) return _to_generation_task_response(retried)
@@ -1,178 +0,0 @@
"""轻量选片接口 POST /generation/variant-plans#1749)。
与正式生成共用同一套选片函数(EditPlanService.ensure_variant_plans →
clone_plan_for_variant / reselect_plan_for_variant → variant_plan_selector),
但**不建任务、不入队、不渲染**
- 仅为 N 个变体创建/选好 EditPlan + clips,返回 plan_id 与片段列表;
- 前端确认后调正式生成接口回传 variant_plan_ids,直接复用这些 plan
不再重复选片(回传后仍按各变体配音幂等重分配段长);
- 配音守卫:voice_library_ids 长度/缺值 → 400variant_voice_resolver),
禁静默 fallback
- 素材不足等选片失败 → 400(与正式生成同口径);除此之外不报错打断。
"""
from __future__ import annotations
import logging
from typing import Any
from app.auth import AuthenticatedUser, get_current_user
from app.dependencies import get_db_session
from fastapi import APIRouter, Depends, HTTPException
from pydantic import BaseModel, Field, model_validator
from sqlalchemy.orm import Session
from packages.domain.variant_voice_resolver import VariantVoiceError, resolve_variant_voice_ids
logger = logging.getLogger(__name__)
router = APIRouter()
class VariantPlanRequest(BaseModel):
"""轻量选片请求体(与前端 variantPlans.ts 契约一致)。"""
template_id: str = Field(default="", description="模板 ID(无 source_edit_plan_id 时用于查找骨架 plan")
asset_ids: list[str] = Field(default_factory=list, description="批次素材池")
count: int = Field(default=1, ge=1, le=50, description="变体数量")
source_edit_plan_id: str = Field(default="", description="源剪辑计划 ID(优先)")
# 配音(可选;传独立配音时严格守卫)
voice_library_id: str = Field(default="", description="统一配音 ID")
voice_library_ids: list[str] = Field(default_factory=list, description="独立配音 ID 列表(长度须=count)")
@model_validator(mode="after")
def _validate(self) -> "VariantPlanRequest":
if not self.template_id.strip() and not self.source_edit_plan_id.strip():
raise ValueError("template_id 与 source_edit_plan_id 至少需要提供一个")
try:
resolve_variant_voice_ids(
count=self.count,
voice_library_id=self.voice_library_id,
voice_library_ids=self.voice_library_ids or None,
)
except VariantVoiceError as exc:
raise ValueError(str(exc)) from exc
return self
class VariantPlanItem(BaseModel):
variant_index: int
plan_id: str
clips: list[dict[str, Any]] = Field(default_factory=list)
class VariantPlanResponse(BaseModel):
items: list[VariantPlanItem]
total: int
@router.post("/variant-plans", response_model=VariantPlanResponse)
def create_variant_plans(
request: VariantPlanRequest,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
db: Session = Depends(get_db_session),
) -> VariantPlanResponse:
"""轻量选片:为 N 个变体创建独立 EditPlan + clips,不建任务/不渲染。
Returns:
200 + {items: [{variant_index, plan_id, clips}], total}
"""
user_id = authenticated_user.user.id
# 配音严格守卫(schema 已校验,此处复用解析取每变体配音)
try:
voices = resolve_variant_voice_ids(
count=request.count,
voice_library_id=request.voice_library_id,
voice_library_ids=request.voice_library_ids or None,
)
except VariantVoiceError as exc:
raise HTTPException(status_code=400, detail=str(exc)) from exc
# 解析源 plan:显式传入优先;否则按 template_id + user 查最新
source_plan_id = request.source_edit_plan_id.strip()
if not source_plan_id and request.template_id.strip():
try:
from packages.adapters.sqlalchemy_impl.models import EditPlanModel
_latest = (
db.query(EditPlanModel)
.filter(
EditPlanModel.template_id == request.template_id.strip(),
EditPlanModel.created_by_user_id == user_id,
)
.order_by(EditPlanModel.created_at.desc())
.first()
)
if _latest:
source_plan_id = _latest.id
except Exception:
logger.exception("[variant-plans] 源 plan 解析失败")
if not source_plan_id:
raise HTTPException(
status_code=400,
detail="缺少剪辑计划:请先完成一次预览生成(或传入 source_edit_plan_id)后再试。",
)
# 配音时长(硬化:异常 → 0.0 不阻断选片)
try:
from app.api.routes.generation_tasks import _query_voice_durations
voice_durations = _query_voice_durations(db, voices)
except Exception:
logger.exception("[variant-plans] 配音时长查询失败(按占位段长选片)")
voice_durations = [0.0] * request.count
from app.services.edit_plan_service import EditPlanService
svc = EditPlanService(db)
try:
plan_ids = svc.ensure_variant_plans(
source_plan_id,
request.count,
list(dict.fromkeys(request.asset_ids or [])),
created_by_user_id=user_id,
voice_durations=voice_durations,
)
except ValueError as ve:
# 素材池为空/时长全未知等可预期错误 → 400(与正式生成同口径)
logger.warning("[variant-plans] 选片失败: %s", ve)
raise HTTPException(status_code=400, detail=f"变体选片失败:{ve}。请增加素材后重试。") from ve
except HTTPException:
raise
except Exception as e:
logger.exception("[variant-plans] 选片异常")
raise HTTPException(status_code=500, detail="选片失败,请稍后重试") from e
# 组装 clips 响应
items: list[VariantPlanItem] = []
for idx, pid in enumerate(plan_ids):
clips = svc.list_clips(pid)
clip_dicts = [
{
"id": c.id,
"order": c.order,
"asset_id": c.asset_id,
"start_time": float(c.start_time or 0.0),
"duration": float(c.duration or 0.0),
"clip_type": c.clip_type,
"transition_effect": c.transition_effect,
"transition_duration": float(c.transition_duration or 0.0),
"playback_speed": float(c.playback_speed or 1.0),
"text_content": c.text_content or "",
"status": c.status or "ready",
}
for c in clips
]
items.append(VariantPlanItem(variant_index=idx, plan_id=pid, clips=clip_dicts))
logger.info(
"[variant-plans] 轻量选片完成: user=%s source=%s count=%d plans=%d",
user_id,
source_plan_id,
request.count,
len(plan_ids),
)
return VariantPlanResponse(items=items, total=len(items))
+1 -7
View File
@@ -43,13 +43,7 @@ def submit_ingest_job(
) )
) )
celery_result = celery_app.send_task("worker.ingest_asset", args=[job.id]) celery_app.send_task("worker.ingest_asset", args=[job.id])
if getattr(celery_result, "id", ""):
try:
job.celery_task_id = celery_result.id
ingest_job_repository.update(job)
except Exception: # noqa: BLE001
pass
return IngestJobResponse( return IngestJobResponse(
id=job.id, id=job.id,
-186
View File
@@ -1,186 +0,0 @@
"""对口型 API 路由 — #1796 MediaKit 对口型, #1809 参数调整.
接口:
POST /api/v1/lipsync/jobs 提交对口型任务
GET /api/v1/lipsync/jobs 任务列表
GET /api/v1/lipsync/jobs/{id} 任务详情
POST /api/v1/lipsync/jobs/{id}/refresh 刷新任务状态
POST /api/v1/lipsync/jobs/{id}/cancel 取消任务
"""
from __future__ import annotations
import logging
from app.auth import AuthenticatedUser, get_current_user
from app.dependencies import (
get_db_session,
get_voice_clone_profile_repository,
)
from app.schemas.lipsync import CreateLipsyncJobRequest, LipsyncJobResponse
from app.services.lipsync_service import LipsyncService
from app.services.mediakit_client import MediaKitError
from fastapi import APIRouter, BackgroundTasks, Depends, HTTPException, Query
from sqlalchemy.orm import Session
logger = logging.getLogger(__name__)
router = APIRouter()
def _get_service(
db: Session = Depends(get_db_session),
voice_clone_repo=Depends(get_voice_clone_profile_repository),
) -> LipsyncService:
# voice_clone_repo 用于克隆音色 profile 解析
# TTS 合成已移至 Celery 异步任务,无需同步注入 cosyvoice_service
return LipsyncService(
db,
voice_clone_repo=voice_clone_repo,
)
# ── POST /jobs — 提交对口型任务 ───────────────────────────────────────────
@router.post("/jobs", response_model=LipsyncJobResponse, status_code=201)
def create_lipsync_job(
body: CreateLipsyncJobRequest,
current_user: AuthenticatedUser = Depends(get_current_user),
svc: LipsyncService = Depends(_get_service),
):
"""提交对口型任务.
#1809/#1822: 前端传 {video_url, voice_id, script_text, speed?, emotion?}
后端创建任务记录(状态 tts_processing),dispatch Celery 异步任务执行 TTS 合成 + MediaKit 提交;
也支持直接传 {video_url, audio_url}(同步提交 MediaKit)。
"""
try:
job = svc.create_job(
user_id=current_user.user.id,
video_url=body.video_url,
audio_url=body.audio_url,
voice_id=body.voice_id,
script_text=body.script_text,
speed=body.speed,
emotion=body.emotion,
enable_video_loop=body.enable_video_loop,
project_id=body.project_id,
)
except ValueError as exc:
# 参数无效(如 voice_id 格式不对、文本过长等)
raise HTTPException(status_code=400, detail=str(exc)) from exc
except MediaKitError as exc:
# 音色无权访问 → 403;参数无效 → 400MediaKit 提交失败 → 502
status_code = 502
if exc.code in ("VoiceForbidden",):
status_code = 403
elif exc.code in ("InvalidInput", "TTSInvalidParam", "VoiceNotReady"):
status_code = 400
raise HTTPException(
status_code=status_code,
detail={
"code": exc.code,
"message": str(exc),
"request_id": getattr(exc, "request_id", ""),
},
) from exc
except Exception as exc:
# 兜底:任何未预期的错误返回 400 而非 500
logger.error("创建对口型任务异常: %s", exc, exc_info=True)
raise HTTPException(
status_code=400,
detail=f"创建对口型任务失败: {exc}",
) from exc
return job
# ── GET /jobs — 任务列表 ─────────────────────────────────────────────────
@router.get("/jobs", response_model=dict)
def list_lipsync_jobs(
project_id: str = Query("", description="项目 ID 过滤"),
status: str = Query("", description="状态过滤"),
offset: int = Query(0, ge=0),
limit: int = Query(20, ge=1, le=100),
current_user: AuthenticatedUser = Depends(get_current_user),
svc: LipsyncService = Depends(_get_service),
):
"""获取对口型任务列表."""
items, total = svc.list_jobs(
user_id=current_user.user.id,
project_id=project_id,
status=status,
offset=offset,
limit=limit,
)
return {
"items": [LipsyncJobResponse.model_validate(j) for j in items],
"total": total,
"offset": offset,
"limit": limit,
}
# ── GET /jobs/{job_id} — 任务详情 ────────────────────────────────────────
@router.get("/jobs/{job_id}", response_model=LipsyncJobResponse)
def get_lipsync_job(
job_id: str,
background: BackgroundTasks,
current_user: AuthenticatedUser = Depends(get_current_user),
svc: LipsyncService = Depends(_get_service),
):
"""获取对口型任务详情.
非终态任务:先返回 DB 缓存,挂后台刷新(下次轮询拿到新状态),
避免 MediaKit 慢响应阻塞前端轮询。
"""
job = svc.get_job(job_id, current_user.user.id)
if job is None:
raise HTTPException(status_code=404, detail="任务不存在")
if job.status not in ("completed", "failed"):
background.add_task(svc.refresh_job_status, job_id, current_user.user.id)
return job
# ── POST /jobs/{job_id}/refresh — 刷新状态 ───────────────────────────────
@router.post("/jobs/{job_id}/refresh", response_model=LipsyncJobResponse)
def refresh_lipsync_job(
job_id: str,
current_user: AuthenticatedUser = Depends(get_current_user),
svc: LipsyncService = Depends(_get_service),
):
"""从 MediaKit 拉取最新状态并更新."""
job = svc.refresh_job_status(job_id, current_user.user.id)
if job is None:
raise HTTPException(status_code=404, detail="任务不存在")
return job
# ── POST /jobs/{job_id}/cancel — 取消任务 ────────────────────────────────
@router.post("/jobs/{job_id}/cancel", response_model=LipsyncJobResponse)
def cancel_lipsync_job(
job_id: str,
current_user: AuthenticatedUser = Depends(get_current_user),
svc: LipsyncService = Depends(_get_service),
):
"""取消对口型任务(仅 pending/tts_processing/submitted 状态可取消)."""
job = svc.cancel_job(job_id, current_user.user.id)
if job is None:
raise HTTPException(status_code=404, detail="任务不存在")
if job.status != "cancelled":
raise HTTPException(
status_code=400,
detail=f"任务状态 {job.status} 不可取消,仅 pending/tts_processing/submitted 可取消",
)
return job
+1 -41
View File
@@ -1,14 +1,13 @@
from typing import Any from typing import Any
from app.auth import AuthenticatedUser, get_current_user from app.auth import AuthenticatedUser, get_current_user
from app.dependencies import get_asset_library_repository, get_project_repository from app.dependencies import get_project_repository
from app.schemas.project import ( from app.schemas.project import (
CreateProjectRequest, CreateProjectRequest,
ListProjectsResponse, ListProjectsResponse,
ProjectResponse, ProjectResponse,
) )
from fastapi import APIRouter, Depends, HTTPException, Response, status from fastapi import APIRouter, Depends, HTTPException, Response, status
from pydantic import BaseModel
from packages.application import ( from packages.application import (
CreateProjectCommand, CreateProjectCommand,
@@ -17,20 +16,10 @@ from packages.application import (
GetProjectUseCase, GetProjectUseCase,
ListProjectsUseCase, ListProjectsUseCase,
) )
from packages.domain import AssetLibraryKind
router = APIRouter() router = APIRouter()
class DefaultContextResponse(BaseModel):
"""幂等默认上下文响应(Issue #1775):默认项目 + 各类型默认素材库 ID。"""
project_id: str
image_library_id: str
video_library_id: str
voice_library_id: str
def _to_project_response(item) -> ProjectResponse: def _to_project_response(item) -> ProjectResponse:
return ProjectResponse( return ProjectResponse(
id=item.id, id=item.id,
@@ -83,35 +72,6 @@ def create_project(
return _to_project_response(project) return _to_project_response(project)
@router.post("/ensure-default", response_model=DefaultContextResponse)
def ensure_default_project_and_libraries(
authenticated_user: AuthenticatedUser = Depends(get_current_user),
project_repository: Any = Depends(get_project_repository),
asset_library_repository: Any = Depends(get_asset_library_repository),
) -> DefaultContextResponse:
"""幂等获取/创建当前用户的默认项目和三类默认素材库(Issue #1775)。
- 同一用户永远只有一个默认项目(部分唯一索引 uq_projects_owner_default
- 同一项目同 kind 永远只有一个默认素材库(唯一约束 uq_asset_libraries_project_kind
- 并发调用/失败重试:唯一约束冲突时返回已存在记录,不报 500
- 项目和素材库的创建各自在仓储事务内幂等,冲突回滚后重查返回同一条
"""
user_id = authenticated_user.user.id
project = project_repository.get_or_create_default_project(user_id)
libraries = {}
for kind in (AssetLibraryKind.VIDEO, AssetLibraryKind.VOICE, AssetLibraryKind.IMAGE):
library = asset_library_repository.get_or_create_default_library(project.id, kind)
libraries[kind] = library.id
return DefaultContextResponse(
project_id=project.id,
image_library_id=libraries[AssetLibraryKind.IMAGE],
video_library_id=libraries[AssetLibraryKind.VIDEO],
voice_library_id=libraries[AssetLibraryKind.VOICE],
)
@router.delete("/{project_id}", status_code=status.HTTP_204_NO_CONTENT, response_model=None, response_class=Response) @router.delete("/{project_id}", status_code=status.HTTP_204_NO_CONTENT, response_model=None, response_class=Response)
def delete_project( def delete_project(
project_id: str, project_id: str,
-123
View File
@@ -1,123 +0,0 @@
"""Script (口播文案库) CRUD routes — Issue #1795."""
from __future__ import annotations
from typing import Optional
from app.auth import AuthenticatedUser, get_current_user
from app.dependencies import get_db_session
from app.schemas.script import (
CreateScriptRequest,
ScriptListResponse,
ScriptResponse,
ScriptSegment,
UpdateScriptRequest,
)
from app.services.script_service import ScriptNotFoundError, ScriptService
from fastapi import APIRouter, Depends, HTTPException, Query, Response, status
from sqlalchemy.orm import Session
router = APIRouter()
def _get_service(session: Session = Depends(get_db_session)) -> ScriptService:
return ScriptService(session)
def _to_response(script) -> ScriptResponse:
segments = script.segments or []
return ScriptResponse(
id=script.id,
user_id=script.user_id,
title=script.title,
content=script.content,
segments=[
ScriptSegment(text=s.get("text", ""), duration=s.get("duration")) if isinstance(s, dict) else s
for s in segments
],
tags=script.tags or [],
created_at=script.created_at,
updated_at=script.updated_at,
)
@router.get("", response_model=ScriptListResponse)
def list_scripts(
skip: int = Query(0, ge=0),
limit: int = Query(50, ge=1, le=200),
tag: Optional[str] = Query(None, description="按标签筛选"),
authenticated_user: AuthenticatedUser = Depends(get_current_user),
svc: ScriptService = Depends(_get_service),
) -> ScriptListResponse:
user_id = authenticated_user.user.id
items, total = svc.list_scripts(user_id, skip=skip, limit=limit, tag=tag)
return ScriptListResponse(
items=[_to_response(i) for i in items],
total=total,
)
@router.post("", response_model=ScriptResponse, status_code=status.HTTP_201_CREATED)
def create_script(
request: CreateScriptRequest,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
svc: ScriptService = Depends(_get_service),
) -> ScriptResponse:
user_id = authenticated_user.user.id
script = svc.create_script(
user_id=user_id,
title=request.title,
content=request.content,
segments=[s.model_dump() for s in request.segments],
tags=request.tags,
)
return _to_response(script)
@router.get("/{script_id}", response_model=ScriptResponse)
def get_script(
script_id: str,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
svc: ScriptService = Depends(_get_service),
) -> ScriptResponse:
user_id = authenticated_user.user.id
try:
script = svc.get_script(script_id, user_id)
except ScriptNotFoundError as exc:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Script not found") from exc
return _to_response(script)
@router.put("/{script_id}", response_model=ScriptResponse)
def update_script(
script_id: str,
request: UpdateScriptRequest,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
svc: ScriptService = Depends(_get_service),
) -> ScriptResponse:
user_id = authenticated_user.user.id
try:
script = svc.update_script(
script_id=script_id,
user_id=user_id,
title=request.title,
content=request.content,
segments=[s.model_dump() for s in request.segments] if request.segments is not None else None,
tags=request.tags,
)
except ScriptNotFoundError as exc:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Script not found") from exc
return _to_response(script)
@router.delete("/{script_id}", status_code=status.HTTP_204_NO_CONTENT, response_model=None, response_class=Response)
def delete_script(
script_id: str,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
svc: ScriptService = Depends(_get_service),
) -> Response:
user_id = authenticated_user.user.id
deleted = svc.delete_script(script_id, user_id)
if not deleted:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Script not found")
return
+1 -7
View File
@@ -375,13 +375,7 @@ def retry_project_task(
storage_key=job.storage_key, storage_key=job.storage_key,
) )
) )
celery_result = celery_app.send_task("worker.ingest_asset", args=[retried.id]) celery_app.send_task("worker.ingest_asset", args=[retried.id])
if getattr(celery_result, "id", ""):
try:
retried.celery_task_id = celery_result.id
ingest_job_repository.update(retried)
except Exception: # noqa: BLE001
pass
return ProjectTaskResponse( return ProjectTaskResponse(
id=f"ingest:{retried.id}", id=f"ingest:{retried.id}",
task_type="ingest", task_type="ingest",
-5
View File
@@ -106,10 +106,6 @@ def list_templates(
tag: str | None = Query(None, description="按标签筛选"), tag: str | None = Query(None, description="按标签筛选"),
keyword: str | None = Query(None, description="按名称关键词搜索"), keyword: str | None = Query(None, description="按名称关键词搜索"),
mode: str | None = Query(None, description="按剪辑模式筛选"), mode: str | None = Query(None, description="按剪辑模式筛选"),
valid_only: bool = Query(
False,
description="仅返回已配置片段的模板(剪辑页传 true;模板编辑器不传,可查看全部模板含草稿)",
),
authenticated_user: AuthenticatedUser = Depends(get_current_user), authenticated_user: AuthenticatedUser = Depends(get_current_user),
template_repository: SQLAlchemyTemplateRepository = Depends(_get_template_repository), template_repository: SQLAlchemyTemplateRepository = Depends(_get_template_repository),
) -> ListTemplatesResponse: ) -> ListTemplatesResponse:
@@ -120,7 +116,6 @@ def list_templates(
tag=tag, tag=tag,
keyword=keyword, keyword=keyword,
mode=mode, mode=mode,
valid_only=valid_only,
) )
use_case = ListTemplatesUseCase(template_repository) use_case = ListTemplatesUseCase(template_repository)
templates = use_case.execute(user_id, skip=skip, limit=limit, filter=tpl_filter) templates = use_case.execute(user_id, skip=skip, limit=limit, filter=tpl_filter)
@@ -65,8 +65,8 @@ def _build_asset_analyses(
if url: if url:
video_urls.append(url) video_urls.append(url)
valid_asset_ids.append(aid) valid_asset_ids.append(aid)
except Exception: except Exception as e:
logger.exception("获取素材URL失败: asset_id=%s", aid) logger.warning("获取素材URL失败: asset_id=%s error=%s", aid, str(e))
if not video_urls: if not video_urls:
logger.info("无可用视频素材,跳过视频理解分析") logger.info("无可用视频素材,跳过视频理解分析")
@@ -108,7 +108,7 @@ def _build_asset_analyses(
return analyses return analyses
except Exception as e: except Exception as e:
logger.exception("MediaKit 视频理解异常,将降级到无分析模式: %s", e) logger.warning("MediaKit 视频理解异常,将降级到无分析模式: %s", str(e))
return {} return {}
@@ -177,7 +177,7 @@ def editor_ai_recommend(
try: try:
db.rollback() db.rollback()
except Exception: except Exception:
logger.exception("db rollback failed in ai_recommend") pass
raise HTTPException( raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail="AI推荐结果保存失败,请稍后重试", detail="AI推荐结果保存失败,请稍后重试",
+161 -398
View File
@@ -23,10 +23,6 @@ import re
from app.auth import AuthenticatedUser, get_current_user from app.auth import AuthenticatedUser, get_current_user
from app.core.storage import get_storage_service from app.core.storage import get_storage_service
from app.dependencies import get_asset_repository, get_db_session from app.dependencies import get_asset_repository, get_db_session
# 默认转场时长(与 worker 端保持一致)
_DEFAULT_TRANSITION_DURATION = 0.5
from app.services.asset_segment_tracker import ( from app.services.asset_segment_tracker import (
REUSE_RATIO_LIMIT, REUSE_RATIO_LIMIT,
SEGMENT_EDGE_GAP, SEGMENT_EDGE_GAP,
@@ -36,19 +32,15 @@ from app.services.asset_segment_tracker import (
remove_used_segment, remove_used_segment,
) )
from app.services.edit_plan_service import EditPlanService from app.services.edit_plan_service import EditPlanService
from app.services.edit_template_service import EditTemplateService, TemplateNotFoundError from app.services.edit_template_service import EditTemplateService
from fastapi import APIRouter, BackgroundTasks, Depends, HTTPException, Query, status from fastapi import APIRouter, BackgroundTasks, Depends, HTTPException, Query, status
from sqlalchemy.orm import Session from sqlalchemy.orm import Session
from packages.adapters.sqlalchemy_impl.asset_repository import SQLAlchemyAssetRepository from packages.adapters.sqlalchemy_impl.asset_repository import SQLAlchemyAssetRepository
from packages.domain.plan_generator_utils import ( from packages.adapters.sqlalchemy_impl.template_repository import (
_calc_random_start_time, SQLAlchemyTemplateRepository,
build_scene_segments,
extract_scene_points_from_metadata,
pick_scene_aware_start,
pick_start_in_scene_segment,
) )
from packages.domain.smart_match import SCORE_RANDOM_NOISE_MAX, score_asset from packages.domain.plan_generator_utils import _calc_random_start_time
from packages.shared.mediakit_client import get_mediakit_client from packages.shared.mediakit_client import get_mediakit_client
from .dependencies import get_draft_plan_id, get_editor_services from .dependencies import get_draft_plan_id, get_editor_services
@@ -132,8 +124,8 @@ def _build_asset_url_map(
result: dict[str, str | None] = {} result: dict[str, str | None] = {}
try: try:
storage = get_storage_service() storage = get_storage_service()
except Exception as e: except Exception:
logger.exception("获取存储服务失败,跳过asset_url生成: %s", e) logger.warning("获取存储服务失败,跳过asset_url生成")
return {aid: None for aid in asset_ids} return {aid: None for aid in asset_ids}
# 批量查询所有 Asset(单次 SQL IN 查询,避免 N+1) # 批量查询所有 Asset(单次 SQL IN 查询,避免 N+1)
@@ -141,7 +133,7 @@ def _build_asset_url_map(
assets = asset_repo.find_by_ids(unique_ids) assets = asset_repo.find_by_ids(unique_ids)
asset_map = {a.id: a for a in assets} asset_map = {a.id: a for a in assets}
except Exception: except Exception:
logger.exception("批量查询素材失败: asset_ids=%s", asset_ids) logger.warning("批量查询素材失败: asset_ids=%s", asset_ids, exc_info=True)
return {aid: None for aid in asset_ids if aid} return {aid: None for aid in asset_ids if aid}
for aid in unique_ids: for aid in unique_ids:
@@ -156,7 +148,7 @@ def _build_asset_url_map(
continue continue
result[aid] = storage.get_download_url(storage_key, expires_seconds=3600) result[aid] = storage.get_download_url(storage_key, expires_seconds=3600)
except Exception: except Exception:
logger.exception("生成素材签名URL失败: asset_id=%s", aid) logger.warning("生成素材签名URL失败: asset_id=%s", aid, exc_info=True)
result[aid] = None result[aid] = None
return result return result
@@ -393,42 +385,50 @@ def _safe_segment_duration(value, default: float) -> float:
def _get_template_segments( def _get_template_segments(
template_id: str, template_id: str,
user_id: str,
tpl_svc: EditTemplateService, tpl_svc: EditTemplateService,
db: Session,
) -> list[tuple[int, float, float]]: ) -> list[tuple[int, float, float]]:
"""获取模板的片段配置(顺序、最短时长、最长时长). """获取模板的片段配置(顺序、最短时长、最长时长).
单一数据源:模板主表为 ``templates``(用户自建,归属 user_id/ 优先从新模板系统(template_clip_configs)查询,
``edit_templates``(全局模板库),片段配置主表为 ``template_clip_configs`` 若不存在则回退到旧模板系统(template_segments)。
(由 ``EditTemplateService.list_clip_configs_for_editor`` 统一读取)。
不再使用"新表抛异常 → 降级直查配置表 → 再降级查 segments"的异常控制流,
也不在正常请求中打印 ``ValueError: 模板不存在`` 堆栈。
Args:
template_id: 模板 ID
user_id: 当前登录用户 ID(用于归属校验)
tpl_svc: 模板编辑器服务
Returns: Returns:
[(segment_order, duration_min, duration_max), ...] 按 order 排序 [(segment_order, duration_min, duration_max), ...] 按 order 排序
模板存在但未配置片段时返回空列表。
Raises:
TemplateNotFoundError: 模板不存在、已删除或不归属于当前用户。
""" """
clip_configs = tpl_svc.list_clip_configs_for_editor(template_id, user_id) # 优先查新模板系统
try:
clip_configs = tpl_svc.list_clip_configs(template_id)
if clip_configs:
result = []
for cc in clip_configs:
dur_min = _safe_segment_duration(cc.min_duration, _DEFAULT_EDITOR_CLIP_DURATION)
dur_max = _safe_segment_duration(
cc.max_duration or cc.min_duration,
_DEFAULT_EDITOR_CLIP_DURATION,
)
dur_min, dur_max = min(dur_min, dur_max), max(dur_min, dur_max)
result.append((cc.order, dur_min, dur_max))
return sorted(result, key=lambda x: x[0])
except Exception:
logger.warning("新模板系统查询clip_configs失败,回退到旧系统", exc_info=True)
result = [] # 回退到旧模板系统(template_segments表)
for cc in clip_configs: try:
dur_min = _safe_segment_duration(cc.min_duration, _DEFAULT_EDITOR_CLIP_DURATION) old_repo = SQLAlchemyTemplateRepository(db)
dur_max = _safe_segment_duration( segments = old_repo.list_segments(template_id)
cc.max_duration or cc.min_duration, if segments:
_DEFAULT_EDITOR_CLIP_DURATION, result = []
) for s in segments:
dur_min, dur_max = min(dur_min, dur_max), max(dur_min, dur_max) dur_min = _safe_segment_duration(s.duration_min, _DEFAULT_EDITOR_CLIP_DURATION)
result.append((cc.order, dur_min, dur_max)) dur_max = _safe_segment_duration(s.duration_max, _DEFAULT_EDITOR_CLIP_DURATION)
return sorted(result, key=lambda x: x[0]) dur_min, dur_max = min(dur_min, dur_max), max(dur_min, dur_max)
result.append((s.segment_order, dur_min, dur_max))
return sorted(result, key=lambda x: x[0])
except Exception:
logger.warning("旧模板系统查询segments失败", exc_info=True)
return []
def _recommended_time_conflicts( def _recommended_time_conflicts(
@@ -449,12 +449,6 @@ def _recommended_time_conflicts(
return False return False
# 向后兼容别名:镜头段构建/段内取点逻辑已下沉到 packages.domain.plan_generator_utils
# 旧测试与历史代码仍按 clips._build_scene_segments / _pick_start_in_scene_segment 导入
_build_scene_segments = build_scene_segments
_pick_start_in_scene_segment = pick_start_in_scene_segment
def _get_mediakit_recommendations( def _get_mediakit_recommendations(
asset_ids: list[str], asset_ids: list[str],
asset_repo, asset_repo,
@@ -486,8 +480,8 @@ def _get_mediakit_recommendations(
if url: if url:
video_urls.append(url) video_urls.append(url)
valid_asset_ids.append(asset_id) valid_asset_ids.append(asset_id)
except Exception: except Exception as e:
logger.exception("获取素材URL失败: asset_id=%s", asset_id) logger.warning("获取素材URL失败: asset_id=%s error=%s", asset_id, e)
if not video_urls: if not video_urls:
return {} return {}
@@ -563,51 +557,10 @@ def _get_mediakit_recommendations(
return recommendations return recommendations
except Exception as e: except Exception as e:
logger.exception("MediaKit 智能选片异常,降级为随机选择: %s", e) logger.warning("MediaKit 智能选片异常,降级为随机选择: %s", e)
return {} return {}
def _calc_plan_internal_duplicate_rate(clips_data: list[dict]) -> float:
"""估算单条成片内部重复率(%.
检查本条成片中同一素材是否有重叠的时间区间。
重叠时长 / 成片总时长 * 100 = 内部重复率。
这是一个轻量估算,不依赖视频指纹;完整查重由 worker 异步完成。
"""
if not clips_data:
return 0.0
# 按素材分组
by_asset: dict[str, list[tuple[float, float]]] = {}
total_duration = 0.0
for c in clips_data:
aid = c.get("asset_id", "")
if not aid:
continue
start = c.get("start_time", 0.0)
end = start + c.get("duration", 0.0)
by_asset.setdefault(aid, []).append((start, end))
total_duration += c.get("duration", 0.0)
if total_duration <= 0:
return 0.0
# 检查同素材内的区间重叠
overlap_duration = 0.0
for segments in by_asset.values():
if len(segments) < 2:
continue
segments_sorted = sorted(segments, key=lambda s: s[0])
last_end = segments_sorted[0][1]
for start, end in segments_sorted[1:]:
overlap = max(0.0, min(end, last_end) - start)
if overlap > 0:
overlap_duration += overlap
last_end = max(last_end, end)
return round(overlap_duration / total_duration * 100, 1)
@router.post("/clips/from-assets", response_model=ClipsFromAssetsResponse) @router.post("/clips/from-assets", response_model=ClipsFromAssetsResponse)
def create_clips_from_assets_editor( def create_clips_from_assets_editor(
template_id: str, template_id: str,
@@ -631,21 +584,13 @@ def create_clips_from_assets_editor(
7. 素材时长为 0 或缺失时报 400,不创建无效片段 7. 素材时长为 0 或缺失时报 400,不创建无效片段
""" """
tpl_svc, plan_svc = services tpl_svc, plan_svc = services
user_id = str(current_user.user.id)
# 1. 查询模板片段配置。模板不存在/已删除/无权限 → 404; # 1. 查询模板 segments
# 模板存在但确实未配置片段 → 422(配置错误,与 404 区分)。 segments = _get_template_segments(template_id, tpl_svc, db)
try:
segments = _get_template_segments(template_id, user_id, tpl_svc)
except TemplateNotFoundError as exc:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail="模板不存在或无权访问",
) from exc
if not segments: if not segments:
raise HTTPException( raise HTTPException(
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY, status_code=status.HTTP_400_BAD_REQUEST,
detail="模板未配置片段", detail="模板没有片段配置,无法创建片段",
) )
# 防御:schema validator 已过滤 null/空串,这里再归一化一次, # 防御:schema validator 已过滤 null/空串,这里再归一化一次,
@@ -660,26 +605,10 @@ def create_clips_from_assets_editor(
# 2. 获取素材实际时长(去重查询) # 2. 获取素材实际时长(去重查询)
unique_asset_ids = list(dict.fromkeys(asset_ids)) unique_asset_ids = list(dict.fromkeys(asset_ids))
asset_durations: dict[str, float] = {} asset_durations: dict[str, float] = {}
asset_smart_scores: dict[str, float] = {}
# 素材 metadata 中缓存的场景切换点(由后台 MediaKit SceneChange 检测写入):
# 有缓存时片段起点从随机镜头段中选取(不同片段来自不同镜头),无缓存回退随机起点
asset_scene_points: dict[str, list[float]] = {}
for asset_id in unique_asset_ids: for asset_id in unique_asset_ids:
asset = asset_repo.get(asset_id) asset = asset_repo.get(asset_id)
if asset and hasattr(asset, "duration"): if asset and hasattr(asset, "duration"):
asset_durations[asset_id] = float(asset.duration or 0.0) asset_durations[asset_id] = float(asset.duration or 0.0)
# 计算 smart_match 综合评分,用于候选排序
smart_score, _ = score_asset(asset)
asset_smart_scores[asset_id] = smart_score
# 读取场景切换点缓存(新素材未检测过时为 None,走随机起点兜底)
cached_points = extract_scene_points_from_metadata(getattr(asset, "metadata", None))
if cached_points:
asset_scene_points[asset_id] = cached_points
logger.info(
"from-assets 场景缓存命中: %d/%d 个素材有场景切换点",
len(asset_scene_points),
len(unique_asset_ids),
)
# 3. 在内存中计算所有片段数据(使用随机起始时间,不调用MediaKit) # 3. 在内存中计算所有片段数据(使用随机起始时间,不调用MediaKit)
# 读取素材 metadata 中持久化的历史已用区间(跨任务/跨调用去重), # 读取素材 metadata 中持久化的历史已用区间(跨任务/跨调用去重),
@@ -712,52 +641,19 @@ def create_clips_from_assets_editor(
return False return False
return reused_durations.get(aid, 0.0) / assigned > REUSE_RATIO_LIMIT return reused_durations.get(aid, 0.0) / assigned > REUSE_RATIO_LIMIT
# 素材耗尽标志:某轮循环中所有素材均被跳过时为 True for i, (_seg_order, dur_min, dur_max) in enumerate(segments):
all_assets_exhausted = False
# 计算转场重叠补偿:每个 clip 需要额外增加的时长
# 目标:渲染后视频总时长 = 模板设定的各片段时长之和
# 公式:每 clip 增加 (n_segments - 1) * td / n_segments
n_segments = len(segments)
if n_segments > 1:
transition_compensation = (n_segments - 1) * _DEFAULT_TRANSITION_DURATION / n_segments
else:
transition_compensation = 0.0
# 打乱 segments 的处理顺序(分配素材的顺序随机化),但最终 clips_data 按原始 order 排序
shuffled_indices = list(range(len(segments)))
random.shuffle(shuffled_indices)
for idx in shuffled_indices:
_seg_order, dur_min, dur_max = segments[idx]
# 在 segment 的 duration_min ~ duration_max 之间随机取值(保留一位小数) # 在 segment 的 duration_min ~ duration_max 之间随机取值(保留一位小数)
raw_duration = random.uniform(dur_min, dur_max) raw_duration = random.uniform(dur_min, dur_max)
# 加上转场补偿,确保最终输出时长 = 模板设定总时长
raw_duration += transition_compensation
# 贪心分配素材:按"已使用次数"升序排列候选素材(使用最少的优先), # 轮询分配素材:跳过时长缺失、复用占比已超 15% 阈值的素材;
# 同次数随机打散,避免"A-B-C-D"的固定组合反复出现。
# 跳过时长缺失、复用占比已超 10% 阈值的素材;
# 选中后计算起点,若该素材可用区间耗尽且复用被闸门拒绝(calc 返回 None), # 选中后计算起点,若该素材可用区间耗尽且复用被闸门拒绝(calc 返回 None),
# 继续尝试下一个素材 # 继续轮询下一个素材
asset_id = "" asset_id = ""
clip_duration = 0.0 clip_duration = 0.0
start_time: float | None = None start_time: float | None = None
# 动态按使用次数排序:优先选使用最少的素材,同次数随机打散 n_assets = len(asset_ids)
asset_use_counts = {aid: len(used_segments.get(aid, [])) for aid in asset_ids} for offset in range(n_assets):
# 排序键:smart_match 评分(注入随机噪声)→ 使用次数 → 纯随机。 candidate = asset_ids[(i + offset) % n_assets]
# 噪声让得分接近的素材排名每次浮动,避免同一批素材反复选出相同组合,
# 从素材组合层面降低成片查重率;分差 > SCORE_RANDOM_NOISE_MAX 时排名稳定,
# 质量差距显著的素材仍保持优先级。
sorted_candidates = sorted(
asset_ids,
key=lambda aid: (
-(asset_smart_scores.get(aid, 0.0) + random.uniform(0.0, SCORE_RANDOM_NOISE_MAX)),
asset_use_counts.get(aid, 0),
random.random(),
),
)
for candidate in sorted_candidates:
candidate_total = asset_durations.get(candidate, 0.0) candidate_total = asset_durations.get(candidate, 0.0)
if candidate_total <= 0: if candidate_total <= 0:
continue continue
@@ -771,30 +667,16 @@ def create_clips_from_assets_editor(
candidate, candidate,
) )
continue continue
# 起始时间选取(不调用 MediaKit,保证接口快速返回) # 随机起始时间(不调用 MediaKit,保证接口快速返回)100 次避不开
# 1) 素材有场景切换点缓存时,优先从随机镜头段中选起点(不同片段来自不同镜头, # 历史区间时走受控复用回调(复用片段累加 reused_durations,回调内部
# 画面内容本质不同),与 used_segments 做冲突避让(含 1.5s 边缘间隙 # 预判复用后占比超 15% 则拒绝并返回 None
# 2) 无缓存 / 镜头段全冲突 → _calc_random_start_time 随机起点兜底; candidate_start = _calc_random_start_time(
# 100 次避不开历史区间时走受控复用回调(复用片段累加 reused_durations candidate,
# 回调内部预判复用后占比超 10% 则拒绝并返回 None) candidate_duration,
candidate_start = None asset_durations,
if candidate in asset_scene_points: used_segments,
candidate_start = pick_scene_aware_start( on_exhausted=reuse_cb,
candidate, )
candidate_duration,
asset_durations,
asset_scene_points,
used_segments,
edge_gap=SEGMENT_EDGE_GAP,
)
if candidate_start is None:
candidate_start = _calc_random_start_time(
candidate,
candidate_duration,
asset_durations,
used_segments,
on_exhausted=reuse_cb,
)
if candidate_start is None: if candidate_start is None:
# 该素材可用区间耗尽且复用被闸门/use_count 上限拒绝 → 尝试下一素材 # 该素材可用区间耗尽且复用被闸门/use_count 上限拒绝 → 尝试下一素材
logger.info( logger.info(
@@ -809,7 +691,6 @@ def create_clips_from_assets_editor(
if not asset_id or start_time is None: if not asset_id or start_time is None:
# 所有素材时长缺失、复用占比超阈值,或区间耗尽且复用被拒 → 素材可切区间不足 # 所有素材时长缺失、复用占比超阈值,或区间耗尽且复用被拒 → 素材可切区间不足
all_assets_exhausted = True
raise HTTPException( raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST, status_code=status.HTTP_400_BAD_REQUEST,
detail="素材可切区间不足,请补充新素材", detail="素材可切区间不足,请补充新素材",
@@ -825,7 +706,7 @@ def create_clips_from_assets_editor(
clips_data.append( clips_data.append(
{ {
"order": _seg_order, "order": i,
"asset_id": asset_id, "asset_id": asset_id,
"start_time": start_time, "start_time": start_time,
"duration": clip_duration, "duration": clip_duration,
@@ -833,9 +714,6 @@ def create_clips_from_assets_editor(
} }
) )
# 按原始 segment order 排序,确保 clips_data 的 order 字段有序(0,1,2,3...
clips_data.sort(key=lambda c: c["order"])
# 4. 事务性替换:清空旧片段 → 创建新片段 → 标记ready(单事务,失败自动回滚) # 4. 事务性替换:清空旧片段 → 创建新片段 → 标记ready(单事务,失败自动回滚)
created_count = plan_svc.replace_all_clips_transactional(plan_id, clips_data) created_count = plan_svc.replace_all_clips_transactional(plan_id, clips_data)
@@ -855,31 +733,11 @@ def create_clips_from_assets_editor(
unique_asset_ids, unique_asset_ids,
) )
# 6. 估算成片内部重复率(本条成片中同一素材的重叠片段时长占比) # 6. 立即返回响应
dup_rate = _calc_plan_internal_duplicate_rate(clips_data)
duplicate_warning = None
if dup_rate > 50:
duplicate_warning = f"查重率 {dup_rate:.1f}% 超过50%,建议更换素材或模板"
logger.exception(
"from-assets 成片查重率超标: plan_id=%s dup_rate=%.1f%%",
plan_id,
dup_rate,
)
# 7. 素材耗尽提示
exhaustion_warning = None
if all_assets_exhausted and created_count < len(segments):
exhaustion_warning = (
"素材可切区间不足,部分片段使用了复用素材。" "建议:1) 补充更多素材到素材库 2) 使用不同的素材组合生成"
)
# 8. 立即返回响应
return ClipsFromAssetsResponse( return ClipsFromAssetsResponse(
created_count=created_count, created_count=created_count,
plan_id=plan_id, plan_id=plan_id,
clip_ids=[], clip_ids=[],
duplicate_warning=duplicate_warning,
exhaustion_warning=exhaustion_warning,
) )
@@ -887,15 +745,7 @@ def _update_mediakit_recommendations_async( # pragma: no cover
plan_id: str, plan_id: str,
asset_ids: list[str], asset_ids: list[str],
) -> None: ) -> None:
"""后台任务:使SceneChange 智能选并更新片段的起始时间. """后台任务:MediaKit 智能选并更新片段的起始时间.
优先使用 SceneChange 策略检测视频镜头切换点,将每个素材按镜头段拆分,
各片段优先从不同镜头段中选取起始时间,实现「不同片段展示不同场景」的效果。
降级策略:
1. SceneChange 优先 → detect_scene_changes 内部已含 TimeInterval 降级
2. 若 detect_scene_changes 仍返回 None → 回退到旧的 analyze_videos 方式
3. 所有方式都失败 → 保持现有随机 start_time,不影响视频生成
此函数在后台异步执行,不影响接口响应时间。 此函数在后台异步执行,不影响接口响应时间。
失败时静默处理,不影响已创建的片段。 失败时静默处理,不影响已创建的片段。
@@ -917,6 +767,12 @@ def _update_mediakit_recommendations_async( # pragma: no cover
asset_repo = SQLAlchemyAssetRepository(db) asset_repo = SQLAlchemyAssetRepository(db)
plan_svc = EditPlanService(db) plan_svc = EditPlanService(db)
# 调用 MediaKit 获取推荐时间
recommendations = _get_mediakit_recommendations(asset_ids, asset_repo)
if not recommendations:
logger.info("后台任务: MediaKit 无推荐结果,跳过更新")
return
# 查询该 plan 的所有片段(分批获取,避免硬编码 limit 截断) # 查询该 plan 的所有片段(分批获取,避免硬编码 limit 截断)
batch_size = 500 batch_size = 500
all_clips = [] all_clips = []
@@ -939,16 +795,15 @@ def _update_mediakit_recommendations_async( # pragma: no cover
unique_asset_ids = list({getattr(c, "asset_id", "") or "" for c in clips} - {""}) unique_asset_ids = list({getattr(c, "asset_id", "") or "" for c in clips} - {""})
assets_map: dict[str, object] = {a.id: a for a in asset_repo.find_by_ids(unique_asset_ids)} assets_map: dict[str, object] = {a.id: a for a in asset_repo.find_by_ids(unique_asset_ids)}
# 按 asset_id 预分组片段对象(按 order 排序,保证按模板顺序分配镜头段 # 按 asset_id 预分组片段时间段(消除 O(N^2) 嵌套循环
clips_by_asset: dict[str, list] = defaultdict(list) clips_by_asset: dict[str, list[tuple[str, float, float]]] = defaultdict(list)
for clip in clips: for clip in clips:
aid = getattr(clip, "asset_id", "") or "" aid = getattr(clip, "asset_id", "") or ""
if aid: if aid and clip.start_time is not None:
clips_by_asset[aid].append(clip) clips_by_asset[aid].append((clip.id, clip.start_time, clip.start_time + clip.duration))
for aid in clips_by_asset:
clips_by_asset[aid].sort(key=lambda c: c.order)
# 读取素材全部历史已用区间(跨任务/跨 plan 持久化记录) # 读取素材全部历史已用区间(跨任务/跨 plan 持久化记录)
# MediaKit 挪点必须与随机选片一样避让历史区间,否则会把片段挪回已用过的画面
historical_segments = get_used_segments(db, unique_asset_ids) historical_segments = get_used_segments(db, unique_asset_ids)
# 已更新的片段ID(用于排除已移动的旧时间段) # 已更新的片段ID(用于排除已移动的旧时间段)
@@ -957,22 +812,16 @@ def _update_mediakit_recommendations_async( # pragma: no cover
updated_segments: dict[str, list[tuple[float, float]]] = {} updated_segments: dict[str, list[tuple[float, float]]] = {}
updated_count = 0 updated_count = 0
# 尝试获取存储服务(用于生成视频 URL) # 遍历片段,按 asset_id 匹配推荐时间
try: for clip in clips:
storage = get_storage_service() asset_id = getattr(clip, "asset_id", "") or ""
except Exception as e: if not asset_id or asset_id not in recommendations:
logger.exception("后台任务: 获取存储服务失败,跳过 SceneChange 更新: %s", e)
return
# 获取 MediaKit 客户端
client = get_mediakit_client()
# 对每个素材,检测场景切换点并分配镜头段
for asset_id in unique_asset_ids:
asset_clips = clips_by_asset.get(asset_id, [])
if not asset_clips:
continue continue
recommended_start = recommendations[asset_id]
clip_duration = clip.duration
# 从预加载字典获取素材(O(1) 查找)
asset = assets_map.get(asset_id) asset = assets_map.get(asset_id)
if not asset: if not asset:
continue continue
@@ -980,181 +829,95 @@ def _update_mediakit_recommendations_async( # pragma: no cover
if asset_total <= 0: if asset_total <= 0:
continue continue
# 获取素材视频 URL # 推荐时间 + 片段时长不能超过素材总时长
video_url: str | None = None if recommended_start + clip_duration > asset_total:
storage_key = getattr(asset, "storage_key", None) or ""
mime = getattr(asset, "mime_type", "") or ""
if storage_key and mime.startswith("video/"):
try:
video_url = storage.get_download_url(storage_key)
except Exception:
logger.exception("后台任务: 获取素材URL失败: asset_id=%s", asset_id)
# 构建该素材的占用区间列表(排除已更新片段)
def _get_other_segments(asset_id_inner, clip_id_inner):
segs: list[tuple[float, float]] = []
for c in clips_by_asset.get(asset_id_inner, []):
cid = c.id
if cid != clip_id_inner and cid not in updated_clip_ids:
segs.append((c.start_time, c.start_time + c.duration))
segs.extend(updated_segments.get(asset_id_inner, []))
# 并入历史已用区间
def _norm(segs_in):
return {(round(float(a), 3), round(float(b), 3)) for a, b in segs_in}
return list(_norm(segs) | _norm(historical_segments.get(asset_id_inner, [])))
# 优先使用 SceneChange 策略
scene_segments: list[tuple[float, float]] = []
# 先查素材 metadata 中的场景点缓存:命中则直接复用,跳过 MediaKit 检测
# (缓存由本任务首次检测后写入,跨任务/跨 plan 复用)
cached_points = extract_scene_points_from_metadata(getattr(asset, "metadata", None))
if cached_points:
scene_segments = build_scene_segments(cached_points, asset_total)
logger.info( logger.info(
"后台任务: 命中场景点缓存: asset_id=%s scenes=%d", "后台任务: 推荐时间越界,跳过: asset_id=%s recommended=%.2f duration=%.1f total=%.1f",
asset_id,
len(scene_segments),
)
if not scene_segments and client.is_available and video_url:
scene_changes = client.detect_scene_changes(video_url)
if scene_changes is not None:
scene_segments = build_scene_segments(scene_changes, asset_total)
logger.info(
"后台任务: 素材场景检测完成: asset_id=%s scenes=%d",
asset_id,
len(scene_segments),
)
# 检测结果写入素材 metadata 缓存:首次生成用随机起点,
# 检测完成后后续生成的渲染前同步路径即可读缓存选镜头段
try:
existing_meta = dict(getattr(asset, "metadata", None) or {})
existing_meta["scene_change_points"] = scene_changes
asset.metadata = existing_meta # type: ignore[attr-defined]
asset_repo.update(asset) # type: ignore[arg-type]
logger.info(
"后台任务: 场景点已写入素材缓存: asset_id=%s points=%d",
asset_id,
len(scene_changes),
)
except Exception:
# 缓存写入失败不影响本次片段更新
logger.exception(
"后台任务: 场景点缓存写入失败: asset_id=%s",
asset_id,
)
# SceneChange 未获得有效结果 → 尝试 analyze_videos 作为 fallback
if not scene_segments and video_url:
fallback_recs = _get_mediakit_recommendations([asset_id], asset_repo)
if fallback_recs and asset_id in fallback_recs:
# analyze_videos 只返回单个推荐点,转为单镜头段
rec_start = fallback_recs[asset_id]
scene_segments = [(rec_start, asset_total)]
logger.info(
"后台任务: 使用 analyze_videos fallback: asset_id=%s start=%.2f",
asset_id,
rec_start,
)
if not scene_segments:
# 所有方式都失败 → 保持现有随机 start_time
logger.info(
"后台任务: SceneChange 与 analyze_videos 均无结果,保持随机起点: asset_id=%s",
asset_id, asset_id,
recommended_start,
clip_duration,
asset_total,
) )
continue continue
# 为每个片段分配不同的镜头段 # 构建排除当前片段及已更新片段后的占用列表(O(M),M=同素材片段数)
scene_segments_pool = list(scene_segments) # 可消费的镜头段池 other_segments: list[tuple[float, float]] = [
for clip in asset_clips: (cs, ce)
clip_duration = clip.duration for cid, cs, ce in clips_by_asset.get(asset_id, [])
recommended_start: float | None = None if cid != clip.id and cid not in updated_clip_ids
]
other_segments.extend(updated_segments.get(asset_id, []))
# 从镜头段池中依次尝试,选一个不冲突的 # 并入该素材全部历史已用区间(含其他 plan/其他任务),set 去重:
for seg_idx, (seg_start, seg_end) in enumerate(scene_segments_pool): # 本 plan 片段创建时已写入历史记录
candidate_start = pick_start_in_scene_segment(seg_start, seg_end, clip_duration) # 并入该素材全部历史已用区间(含其他 plan/其他任务)。
if candidate_start is None: # set 去重前先归一化精度(round 3 位),避免浮点尾差导致逻辑相同的
continue # 镜头段太短,跳过 # 区间(如 1.0 与 1.0000000001)被误判为不同区间
def _norm(segs):
return {(round(float(a), 3), round(float(b), 3)) for a, b in segs}
# 检查越界 other_segments = list(_norm(other_segments) | _norm(historical_segments.get(asset_id, [])))
if candidate_start + clip_duration > asset_total:
continue
# 检查与已用区间冲突 # 检查推荐时间是否与同 plan 片段或历史已用区间冲突(含 0.3s 边缘间隙):
other_segs = _get_other_segments(asset_id, clip.id) # 冲突时放弃该推荐、保留原随机起点(不硬挪到已用过的画面)
if _recommended_time_conflicts(candidate_start, clip_duration, other_segs): if _recommended_time_conflicts(recommended_start, clip_duration, other_segments):
continue logger.info(
"后台任务: 推荐时间与同片/历史区间冲突,保留原起点: asset_id=%s recommended=%.2f",
asset_id,
recommended_start,
)
continue
recommended_start = candidate_start # 逐个更新并捕获异常(单点失败不影响其他片段)
# 消费该镜头段(从池中移除,下一个片段用不同镜头段) try:
scene_segments_pool.pop(seg_idx) old_start = clip.start_time
break old_end = old_start + clip_duration
# MediaKit 移动片段起点 + 同步素材 metadata 区间记录放在同一事务:
if recommended_start is None: # 删旧区间记录(按 plan_id + 旧 start 匹配,兼容无 plan_id 的旧数据)、
# 镜头段用完或都冲突 → 尝试 _calc_random_start_time 兜底 # 写新区间,最后统一 commit;任一步失败整体 rollback
used_segs_for_calc: dict[str, list[tuple[float, float]]] = { # 保证 clip.start_time 与 metadata.used_time_ranges 不出现不一致。
asset_id: _get_other_segments(asset_id, clip.id) plan_svc.update_clip(clip.id, start_time=recommended_start)
}
fallback_start = _calc_random_start_time(
asset_id,
clip_duration,
{asset_id: asset_total},
used_segs_for_calc,
)
if fallback_start is None:
continue # 完全无法分配,保持原起点
recommended_start = fallback_start
# 更新片段起始时间
try: try:
old_start = clip.start_time if remove_used_segment(db, asset_id, old_start, old_end, plan_id=plan_id):
old_end = old_start + clip_duration record_used_segments(
db,
plan_svc.update_clip(clip.id, start_time=recommended_start) asset_id,
try: recommended_start,
if remove_used_segment(db, asset_id, old_start, old_end, plan_id=plan_id): recommended_start + clip_duration,
record_used_segments( plan_id,
db,
asset_id,
recommended_start,
recommended_start + clip_duration,
plan_id,
)
except Exception:
logger.exception(
"后台任务: 同步素材区间记录失败,回滚本次片段更新: clip_id=%s",
clip.id,
) )
db.rollback() except Exception as me:
continue logger.warning(
db.commit() "后台任务: 同步素材区间记录失败,回滚本次片段更新: clip_id=%s error=%s",
updated_count += 1
updated_clip_ids.add(clip.id)
updated_segments.setdefault(asset_id, []).append(
(recommended_start, recommended_start + clip_duration)
)
logger.info(
"后台任务: 更新片段起始时间(场景选帧): clip_id=%s asset_id=%s start_time=%.2f",
clip.id, clip.id,
asset_id, me,
recommended_start,
) )
except Exception: db.rollback()
logger.exception("后台任务: 单个片段更新失败: clip_id=%s", clip.id)
try:
db.rollback()
except Exception:
pass
continue continue
db.commit()
updated_count += 1
updated_clip_ids.add(clip.id)
except Exception as ue:
logger.warning("后台任务: 单个片段更新失败: clip_id=%s error=%s", clip.id, ue)
try:
db.rollback()
except Exception:
pass
continue
updated_segments.setdefault(asset_id, []).append((recommended_start, recommended_start + clip_duration))
logger.info(
"后台任务: 更新片段起始时间: clip_id=%s asset_id=%s start_time=%.2f",
clip.id,
asset_id,
recommended_start,
)
logger.info("后台任务完成: plan_id=%s 成功更新 %d 个片段", plan_id, updated_count) logger.info("后台任务完成: plan_id=%s 成功更新 %d 个片段", plan_id, updated_count)
except Exception: except Exception as e:
# 后台任务失败不影响已创建的片段,静默处理 # 后台任务失败不影响已创建的片段,静默处理
logger.exception("后台任务异常: plan_id=%s", plan_id) logger.warning("后台任务异常: plan_id=%s error=%s", plan_id, e, exc_info=True)
if db: if db:
try: try:
db.rollback() db.rollback()
@@ -41,33 +41,29 @@ def get_draft_plan_id(
这是模板编辑器路由的核心依赖——所有编辑器端点都先经过这里, 这是模板编辑器路由的核心依赖——所有编辑器端点都先经过这里,
确保 template_id → plan_id 的映射始终存在。 确保 template_id → plan_id 的映射始终存在。
模板读取遵循单一数据源、显式判定(不使用异常降级): 兼容策略:优先从新模板系统(edit_templates 表)查找,
- 用户自建模板在旧表 ``templates``(归属 user_idis_active=True); 若不存在则回退到旧模板系统(templates 表),确保用户自建模板可用。
- 全局模板在新表 ``edit_templates``(无 user_id,全局可读)。
模板不存在、已删除或不归属于当前用户时,一律返回 404。
""" """
tpl_svc, plan_svc = services tpl_svc, plan_svc = services
user_id = str(current_user.user.id) user_id = str(current_user.user.id)
# 0. 门禁:校验模板存在且可访问(即使草稿已缓存命中也要校验,
# 避免模板被删除/无权访问后仍可通过既有草稿 plan 继续操作)。
old_repo = SQLAlchemyTemplateRepository(db)
old_template = old_repo.get_active(template_id, user_id)
is_global_template = tpl_svc.get_template(template_id) is not None
if old_template is None and not is_global_template:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="模板不存在")
# 1. 草稿已存在 → 直接返回 # 1. 草稿已存在 → 直接返回
draft = tpl_svc.get_template_draft(template_id) draft = tpl_svc.get_template_draft(template_id)
if draft is not None: if draft is not None:
return draft.id return draft.id
# 2. 全局模板(新系统)→ 用新服务创建草稿 # 2. 新系统有模板 → 用新服务创建草稿
if is_global_template: if tpl_svc.get_template(template_id) is not None:
draft = tpl_svc.create_template_draft(template_id, user_id=user_id) draft = tpl_svc.create_template_draft(template_id, user_id=user_id)
return draft.id return draft.id
# 3. 旧模板(templates 表)→ 基于旧模板创建草稿计划 # 3. 回退到旧模板系统templates 表)
old_repo = SQLAlchemyTemplateRepository(db)
old_template = old_repo.get(template_id, user_id=user_id)
if old_template is None:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="模板不存在")
# 4. 基于旧模板创建草稿计划
from app.services.plan_generator_service import PlanGeneratorService from app.services.plan_generator_service import PlanGeneratorService
from packages.domain.edit_template import EditTemplate, EditTemplateStatus from packages.domain.edit_template import EditTemplate, EditTemplateStatus
@@ -150,7 +150,7 @@ def rollback_template(
try: try:
tpl = tpl_svc.rollback_to_version(template_id, request.version) tpl = tpl_svc.rollback_to_version(template_id, request.version)
except ValueError as exc: except ValueError as exc:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=str(exc)) from exc raise HTTPException(status_code=400, detail=str(exc)) from exc
clip_configs = tpl_svc.list_clip_configs(template_id) clip_configs = tpl_svc.list_clip_configs(template_id)
return EditorRollbackResponse( return EditorRollbackResponse(
@@ -41,17 +41,17 @@ def list_editor_transition_presets(
_: AuthenticatedUser = Depends(get_current_user), _: AuthenticatedUser = Depends(get_current_user),
) -> TransitionPresetListResponse: ) -> TransitionPresetListResponse:
"""获取转场预设列表""" """获取转场预设列表"""
from packages.domain.transition_presets import TRANSITION_PRESET_LIBRARY from packages.domain.transition_presets import TRANSITION_PRESETS
items = [ items = [
{ {
"id": p.id, "id": p["id"],
"name": p.name, "name": p["name"],
"category": p.category, "category": p.get("category", "通用"),
"duration": p.default_duration, "duration": p.get("default_duration", 0.5),
"description": p.description, "description": p.get("description", ""),
} }
for p in TRANSITION_PRESET_LIBRARY for p in TRANSITION_PRESETS
] ]
return TransitionPresetListResponse(items=items, total=len(items)) return TransitionPresetListResponse(items=items, total=len(items))
@@ -123,17 +123,17 @@ def list_editor_filter_presets(
_: AuthenticatedUser = Depends(get_current_user), _: AuthenticatedUser = Depends(get_current_user),
) -> FilterPresetListResponse: ) -> FilterPresetListResponse:
"""获取滤镜预设列表""" """获取滤镜预设列表"""
from packages.domain.filter_presets import FILTER_PRESET_LIBRARY from packages.domain.filter_presets import FILTER_PRESETS
items = [ items = [
{ {
"id": p.id, "id": p["id"],
"name": p.name, "name": p["name"],
"category": p.category, "category": p.get("category", "通用"),
"thumbnail": p.lut_url, "thumbnail": p.get("thumbnail", ""),
"description": p.description, "description": p.get("description", ""),
} }
for p in FILTER_PRESET_LIBRARY for p in FILTER_PRESETS
] ]
return FilterPresetListResponse(items=items, total=len(items)) return FilterPresetListResponse(items=items, total=len(items))
@@ -189,8 +189,6 @@ class ClipsFromAssetsResponse(BaseModel):
plan_id: str = "" plan_id: str = ""
message: str = "" message: str = ""
clip_ids: List[str] = Field(default_factory=list, description="创建的片段ID列表") clip_ids: List[str] = Field(default_factory=list, description="创建的片段ID列表")
duplicate_warning: Optional[str] = Field(default=None, description="查重率超标警告")
exhaustion_warning: Optional[str] = Field(default=None, description="素材耗尽警告")
# ── 封面配置 ──────────────────────────────────────────────────────────────── # ── 封面配置 ────────────────────────────────────────────────────────────────
+2 -32
View File
@@ -2,9 +2,7 @@
from __future__ import annotations from __future__ import annotations
import json
import logging import logging
import subprocess
import tempfile import tempfile
from pathlib import Path from pathlib import Path
from typing import Any, Optional from typing import Any, Optional
@@ -173,14 +171,6 @@ def synthesize(
# job.voice_id 统一存解析后的 CosyVoice voice_id # job.voice_id 统一存解析后的 CosyVoice voice_id
actual_voice_id = resolved_profile.voice_id actual_voice_id = resolved_profile.voice_id
# 语速/情绪等合成参数随 metadata 落库,workflow 提交 CosyVoice 时读取透传
synthesis_meta = {
"speed": request.speed,
"emotion": request.emotion or "",
}
if request.metadata_:
synthesis_meta.update(request.metadata_)
use_case = CreateTTSJobUseCase(repository) use_case = CreateTTSJobUseCase(repository)
job = use_case.execute( job = use_case.execute(
user_id=user_id, user_id=user_id,
@@ -188,7 +178,7 @@ def synthesize(
voice_id=actual_voice_id, voice_id=actual_voice_id,
voice_model=request.voice_model, voice_model=request.voice_model,
voice_clone_profile_id=voice_clone_profile_id, voice_clone_profile_id=voice_clone_profile_id,
metadata=synthesis_meta, metadata=request.metadata_,
) )
# 提交 CosyVoice 合成任务 # 提交 CosyVoice 合成任务
@@ -440,8 +430,6 @@ def save_tts_job_to_library(
storage_key = f"uploads/voice/tts/{job.id}.{audio_format}" storage_key = f"uploads/voice/tts/{job.id}.{audio_format}"
tmp_path: Path | None = None tmp_path: Path | None = None
audio_duration: float | None = None
file_size = 0
try: try:
with tempfile.NamedTemporaryFile(suffix=f".{audio_format}", delete=False) as tmp: with tempfile.NamedTemporaryFile(suffix=f".{audio_format}", delete=False) as tmp:
tmp_path = Path(tmp.name) tmp_path = Path(tmp.name)
@@ -457,23 +445,6 @@ def save_tts_job_to_library(
) )
file_size = tmp_path.stat().st_size file_size = tmp_path.stat().st_size
storage_service.upload_file(tmp_path, storage_key, content_type=content_type) storage_service.upload_file(tmp_path, storage_key, content_type=content_type)
# 从音频文件提取时长(ffprobe),作为 job.duration 的兜底
try:
proc = subprocess.run(
[
"ffprobe", "-v", "quiet", "-print_format", "json",
"-show_format", str(tmp_path),
],
capture_output=True, text=True, timeout=10,
)
if proc.returncode == 0:
fmt = json.loads(proc.stdout).get("format", {})
dur = float(fmt.get("duration", 0))
if dur > 0:
audio_duration = dur
except Exception:
logger.warning("ffprobe 提取时长失败: job_id=%s", job.id, exc_info=True)
except HTTPException: except HTTPException:
raise raise
except Exception as e: except Exception as e:
@@ -511,7 +482,7 @@ def save_tts_job_to_library(
mime_type=content_type, mime_type=content_type,
metadata=metadata_, metadata=metadata_,
file_size=file_size, file_size=file_size,
duration=job.duration or audio_duration or None, duration=job.duration or None,
status=AssetStatus.READY, status=AssetStatus.READY,
classification_status=ClassificationStatus.PENDING, # 音频不参与内容分类,保持 pending 与 ingest 链路一致 classification_status=ClassificationStatus.PENDING, # 音频不参与内容分类,保持 pending 与 ingest 链路一致
uploaded_by_user_id=user_id, uploaded_by_user_id=user_id,
@@ -575,7 +546,6 @@ def preview_tts(
text=request.text, text=request.text,
voice_id=actual_voice_id, voice_id=actual_voice_id,
speed=request.speed, speed=request.speed,
emotion=request.emotion,
) )
except CosyVoiceError as e: except CosyVoiceError as e:
raise HTTPException( raise HTTPException(
+49 -341
View File
@@ -23,7 +23,6 @@ from app.schemas.upload import (
from fastapi import APIRouter, Depends, File, Form, HTTPException, UploadFile, status from fastapi import APIRouter, Depends, File, Form, HTTPException, UploadFile, status
from packages.application import SubmitIngestJobCommand, SubmitIngestJobUseCase from packages.application import SubmitIngestJobCommand, SubmitIngestJobUseCase
from packages.domain import Asset, AssetStatus
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -81,199 +80,12 @@ def _validate_mime_type(content_type: str | None) -> str:
return base_type return base_type
def _infer_mime_type_from_storage_key(storage_key: str) -> str:
"""从 storage_key 推断 MIME 类型(与 worker 端保持一致)。"""
lower_filename = storage_key.rsplit("/", 1)[-1].lower()
_MIME_MAP = {
".mov": "video/quicktime",
".mp4": "video/mp4",
".avi": "video/x-msvideo",
".mkv": "video/x-matroska",
".webm": "video/webm",
".png": "image/png",
".gif": "image/gif",
".bmp": "image/bmp",
".svg": "image/svg+xml",
".jpg": "image/jpeg",
".jpeg": "image/jpeg",
".mp3": "audio/mpeg",
".wav": "audio/wav",
".ogg": "audio/ogg",
".flac": "audio/flac",
".m4a": "audio/x-m4a",
}
for ext, mime in _MIME_MAP.items():
if lower_filename.endswith(ext):
return mime
return "video/mp4" # default
# 兜底去重:无 file_hash / client_upload_id 且大小已知时,同库同名同大小近期活动记录视为重复
FALLBACK_DEDUP_WINDOW_MINUTES = 30
def _find_duplicate_asset(
asset_repository: Any,
*,
library_id: str,
file_hash: str,
client_upload_id: str,
filename: str,
file_size: int = 0,
) -> Any:
"""complete/上传幂等去重,按优先级查找已存在的素材。
1. client_upload_id(客户端幂等 token,同一次上传的重试保持一致)
2. file_hash(内容哈希,不同上传只要内容相同即去重)
3. 兜底(严格模式,宁可漏判不可误杀):file_hash 与 client_upload_id
均缺失、且 file_size > 0 时,同库 + 同文件名 + **同大小** 且 30 分钟内
仍处 uploading/processing 的记录才判重。
- file_hash 非空时跳过兜底(hash 已代表内容;同名但内容全新的视频
如 iPhone 的 IMG_xxxx.MOV 绝不能被同名占位误杀)
- file_size=0(未知)时不允许仅凭同名 + processing 判重,直接放行
全部为鸭子类型调用:旧仓储无对应方法时静默跳过,不破坏既有实现。
"""
if client_upload_id:
find = getattr(asset_repository, "find_by_library_and_client_upload_id", None)
if callable(find):
existing = find(library_id=library_id, client_upload_id=client_upload_id)
if existing is not None:
logger.info(
"素材幂等命中(client_upload_id): library=%s token=%s asset=%s",
library_id,
client_upload_id,
getattr(existing, "id", "?"),
)
return existing
if file_hash:
existing = asset_repository.find_by_library_and_file_hash(
library_id=library_id,
file_hash=file_hash,
)
if existing is not None:
logger.info(
"素材去重命中(file_hash): library=%s hash=%s asset=%s",
library_id,
file_hash,
existing.id,
)
return existing
# 同名兜底去重(最后防线,严格模式):
# - 仅当 file_hash / client_upload_id 均缺失时启用(hash 能代表内容时不靠同名猜)
# - file_size 必须 > 0 且与记录大小严格一致;大小未知(0)直接放行
# - 只命中近期 UPLOADING/PROCESSING 活动记录(READY 历史素材不拦)
if filename and not file_hash and not client_upload_id and file_size and file_size > 0:
find_recent = getattr(asset_repository, "find_recent_active_by_library_and_name", None)
if callable(find_recent):
existing = find_recent(
library_id=library_id,
name=filename,
within_minutes=FALLBACK_DEDUP_WINDOW_MINUTES,
file_size=file_size,
)
if existing is not None:
logger.info(
"素材幂等兜底命中(近期同名同大小活动记录): library=%s name=%s asset=%s status=%s size=%s",
library_id,
filename,
getattr(existing, "id", "?"),
getattr(existing, "status", None),
file_size,
)
return existing
elif filename and not file_hash and not client_upload_id and not file_size:
logger.debug(
"同名兜底去重跳过(file_size 未知,宁可放行不可误杀): library=%s name=%s",
library_id,
filename,
)
return None
def _create_pending_asset(
asset_repository,
project_id,
library_id,
storage_key,
filename,
mime_type,
user_id,
file_hash="",
client_upload_id="",
file_size: int = 0,
):
"""立即创建或复用一条 PROCESSING 状态的 Asset 记录。
find-or-createprepare 阶段已按 file_hash/client_upload_id 预建的占位记录
会被 find_by_library_and_file_hash/find_by_library_and_client_upload_id 命中,
直接复用并补齐字段(避免 pre-create + complete 重复建两条)。
Issue #1776: 素材库计数由 asset_repository.create() 自动维护。
"""
# 1. 按 client_upload_id / file_hash 查找现有记录
existing = None
if client_upload_id:
find_by_cuid = getattr(asset_repository, "find_by_library_and_client_upload_id", None)
if callable(find_by_cuid):
existing = find_by_cuid(library_id=library_id, client_upload_id=client_upload_id)
if existing is None and file_hash:
existing = asset_repository.find_by_library_and_file_hash(library_id=library_id, file_hash=file_hash)
if existing is not None:
# 补齐字段(幂等:避免重复建记录,前端已拿到 asset_id)
changed = False
if file_hash and not existing.file_hash:
existing.file_hash = file_hash
changed = True
if client_upload_id and not existing.client_upload_id:
existing.client_upload_id = client_upload_id
changed = True
if file_size and not existing.file_size:
existing.file_size = file_size
changed = True
if existing.status not in (AssetStatus.PROCESSING, AssetStatus.UPLOADING):
existing.status = AssetStatus.PROCESSING
changed = True
if changed:
try:
asset_repository.update(existing)
except Exception: # noqa: BLE001 — 字段补齐失败不阻塞主流程
pass
return existing
asset = Asset.create(
project_id=project_id,
library_id=library_id,
name=filename,
storage_key=storage_key,
mime_type=mime_type,
status=AssetStatus.PROCESSING,
uploaded_by_user_id=user_id,
file_hash=file_hash,
client_upload_id=client_upload_id,
file_size=file_size,
)
return asset_repository.create(asset)
def _persist_celery_task_id(repo: Any, job: Any, celery_task_id: str) -> None:
"""记录 celery 消息 ID 到任务行,供孤儿清理时 revoke/清除队列消息(#1714)。"""
if not celery_task_id:
return
try:
job.celery_task_id = celery_task_id
repo.update(job)
except Exception: # noqa: BLE001 — 记录失败不影响主流程(执行前状态守卫兜底)
pass
def _submit_ingest_job( def _submit_ingest_job(
project_id: str, project_id: str,
library_id: str, library_id: str,
storage_key: str, storage_key: str,
ingest_job_repository: Any, ingest_job_repository: Any,
file_hash: str = "", file_hash: str = "",
asset_id: str = "",
) -> Any: ) -> Any:
use_case = SubmitIngestJobUseCase(ingest_job_repository) use_case = SubmitIngestJobUseCase(ingest_job_repository)
job = use_case.execute( job = use_case.execute(
@@ -282,11 +94,9 @@ def _submit_ingest_job(
library_id=library_id, library_id=library_id,
storage_key=storage_key, storage_key=storage_key,
file_hash=file_hash, file_hash=file_hash,
asset_id=asset_id,
) )
) )
celery_result = celery_app.send_task("worker.ingest_asset", args=[job.id]) celery_app.send_task("worker.ingest_asset", args=[job.id])
_persist_celery_task_id(ingest_job_repository, job, getattr(celery_result, "id", ""))
return job return job
@@ -296,15 +106,9 @@ async def prepare_direct_upload(
authenticated_user: AuthenticatedUser = Depends(get_current_user), authenticated_user: AuthenticatedUser = Depends(get_current_user),
project_repository: Any = Depends(get_project_repository), project_repository: Any = Depends(get_project_repository),
asset_library_repository: Any = Depends(get_asset_library_repository), asset_library_repository: Any = Depends(get_asset_library_repository),
asset_repository: Any = Depends(get_asset_repository),
storage_service: OSSStorageService = Depends(get_storage_service), storage_service: OSSStorageService = Depends(get_storage_service),
) -> DirectUploadPrepareResponse: ) -> DirectUploadPrepareResponse:
"""创建浏览器直传 OSS 的短期表单签名,并在签名前按 file_hash/client_upload_id 去重。 """创建浏览器直传 OSS 的短期表单签名"""
命中去重:直接返回 duplicated=True + skip_transfer=True(前端跳过 OSS 直传),
未命中:正常签名 OSS 并立即预建一条 PROCESSING 状态的 asset 记录占住
file_hash 闸门,响应带 asset_id 供前端/后续 complete 关联。
"""
settings = get_settings() settings = get_settings()
max_size_bytes = settings.OSS_DIRECT_UPLOAD_MAX_MB * 1024 * 1024 max_size_bytes = settings.OSS_DIRECT_UPLOAD_MAX_MB * 1024 * 1024
if request.file_size > max_size_bytes: if request.file_size > max_size_bytes:
@@ -323,39 +127,8 @@ async def prepare_direct_upload(
asset_library_repository, asset_library_repository,
) )
safe_filename = request.filename.replace("/", "_").replace("\\", "_")
# ── prepare 阶段去重:OSS 签名之前先查已存在素材 ──
if request.file_hash or request.client_upload_id:
existing = _find_duplicate_asset(
asset_repository,
library_id=request.library_id,
file_hash=request.file_hash,
client_upload_id=request.client_upload_id,
filename=request.filename,
file_size=request.file_size,
)
if existing is not None:
logger.info(
"prepare 命中去重: library=%s hash=%s cuid=%s existing_asset=%s",
request.library_id,
request.file_hash,
request.client_upload_id,
existing.id,
)
return DirectUploadPrepareResponse(
upload_url="",
method="",
storage_key=existing.storage_key,
expires_at="",
fields={},
max_size_bytes=0,
duplicated=True,
skip_transfer=True,
asset_id=existing.id,
)
file_id = uuid4().hex[:8] file_id = uuid4().hex[:8]
safe_filename = request.filename.replace("/", "_").replace("\\", "_")
storage_key = f"uploads/{file_id}/{safe_filename}" storage_key = f"uploads/{file_id}/{safe_filename}"
try: try:
payload = storage_service.create_direct_upload_post( payload = storage_service.create_direct_upload_post(
@@ -374,28 +147,6 @@ async def prepare_direct_upload(
detail=f"Failed to prepare upload: {type(error).__name__}", detail=f"Failed to prepare upload: {type(error).__name__}",
) from error ) from error
# ── 预建 asset 占位:占住 file_hash/client_upload_id 闸门,避免并发重复上传 ──
pending_asset_id = ""
if request.file_hash or request.client_upload_id:
try:
pending = _create_pending_asset(
asset_repository=asset_repository,
project_id=request.project_id,
library_id=request.library_id,
storage_key=storage_key,
filename=safe_filename,
mime_type=validated_content_type,
user_id=authenticated_user.user.id,
file_hash=request.file_hash,
client_upload_id=request.client_upload_id,
file_size=request.file_size,
)
pending_asset_id = pending.id
# Issue #1776: 计数由 asset_repository.create() 自动维护
except Exception as error:
# 预建失败不阻塞签名:complete 仍可按 OSS 文件 + hash 兜底去重
logger.warning("预建 asset 占位失败,降级走 old flow: %s", error)
return DirectUploadPrepareResponse( return DirectUploadPrepareResponse(
upload_url=str(payload["url"]), upload_url=str(payload["url"]),
method=str(payload["method"]), method=str(payload["method"]),
@@ -403,9 +154,6 @@ async def prepare_direct_upload(
expires_at=str(payload["expires_at"]), expires_at=str(payload["expires_at"]),
fields={str(key): str(value) for key, value in dict(payload["fields"]).items()}, fields={str(key): str(value) for key, value in dict(payload["fields"]).items()},
max_size_bytes=max_size_bytes, max_size_bytes=max_size_bytes,
duplicated=False,
skip_transfer=False,
asset_id=pending_asset_id,
) )
@@ -419,7 +167,7 @@ async def complete_direct_upload(
asset_repository: Any = Depends(get_asset_repository), asset_repository: Any = Depends(get_asset_repository),
storage_service: OSSStorageService = Depends(get_storage_service), storage_service: OSSStorageService = Depends(get_storage_service),
) -> DirectUploadCompleteResponse: ) -> DirectUploadCompleteResponse:
"""确认浏览器直传完成并创建导入任务(幂等:重复 complete 返回同一素材)""" """确认浏览器直传完成并创建导入任务。"""
require_project_and_library( require_project_and_library(
request.project_id, request.project_id,
request.library_id, request.library_id,
@@ -429,29 +177,6 @@ async def complete_direct_upload(
normalized_key = storage_service._normalize_storage_key(request.storage_key) normalized_key = storage_service._normalize_storage_key(request.storage_key)
if not normalized_key.startswith("uploads/"): if not normalized_key.startswith("uploads/"):
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="Invalid upload key") raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="Invalid upload key")
filename = normalized_key.rsplit("/", 1)[-1]
# ── 幂等去重(放在 OSS 检查之前):complete 超时后前端重试时,
# 第一次 complete 可能已建好占位记录,此时即使 OSS 检查失败也必须返回
# 已存在记录,绝不能再建第二条。─
existing = _find_duplicate_asset(
asset_repository,
library_id=request.library_id,
file_hash=request.file_hash,
client_upload_id=request.client_upload_id,
filename=filename,
file_size=request.file_size,
)
if existing is not None:
return DirectUploadCompleteResponse(
storage_key=existing.storage_key,
ingest_job_id="",
duplicated=True,
asset_id=existing.id,
url=storage_service.get_url(existing.storage_key),
)
try: try:
file_exists = storage_service.file_exists(normalized_key) file_exists = storage_service.file_exists(normalized_key)
except Exception as error: except Exception as error:
@@ -463,21 +188,26 @@ async def complete_direct_upload(
if not file_exists: if not file_exists:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Uploaded file not found") raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Uploaded file not found")
# 立即创建 Asset 记录(PROCESSING 状态),使前端刷新后即可看到新素材 # ── 素材去重检测:同素材库 + 同 file_hash 视为重复 ──
mime_type = _infer_mime_type_from_storage_key(normalized_key) if request.file_hash:
pending_asset = _create_pending_asset( existing = asset_repository.find_by_library_and_file_hash(
asset_repository=asset_repository, library_id=request.library_id,
project_id=request.project_id, file_hash=request.file_hash,
library_id=request.library_id, )
storage_key=normalized_key, if existing is not None:
filename=filename, logger.info(
mime_type=mime_type, "素材去重命中: library=%s hash=%s existing_asset=%s",
user_id=authenticated_user.user.id, request.library_id,
file_hash=request.file_hash, request.file_hash,
client_upload_id=request.client_upload_id, existing.id,
file_size=request.file_size, )
) return DirectUploadCompleteResponse(
# Issue #1776: 计数由 asset_repository.create() 自动维护 storage_key=normalized_key,
ingest_job_id="",
duplicated=True,
asset_id=existing.id,
url=storage_service.get_url(normalized_key),
)
job = _submit_ingest_job( job = _submit_ingest_job(
project_id=request.project_id, project_id=request.project_id,
@@ -485,14 +215,8 @@ async def complete_direct_upload(
storage_key=normalized_key, storage_key=normalized_key,
ingest_job_repository=ingest_job_repository, ingest_job_repository=ingest_job_repository,
file_hash=request.file_hash, file_hash=request.file_hash,
asset_id=pending_asset.id,
)
return DirectUploadCompleteResponse(
storage_key=normalized_key,
ingest_job_id=job.id,
asset_id=pending_asset.id,
url=storage_service.get_url(normalized_key),
) )
return DirectUploadCompleteResponse(storage_key=normalized_key, ingest_job_id=job.id, url=storage_service.get_url(normalized_key))
@router.post( @router.post(
@@ -505,8 +229,7 @@ async def upload_asset(
project_id: str = Form(..., min_length=1, description="项目 ID"), project_id: str = Form(..., min_length=1, description="项目 ID"),
library_id: str = Form(..., min_length=1, description="素材库 ID"), library_id: str = Form(..., min_length=1, description="素材库 ID"),
file: UploadFile = File(..., description="要上传的文件(视频、音频、图片等)"), file: UploadFile = File(..., description="要上传的文件(视频、音频、图片等)"),
file_hash: str = Form(default="", description="文件哈希,用于去重检测"), file_hash: str = Form(default="", description="文件 MD5 哈希,用于去重检测"),
client_upload_id: str = Form(default="", description="客户端幂等 token(同一次上传的重试保持一致)"),
authenticated_user: AuthenticatedUser = Depends(get_current_user), authenticated_user: AuthenticatedUser = Depends(get_current_user),
ingest_job_repository: Any = Depends(get_ingest_job_repository), ingest_job_repository: Any = Depends(get_ingest_job_repository),
project_repository: Any = Depends(get_project_repository), project_repository: Any = Depends(get_project_repository),
@@ -517,31 +240,32 @@ async def upload_asset(
"""上传素材文件并触发导入流水线。""" """上传素材文件并触发导入流水线。"""
require_project_and_library(project_id, library_id, project_repository, asset_library_repository) require_project_and_library(project_id, library_id, project_repository, asset_library_repository)
# P2-5: 服务端验证 MIME 类型(先验证,再幂等去重,避免非法类型绕过) # ── 素材去重检测:上传前检查同素材库 + 同 file_hash ──
if file_hash:
existing = asset_repository.find_by_library_and_file_hash(
library_id=library_id,
file_hash=file_hash,
)
if existing is not None:
logger.info(
"素材去重命中(multipart): library=%s hash=%s existing_asset=%s",
library_id,
file_hash,
existing.id,
)
return UploadAssetResponse(
storage_key=existing.storage_key,
ingest_job_id="",
url="",
duplicated=True,
asset_id=existing.id,
)
# P2-5: 服务端验证 MIME 类型
validated_content_type = _validate_mime_type(file.content_type) validated_content_type = _validate_mime_type(file.content_type)
safe_filename = file.filename.replace("/", "_").replace("\\", "_") if file.filename else "unknown"
# ── 幂等去重:client_upload_id → file_hash → 近期活动同名记录兜底 ──
# 放在 OSS 上传之前:重复提交直接返回,不占 OSS 流量、不建新记录。
existing = _find_duplicate_asset(
asset_repository,
library_id=library_id,
file_hash=file_hash,
client_upload_id=client_upload_id,
filename=safe_filename,
file_size=0,
)
if existing is not None:
return UploadAssetResponse(
storage_key=existing.storage_key,
ingest_job_id="",
url="",
duplicated=True,
asset_id=existing.id,
)
file_id = uuid4().hex[:8] file_id = uuid4().hex[:8]
safe_filename = file.filename.replace("/", "_").replace("\\", "_") if file.filename else "unknown"
storage_key = f"uploads/{file_id}/{safe_filename}" storage_key = f"uploads/{file_id}/{safe_filename}"
try: try:
@@ -560,32 +284,16 @@ async def upload_asset(
detail=f"Failed to upload file: {type(error).__name__}", detail=f"Failed to upload file: {type(error).__name__}",
) from error ) from error
# 立即创建 Asset 记录(PROCESSING 状态),使前端刷新后即可看到新素材
pending_asset = _create_pending_asset(
asset_repository=asset_repository,
project_id=project_id,
library_id=library_id,
storage_key=storage_key,
filename=safe_filename,
mime_type=validated_content_type,
user_id=authenticated_user.user.id,
file_hash=file_hash,
client_upload_id=client_upload_id,
)
# Issue #1776: 计数由 asset_repository.create() 自动维护
job = _submit_ingest_job( job = _submit_ingest_job(
project_id=project_id, project_id=project_id,
library_id=library_id, library_id=library_id,
storage_key=storage_key, storage_key=storage_key,
ingest_job_repository=ingest_job_repository, ingest_job_repository=ingest_job_repository,
file_hash=file_hash, file_hash=file_hash,
asset_id=pending_asset.id,
) )
return UploadAssetResponse( return UploadAssetResponse(
storage_key=storage_key, storage_key=storage_key,
ingest_job_id=job.id, ingest_job_id=job.id,
asset_id=pending_asset.id,
url=file_url, url=file_url,
) )
-74
View File
@@ -15,7 +15,6 @@ from app.schemas.video_center import (
VideoItemResponse, VideoItemResponse,
) )
from fastapi import APIRouter, Depends, HTTPException, Query, Response from fastapi import APIRouter, Depends, HTTPException, Query, Response
from pydantic import BaseModel, Field
from packages.application import ( from packages.application import (
GetGeneratedVideoUseCase, GetGeneratedVideoUseCase,
@@ -54,8 +53,6 @@ def _to_video_response(item, storage: OSSStorageService | None = None) -> VideoI
download_url=download_url, download_url=download_url,
generated_at=format_utc_datetime(item.generated_at) if hasattr(item, "generated_at") else "", generated_at=format_utc_datetime(item.generated_at) if hasattr(item, "generated_at") else "",
duplicate_rate=getattr(item, "duplicate_rate", None), duplicate_rate=getattr(item, "duplicate_rate", None),
visual_similarity=getattr(item, "visual_similarity", None),
match_count=getattr(item, "match_count", None),
) )
@@ -240,74 +237,3 @@ def get_batch_download_status(
status=api_status, status=api_status,
download_url=download_url, download_url=download_url,
) )
# ── 重新计算查重率 ─────────────────────────────────────────────────
class RecomputeDedupRequest(BaseModel):
"""重新计算查重率请求。"""
video_ids: list[str] | None = Field(
None,
description="指定视频 ID 列表。为空则对当前用户所有缺少查重数据的视频重新计算。",
)
force: bool = Field(
False,
description="强制重算:即使视频已有查重数据也重新入队(#1702 查重算法升级后用于存量视频重算)。",
)
class RecomputeDedupResponse(BaseModel):
"""重新计算查重率响应。"""
enqueued: int = Field(..., description="已入队的任务数量")
total_scanned: int = Field(..., description="扫描的视频总数")
skipped: int = Field(..., description="已有查重数据跳过的数量")
message: str = ""
@router.post("/videos/recompute-dedup", response_model=RecomputeDedupResponse)
def recompute_dedup(
request: RecomputeDedupRequest = RecomputeDedupRequest(),
repo=Depends(get_generated_video_repository),
current_user: AuthenticatedUser = Depends(get_current_user),
):
"""重新计算视频的查重率/视觉相似度。
对于已存在但缺少 duplicate_rate / video_fingerprint 的视频,
触发异步 Celery 任务重新下载并计算指纹 + 查重率。
不传 video_ids 时,对当前用户所有视频进行检查。
"""
user_id = current_user.user.id
# 获取目标视频列表
if request.video_ids:
all_videos = repo.get_by_ids(request.video_ids)
# 安全校验:只处理当前用户的视频
target_videos = [v for v in all_videos if v.user_id == user_id]
else:
target_videos = repo.list_by_user(user_id)
total_scanned = len(target_videos)
enqueued = 0
skipped = 0
for video in target_videos:
# 已有完整查重数据的跳过(force=True 时强制重算,#1702 算法升级后存量视频需要重算指纹/分片)
if not request.force and video.duplicate_rate is not None and video.video_fingerprint:
skipped += 1
continue
# 触发异步查重任务
celery_app.send_task("worker.check_duplicate", args=[video.id])
enqueued += 1
logger.info("Enqueued re-dedup for video %s (user=%s, force=%s)", video.id, user_id, request.force)
return RecomputeDedupResponse(
enqueued=enqueued,
total_scanned=total_scanned,
skipped=skipped,
message=f"已入队 {enqueued} 个查重任务" if enqueued > 0 else "所有视频查重数据已完整",
)
+6 -6
View File
@@ -163,12 +163,12 @@ def create_voice_clone(
celery_app.send_task("worker.process_voice_clone", args=[profile.id]) celery_app.send_task("worker.process_voice_clone", args=[profile.id])
logger.info(f"Celery task dispatched for voice clone {profile.id}") logger.info(f"Celery task dispatched for voice clone {profile.id}")
except Exception as e: except Exception as e:
logger.exception("Failed to dispatch Celery task") logger.error(f"Failed to dispatch Celery task: {e}")
# P2-3: Celery 调度失败时标记 profile 为 failed,避免永久卡在 processing # P2-3: Celery 调度失败时标记 profile 为 failed,避免永久卡在 processing
try: try:
workflow.process_clone_failure(profile.id, f"Celery 任务调度失败: {e}") workflow.process_clone_failure(profile.id, f"Celery 任务调度失败: {e}")
except Exception: except Exception as inner_e:
logger.exception("Failed to mark profile as failed after dispatch error") logger.error(f"Failed to mark profile as failed after dispatch error: {inner_e}")
return _to_response(profile) return _to_response(profile)
@@ -277,12 +277,12 @@ def retry_voice_clone(
celery_app.send_task("worker.process_voice_clone", args=[profile.id]) celery_app.send_task("worker.process_voice_clone", args=[profile.id])
logger.info(f"Celery task dispatched for voice clone retry {profile.id}") logger.info(f"Celery task dispatched for voice clone retry {profile.id}")
except Exception as e: except Exception as e:
logger.exception("Failed to dispatch Celery task") logger.error(f"Failed to dispatch Celery task: {e}")
# P2-3: Celery 调度失败时标记 profile 为 failed,避免永久卡在 processing # P2-3: Celery 调度失败时标记 profile 为 failed,避免永久卡在 processing
try: try:
workflow.process_clone_failure(profile.id, f"Celery 任务调度失败: {e}") workflow.process_clone_failure(profile.id, f"Celery 任务调度失败: {e}")
except Exception: except Exception as inner_e:
logger.exception("Failed to mark profile as failed after dispatch error") logger.error(f"Failed to mark profile as failed after dispatch error: {inner_e}")
return _to_response(profile) return _to_response(profile)
+4 -263
View File
@@ -6,26 +6,12 @@
from __future__ import annotations from __future__ import annotations
import logging import logging
import shutil
import subprocess
import tempfile
import time import time
from pathlib import Path
from typing import Literal, Optional from typing import Literal, Optional
from uuid import uuid4
from app.api.routes._helpers import get_user_plan from app.api.routes._helpers import get_user_plan
from app.auth import AuthenticatedUser, get_current_user from app.auth import AuthenticatedUser, get_current_user
from app.core.storage import get_storage_service from app.dependencies import get_audio_url_signer, get_cosyvoice_service, get_db_session, get_user_repository
from app.dependencies import (
get_asset_library_repository,
get_asset_repository,
get_audio_url_signer,
get_cosyvoice_service,
get_db_session,
get_project_repository,
get_user_repository,
)
from app.schemas.voice import ( from app.schemas.voice import (
PresetVoiceItemResponse, PresetVoiceItemResponse,
PresetVoiceListResponse, PresetVoiceListResponse,
@@ -38,7 +24,7 @@ from app.schemas.voice_library import (
UpdateVoiceLibraryRequest, UpdateVoiceLibraryRequest,
VoiceLibraryItemResponse, VoiceLibraryItemResponse,
) )
from fastapi import APIRouter, Depends, File, Form, HTTPException, Query, Response, UploadFile, status from fastapi import APIRouter, Depends, HTTPException, Query, Response, status
from sqlalchemy.orm import Session from sqlalchemy.orm import Session
from packages.adapters.sqlalchemy_impl.voice_clone_profile_repository import SQLAlchemyVoiceCloneProfileRepository from packages.adapters.sqlalchemy_impl.voice_clone_profile_repository import SQLAlchemyVoiceCloneProfileRepository
@@ -54,12 +40,8 @@ from packages.application.voice_library.use_cases import (
QuotaExceededError, QuotaExceededError,
UpdateVoiceLibraryUseCase, UpdateVoiceLibraryUseCase,
) )
from packages.domain import Asset, AssetStatus
from packages.domain.classification import AssetLibraryKind, ClassificationStatus
from packages.domain.entities import AssetLibrary
from packages.domain.preset_voices import PRESET_VOICES, get_preset_voice_by_id from packages.domain.preset_voices import PRESET_VOICES, get_preset_voice_by_id
from packages.ports.user_repository import UserRepository from packages.ports.user_repository import UserRepository
from packages.shared.storage import SharedStorageService
router = APIRouter() router = APIRouter()
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -105,8 +87,8 @@ def _resolve_preset_preview_url(
_preset_preview_cache[voice_id] = (audio_url, time.time()) _preset_preview_cache[voice_id] = (audio_url, time.time())
logger.info("Preset voice preview generated: %s", voice_id) logger.info("Preset voice preview generated: %s", voice_id)
return audio_url return audio_url
except Exception: except Exception as e:
logger.exception("Failed to generate preset voice preview: voice_id=%s", voice_id) logger.warning("Failed to generate preview for %s, using fallback: %s", voice_id, e)
return fallback_url return fallback_url
@@ -127,7 +109,6 @@ def _resolve_all_preset_preview_urls(
try: try:
result_map[p.voice_id] = _resolve_preset_preview_url(p.voice_id, p.preview_url, cosyvoice) result_map[p.voice_id] = _resolve_preset_preview_url(p.voice_id, p.preview_url, cosyvoice)
except Exception: except Exception:
logger.exception("Failed to resolve preset preview URL: voice_id=%s", p.voice_id)
result_map[p.voice_id] = p.preview_url result_map[p.voice_id] = p.preview_url
return result_map return result_map
@@ -526,243 +507,3 @@ def delete_voice(
if not deleted: if not deleted:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Voice not found") raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Voice not found")
return return
# ── 提取视频配音 ─────────────────────────────────────────────────────
# 支持的视频格式
EXTRACT_VIDEO_MIMES = frozenset({"video/mp4", "video/quicktime", "video/webm", "video/x-msvideo"})
MAX_EXTRACT_SIZE = 500 * 1024 * 1024 # 500MB
@router.post(
"/extract-voice",
status_code=status.HTTP_201_CREATED,
)
def extract_voice_from_video(
file: UploadFile = File(...),
project_id: str = Form(...),
authenticated_user: AuthenticatedUser = Depends(get_current_user),
project_repository=Depends(get_project_repository),
asset_library_repository=Depends(get_asset_library_repository),
asset_repository=Depends(get_asset_repository),
storage_service: SharedStorageService = Depends(get_storage_service),
sign_url=Depends(get_audio_url_signer),
):
"""从上传的视频中提取人声配音。
流程:
1. 接收视频文件(mp4/mov/webm
2. ffmpeg 提取音频 + 降噪 + 编码为 mp3
3. 上传到 OSS,创建 Asset 记录到配音素材库
4. 返回素材信息(时长、文件大小、URL)
"""
user_id = authenticated_user.user.id
# 校验文件类型
content_type = file.content_type or ""
if content_type and content_type not in EXTRACT_VIDEO_MIMES:
# 兜底:按扩展名判断
ext = (file.filename or "").rsplit(".", 1)[-1].lower()
ext_to_mime = {"mp4": "video/mp4", "mov": "video/quicktime", "webm": "video/webm", "avi": "video/x-msvideo"}
if ext not in ext_to_mime:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="仅支持 mp4/mov/webm/avi 格式的视频文件",
)
content_type = ext_to_mime[ext]
# 找到(或自动创建)用户 voice 素材库(复用 TTS 的逻辑)
library = _find_or_create_voice_library_for_extract(
user_id=user_id,
project_repository=project_repository,
asset_library_repository=asset_library_repository,
)
tmp_dir = None
try:
tmp_dir = Path(tempfile.mkdtemp(prefix="voice_extract_"))
video_path = tmp_dir / f"input_{uuid4().hex[:8]}_{file.filename or 'video.mp4'}"
audio_path = tmp_dir / f"output_{uuid4().hex[:8]}.mp3"
# 保存上传的视频到临时文件
with open(video_path, "wb") as f:
total = 0
while chunk := file.file.read(1024 * 1024): # 1MB chunks
total += len(chunk)
if total > MAX_EXTRACT_SIZE:
raise HTTPException(
status_code=status.HTTP_413_REQUEST_ENTITY_TOO_LARGE,
detail="视频文件过大,最大支持 500MB",
)
f.write(chunk)
if video_path.stat().st_size == 0:
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="视频文件为空")
# ffmpeg: 提取音频 + 降噪 + 编码 mp3
# 滤镜链:highpass(去低频噪声) → afftdn(FFT降噪) → lowpass(去高频噪声)
ffmpeg_cmd = [
"ffmpeg",
"-y",
"-i",
str(video_path),
"-vn", # 不要视频
"-af",
"highpass=f=80,afftdn=nf=-25:tn=1,lowpass=f=8000",
"-acodec",
"libmp3lame",
"-ab",
"192k",
"-ar",
"44100",
"-ac",
"1", # 单声道(人声足够)
str(audio_path),
]
result = subprocess.run(
ffmpeg_cmd,
capture_output=True,
timeout=300, # 5 分钟超时
)
if result.returncode != 0:
stderr_text = result.stderr.decode("utf-8", errors="replace")[-500:]
logger.error("ffmpeg 提取配音失败: %s", stderr_text)
raise HTTPException(
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
detail="视频音频提取失败,可能该视频没有音轨或格式不支持",
)
if not audio_path.exists() or audio_path.stat().st_size == 0:
raise HTTPException(
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
detail="音频提取结果为空",
)
# 获取音频时长
duration = _get_audio_duration(audio_path)
file_size = audio_path.stat().st_size
# 上传到 OSS
audio_ext = "mp3"
storage_key = f"uploads/voice/extracted/{uuid4().hex}.{audio_ext}"
storage_service.upload_file(audio_path, storage_key, content_type="audio/mpeg")
# 创建 Asset 记录
original_name = (file.filename or "video").rsplit(".", 1)[0]
asset_name = f"{original_name}-配音"
asset = Asset.create(
project_id=library.project_id,
library_id=library.id,
name=asset_name,
storage_key=storage_key,
mime_type="audio/mpeg",
metadata={
"source": "video_extract",
"original_video": file.filename or "unknown",
},
file_size=file_size,
duration=duration,
status=AssetStatus.READY,
classification_status=ClassificationStatus.PENDING,
uploaded_by_user_id=user_id,
)
asset = asset_repository.create(asset)
return {
"id": asset.id,
"name": asset.name,
"audio_url": sign_url(storage_key),
"duration": duration,
"file_size": file_size,
"status": "completed",
"source": "video_extract",
}
except HTTPException:
raise
except subprocess.TimeoutExpired:
raise HTTPException(
status_code=status.HTTP_504_GATEWAY_TIMEOUT,
detail="视频处理超时,请尝试较短的视频",
) from None
except Exception as e:
logger.exception("提取视频配音失败: %s", e)
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail="提取配音失败,请稍后重试",
) from e
finally:
# 清理临时文件
if tmp_dir and Path(tmp_dir).exists():
shutil.rmtree(tmp_dir, ignore_errors=True)
def _find_or_create_voice_library_for_extract(*, user_id, project_repository, asset_library_repository):
"""为用户找到或创建 voice 素材库(与 TTS 保存逻辑一致)。"""
projects = project_repository.find_accessible_projects(user_id)
if not projects:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="没有可用的项目,请先创建项目",
)
for project in projects:
for lib in asset_library_repository.find_by_project(project.id):
kind = lib.kind.value if hasattr(lib.kind, "value") else lib.kind
if kind == AssetLibraryKind.VOICE.value:
return lib
# 自动创建
from sqlalchemy.exc import IntegrityError
project = projects[0]
library = AssetLibrary.create(
project_id=project.id,
name="配音素材库",
kind=AssetLibraryKind.VOICE,
)
try:
return asset_library_repository.create(library)
except IntegrityError:
session = getattr(asset_library_repository, "session", None)
if session is not None:
try:
session.rollback()
except Exception:
logger.exception("session rollback failed in _find_or_create_voice_library")
for lib in asset_library_repository.find_by_project(project.id):
kind = lib.kind.value if hasattr(lib.kind, "value") else lib.kind
if kind == AssetLibraryKind.VOICE.value:
return lib
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail="配音素材库创建失败",
) from None
def _get_audio_duration(audio_path: Path) -> float:
"""用 ffprobe 获取音频时长(秒)。"""
try:
result = subprocess.run(
[
"ffprobe",
"-v",
"quiet",
"-show_entries",
"format=duration",
"-of",
"csv=p=0",
str(audio_path),
],
capture_output=True,
timeout=10,
)
if result.returncode == 0 and result.stdout.strip():
return float(result.stdout.strip())
except (ValueError, subprocess.TimeoutExpired):
pass
return 0.0
View File
-8
View File
@@ -5,11 +5,3 @@ settings = get_settings()
celery_app = Celery("xiaoxia-saas-api") celery_app = Celery("xiaoxia-saas-api")
celery_app.conf.broker_url = settings.CELERY_BROKER_URL celery_app.conf.broker_url = settings.CELERY_BROKER_URL
celery_app.conf.result_backend = settings.CELERY_RESULT_BACKEND celery_app.conf.result_backend = settings.CELERY_RESULT_BACKEND
# #1714 队列隔离:视频生成走 generation 队列,素材转码走 transcode 队列
try:
from packages.shared.celery_queues import apply_queue_settings
apply_queue_settings(celery_app)
except Exception: # noqa: BLE001 — 队列配置失败不阻断 API 启动
pass
+3 -131
View File
@@ -8,145 +8,27 @@ logger = logging.getLogger(__name__)
# ── 限流阈值常量(全系统统一管理,不要在业务代码里硬编码) ── # ── 限流阈值常量(全系统统一管理,不要在业务代码里硬编码) ──
USER_PENDING_LIMIT = 3 # 单用户 pending 上限 USER_PENDING_LIMIT = 3 # 单用户 pending 上限
GLOBAL_PENDING_LIMIT = 20 # 全局 pending 上限 GLOBAL_PENDING_LIMIT = 20 # 全局 pending 上限
WORKER_CONCURRENCY = 4 # worker 渲染并发数(infra/docker/compose.yml WORKER_CONCURRENCY 默认值)
# 限流错误码:前端据此区分"排队等待"与"创建失败"
ERROR_CODE_USER_QUEUE_FULL = "USER_QUEUE_FULL" # 429:用户自己的任务排队中
ERROR_CODE_SYSTEM_QUEUE_FULL = "SYSTEM_QUEUE_FULL" # 503:系统整体繁忙
class UserPendingLimitExceeded(Exception): class UserPendingLimitExceeded(Exception):
"""用户 pending 任务数超限,返回 429。""" """用户 pending 任务数超限,返回 429。"""
def __init__( def __init__(self, user_id: str, pending_count: int, limit: int):
self,
user_id: str,
pending_count: int,
limit: int,
*,
running_count: int = 0,
requested_count: int = 1,
queue_ahead: int = 0,
estimated_wait_seconds: int = 0,
):
self.user_id = user_id self.user_id = user_id
self.pending_count = pending_count self.pending_count = pending_count
self.limit = limit self.limit = limit
# 排队上下文(用于 429 结构化提示,前端展示"排队中"而非"创建失败"
self.running_count = running_count
self.requested_count = requested_count
self.queue_ahead = queue_ahead
self.estimated_wait_seconds = estimated_wait_seconds
super().__init__(f"用户 {user_id} pending 任务数 {pending_count} 超过上限 {limit}") super().__init__(f"用户 {user_id} pending 任务数 {pending_count} 超过上限 {limit}")
class GlobalQueueFull(Exception): class GlobalQueueFull(Exception):
"""全局限流,返回 503。""" """全局限流,返回 503。"""
def __init__( def __init__(self, pending_count: int, limit: int):
self,
pending_count: int,
limit: int,
*,
running_count: int = 0,
queue_ahead: int = 0,
estimated_wait_seconds: int = 0,
):
self.pending_count = pending_count self.pending_count = pending_count
self.limit = limit self.limit = limit
self.running_count = running_count
self.queue_ahead = queue_ahead
self.estimated_wait_seconds = estimated_wait_seconds
super().__init__(f"系统 pending 任务数 {pending_count} 超过上限 {limit}") super().__init__(f"系统 pending 任务数 {pending_count} 超过上限 {limit}")
def _estimate_wait_seconds(queue_ahead: int, generation_task_repository: Any) -> int:
"""根据排队任务数 + worker 并发数 + 历史平均任务耗时估算等待秒数。
估算公式:ceil(排队任务数 / 并发数) × 平均单任务耗时。
拿不到历史数据时仓储层返回默认 120 秒。
"""
import math
if queue_ahead <= 0:
return 0
try:
estimator = getattr(generation_task_repository, "estimate_avg_duration_seconds", None)
avg_seconds = estimator() if estimator is not None else 120.0
except Exception:
avg_seconds = 120.0
return int(math.ceil(queue_ahead / WORKER_CONCURRENCY) * avg_seconds)
def build_rate_limit_detail(
exc: Exception,
generation_task_repository: Any,
*,
scope: str = "user",
) -> dict:
"""构造结构化限流响应体(HTTPException 的 detail)。
前端按 detail.code 判断场景:
- USER_QUEUE_FULL (429):用户自己的任务在排队,应提示"等待/继续排队",不是创建失败
- SYSTEM_QUEUE_FULL (503):系统繁忙,稍后重试
detail 字段:
- code: 错误码
- message: 可读中文提示(可直接展示)
- queued_count: 当前排队(pending)任务数
- running_count: 当前渲染中(running)任务数
- queue_ahead: 前方排队任务数(预计等待批次依据)
- estimated_wait_seconds: 预计等待秒数
- limit: 对应限流上限
"""
if scope == "user" and isinstance(exc, UserPendingLimitExceeded):
running = exc.running_count
if not running:
try:
counter = getattr(generation_task_repository, "count_running_by_user", None)
running = counter(exc.user_id) if counter is not None else 0
except Exception:
running = 0
queue_ahead = exc.queue_ahead or max(exc.pending_count, 0)
wait = exc.estimated_wait_seconds or _estimate_wait_seconds(queue_ahead, generation_task_repository)
wait_minutes = max(1, round(wait / 60))
message = (
f"您有 {exc.pending_count} 个任务正在排队、{running} 个正在渲染,"
f"同一时间最多提交 {exc.limit} 个任务。请等待约 {wait_minutes} 分钟后再提交"
)
return {
"code": ERROR_CODE_USER_QUEUE_FULL,
"message": message,
"queued_count": exc.pending_count,
"running_count": running,
"queue_ahead": queue_ahead,
"estimated_wait_seconds": wait,
"limit": exc.limit,
}
# 全局繁忙
pending = getattr(exc, "pending_count", 0)
running = getattr(exc, "running_count", 0)
if not running:
try:
counter = getattr(generation_task_repository, "count_running_total", None)
running = counter() if counter is not None else 0
except Exception:
running = 0
queue_ahead = getattr(exc, "queue_ahead", 0) or pending
wait = getattr(exc, "estimated_wait_seconds", 0) or _estimate_wait_seconds(queue_ahead, generation_task_repository)
wait_minutes = max(1, round(wait / 60))
return {
"code": ERROR_CODE_SYSTEM_QUEUE_FULL,
"message": f"系统繁忙:当前 {pending} 个任务排队中、{running} 个渲染中,预计等待约 {wait_minutes} 分钟,请稍后再试",
"queued_count": pending,
"running_count": running,
"queue_ahead": queue_ahead,
"estimated_wait_seconds": wait,
"limit": getattr(exc, "limit", GLOBAL_PENDING_LIMIT),
}
def check_queue_limits( def check_queue_limits(
user_id: str, user_id: str,
generation_task_repository: Any, generation_task_repository: Any,
@@ -279,17 +161,7 @@ def safe_enqueue_generation_task(
# ── 发送 Celery 任务 ── # ── 发送 Celery 任务 ──
try: try:
celery_result = celery_app.send_task("worker.generate_video", args=[task.id]) celery_app.send_task("worker.generate_video", args=[task.id])
# 记录 celery 消息 ID:孤儿清理/超时作废时据此 revoke + 清除队列消息(#1714
celery_task_id = getattr(celery_result, "id", "")
if celery_task_id:
try:
task.celery_task_id = celery_task_id
generation_task_repository.update(task)
except Exception as persist_err: # noqa: BLE001
logger.warning(
"%s 持久化 celery_task_id 失败(不影响主流程): task_id=%s err=%s", log_prefix, task.id, persist_err
)
except Exception as e: except Exception as e:
logger.error( logger.error(
"%s 入队失败,标记为失败: task_id=%s error=%s", "%s 入队失败,标记为失败: task_id=%s error=%s",
-124
View File
@@ -1,124 +0,0 @@
"""AI数字人渲染合成管线 API Schema — #1798."""
from __future__ import annotations
from datetime import datetime
from typing import Any, Optional
from pydantic import BaseModel, Field, field_validator
class BRollSegment(BaseModel):
"""B-roll 片段配置."""
script_segment_index: int = Field(..., ge=0, description="对应文案片段索引")
asset_url: str = Field(..., description="B-roll 素材 URL")
mode: str = Field(..., description="插入模式: fullscreen 或 pip")
start_time: float = Field(..., ge=0.0, description="在对口型视频中的起始时间(秒)")
end_time: float = Field(..., ge=0.0, description="在对口型视频中的结束时间(秒)")
pip_position: Optional[str] = Field("bottom_right", description="pip 模式位置")
pip_scale: Optional[float] = Field(0.3, ge=0.05, le=1.0, description="pip 模式缩放比例")
@field_validator("mode")
@classmethod
def validate_mode(cls, v: str) -> str:
v = v.strip().lower()
if v not in ("fullscreen", "pip"):
raise ValueError("mode 必须为 fullscreen 或 pip")
return v
@field_validator("asset_url")
@classmethod
def validate_asset_url(cls, v: str) -> str:
v = v.strip()
if not v:
raise ValueError("asset_url 不能为空")
if not v.startswith(("http://", "https://")):
raise ValueError("asset_url 必须是 HTTP/HTTPS URL")
return v
@field_validator("end_time")
@classmethod
def validate_end_time(cls, v: float, info: Any) -> float:
start = info.data.get("start_time", 0.0)
if v <= start:
raise ValueError("end_time 必须大于 start_time")
return v
class CreateAiAvatarRenderRequest(BaseModel):
"""创建渲染任务请求."""
lipsync_job_id: str = Field(..., description="对口型任务 ID")
script_id: str = Field("", description="文案 ID(选自文案库时传;手动输入文案直生场景可留空)")
b_roll_segments: list[BRollSegment] = Field(default_factory=list, description="B-roll 片段列表")
title_config: dict[str, Any] = Field(default_factory=dict, description="标题配置")
cover_config: dict[str, Any] = Field(default_factory=dict, description="封面配置")
project_id: str = Field("", description="项目 ID")
@field_validator("lipsync_job_id")
@classmethod
def validate_lipsync_job_id(cls, v: str) -> str:
v = v.strip()
if not v:
raise ValueError("lipsync_job_id 不能为空")
return v
@field_validator("script_id")
@classmethod
def validate_script_id(cls, v: str) -> str:
# script_id 可选:手动输入文案(TTS 直生)场景不关联文案库条目
return (v or "").strip()
class AiAvatarRenderJobResponse(BaseModel):
"""渲染任务响应."""
id: str
user_id: str
project_id: str
lipsync_job_id: str
script_id: str = ""
b_roll_segments: list[dict[str, Any]]
title_config: dict[str, Any]
cover_config: dict[str, Any]
status: str
progress: int
output_video_url: str
output_cover_url: str
output_duration: float
error_message: str
submitted_at: Optional[datetime] = None
started_at: Optional[datetime] = None
completed_at: Optional[datetime] = None
created_at: datetime
updated_at: datetime
class Config:
from_attributes = True
class AiAvatarRenderProgressResponse(BaseModel):
"""渲染进度响应."""
status: str
progress: int
output_video_url: str
output_cover_url: str
output_duration: float
error_message: str
class SmartCoverRequest(BaseModel):
"""智能封面请求 — MediaKit 抽帧 + 质量评分选最佳帧."""
video_url: str = Field(..., description="数字人视频 URL(对口型/渲染成片)")
max_frames: int = Field(5, ge=1, le=10, description="抽帧数量(默认 5")
class SmartCoverResponse(BaseModel):
"""智能封面响应."""
cover_url: str = Field("", description="封面图公网 URL(OSS,非临时);失败为空")
status: str = Field("completed", description="completed / fallback_failed")
message: str = Field("", description="失败原因(如有)")
-3
View File
@@ -28,9 +28,6 @@ class DuplicationRecordResponse(BaseModel):
status: str = "pending" status: str = "pending"
duplicate_rate: float | None = None duplicate_rate: float | None = None
duplicate_count: int = 0 duplicate_count: int = 0
# #1661 视觉相似度(归一化 0~1)/ 匹配视频数
visual_similarity: float | None = None
match_count: int | None = None
created_at: str created_at: str
updated_at: str updated_at: str
-4
View File
@@ -25,10 +25,6 @@ class GeneratedVideoResponse(BaseModel):
review_status: str = "pending_review" review_status: str = "pending_review"
generation_params: dict = Field(default_factory=dict) generation_params: dict = Field(default_factory=dict)
download_url: str | None = None download_url: str | None = None
# #1660 查重率(百分比 0~100)/ 视觉相似度(0~1)/ 匹配帧数
duplicate_rate: float | None = None
visual_similarity: float | None = None
match_count: int | None = None
class GeneratedVideoDownloadUrlResponse(BaseModel): class GeneratedVideoDownloadUrlResponse(BaseModel):
+8 -109
View File
@@ -25,24 +25,14 @@ class CreateGenerationTaskRequest(BaseModel):
asset_library_id: str = "" asset_library_id: str = ""
strategy_id: str = "" strategy_id: str = ""
voice_library_id: str = "" voice_library_id: str = ""
# ── 多变体独立配音(批量生成)──
# 长度 1 = 所有变体共用;长度 = count = 每个变体独立配音;空数组 = 回退 voice_library_id
voice_library_ids: list[str] = Field(
default_factory=list,
description="各变体独立配音素材库ID数组:长度1=共用,长度=count=独立。为空时回退 voice_library_id",
)
created_by_user_id: str = "" created_by_user_id: str = ""
# ── 模板模式新增字段 ── # ── 模板模式新增字段 ──
template_id: str = "" template_id: str = ""
asset_ids: list[str] = Field(default_factory=list) asset_ids: list[str] = Field(default_factory=list)
title_ids: list[str] = Field(default_factory=list) title_ids: list[str] = Field(default_factory=list)
voice_ids: list[str] = Field(default_factory=list)
# ── 来源剪辑计划 ── # ── 来源剪辑计划 ──
source_edit_plan_id: str = "" source_edit_plan_id: str = ""
# ── variant-plans 轻量选片回传(#1749):正式生成直接复用,不再重选 ──
variant_plan_ids: list[str] = Field(
default_factory=list,
description="POST /generation/variant-plans 返回的各变体 plan_id(长度须=count);为空则走服务端选片",
)
# ── 标题配置(结构化)── # ── 标题配置(结构化)──
title_config: dict | None = Field( title_config: dict | None = Field(
default=None, default=None,
@@ -85,47 +75,6 @@ class CreateGenerationTaskRequest(BaseModel):
output_width: int = Field(default=1280, description="输出视频宽度") output_width: int = Field(default=1280, description="输出视频宽度")
output_height: int = Field(default=720, description="输出视频高度") output_height: int = Field(default=720, description="输出视频高度")
cover_url: str = Field(default="", description="封面图片 URL") cover_url: str = Field(default="", description="封面图片 URL")
# ── 多变体独立封面(批量生成)──
# 长度 1 = 所有变体共用;长度 = count = 每个变体独立封面;空数组 = 回退 cover_url
cover_urls: list[str] = Field(
default_factory=list,
description="各变体独立封面URL数组:长度1=共用,长度=count=独立。为空时回退 cover_url",
)
# ── 多变体独立标题文字(批量生成)──
# 长度 1 = 所有变体共用;长度 = count = 每个变体独立标题文字;空数组 = 使用 title_config.text
titles: list[str] = Field(
default_factory=list,
description="各变体独立标题文字数组:长度1=共用,长度=count=独立。为空时使用 title_config.text",
)
@model_validator(mode="after")
def _check_variant_arrays(self) -> "CreateGenerationTaskRequest":
"""变体数组字段长度校验 + #1749 配音严格守卫。
- cover_urls/titles:空(回退单值)、长度 1(共用)或长度 = count(独立);
- voice_library_ids:独立配音长度必须恰好 = count 且逐项非空,禁止静默 fallback
(长度 1 的"共用"场景请用 voice_library_id 单值字段);
- variant_plan_ids:非空时长度必须 = count。
"""
for name in ("cover_urls", "titles"):
arr = getattr(self, name)
if arr and len(arr) != 1 and len(arr) != self.count:
raise ValueError(f"{name} 长度必须为 1(共用)或 {self.count}(与 count 一致),当前为 {len(arr)}")
from packages.domain.variant_voice_resolver import VariantVoiceError, resolve_variant_voice_ids
try:
resolve_variant_voice_ids(
count=self.count,
voice_library_id=self.voice_library_id,
voice_library_ids=self.voice_library_ids or None,
)
except VariantVoiceError as exc:
raise ValueError(str(exc)) from exc
if self.variant_plan_ids and len(self.variant_plan_ids) != self.count:
raise ValueError(f"variant_plan_ids 长度({len(self.variant_plan_ids)})必须与 count({self.count})一致")
return self
@model_validator(mode="after") @model_validator(mode="after")
def _check_at_least_one_mode(self) -> "CreateGenerationTaskRequest": def _check_at_least_one_mode(self) -> "CreateGenerationTaskRequest":
@@ -134,9 +83,9 @@ class CreateGenerationTaskRequest(BaseModel):
if not has_project and not has_template: if not has_project and not has_template:
raise ValueError("project_id 或 template_id 至少需要提供一个") raise ValueError("project_id 或 template_id 至少需要提供一个")
has_library = bool(self.asset_library_id.strip()) has_library = bool(self.asset_library_id.strip())
has_assets = bool(self.asset_ids or self.title_ids) has_assets = bool(self.asset_ids or self.title_ids or self.voice_ids)
if not has_library and not has_assets: if not has_library and not has_assets:
raise ValueError("asset_library_id 或 asset_ids/title_ids 至少需要提供一个") raise ValueError("asset_library_id 或 asset_ids/title_ids/voice_ids 至少需要提供一个")
return self return self
@@ -213,6 +162,7 @@ class CreatePreviewGenerationTaskRequest(BaseModel):
template_id: str template_id: str
asset_ids: list[str] = Field(default_factory=list) asset_ids: list[str] = Field(default_factory=list)
title_ids: list[str] = Field(default_factory=list) title_ids: list[str] = Field(default_factory=list)
voice_ids: list[str] = Field(default_factory=list)
voice_library_id: str = Field( voice_library_id: str = Field(
default="", description="配音素材库ID(用户上传的音频或AI配音),对应配音选择页面选择的配音素材" default="", description="配音素材库ID(用户上传的音频或AI配音),对应配音选择页面选择的配音素材"
) )
@@ -235,44 +185,8 @@ class CreatePreviewGenerationTaskRequest(BaseModel):
) )
title_config: dict = Field( title_config: dict = Field(
default_factory=dict, default_factory=dict,
description="标题配置(可选),渲染时烧录到预览视频中。支持字段: text/font/font_size/font_color/position/bold/stroke/shadow。N个变体时样式全局共用", description="标题配置(可选),渲染时烧录到预览视频中。支持字段: text/font/font_size/font_color/position/bold/stroke/shadow",
) )
# ── 多变体独立配置(preview_count > 1)──
# 长度 1 = 所有变体共用;长度 = preview_count = 每个变体独立;空数组 = 回退单值字段
titles: list[str] = Field(
default_factory=list,
description="各变体独立标题文字数组:长度1=共用,长度=preview_count=独立。为空时使用 title_config.text",
)
voice_library_ids: list[str] = Field(
default_factory=list,
description="各变体独立配音素材库ID数组:长度1=共用,长度=preview_count=独立。为空时回退 voice_library_id",
)
cover_urls: list[str] = Field(
default_factory=list,
description="各变体独立封面URL数组:长度1=共用,长度=preview_count=独立(预览阶段通常为空)",
)
@model_validator(mode="after")
def _check_variant_arrays(self) -> "CreatePreviewGenerationTaskRequest":
"""变体数组字段长度校验 + #1749 配音严格守卫。"""
for name in ("titles", "cover_urls"):
arr = getattr(self, name)
if arr and len(arr) != 1 and len(arr) != self.preview_count:
raise ValueError(
f"{name} 长度必须为 1(共用)或 {self.preview_count}(与 preview_count 一致),当前为 {len(arr)}"
)
from packages.domain.variant_voice_resolver import VariantVoiceError, resolve_variant_voice_ids
try:
resolve_variant_voice_ids(
count=self.preview_count,
voice_library_id=self.voice_library_id,
voice_library_ids=self.voice_library_ids or None,
)
except VariantVoiceError as exc:
raise ValueError(str(exc)) from exc
return self
@model_validator(mode="after") @model_validator(mode="after")
def _check_template_id(self) -> "CreatePreviewGenerationTaskRequest": def _check_template_id(self) -> "CreatePreviewGenerationTaskRequest":
@@ -282,13 +196,13 @@ class CreatePreviewGenerationTaskRequest(BaseModel):
@model_validator(mode="after") @model_validator(mode="after")
def _check_asset_ids(self) -> "CreatePreviewGenerationTaskRequest": def _check_asset_ids(self) -> "CreatePreviewGenerationTaskRequest":
if not self.asset_ids and not self.title_ids: if not self.asset_ids and not self.title_ids and not self.voice_ids:
raise ValueError("asset_ids/title_ids 至少需要提供一个") raise ValueError("asset_ids/title_ids/voice_ids 至少需要提供一个")
return self return self
class PreviewGenerationTaskResponse(BaseModel): class PreviewGenerationTaskResponse(BaseModel):
"""单个预览变体任务响应。 """预览生成任务响应。
包含任务状态、进度、分辨率、生成结果 URL 等关键字段。 包含任务状态、进度、分辨率、生成结果 URL 等关键字段。
""" """
@@ -297,7 +211,6 @@ class PreviewGenerationTaskResponse(BaseModel):
status: str status: str
progress: float progress: float
is_preview: bool = True is_preview: bool = True
variant_index: int = 0
resolution: str = "" resolution: str = ""
video_url: str = "" video_url: str = ""
duration: float = 0.0 duration: float = 0.0
@@ -306,21 +219,7 @@ class PreviewGenerationTaskResponse(BaseModel):
transition_count: int = 0 transition_count: int = 0
material_usage: dict = Field(default_factory=dict) material_usage: dict = Field(default_factory=dict)
error_message: str = "" error_message: str = ""
title_text: str = ""
voice_library_id: str = ""
created_at: datetime | None = None created_at: datetime | None = None
started_at: datetime | None = None started_at: datetime | None = None
finished_at: datetime | None = None finished_at: datetime | None = None
generate_duration: float = 0.0 generate_duration: float = 0.0
class BatchPreviewGenerationTaskResponse(BaseModel):
"""批量预览任务响应:preview_count=N 时返回 N 个独立变体任务。
- items: 变体任务数组,按 variant_index 顺序排列,每个含独立 task_id/状态/预览视频URL
- total: 变体总数(= preview_count
- 前端按 items[i].task_id 分别轮询 GET /preview/{task_id} 获取进度与结果
"""
items: list[PreviewGenerationTaskResponse]
total: int
-100
View File
@@ -1,100 +0,0 @@
"""对口型 API Schema 定义 — #1796 / #1809 / #1822.
支持两种输入模式(二选一):
1. TTS 直生模式(推荐):传 voice_id + script_text+ speed/emotion),
后端内部先调 CosyVoice 合成音频,再提交 MediaKit 对口型。
2. 直接音频模式:传 video_url + audio_url(音频已由调用方准备好)。
"""
from __future__ import annotations
from datetime import datetime
from typing import Optional
from pydantic import BaseModel, Field, model_validator
class LipsyncJobResponse(BaseModel):
"""对口型任务响应."""
id: str
user_id: str
project_id: str
video_url: str
audio_url: str
enable_video_loop: bool
voice_id: str = ""
script_text: str = ""
speed: float = 1.0
emotion: str = ""
mediakit_task_id: str
status: str
output_video_url: str
output_duration: float
error_message: str
error_code: str
submitted_at: Optional[datetime] = None
completed_at: Optional[datetime] = None
created_at: datetime
updated_at: datetime
class Config:
from_attributes = True
class CreateLipsyncJobRequest(BaseModel):
"""创建对口型任务请求.
两种模式(二选一):
- TTS 直生:voice_id + script_text 必填(+ 可选 speed/emotion);audio_url 留空。
- 直接音频:video_url + audio_url 必填。
"""
video_url: str = Field(..., description="人物视频 URL(MP4,≤30min,单人真人)")
# 模式 2:直接音频
audio_url: str = Field("", description="驱动音频 URLmp3/aac/wav/m4a/flac);直生模式留空")
# 模式 1TTS 直生
voice_id: str = Field("", description="音色 ID(预置音色或克隆音色 profile UUID")
script_text: str = Field("", description="要合成的文案(直生模式必填,最长 5000 字符)")
speed: float = Field(1.0, ge=0.5, le=2.0, description="语速(0.5-2.0),默认 1.0")
emotion: str = Field("", description="情绪(natural/excited/calm/friendly 或中文 自然/兴奋/沉稳/亲切)")
enable_video_loop: bool = Field(False, description="音频长于视频时是否循环画面")
project_id: str = Field("", description="项目 ID(可选)")
@model_validator(mode="after")
def _validate_input_mode(self) -> "CreateLipsyncJobRequest":
video = (self.video_url or "").strip()
if not video:
raise ValueError("video_url 不能为空")
if not video.startswith(("http://", "https://")):
raise ValueError("video_url 必须是 HTTP/HTTPS URL")
lower = video.lower().split("?")[0]
if not lower.endswith(".mp4"):
raise ValueError("video_url 仅支持 MP4 格式")
has_audio = bool((self.audio_url or "").strip())
has_tts = bool((self.voice_id or "").strip()) and bool((self.script_text or "").strip())
if not has_audio and not has_tts:
raise ValueError(
"必须提供驱动音频:要么传 audio_url(直接音频模式),"
"要么同时传 voice_id + script_textTTS 直生模式)"
)
if has_tts and len(self.script_text) > 5000:
raise ValueError("script_text 最长 5000 字符")
if has_audio:
au = self.audio_url.strip()
if not au.startswith(("http://", "https://")):
raise ValueError("audio_url 必须是 HTTP/HTTPS URL")
au_lower = au.lower().split("?")[0]
allowed = (".mp3", ".aac", ".wav", ".m4a", ".flac")
if not any(au_lower.endswith(ext) for ext in allowed):
raise ValueError(f"audio_url 格式不支持,仅支持: {', '.join(allowed)}")
self.audio_url = au
return self
-45
View File
@@ -1,45 +0,0 @@
"""Script (口播文案库) Pydantic schemas — Issue #1795."""
from __future__ import annotations
from datetime import datetime
from typing import List, Optional
from pydantic import BaseModel, Field
class ScriptSegment(BaseModel):
"""单段文案."""
text: str
duration: Optional[float] = None
class ScriptResponse(BaseModel):
id: str
user_id: str
title: str
content: str
segments: List[ScriptSegment] = Field(default_factory=list)
tags: List[str] = Field(default_factory=list)
created_at: datetime
updated_at: datetime
class ScriptListResponse(BaseModel):
items: list[ScriptResponse]
total: int = 0
class CreateScriptRequest(BaseModel):
title: str = Field(..., min_length=1, max_length=255)
content: str = ""
segments: List[ScriptSegment] = Field(default_factory=list)
tags: List[str] = Field(default_factory=list)
class UpdateScriptRequest(BaseModel):
title: Optional[str] = Field(None, min_length=1, max_length=255)
content: Optional[str] = None
segments: Optional[List[ScriptSegment]] = None
tags: Optional[List[str]] = None
-2
View File
@@ -16,7 +16,6 @@ class TTSSynthesizeRequest(BaseModel):
output_name: str = Field("", description="输出文件名") output_name: str = Field("", description="输出文件名")
language: str = Field("zh-CN", description="语言") language: str = Field("zh-CN", description="语言")
speed: float = Field(1.0, ge=0.5, le=2.0, description="语速") speed: float = Field(1.0, ge=0.5, le=2.0, description="语速")
emotion: str = Field("", description="情绪(natural/excited/calm/friendly,或中文 自然/兴奋/沉稳/亲切)")
voice_model: str = Field("", description="语音模型名称") voice_model: str = Field("", description="语音模型名称")
voice_clone_profile_id: str = Field("", description="关联的音色克隆档案 ID") voice_clone_profile_id: str = Field("", description="关联的音色克隆档案 ID")
format: str = Field("mp3", description="输出格式(mp3/wav/pcm") format: str = Field("mp3", description="输出格式(mp3/wav/pcm")
@@ -110,7 +109,6 @@ class TTSPreviewRequest(BaseModel):
text: str = Field(..., min_length=1, max_length=200, description="合成文本,限制 200 字") text: str = Field(..., min_length=1, max_length=200, description="合成文本,限制 200 字")
voice_id: str = Field(..., min_length=1, description="音色 ID") voice_id: str = Field(..., min_length=1, description="音色 ID")
speed: float = Field(1.0, ge=0.5, le=2.0, description="语速") speed: float = Field(1.0, ge=0.5, le=2.0, description="语速")
emotion: str = Field("", description="情绪(natural/excited/calm/friendly,或中文)")
pitch: float = Field(1.0, ge=0.5, le=2.0, description="音调(预留,当前未使用)") pitch: float = Field(1.0, ge=0.5, le=2.0, description="音调(预留,当前未使用)")
+5 -11
View File
@@ -16,7 +16,6 @@ class DirectUploadPrepareRequest(BaseModel):
content_type: str = Field(default="application/octet-stream", min_length=1, max_length=100) content_type: str = Field(default="application/octet-stream", min_length=1, max_length=100)
file_size: int = Field(..., gt=0) file_size: int = Field(..., gt=0)
file_hash: str = Field(default="", max_length=64, description="文件 MD5 哈希,用于去重检测") file_hash: str = Field(default="", max_length=64, description="文件 MD5 哈希,用于去重检测")
client_upload_id: str = Field(default="", max_length=64, description="客户端幂等 token(同一次上传的重试保持一致)")
class DirectUploadPrepareResponse(BaseModel): class DirectUploadPrepareResponse(BaseModel):
@@ -26,25 +25,20 @@ class DirectUploadPrepareResponse(BaseModel):
expires_at: str expires_at: str
fields: dict[str, str] fields: dict[str, str]
max_size_bytes: int max_size_bytes: int
duplicated: bool = False
skip_transfer: bool = False
asset_id: str = ""
class DirectUploadCompleteRequest(BaseModel): class DirectUploadCompleteRequest(BaseModel):
project_id: str = Field(..., min_length=1) project_id: str = Field(..., min_length=1)
library_id: str = Field(..., min_length=1) library_id: str = Field(..., min_length=1)
storage_key: str = Field(..., min_length=1, max_length=255) storage_key: str = Field(..., min_length=1, max_length=255)
file_hash: str = Field(default="", max_length=64, description="文件哈希,用于去重检测") file_hash: str = Field(default="", max_length=64, description="文件 MD5 哈希,用于去重检测")
client_upload_id: str = Field(default="", max_length=64, description="客户端幂等 token(同一次上传的重试保持一致)")
file_size: int = Field(default=0, ge=0, description="文件大小(字节),用于无 hash 时的兜底去重")
class DirectUploadCompleteResponse(BaseModel): class DirectUploadCompleteResponse(BaseModel):
storage_key: str storage_key: str
ingest_job_id: str ingest_job_id: str
duplicated: bool = Field(default=False, description="是否为重复素材/重复 complete(命中幂等去重)") duplicated: bool = Field(default=False, description="是否为重复素材(命中去重)")
asset_id: str = Field(default="", description="素材 asset_id重复 complete 时返回已存在记录") asset_id: str = Field(default="", description="重复素材 asset_idduplicated=true 时返回")
url: str = Field(default="", description="Public URL of uploaded file") url: str = Field(default="", description="Public URL of uploaded file")
@@ -52,5 +46,5 @@ class UploadAssetResponse(BaseModel):
storage_key: str storage_key: str
ingest_job_id: str ingest_job_id: str
url: str = Field(..., description="Public URL of uploaded file") url: str = Field(..., description="Public URL of uploaded file")
duplicated: bool = Field(default=False, description="是否为重复素材/重复提交(命中幂等去重)") duplicated: bool = Field(default=False, description="是否为重复素材(命中去重)")
asset_id: str = Field(default="", description="素材 asset_id重复提交时返回已存在记录") asset_id: str = Field(default="", description="重复素材 asset_idduplicated=true 时返回")
-3
View File
@@ -22,10 +22,7 @@ class VideoItemResponse(BaseModel):
generation_params: dict = Field(default_factory=dict) generation_params: dict = Field(default_factory=dict)
download_url: str | None = None download_url: str | None = None
generated_at: str = "" generated_at: str = ""
# #1660 查重率(百分比 0~100)/ 视觉相似度(0~1)/ 匹配帧数
duplicate_rate: float | None = None duplicate_rate: float | None = None
visual_similarity: float | None = None
match_count: int | None = None
class ListVideosResponse(BaseModel): class ListVideosResponse(BaseModel):
@@ -1,222 +0,0 @@
"""AI 数字人封面服务 — 复用智能剪辑的 MediaKit 抽帧 + 质量评分选最佳帧.
与 generation_cover.py 的智能选帧能力对齐(不再用 FFmpeg 简单截帧):
1. MediaKit extract_frames 抽取多帧(默认 5 帧,SpecifiedFrames 策略)
2. cover_frame_scorer.score_frames 按清晰度/亮度/色彩评分选最佳
3. 下载最佳帧并转存 OSS,返回公网封面 URL
降级:MediaKit 不可用或抽帧失败时返回空字符串,由调用方决定回退策略。
"""
from __future__ import annotations
import logging
import tempfile
import uuid
from pathlib import Path
from typing import Optional
from urllib.parse import urlparse
logger = logging.getLogger(__name__)
# MediaKit 抽帧轮询参数(与 MediaKit API timeout=60s 对齐)
COVER_POLL_INTERVAL = 3.0
COVER_MAX_POLL_ATTEMPTS = 20 # 最多等 60 秒
# 帧图片下载超时(秒)
FRAME_DOWNLOAD_TIMEOUT = 20
# 最佳帧下载超时(用于 persist)
BEST_FRAME_DOWNLOAD_TIMEOUT = 30
# 自家 OSS 私有桶 URL 重签有效期(供 MediaKit GPU worker 拉取)
MEDIAKIT_URL_TTL_SECONDS = 7 * 24 * 3600
def _sign_video_url_for_mediakit(video_url: str) -> str:
"""如果 video_url 是自家 OSS 私有桶 URL,重新签名为长有效期预签名 URL。
MediaKit GPU worker 需要能公网访问 video_url,裸 public_url 在私有桶下会 403。
"""
if not video_url:
return video_url
try:
from packages.shared.storage import get_shared_storage_service
storage = get_shared_storage_service()
public_base = getattr(storage, "public_url", "")
if not isinstance(public_base, str) or not public_base:
return video_url
own_host = urlparse(public_base).netloc.lower()
url_host = urlparse(video_url).netloc.lower()
if own_host and url_host == own_host:
# 是自家 OSS URL,重签 7 天有效期供 MediaKit 拉取
signed = storage.get_download_url(video_url, expires_seconds=MEDIAKIT_URL_TTL_SECONDS)
if signed:
logger.info("[数字人封面] video_url 已重签(自家 OSS 私有桶)")
return signed
except Exception:
logger.warning("[数字人封面] video_url 重签失败,使用原始 URL", exc_info=True)
return video_url
def select_best_cover_frame(video_url: str, *, max_frames: int = 5) -> str:
"""从视频抽取多帧并评分选最佳帧,返回最佳帧的临时 URL.
Args:
video_url: 可公网访问的视频 URL
max_frames: 抽帧数量
Returns:
最佳帧图片 URL;失败返回空字符串
"""
if not video_url:
return ""
# 确保 MediaKit 能访问 video_url(自家 OSS 私有桶需重签)
video_url = _sign_video_url_for_mediakit(video_url)
try:
from packages.shared.cover_frame_scorer import score_frames
from packages.shared.mediakit_client import get_mediakit_client
mk = get_mediakit_client()
if not mk.is_available:
logger.warning("[数字人封面] MediaKit 未配置,无法智能抽帧")
return ""
logger.info(
"[数字人封面] 开始抽帧: video_url=%s max_frames=%d poll_interval=%.1f max_poll=%d",
video_url[:80],
max_frames,
COVER_POLL_INTERVAL,
COVER_MAX_POLL_ATTEMPTS,
)
snapshots = mk.extract_frames(
video_url=video_url,
strategy="SpecifiedFrames",
max_frames=max_frames,
poll_interval=COVER_POLL_INTERVAL,
max_poll_attempts=COVER_MAX_POLL_ATTEMPTS,
max_retries=1,
)
if not snapshots:
logger.warning("[数字人封面] MediaKit 未返回帧: %s", video_url[:80])
return ""
if len(snapshots) == 1:
return snapshots[0].get("image_url") or snapshots[0].get("url") or ""
# 使用连接池下载各帧(复用 TCP 连接,减少延迟)
import httpx
candidates = []
with httpx.Client(timeout=FRAME_DOWNLOAD_TIMEOUT, follow_redirects=True) as client:
for snap in snapshots:
url = snap.get("image_url") or snap.get("url") or ""
if not url:
continue
tmp_path: Optional[str] = None
try:
resp = client.get(url)
resp.raise_for_status()
with tempfile.NamedTemporaryFile(suffix=".jpg", delete=False) as tmp:
tmp.write(resp.content)
tmp_path = tmp.name
candidates.append({"image_path": tmp_path, "url": url})
except Exception as e:
logger.warning("[数字人封面] 帧下载失败,跳过: url=%s err=%s", url[:80], e)
candidates.append({"image_path": None, "url": url, "score": 0.0})
if not candidates:
return snapshots[0].get("image_url") or snapshots[0].get("url") or ""
scored = score_frames(candidates)
best = scored[0] if scored else None
best_url = best.get("url", "") if best else ""
# 清理临时文件
for c in candidates:
p = c.get("image_path")
if p:
try:
Path(p).unlink(missing_ok=True)
except Exception:
pass
logger.info(
"[数字人封面] 智能选帧完成: candidates=%d best_score=%s",
len(candidates),
best.get("score") if best else "n/a",
)
return best_url
except Exception:
logger.warning("[数字人封面] 智能选帧失败", exc_info=True)
return ""
def persist_cover_to_oss(frame_url: str, *, job_id: str = "", prefix: str = "ai-avatar/covers") -> str:
"""下载帧图并转存到 OSS,返回公网封面 URL.
Args:
frame_url: MediaKit 返回的临时帧图 URL
job_id: 关联任务 ID(用于 OSS key 命名)
prefix: OSS key 前缀
Returns:
OSS 公网 URL;失败回退原始 frame_url
"""
if not frame_url:
return ""
tmp_path: Optional[str] = None
try:
import httpx
with httpx.Client(timeout=BEST_FRAME_DOWNLOAD_TIMEOUT, follow_redirects=True) as client:
resp = client.get(frame_url)
resp.raise_for_status()
if not resp.content:
logger.warning("[数字人封面] 帧图内容为空: %s", frame_url[:80])
return frame_url
with tempfile.NamedTemporaryFile(suffix=".jpg", delete=False) as tmp:
tmp.write(resp.content)
tmp_path = tmp.name
from packages.shared.storage import get_shared_storage_service
storage = get_shared_storage_service()
token = job_id or uuid.uuid4().hex[:12]
cover_key = f"{prefix}/{token}/cover_{uuid.uuid4().hex[:8]}.jpg"
public_url = storage.upload_file(
file_or_path=tmp_path,
storage_key=cover_key,
content_type="image/jpeg",
)
logger.info("[数字人封面] 封面已转存 OSS: key=%s", cover_key)
# 私有桶:返回预签名 URL(前端才能加载)
if public_url:
signed = storage.get_download_url(cover_key, expires_seconds=86400)
return signed
return frame_url
except Exception:
logger.warning("[数字人封面] 封面转存 OSS 失败,返回原始 URL", exc_info=True)
return frame_url
finally:
if tmp_path:
try:
Path(tmp_path).unlink(missing_ok=True)
except Exception:
pass
def generate_smart_cover(video_url: str, *, job_id: str = "", max_frames: int = 5) -> str:
"""一站式:MediaKit 智能抽帧选最佳 → 转存 OSS,返回封面公网 URL.
供独立封面接口与渲染管线复用。失败返回空字符串。
"""
best_frame = select_best_cover_frame(video_url, max_frames=max_frames)
if not best_frame:
return ""
return persist_cover_to_oss(best_frame, job_id=job_id)
@@ -1,426 +0,0 @@
"""AI数字人渲染合成 Service — #1798.
职责:
- 创建/查询/取消渲染任务
- 调用 Celery 异步任务执行渲染
- B-roll 合成 + 标题叠加 + 封面提取
- 用户隔离
"""
from __future__ import annotations
import logging
import os
import tempfile
import uuid
from datetime import datetime, timezone
from typing import Any, Optional
from sqlalchemy.orm import Session
from packages.adapters.sqlalchemy_impl.models import (
AiAvatarRenderJob,
LipsyncJobModel,
ScriptModel,
)
from packages.domain.video_filter_builder import (
build_cover_extract_command,
build_title_drawtext_filter,
)
from packages.shared.storage import get_shared_storage_service
logger = logging.getLogger(__name__)
class AiAvatarRenderError(Exception):
"""渲染服务异常."""
def __init__(self, message: str, code: str = "RenderError"):
self.code = code
super().__init__(message)
class AiAvatarRenderService:
"""AI数字人渲染合成 Service."""
def __init__(self, db: Session):
self.db = db
# ── 创建任务 ──────────────────────────────────────────────────────────
def create_render_job(
self,
*,
user_id: str,
lipsync_job_id: str,
script_id: str = "",
b_roll_segments: list[dict[str, Any]] | None = None,
title_config: dict[str, Any],
cover_config: dict[str, Any],
project_id: str = "",
) -> AiAvatarRenderJob:
"""创建渲染任务.
Raises:
AiAvatarRenderError: 校验失败
"""
# 1. 验证对口型任务
lipsync_job = (
self.db.query(LipsyncJobModel)
.filter(
LipsyncJobModel.id == lipsync_job_id,
LipsyncJobModel.user_id == user_id,
)
.first()
)
if lipsync_job is None:
raise AiAvatarRenderError("对口型任务不存在", code="LipsyncJobNotFound")
if lipsync_job.status != "completed":
raise AiAvatarRenderError(
f"对口型任务状态为 {lipsync_job.status},仅 completed 状态可渲染",
code="LipsyncJobNotCompleted",
)
if not lipsync_job.output_video_url:
raise AiAvatarRenderError("对口型任务输出视频 URL 为空", code="LipsyncJobNoOutput")
# 2. 验证文案归属(仅当选了文案库条目时;手动输入文案直生场景 script_id 可空)
script_id = (script_id or "").strip()
if script_id:
script = (
self.db.query(ScriptModel)
.filter(
ScriptModel.id == script_id,
ScriptModel.user_id == user_id,
)
.first()
)
if script is None:
raise AiAvatarRenderError("文案不存在或无权访问", code="ScriptNotFound")
# 3. 创建渲染任务
job_id = str(uuid.uuid4())
job = AiAvatarRenderJob(
id=job_id,
user_id=user_id,
project_id=project_id,
lipsync_job_id=lipsync_job_id,
script_id=script_id,
b_roll_segments=[s if isinstance(s, dict) else s.model_dump() for s in (b_roll_segments or [])],
title_config=title_config,
cover_config=cover_config,
status="pending",
)
self.db.add(job)
self.db.flush()
job.submitted_at = datetime.now(timezone.utc)
self.db.commit()
self.db.refresh(job)
return job
# ── 查询任务 ──────────────────────────────────────────────────────────
def get_render_job(self, job_id: str, user_id: str) -> Optional[AiAvatarRenderJob]:
"""获取渲染任务详情(用户隔离)."""
return (
self.db.query(AiAvatarRenderJob)
.filter(
AiAvatarRenderJob.id == job_id,
AiAvatarRenderJob.user_id == user_id,
)
.first()
)
def list_render_jobs(
self,
*,
user_id: str,
project_id: str = "",
status: str = "",
offset: int = 0,
limit: int = 20,
) -> tuple[list[AiAvatarRenderJob], int]:
"""获取渲染任务列表(分页 + 用户隔离)."""
query = self.db.query(AiAvatarRenderJob).filter(AiAvatarRenderJob.user_id == user_id)
if project_id:
query = query.filter(AiAvatarRenderJob.project_id == project_id)
if status:
query = query.filter(AiAvatarRenderJob.status == status)
total = query.count()
items = query.order_by(AiAvatarRenderJob.created_at.desc()).offset(offset).limit(limit).all()
return items, total
# ── 取消任务 ──────────────────────────────────────────────────────────
def cancel_render_job(self, job_id: str, user_id: str) -> Optional[AiAvatarRenderJob]:
"""取消渲染任务(仅 pending 状态可取消)."""
job = self.get_render_job(job_id, user_id)
if job is None:
return None
if job.status in ("pending", "submitted"):
job.status = "cancelled"
job.updated_at = datetime.now(timezone.utc)
self.db.commit()
self.db.refresh(job)
return job
# ── 重试任务 ──────────────────────────────────────────────────────────
def retry_render_job(self, job_id: str, user_id: str) -> Optional[AiAvatarRenderJob]:
"""重试失败的渲染任务."""
job = self.get_render_job(job_id, user_id)
if job is None:
return None
if job.status != "failed":
return None
job.status = "pending"
job.progress = 0
job.error_message = ""
job.output_video_url = ""
job.output_cover_url = ""
job.output_duration = 0.0
job.started_at = None
job.completed_at = None
job.updated_at = datetime.now(timezone.utc)
self.db.commit()
self.db.refresh(job)
return job
# ── 执行渲染(Celery 异步调用) ──────────────────────────────────────
def execute_render(self, job_id: str) -> None:
"""执行渲染管线.
由 Celery 异步任务调用,流程:
1. 下载对口型输出视频 (20%)
2. 构建 FFmpeg 滤镜链 (40%)
3. 执行 FFmpeg 渲染 (80%)
4. 提取封面 (90%)
5. 上传到 OSS (95%)
6. 更新任务状态 (100%)
"""
job = self.db.query(AiAvatarRenderJob).filter(AiAvatarRenderJob.id == job_id).first()
if job is None:
logger.error("渲染任务不存在: %s", job_id)
return
if job.status == "cancelled":
logger.info("渲染任务已取消: %s", job_id)
return
try:
# 更新状态为 processing
job.status = "processing"
job.started_at = datetime.now(timezone.utc)
job.progress = 5
job.updated_at = datetime.now(timezone.utc)
self.db.commit()
# 获取对口型任务信息
lipsync_job = self.db.query(LipsyncJobModel).filter(LipsyncJobModel.id == job.lipsync_job_id).first()
if lipsync_job is None:
raise AiAvatarRenderError("关联的对口型任务不存在", code="LipsyncJobNotFound")
# 1. 下载对口型输出视频 (20%)
input_video_path = self._download_video(lipsync_job.output_video_url)
job.progress = 20
self.db.commit()
# 2. 构建 FFmpeg 滤镜链 (40%)
from packages.domain.video_filter_builder import build_broll_overlay_filter
filter_complex = build_broll_overlay_filter(
b_roll_segments=job.b_roll_segments,
video_duration=lipsync_job.output_duration,
)
# 标题叠加
title_filter = build_title_drawtext_filter(job.title_config)
if title_filter:
if filter_complex:
filter_complex += f"[vout]{title_filter}[vout_titled];"
else:
filter_complex = f"[0:v]{title_filter}[vout_titled];"
# 清理末尾分号
if filter_complex.endswith(";"):
filter_complex = filter_complex[:-1]
# 最终输出标签
final_label = "vout_titled" if title_filter else ("vout" if filter_complex else None)
job.progress = 40
self.db.commit()
# 3. 执行 FFmpeg 渲染 (80%)
with tempfile.TemporaryDirectory() as tmpdir:
output_video_path = os.path.join(tmpdir, "output.mp4")
cmd = self._build_ffmpeg_command(
input_video=input_video_path,
b_roll_segments=job.b_roll_segments,
filter_complex=filter_complex,
final_label=final_label,
output_path=output_video_path,
)
exit_code = os.system(cmd)
if exit_code != 0:
raise AiAvatarRenderError(f"FFmpeg 渲染失败,退出码: {exit_code}", code="FFmpegFailed")
job.progress = 80
self.db.commit()
# 4. 提取封面 (90%)
cover_path = ""
if job.cover_config:
cover_path = os.path.join(tmpdir, "cover.jpg")
cover_cmd = build_cover_extract_command(job.cover_config, cover_path)
cover_cmd = cover_cmd.replace("INPUT_VIDEO", output_video_path)
cover_exit = os.system(cover_cmd)
if cover_exit != 0:
logger.warning("封面提取失败,跳过: %s", cover_cmd)
cover_path = ""
job.progress = 90
self.db.commit()
# 5. 上传到 OSS (95%)
output_video_url = self._upload_to_oss(output_video_path, f"ai-avatar/{job_id}/output.mp4")
job.output_video_url = output_video_url
# 封面:优先复用智能剪辑的 MediaKit 抽帧 + 质量评分选最佳帧;
# MediaKit 不可用时回退到 FFmpeg 已按 cover_config 抽取的 cover_path
smart_cover_url = ""
if output_video_url:
try:
from app.services.ai_avatar_cover_service import (
generate_smart_cover,
)
smart_cover_url = generate_smart_cover(output_video_url, job_id=job_id, max_frames=5)
except Exception:
logger.warning("智能封面(MediaKit)失败,回退 FFmpeg 封面 job_id=%s", job_id, exc_info=True)
if smart_cover_url:
job.output_cover_url = smart_cover_url
elif cover_path:
output_cover_url = self._upload_to_oss(cover_path, f"ai-avatar/{job_id}/cover.jpg")
job.output_cover_url = output_cover_url
# 获取输出视频时长
job.output_duration = lipsync_job.output_duration
job.progress = 95
self.db.commit()
# 6. 完成
job.status = "completed"
job.progress = 100
job.completed_at = datetime.now(timezone.utc)
job.updated_at = datetime.now(timezone.utc)
self.db.commit()
logger.info("渲染任务完成: %s", job_id)
# 7. 自动保存成片记录到成片库
if job.output_video_url:
try:
from packages.adapters.sqlalchemy_impl.generated_video_repository import (
SQLAlchemyGeneratedVideoRepository,
)
from packages.domain.generated_video import GeneratedVideo
clip_name = f"AI数字人_{job_id[:8]}"
clip = GeneratedVideo.create(
project_id=job.project_id,
generation_task_id=job.lipsync_job_id,
name=clip_name,
file_url=job.output_video_url,
user_id=job.user_id,
duration=job.output_duration or 0.0,
thumbnail_url=job.output_cover_url or None,
generation_params={
"source": "ai_avatar_render",
"render_job_id": job.id,
},
)
video_repo = SQLAlchemyGeneratedVideoRepository(self.db)
video_repo.create(clip)
logger.info("成片记录已保存到成片库: clip_id=%s, render_job=%s", clip.id, job_id)
except Exception as clip_err:
logger.warning(
"自动保存成片记录失败(不影响渲染任务状态): render_job=%s, error=%s",
job_id,
clip_err,
)
except AiAvatarRenderError as exc:
job.status = "failed"
job.error_message = str(exc)
job.updated_at = datetime.now(timezone.utc)
self.db.commit()
logger.error("渲染任务失败 [%s]: %s", job_id, exc)
except Exception as exc:
job.status = "failed"
job.error_message = f"渲染异常: {str(exc)}"
job.updated_at = datetime.now(timezone.utc)
self.db.commit()
logger.exception("渲染任务异常 [%s]", job_id)
def _download_video(self, url: str) -> str:
"""下载视频到临时文件."""
import httpx
tmp = tempfile.NamedTemporaryFile(suffix=".mp4", delete=False)
try:
with httpx.Client(timeout=120) as client:
resp = client.get(url)
resp.raise_for_status()
tmp.write(resp.content)
return tmp.name
except Exception:
if os.path.exists(tmp.name):
os.unlink(tmp.name)
raise
def _build_ffmpeg_command(
self,
*,
input_video: str,
b_roll_segments: list[dict[str, Any]],
filter_complex: str,
final_label: Optional[str],
output_path: str,
) -> str:
"""构建 FFmpeg 命令."""
# 输入文件
inputs = f"-i {input_video}"
for seg in b_roll_segments:
asset_url = seg.get("asset_url", "")
if asset_url:
inputs += f" -i {asset_url}"
# 滤镜
if filter_complex and final_label:
filter_arg = f'-filter_complex "{filter_complex}" -map "[{final_label}]"'
elif filter_complex:
filter_arg = f'-filter_complex "{filter_complex}"'
else:
filter_arg = ""
return f"ffmpeg {inputs} {filter_arg} -c:v libx264 -preset veryfast -crf 23 -y {output_path}"
def _upload_to_oss(self, local_path: str, oss_key: str) -> str:
"""上传文件到 OSS,返回 URL.
使用 SharedStorageService 统一存储服务。
"""
storage = get_shared_storage_service()
url = storage.upload_file_smart(local_path, oss_key)
if url is None:
raise AiAvatarRenderError(
f"上传文件到 OSS 失败: {oss_key}",
code="OSSUploadFailed",
)
logger.info("上传文件到 OSS 成功: %s -> %s", local_path, url)
return url
@@ -3,7 +3,7 @@
在素材 metadataassets.classification_result JSON)中持久化已使用的片段时间区间, 在素材 metadataassets.classification_result JSON)中持久化已使用的片段时间区间,
供 from-assets 创建片段时避开历史区间,实现跨任务/跨调用的片段去重; 供 from-assets 创建片段时避开历史区间,实现跨任务/跨调用的片段去重;
素材可用区间耗尽后进入受控复用:允许有限次数(MAX_RANGE_USE_COUNT)复用最久未用 素材可用区间耗尽后进入受控复用:允许有限次数(MAX_RANGE_USE_COUNT)复用最久未用
的历史区间,配合调用方的成片复用占比控制(MAX_REUSE_RATIO = 10%),把任意两条 的历史区间,配合调用方的成片复用占比控制(MAX_REUSE_RATIO = 15%),把任意两条
成片的画面重复率控制在阈值内。 成片的画面重复率控制在阈值内。
metadata 中的记录字段 ``used_time_ranges``:: metadata 中的记录字段 ``used_time_ranges``::
@@ -40,14 +40,14 @@ logger = logging.getLogger(__name__)
USED_RANGES_KEY = "used_time_ranges" USED_RANGES_KEY = "used_time_ranges"
# ── 受控复用配置常量 ───────────────────────────────────────────────────────── # ── 受控复用配置常量 ─────────────────────────────────────────────────────────
MAX_RANGE_USE_COUNT = 2 MAX_RANGE_USE_COUNT = 3
"""单条历史区间最多被使用次数(含首次),达到后不再参与复用。""" """单条历史区间最多被使用次数(含首次),达到后不再参与复用。"""
REUSE_RATIO_LIMIT = 0.10 REUSE_RATIO_LIMIT = 0.15
"""单条成片中,单个素材的复用片段累计时长 / 该素材在成片中的总时长上限(10%)。 """单条成片中,单个素材的复用片段累计时长 / 该素材在成片中的总时长上限(15%)。
超过则该素材不再分配新片段(调用方在轮询分配时跳过)。""" 超过则该素材不再分配新片段(调用方在轮询分配时跳过)。"""
SEGMENT_EDGE_GAP = 1.5 SEGMENT_EDGE_GAP = 0.3
"""冲突判定边缘间隙(秒):历史区间按 [start-gap, end+gap] 扩边后参与冲突检测, """冲突判定边缘间隙(秒):历史区间按 [start-gap, end+gap] 扩边后参与冲突检测,
避免两条片段首尾紧贴导致画面观感重复;记录仍存实际值。""" 避免两条片段首尾紧贴导致画面观感重复;记录仍存实际值。"""
@@ -397,12 +397,12 @@ def make_reuse_callback(
db: SQLAlchemy session db: SQLAlchemy session
asset_durations: 素材 ID -> 总时长(回调需要素材总时长做边界约束) asset_durations: 素材 ID -> 总时长(回调需要素材总时长做边界约束)
reused_tracker: 可选的 ``{asset_id: 累计复用时长}``,回调成功返回复用区间时 reused_tracker: 可选的 ``{asset_id: 累计复用时长}``,回调成功返回复用区间时
会把本次片段时长累加进去,供调用方统计成片复用占比(10% 阈值)。 会把本次片段时长累加进去,供调用方统计成片复用占比(15% 阈值)。
assigned_tracker: 可选的 ``{asset_id: 已分配片段总时长}``,配合 ratio_limit assigned_tracker: 可选的 ``{asset_id: 已分配片段总时长}``,配合 ratio_limit
在复用前预判:若复用本片段后占比 (reused + clip_duration) / 在复用前预判:若复用本片段后占比 (reused + clip_duration) /
(assigned + clip_duration) 超过 ratio_limit,则拒绝复用、返回 None (assigned + clip_duration) 超过 ratio_limit,则拒绝复用、返回 None
(保证成片复用占比不超阈值)。 (保证成片复用占比不超阈值)。
ratio_limit: 单条成片复用时长占比上限,默认 10% ratio_limit: 单条成片复用时长占比上限,默认 15%
Returns: Returns:
回调函数 ``(asset_id, clip_duration) -> (start, end) | None``。 回调函数 ``(asset_id, clip_duration) -> (start, end) | None``。
+3 -423
View File
@@ -409,13 +409,8 @@ class EditPlanService:
clip_type=clip_item.get("clip_type", "main"), clip_type=clip_item.get("clip_type", "main"),
order=order, order=order,
asset_id=clip_item.get("asset_id", ""), asset_id=clip_item.get("asset_id", ""),
text_content=clip_item.get("text_content", ""),
start_time=clip_item.get("start_time", 0.0), start_time=clip_item.get("start_time", 0.0),
duration=clip_item.get("duration", 0.0), duration=clip_item.get("duration", 0.0),
transition_effect=clip_item.get("transition_effect", "cut"),
transition_duration=clip_item.get("transition_duration", 0.0),
playback_speed=clip_item.get("playback_speed", 1.0),
config=clip_item.get("config") or None,
) )
model = EditPlanClipModel( model = EditPlanClipModel(
id=clip.id, id=clip.id,
@@ -464,159 +459,6 @@ class EditPlanService:
logger.exception("事务性替换片段失败: plan_id=%s", plan_id) logger.exception("事务性替换片段失败: plan_id=%s", plan_id)
raise raise
def reselect_plan_for_variant(
self,
source_plan_id: str,
candidate_asset_ids: list[str],
*,
created_by_user_id: str = "",
name_suffix: str = "变体",
voice_duration: float = 0.0,
rng=None,
) -> EditPlan:
"""为批量变体生成独立 plan:完整重跑单视频选片流程(#1743)。
与 clone_plan_for_variant(只重算起点、素材/顺序不变)不同,本方法:
- 源 plan 片段骨架(clip_type/order/duration/文案/转场)保留;
- 素材池 shuffle 随机分配 + main 片段顺序洗牌;
- 起点走场景镜头洗牌/随机起点/历史区间避让(与单视频同一入口);
- 批次内同素材区间重叠 >20% 自动重选起点;
- 新片段区间 record_used_segments 写回素材 metadata(跨变体/跨任务避让)。
Args:
source_plan_id: 源 plan(任务 0 / 预览源)。
candidate_asset_ids: 素材池(源 plan 素材 ∪ 批次素材)。
created_by_user_id: 新 plan 归属用户。
name_suffix: plan 名后缀。
rng: 可选随机数(测试注入种子)。
Raises:
ValueError: 源 plan 不存在/无片段、素材池为空或时长全未知。
"""
from packages.adapters.sqlalchemy_impl.models import AssetModel
from packages.domain.plan_generator_utils import extract_scene_points_from_metadata
from packages.domain.variant_plan_selector import reselect_clips_for_variant
source = self.get_plan_or_raise(source_plan_id)
# 分页读取源 plan 全部片段
clips: List[EditPlanClip] = []
skip, page = 0, 500
while True:
batch = self._clip_repo.list_by_plan(source_plan_id, skip=skip, limit=page)
if not batch:
break
clips.extend(batch)
if len(batch) < page:
break
skip += page
if not clips:
raise ValueError(f"源 plan 无片段,无法生成变体: {source_plan_id}")
source_clips_data: list[dict[str, Any]] = [
{
"order": c.order if c.order is not None else i,
"asset_id": c.asset_id,
"start_time": float(c.start_time or 0.0),
"duration": float(c.duration or 0.0),
"clip_type": c.clip_type,
"playback_speed": float(c.playback_speed or 1.0),
"transition_effect": c.transition_effect,
"transition_duration": float(c.transition_duration or 0.0),
"text_content": c.text_content or "",
"config": c.config or {},
}
for i, c in enumerate(clips)
]
db = self._clip_repo.session
# #1749:配音时长 → 每段目标段长(片段数=模板片段数定死;素材不足由渲染末帧冻结铺满)
target_durations: list[float] | None = None
try:
voice = float(voice_duration or 0.0)
except (TypeError, ValueError):
voice = 0.0
if voice > 0 and source_clips_data:
from packages.domain.voice_duration_planner import plan_clip_durations
_effects: list[str | None] = [c.get("transition_effect") for c in source_clips_data]
_tdurs: list[float] = [float(c.get("transition_duration") or 0.0) for c in source_clips_data]
target_durations = plan_clip_durations(
len(source_clips_data),
voice,
transition_effects=_effects,
transition_durations=_tdurs,
)
if target_durations:
for _c, _d in zip(source_clips_data, target_durations, strict=False):
_c["duration"] = _d
# 素材池 = 源 plan 素材 ∪ 调用方传入素材(去重保序)
pool_ids: list[str] = []
seen = set()
for aid in [c.asset_id for c in clips if c.asset_id] + list(candidate_asset_ids or []):
if aid and aid not in seen:
seen.add(aid)
pool_ids.append(aid)
# 时长 + 场景点
durations: dict[str, float] = {}
scene_points: dict[str, list[float]] = {}
if pool_ids:
for m in db.query(AssetModel).filter(AssetModel.id.in_(pool_ids)).all():
durations[m.id] = float(getattr(m, "duration", 0.0) or 0.0)
pts = extract_scene_points_from_metadata(getattr(m, "metadata", None))
if pts:
scene_points[m.id] = pts
historical = get_used_segments(db, pool_ids)
# 创建新 plan(复制模板归属与 config
new_plan = self.create_plan(
template_id=source.template_id,
name=f"{source.name or '剪辑计划'} · {name_suffix}",
config=dict(source.config or {}),
total_duration=source.total_duration,
project_id=source.project_id or "",
created_by_user_id=created_by_user_id or (source.created_by_user_id or ""),
)
# 批次内区间:以源 plan(变体 0)片段为初始避让对象
batch_segments: dict[str, list[tuple[float, float]]] = {}
for c in clips:
if c.asset_id and float(c.duration or 0) > 0:
st = float(c.start_time or 0.0)
batch_segments.setdefault(c.asset_id, []).append((st, st + float(c.duration)))
clips_data = reselect_clips_for_variant(
source_clips_data,
pool_ids,
asset_durations=durations,
asset_scene_points=scene_points,
historical_used_segments=historical,
batch_segments=batch_segments,
target_durations=target_durations,
rng=rng,
)
# 片段区间写回素材 metadata(与落库同事务;replace_all_clips_transactional 内 commit
for item in clips_data:
aid = item.get("asset_id", "")
if aid:
st = float(item.get("start_time", 0.0))
record_used_segments(db, aid, st, st + float(item.get("duration", 0.0)), new_plan.id)
self.replace_all_clips_transactional(new_plan.id, clips_data)
logger.info(
"变体独立选片完成: source=%s new=%s clips=%d assets=%d",
source_plan_id,
new_plan.id,
len(clips_data),
len(pool_ids),
)
return new_plan
def clone_plan_for_variant( def clone_plan_for_variant(
self, self,
source_plan_id: str, source_plan_id: str,
@@ -727,268 +569,6 @@ class EditPlanService:
) )
return new_plan return new_plan
# ── #1749 配音时长分配 / 素材时长查询 / 批量变体 plan 确保 ──────────────
def get_asset_durations(self, asset_ids: list[str]) -> dict[str, float]:
"""批量查询素材时长(秒),O(N) 单查;缺失/异常返回 0.0。"""
from packages.adapters.sqlalchemy_impl.models import AssetModel
ids = [a for a in dict.fromkeys(asset_ids or []) if a]
if not ids:
return {}
db = self._clip_repo.session
out: dict[str, float] = {}
for m in db.query(AssetModel).filter(AssetModel.id.in_(ids)).all():
try:
out[m.id] = float(getattr(m, "duration", 0.0) or 0.0)
except (TypeError, ValueError):
out[m.id] = 0.0
return out
def apply_voice_duration_to_plan(self, plan_id: str, voice_duration: float) -> Optional[EditPlan]:
"""把配音时长分配到 plan 的每段(#1749)。
- 片段数保持不变(= 模板片段数,定死);
- 每段 duration 按 voice_duration_planner 分配(含转场重叠扣减);
- 素材短于段长 → start_time 钳制为 0(末帧冻结由渲染侧 tpad/apad 铺满);
- plan.total_duration 回写为成片净时长(≈ 配音时长);
- 幂等:配音时长相同则分配结果不变,可重复调用。
无配音(<=0)或无片段时直接返回 None,不报错。
"""
try:
voice = float(voice_duration or 0.0)
except (TypeError, ValueError):
return None
if voice <= 0:
return None
plan = self.get_plan(plan_id)
if plan is None:
return None
clips: List[EditPlanClip] = []
skip, page = 0, 500
while True:
batch = self._clip_repo.list_by_plan(plan_id, skip=skip, limit=page)
if not batch:
break
clips.extend(batch)
if len(batch) < page:
break
skip += page
if not clips:
return None
clips.sort(key=lambda c: (c.order if c.order is not None else 0))
from packages.domain.voice_duration_planner import plan_clip_durations, total_output_duration
# #1764:从 plan config 读取节奏模板
rhythm_template = None
if plan and hasattr(plan, "config") and plan.config:
rhythm_template = plan.config.get("rhythm_template")
# #1768:先获取素材时长,传入 plan_clip_durations 用于最大片段钳制
asset_ids = [c.asset_id for c in clips if c.asset_id]
durations = self.get_asset_durations(asset_ids)
asset_durations_for_plan = [durations.get(c.asset_id, 0.0) for c in clips]
target = plan_clip_durations(
len(clips),
voice,
transition_effects=[c.transition_effect for c in clips],
transition_durations=[float(c.transition_duration or 0.0) for c in clips],
rhythm_template=rhythm_template,
asset_durations=asset_durations_for_plan,
)
if not target:
return None
clips_data: list[dict] = []
for i, c in enumerate(clips):
dur = float(target[i])
total = durations.get(c.asset_id, 0.0)
start = float(c.start_time or 0.0)
if c.asset_id and total > 0:
# 素材短于段长:起点钳 0,段长超出部分渲染侧末帧冻结
max_start = max(0.0, total - min(dur, total))
start = min(start, max_start)
clips_data.append(
{
"order": c.order if c.order is not None else i,
"asset_id": c.asset_id or "",
"start_time": round(start, 3),
"duration": dur,
"clip_type": c.clip_type,
"playback_speed": float(c.playback_speed or 1.0),
"transition_effect": c.transition_effect,
"transition_duration": float(c.transition_duration or 0.0),
"text_content": c.text_content or "",
"config": c.config or {},
}
)
self.replace_all_clips_transactional(plan_id, clips_data)
net = total_output_duration(
target,
transition_effects=[c.transition_effect for c in clips],
transition_durations=[float(c.transition_duration or 0.0) for c in clips],
)
try:
plan.total_duration = net
db = self._clip_repo.session
db.commit()
except Exception:
db.rollback()
logger.exception("回写 plan.total_duration 失败(不阻断): plan_id=%s", plan_id)
logger.info(
"配音时长分配完成: plan=%s clips=%d voice=%.2fs 成片净时长=%.2fs",
plan_id,
len(clips),
voice,
net,
)
return plan
def ensure_variant_plans(
self,
source_plan_id: str,
count: int,
candidate_asset_ids: list[str],
*,
created_by_user_id: str = "",
voice_durations: Optional[list[float]] = None,
rng=None,
) -> list[str]:
"""确保批量 N 个变体各自拥有独立 plan(#1749 批量正式生成/预览共用)。
- 变体 0clone 源 plan(不污染源 plan,片段独立可改),并按配音分配段长;
- 变体 1..N-1reselect_plan_for_variant 完整重跑选片(素材级去重);
- voice_durations:每个变体的配音时长(独立配音各自时长;统一配音同值);
缺省/为 0 时不分配(段长保持骨架/模板值)。
Returns:
plan_id 列表,长度 == countindex 即 variant_index。
"""
import random as _random
rng = rng or _random.Random()
plan_ids: list[str] = []
# 变体 0:clone(片段结构同源 plan,起点重算),不污染源 plan
plan0 = self.clone_plan_for_variant(
source_plan_id,
created_by_user_id=created_by_user_id,
name_suffix="变体1",
)
v0_voice = 0.0
if voice_durations and len(voice_durations) > 0:
try:
v0_voice = float(voice_durations[0] or 0.0)
except (TypeError, ValueError):
v0_voice = 0.0
if v0_voice > 0:
try:
self.apply_voice_duration_to_plan(plan0.id, v0_voice)
except Exception:
logger.exception("变体0 配音分配失败(不阻断): plan=%s", plan0.id)
plan_ids.append(plan0.id)
# 变体 1..N-1:独立选片
for i in range(1, count):
voice = 0.0
if voice_durations and i < len(voice_durations):
try:
voice = float(voice_durations[i] or 0.0)
except (TypeError, ValueError):
voice = 0.0
variant = self.reselect_plan_for_variant(
source_plan_id,
candidate_asset_ids,
created_by_user_id=created_by_user_id,
name_suffix=f"变体{i + 1}",
voice_duration=voice,
rng=rng,
)
plan_ids.append(variant.id)
# #1764:为每个变体生成独立节奏模板(让批量视频片段时长分布不同)
from packages.domain.voice_duration_planner import RHYTHM_TEMPLATES, adapt_template_length
clip_count = 0
if voice_durations and len(voice_durations) > 0:
# 从源 plan 获取片段数
source_plan = self.get_plan(source_plan_id)
if source_plan and hasattr(source_plan, "clips"):
clip_count = len(list(source_plan.clips)) if source_plan.clips else 0
rhythm_templates_for_variants = []
if clip_count > 0:
for idx in range(len(plan_ids)):
# 每个变体用不同的 seed 选择节奏模板
variant_seed = rng.randint(0, 999999)
template = adapt_template_length(RHYTHM_TEMPLATES[variant_seed % len(RHYTHM_TEMPLATES)], clip_count)
rhythm_templates_for_variants.append(template)
logger.info("变体 %d 节奏模板: plan=%s template=%s", idx, plan_ids[idx], template)
# #1767:BGM 池差异化分配(让批量变体使用不同 BGM / 段落 / 音量)
from packages.domain.bgm_pool import allocate_bgm_pool_for_variants
source_bgm_config = {}
source_plan = self.get_plan(source_plan_id)
if source_plan and source_plan.config:
source_bgm_config = source_plan.config.get("bgm", {}) or {}
variant_seeds_for_bgm = [rng.randint(0, 999999) for _ in plan_ids]
bgm_pool_assignments = allocate_bgm_pool_for_variants(source_bgm_config, variant_seeds_for_bgm)
# 为每个变体生成独立视觉扰动参数(让批量视频画面本身更不同)
from packages.domain.variant_plan_selector import generate_visual_perturbation
for idx, pid in enumerate(plan_ids):
try:
perturbation = generate_visual_perturbation(rng)
# 变体 0 不做 hflip(保持预览 plan 原始画面方向)
if idx == 0:
perturbation["hflip"] = False
config_update = {"visual_perturbation": perturbation}
# #1764:写入节奏模板
if idx < len(rhythm_templates_for_variants):
config_update["rhythm_template"] = rhythm_templates_for_variants[idx]
# #1765:写入像素级扰动滤镜
from packages.domain.variant_plan_selector import generate_pixel_perturbation
pixel_pert = generate_pixel_perturbation(rng)
config_update["pixel_perturbation"] = pixel_pert
# #1767:写入 BGM 池分配(覆盖 bgm 配置中的 preset_id / audio_offset / volume_adjust_db
if idx < len(bgm_pool_assignments):
existing_bgm = dict((source_plan.config or {}).get("bgm", {}) or {})
existing_bgm.update(bgm_pool_assignments[idx])
config_update["bgm"] = existing_bgm
self.update_plan_config(pid, config_update)
logger.info(
"变体 %d 视觉扰动+像素扰动+BGM池: plan=%s vis=%s pix=%s bgm=%s",
idx,
pid,
perturbation,
pixel_pert,
bgm_pool_assignments[idx] if idx < len(bgm_pool_assignments) else None,
)
except Exception:
logger.exception("变体 %d 视觉扰动生成失败(不阻断): plan=%s", idx, pid)
# 标记所有变体 plan 的 clips 为 ready(已分配素材+起点,语义上就是 ready)
for pid in plan_ids:
try:
self.mark_clips_ready(pid)
except Exception:
logger.exception("标记 clips ready 失败(不阻断): plan=%s", pid)
return plan_ids
# ── 片段分割与合并 ────────────────────────────────────────────────────── # ── 片段分割与合并 ──────────────────────────────────────────────────────
def split_clip(self, clip_id: str, split_time: float) -> Dict[str, Any]: def split_clip(self, clip_id: str, split_time: float) -> Dict[str, Any]:
@@ -1212,7 +792,7 @@ class EditPlanService:
config_asset_ids_count = len((plan.config or {}).get("asset_ids", [])) config_asset_ids_count = len((plan.config or {}).get("asset_ids", []))
clips_with_asset_count = sum(1 for c in clips if c.asset_id) clips_with_asset_count = sum(1 for c in clips if c.asset_id)
logger.info( logger.info(
"can_generate 诊断: plan=%s status=%s total_clips=%d clips_with_asset=%d config_asset_ids_count=%d", "can_generate 诊断: plan=%s status=%s total_clips=%d " "clips_with_asset=%d config_asset_ids_count=%d",
plan_id, plan_id,
plan.status, plan.status,
len(clips), len(clips),
@@ -1224,7 +804,7 @@ class EditPlanService:
config_asset_ids = (plan.config or {}).get("asset_ids", []) config_asset_ids = (plan.config or {}).get("asset_ids", [])
if config_asset_ids: if config_asset_ids:
logger.warning( logger.warning(
"can_generate 最后防线触发: plan=%s clips=%d 均无素材,从 config.asset_ids(%d个) 自动分配", "can_generate 最后防线触发: plan=%s clips=%d 均无素材," "从 config.asset_ids(%d个) 自动分配",
plan_id, plan_id,
len(clips), len(clips),
len(config_asset_ids), len(config_asset_ids),
@@ -1256,7 +836,7 @@ class EditPlanService:
return False, "没有可渲染的就绪片段,自动修复后仍未分配素材" return False, "没有可渲染的就绪片段,自动修复后仍未分配素材"
else: else:
logger.warning( logger.warning(
"can_generate 失败: plan=%s clips=%d 均无素材,且 config.asset_ids 为空,无法自动修复", "can_generate 失败: plan=%s clips=%d 均无素材," "且 config.asset_ids 为空,无法自动修复",
plan_id, plan_id,
len(clips), len(clips),
) )
+1 -65
View File
@@ -34,17 +34,6 @@ from packages.domain.template_clip_converter import (
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
class TemplateNotFoundError(Exception):
"""模板不存在、已删除或当前用户无权访问.
"模板存在但无片段配置"区分:路由层应映射为 HTTP 404。
"""
def __init__(self, template_id: str) -> None:
self.template_id = template_id
super().__init__(f"模板不存在: {template_id}")
class EditTemplateService: class EditTemplateService:
"""模板管理服务 """模板管理服务
@@ -228,14 +217,7 @@ class EditTemplateService:
skip: int = 0, skip: int = 0,
limit: int = 100, limit: int = 100,
) -> List[TemplateClipConfig]: ) -> List[TemplateClipConfig]:
"""列出模板的片段配置 """列出模板的片段配置"""
注意:本方法要求模板存在于新表 ``edit_templates``(全局模板库),
主要服务于新模板系统的写入/发布路径。用户自建模板存放在旧表
``templates``,不在 ``edit_templates`` 中,读取其片段配置请改用
:meth:`list_clip_configs_for_editor`,后者直接读取片段配置主表
``template_clip_configs``,不依赖新模板主表、也不靠异常降级。
"""
# 确保模板存在 # 确保模板存在
self.get_template_or_raise(template_id) self.get_template_or_raise(template_id)
return self._clip_config_repo.list_by_template( return self._clip_config_repo.list_by_template(
@@ -245,52 +227,6 @@ class EditTemplateService:
limit=limit, limit=limit,
) )
def list_clip_configs_for_editor(
self,
template_id: str,
user_id: str,
*,
clip_type: Optional[ClipType] = None,
skip: int = 0,
limit: int = 100,
) -> List[TemplateClipConfig]:
"""编辑器读取模板片段配置的单一数据源入口.
片段配置主表是 ``template_clip_configs``(直接读取,不抛异常、不降级)。
模板主表按双表现状显式判定,不使用 try/except 控制流:
1. 用户自建模板在旧表 ``templates``(归属 user_id)→ 校验归属与未删除后直接读;
2. 全局模板在新表 ``edit_templates``(无 user_id,全局可读)→ 直接读;
3. 两者都没有 → 模板不存在/无权限,抛 :class:`TemplateNotFoundError`。
Args:
template_id: 模板 ID
user_id: 当前登录用户 ID(用于旧表模板归属校验)
Raises:
TemplateNotFoundError: 模板不存在、已删除或不归属于当前用户。
"""
# 1) 用户自建模板(旧表 templates,归属 user_id
if self._clip_config_repo.template_owned_by(template_id, user_id):
return self._clip_config_repo.list_by_template(
template_id,
clip_type=clip_type,
skip=skip,
limit=limit,
)
# 2) 全局模板(新表 edit_templates,无 user_id,全局可读)
if self._template_repo.get(template_id) is not None:
return self._clip_config_repo.list_by_template(
template_id,
clip_type=clip_type,
skip=skip,
limit=limit,
)
# 3) 两表都没有:不存在 / 已删除 / 无权限
raise TemplateNotFoundError(template_id)
def get_clip_config(self, config_id: str) -> Optional[TemplateClipConfig]: def get_clip_config(self, config_id: str) -> Optional[TemplateClipConfig]:
"""获取片段配置详情""" """获取片段配置详情"""
return self._clip_config_repo.get(config_id) return self._clip_config_repo.get(config_id)
-398
View File
@@ -1,398 +0,0 @@
"""对口型 Service — #1796 MediaKit 对口型业务逻辑, #1809 参数调整.
职责:
- 创建/查询对口型任务
- 双输入模式:TTS 直生(voice_id + script_text,内部先合成音频转存 OSS)或直接音频(audio_url
- 调用 MediaKit 客户端提交异步任务
- 轮询更新任务状态(中间状态同步 DB,成片转存自家 OSS)
- 用户隔离(每个用户只能操作自己的任务)
"""
from __future__ import annotations
import io
import logging
import uuid
from datetime import datetime, timezone
from typing import Optional
from urllib.parse import urlparse
from app.services.mediakit_client import (
STATUS_COMPLETED,
STATUS_FAILED,
STATUS_RUNNING,
MediaKitClient,
MediaKitError,
get_mediakit_client,
)
# Celery 异步任务:TTS 合成 + MediaKit 提交(#lipsync-speed-optimization
from app.tasks.lipsync_tts import tts_synthesize_and_submit
from sqlalchemy.orm import Session
from packages.adapters.sqlalchemy_impl.models import LipsyncJobModel
from packages.application.cosyvoice_service import CosyVoiceError, normalize_emotion
from packages.shared.storage import get_shared_storage_service
from packages.shared.url_security import ALLOWED_AUDIO_MIME_TYPES, safe_download_bytes
logger = logging.getLogger(__name__)
# 传给 MediaKit GPU worker / 回给前端播放的 OSS 预签名有效期:7 天。
# MediaKit 排队 + 拉取可能延迟,私有桶裸 URL 或 1 小时短预签名都会 403,故统一重签长有效期。
MEDIAKIT_URL_TTL_SECONDS = 7 * 24 * 3600
class LipsyncService:
"""对口型任务 Service."""
def __init__(
self,
db: Session,
client: Optional[MediaKitClient] = None,
cosyvoice_service=None,
voice_clone_repo=None,
):
self.db = db
self.client = client or get_mediakit_client()
self._cosyvoice = cosyvoice_service
self._voice_clone_repo = voice_clone_repo
def _get_cosyvoice(self):
"""延迟获取 CosyVoiceService(与 tts 路由一致,含 OSS 预签名配置)."""
if self._cosyvoice is None:
from app.dependencies import get_cosyvoice_service
self._cosyvoice = get_cosyvoice_service()
return self._cosyvoice
def _resolve_voice_id(self, voice_id: str, user_id: str) -> str:
"""将克隆音色 profile UUID 解析为 CosyVoice voice_id。
与 /tts/synthesize 保持一致:命中 profile → 校验归属 → 返回其 voice_id;
未命中(预置音色 ID 或克隆 CosyVoice voice_id)原样返回。
"""
if not voice_id:
return ""
if self._voice_clone_repo is None:
try:
from app.dependencies import get_voice_clone_profile_repository
self._voice_clone_repo = get_voice_clone_profile_repository(self.db)
except Exception:
return voice_id
try:
profile = self._voice_clone_repo.get(voice_id)
except Exception:
return voice_id
if profile is None:
return voice_id
if getattr(profile, "user_id", "") != user_id:
raise MediaKitError("无权访问该音色", code="VoiceForbidden")
if not getattr(profile, "voice_id", ""):
raise MediaKitError("音色克隆尚未完成,请稍后再试", code="VoiceNotReady")
return profile.voice_id
def _synthesize_and_persist_audio(
self,
*,
user_id: str,
job_id: str,
voice_id: str,
script_text: str,
speed: float,
emotion: str,
) -> str:
"""TTS 直生:调 CosyVoice 合成音频并转存 OSS,返回可公网访问的音频 URL.
Raises:
MediaKitError: 合成失败
"""
actual_voice_id = self._resolve_voice_id(voice_id, user_id)
cosyvoice = self._get_cosyvoice()
try:
result = cosyvoice.submit_synthesize_task(
text=script_text,
voice_id=actual_voice_id,
speed=speed,
emotion=normalize_emotion(emotion),
)
except CosyVoiceError as exc:
raise MediaKitError(f"TTS 合成失败: {exc}", code="TTSSynthesisFailed") from exc
except ValueError as exc:
raise MediaKitError(f"TTS 参数错误: {exc}", code="TTSInvalidParam") from exc
temp_url = result.get("audio_url", "")
if not temp_url:
raise MediaKitError("TTS 未返回音频 URL", code="TTSNoAudio")
# 转存到自家 OSS,避免临时 URL 过期导致 MediaKit 拉取失败
try:
audio_data = safe_download_bytes(
temp_url,
purpose="lipsync_tts_audio",
allowed_mime_types=ALLOWED_AUDIO_MIME_TYPES,
timeout=60.0,
)
storage = get_shared_storage_service()
storage_key = f"lipsync-tts/{user_id}/{job_id}.mp3"
permanent_url = storage.upload_file(io.BytesIO(audio_data), storage_key, content_type="audio/mpeg")
logger.info("对口型 TTS 音频已转存 OSS: job_id=%s key=%s", job_id, storage_key)
return permanent_url
except Exception as exc:
logger.warning("TTS 音频转存 OSS 失败,回退临时 URL: job_id=%s err=%s", job_id, exc)
return temp_url
# ── 创建任务 ──────────────────────────────────────────────────────────
def create_job(
self,
*,
user_id: str,
video_url: str,
audio_url: str = "",
voice_id: str = "",
script_text: str = "",
speed: float = 1.0,
emotion: str = "",
enable_video_loop: bool = False,
project_id: str = "",
) -> LipsyncJobModel:
"""创建对口型任务.
两种输入模式:
- TTS 直生:voice_id + script_textaudio_url 留空)
→ 先创建 DB 记录(状态 tts_processing),再 dispatch Celery 异步任务
执行 TTS 合成 + MediaKit 提交。API 响应 <1s。
- 直接音频:提供 audio_url
→ 同步提交 MediaKit,状态直接设为 submitted。
Raises:
MediaKitError: 参数校验失败或 MediaKit 提交失败(仅直接音频模式)
"""
# 0. 输入校验
if not audio_url:
if not (voice_id and script_text):
raise MediaKitError(
"必须提供 audio_url 或 voice_id+script_text",
code="InvalidInput",
)
# TTS 模式:在 HTTP 请求中同步校验音色归属,快速失败
self._resolve_voice_id(voice_id, user_id)
# 1. 创建数据库记录
job_id = str(uuid.uuid4())
is_tts_mode = not bool(audio_url)
job = LipsyncJobModel(
id=job_id,
user_id=user_id,
project_id=project_id,
video_url=video_url,
audio_url=audio_url,
enable_video_loop=enable_video_loop,
voice_id=voice_id or "",
script_text=script_text or "",
speed=speed,
emotion=normalize_emotion(emotion),
status="tts_processing" if is_tts_mode else "pending",
)
self.db.add(job)
self.db.flush()
if is_tts_mode:
# 2a. TTS 模式:dispatch Celery 异步任务处理 TTS 合成 + MediaKit 提交
try:
tts_synthesize_and_submit.apply_async(
args=(
job_id,
user_id,
voice_id,
script_text,
speed,
normalize_emotion(emotion),
)
)
except Exception as exc:
# 投递失败时立即把 job 标成 failed 并写入 error_message
# 前端轮询时能直接看到失败原因,不会无限卡在 tts_processing。
logger.exception(
"Celery 任务提交失败,TTS 任务已创建但未触发执行: job_id=%s err=%s",
job_id,
exc,
)
job.status = "failed"
job.error_message = f"Celery 任务投递失败: {exc}"
job.error_code = "AsyncDispatchFailed"
job.updated_at = datetime.now(timezone.utc)
else:
# 2b. 直接音频模式:同步签名并提交 MediaKit
video_url = self._sign_media_url(video_url)
if audio_url:
audio_url = self._sign_media_url(audio_url)
job.audio_url = audio_url
try:
result = self.client.submit_lipsync(
video_url=video_url,
audio_url=audio_url,
enable_video_loop=enable_video_loop,
client_token=job_id,
)
job.mediakit_task_id = result["task_id"]
job.status = "submitted"
job.submitted_at = datetime.now(timezone.utc)
except MediaKitError as exc:
job.status = "failed"
job.error_message = str(exc)
job.error_code = exc.code
logger.error("提交对口型任务失败: %s", exc)
raise
self.db.commit()
self.db.refresh(job)
return job
# ── 查询任务 ──────────────────────────────────────────────────────────
def get_job(self, job_id: str, user_id: str) -> Optional[LipsyncJobModel]:
"""获取任务详情(用户隔离)."""
return (
self.db.query(LipsyncJobModel)
.filter(LipsyncJobModel.id == job_id, LipsyncJobModel.user_id == user_id)
.first()
)
def list_jobs(
self,
*,
user_id: str,
project_id: str = "",
status: str = "",
offset: int = 0,
limit: int = 20,
) -> tuple[list[LipsyncJobModel], int]:
"""获取任务列表(分页 + 用户隔离)."""
query = self.db.query(LipsyncJobModel).filter(LipsyncJobModel.user_id == user_id)
if project_id:
query = query.filter(LipsyncJobModel.project_id == project_id)
if status:
query = query.filter(LipsyncJobModel.status == status)
total = query.count()
items = query.order_by(LipsyncJobModel.created_at.desc()).offset(offset).limit(limit).all()
return items, total
# ── 更新任务状态(轮询) ──────────────────────────────────────────────
def refresh_job_status(self, job_id: str, user_id: str) -> Optional[LipsyncJobModel]:
"""从 MediaKit 拉取最新状态并更新本地记录.
Returns:
更新后的 Job,或 None(任务不存在/不属于该用户)
"""
job = self.get_job(job_id, user_id)
if job is None:
return None
# 终态不需要再轮询
if job.status in (STATUS_COMPLETED, "failed"):
return job
# 未提交的任务不轮询
if not job.mediakit_task_id:
return job
try:
status_data = self.client.get_task_status(job.mediakit_task_id)
except MediaKitError as exc:
logger.error("轮询对口型任务状态失败 [%s]: %s", job_id, exc)
return job
mk_status = status_data.get("status", STATUS_RUNNING)
logger.info("MediaKit 对口型状态 [%s]: %s", job_id, mk_status)
if mk_status == STATUS_COMPLETED:
result = status_data.get("result", {})
job.status = STATUS_COMPLETED
output_url = result.get("video_url", "")
# MediaKit 输出为临时 URL,转存自家 OSS 防止过期(失败则回退临时 URL)
job.output_video_url = self._persist_output_video(output_url, job_id, user_id)
job.output_duration = result.get("duration", 0.0)
job.completed_at = datetime.now(timezone.utc)
elif mk_status == STATUS_FAILED:
error = status_data.get("error", {})
job.status = "failed"
job.error_message = error.get("message", "任务执行失败")
job.error_code = error.get("code", "TaskFailed")
job.completed_at = datetime.now(timezone.utc)
else:
# 中间状态(running/processing/queued 等)同步到 DB,避免前端永远卡在 submitted
if isinstance(mk_status, str) and mk_status:
job.status = mk_status
job.updated_at = datetime.now(timezone.utc)
self.db.commit()
self.db.refresh(job)
return job
def _persist_output_video(self, temp_url: str, job_id: str, user_id: str) -> str:
"""将 MediaKit 输出的临时视频 URL 转存到自家 OSS.
失败时回退返回原始临时 URL,不影响任务完成。
"""
if not temp_url:
return ""
try:
import httpx
with httpx.Client(timeout=180.0, follow_redirects=True) as client:
resp = client.get(temp_url)
resp.raise_for_status()
data = resp.content
storage = get_shared_storage_service()
storage_key = f"lipsync-outputs/{user_id}/{job_id}.mp4"
permanent_url = storage.upload_file(io.BytesIO(data), storage_key, content_type="video/mp4")
logger.info("对口型输出视频已转存 OSS: job_id=%s key=%s", job_id, storage_key)
return self._sign_media_url(permanent_url) or temp_url
except Exception as exc:
logger.warning("对口型输出视频转存 OSS 失败,回退临时 URL: job_id=%s err=%s", job_id, exc)
return temp_url
def _sign_media_url(self, url: str) -> str:
"""对自家 OSS 私有桶 URL 重签长有效期预签名,供 MediaKit 拉取 / 前端播放。
- 裸 public_urlupload_file 返回,不带签名)→ 私有桶匿名访问 403,重签。
- 已带签名但即将过期的 URL(如前端 1h 预签名)→ 抽 storage_key 后重签。
- 外部 URLCosyVoice/MediaKit 临时链接,非本桶 host)→ 原样透传。
- 任何异常都降级原样返回,不阻断主流程。
"""
if not url:
return url
try:
storage = get_shared_storage_service()
public_base = getattr(storage, "public_url", "")
if not isinstance(public_base, str) or not public_base:
return url # 无法判定归属,保守透传
own_host = urlparse(public_base).netloc.lower()
host = urlparse(url).netloc.lower()
if not own_host or host != own_host:
return url # 非自家 OSS(外部临时链接),不处理
signed = storage.get_download_url(url, expires_seconds=MEDIAKIT_URL_TTL_SECONDS)
return signed or url
except Exception as exc: # noqa: BLE001 - 签名失败不阻断,降级原 URL
logger.warning("对口型 URL 重签失败,原样返回: url_prefix=%s err=%s", url[:80], exc)
return url
# ── 取消任务 ──────────────────────────────────────────────────────────
def cancel_job(self, job_id: str, user_id: str) -> Optional[LipsyncJobModel]:
"""取消任务(仅 pending/tts_processing/submitted 状态可取消)."""
job = self.get_job(job_id, user_id)
if job is None:
return None
if job.status in ("pending", "tts_processing", "submitted"):
job.status = "cancelled"
job.updated_at = datetime.now(timezone.utc)
self.db.commit()
self.db.refresh(job)
return job
-243
View File
@@ -1,243 +0,0 @@
"""MediaKit 客户端 — 封装火山引擎 AI MediaKit 对口型 API.
接口文档:https://docs.volcengine.com/docs/6448/2656064
异步任务流程:
1. POST /api/v1/tools/lip-sync 提交对口型任务 → 返回 task_id
2. GET /api/v1/tasks/{task_id} 轮询任务状态 → running/completed/failed
3. completed 时 result.video_url 为口型对齐视频(临时链接 24h 有效)
设计原则:
- API Key 从配置读取(settings.mediakit_api_key
- 未配置 API Key 时所有方法返回降级响应,不阻塞主流程
- HTTP 超时/网络异常统一包装为 MediaKitError
"""
from __future__ import annotations
import logging
from typing import Any, Optional
import httpx
from packages.config import get_api_settings
logger = logging.getLogger(__name__)
# ── 任务状态常量 ──────────────────────────────────────────────────────────
STATUS_RUNNING = "running"
STATUS_COMPLETED = "completed"
STATUS_FAILED = "failed"
class MediaKitError(Exception):
"""MediaKit API 调用异常."""
def __init__(self, message: str, code: str = "", request_id: str = ""):
self.code = code
self.request_id = request_id
super().__init__(message)
class MediaKitClient:
"""火山引擎 AI MediaKit 对口型 API 客户端.
用法:
client = get_mediakit_client()
result = client.submit_lipsync(video_url="...", audio_url="...")
task_id = result["task_id"]
status = client.get_task_status(task_id)
# {"status": "completed", "result": {"video_url": "...", "duration": 60.5}}
"""
def __init__(self) -> None:
settings = get_api_settings()
self._api_key = settings.mediakit_api_key
self._base_url = settings.mediakit_base_url.rstrip("/")
self._timeout = settings.mediakit_timeout
@property
def is_available(self) -> bool:
"""是否已配置 API Key(未配置时自动降级)."""
return bool(self._api_key)
def _headers(self) -> dict[str, str]:
return {
"Authorization": f"Bearer {self._api_key}",
"Content-Type": "application/json",
}
# ── 提交对口型任务 ────────────────────────────────────────────────────
def submit_lipsync(
self,
*,
video_url: str,
audio_url: str,
enable_video_loop: bool = False,
callback_url: Optional[str] = None,
callback_args: Optional[str] = None,
client_token: Optional[str] = None,
) -> dict[str, Any]:
"""提交视频口型对齐任务.
Args:
video_url: 人物视频 URLMP4,≤30min,单人真人)
audio_url: 驱动音频 URLmp3/aac/wav/m4a/flac
enable_video_loop: 音频长于视频时是否循环画面
callback_url: 任务完成回调 URL
callback_args: 回调时原样返回的自定义参数
client_token: 幂等控制 token
Returns:
{"success": True, "task_id": "...", "request_id": "..."}
Raises:
MediaKitError: API 调用失败
"""
if not self.is_available:
raise MediaKitError("MediaKit API Key 未配置", code="NotConfigured")
payload: dict[str, Any] = {
"video_url": video_url,
"audio_url": audio_url,
}
if enable_video_loop:
payload["enable_video_loop"] = True
if callback_url:
payload["callback_url"] = callback_url
if callback_args:
payload["callback_args"] = callback_args[:512] # API 限制 512 字节
if client_token:
payload["client_token"] = client_token[:64] # API 限制 64 字符
try:
with httpx.Client(timeout=self._timeout) as client:
resp = client.post(
f"{self._base_url}/tools/lip-sync",
headers=self._headers(),
json=payload,
)
resp.raise_for_status()
data = resp.json()
except httpx.TimeoutException as exc:
raise MediaKitError(f"MediaKit API 超时 ({self._timeout}s)", code="Timeout") from exc
except httpx.HTTPStatusError as exc:
body = exc.response.text[:500]
raise MediaKitError(
f"MediaKit API HTTP {exc.response.status_code}: {body}",
code="HttpError",
) from exc
except httpx.RequestError as exc:
raise MediaKitError(f"MediaKit API 网络错误: {exc}", code="NetworkError") from exc
except Exception as exc:
raise MediaKitError(f"MediaKit API 未知错误: {exc}", code="UnknownError") from exc
if not data.get("success"):
error = data.get("error", {})
raise MediaKitError(
error.get("message", "提交任务失败"),
code=error.get("code", "SubmitFailed"),
request_id=data.get("request_id", ""),
)
return {
"success": True,
"task_id": data["task_id"],
"request_id": data.get("request_id", ""),
}
# ── 查询任务状态 ──────────────────────────────────────────────────────
def get_task_status(self, task_id: str) -> dict[str, Any]:
"""查询异步任务状态和结果.
Args:
task_id: 提交任务时返回的任务 ID
Returns:
{
"success": True,
"task_id": "...",
"status": "running" | "completed" | "failed",
"result": {"video_url": "...", "duration": 60.5} | None,
"error": {"code": "...", "message": "..."} | None,
"created_at": 1777291767,
"finished_at": 1777291851 | None,
"expires_at": 1777464650 | None,
}
Raises:
MediaKitError: API 调用失败
"""
if not self.is_available:
raise MediaKitError("MediaKit API Key 未配置", code="NotConfigured")
try:
with httpx.Client(timeout=self._timeout) as client:
resp = client.get(
f"{self._base_url}/tasks/{task_id}",
headers=self._headers(),
)
resp.raise_for_status()
data = resp.json()
except httpx.TimeoutException as exc:
raise MediaKitError(f"MediaKit API 超时 ({self._timeout}s)", code="Timeout") from exc
except httpx.HTTPStatusError as exc:
body = exc.response.text[:500]
raise MediaKitError(
f"MediaKit API HTTP {exc.response.status_code}: {body}",
code="HttpError",
) from exc
except httpx.RequestError as exc:
raise MediaKitError(f"MediaKit API 网络错误: {exc}", code="NetworkError") from exc
except Exception as exc:
raise MediaKitError(f"MediaKit API 未知错误: {exc}", code="UnknownError") from exc
if not data.get("success"):
error = data.get("error", {})
raise MediaKitError(
error.get("message", "查询任务失败"),
code=error.get("code", "QueryFailed"),
request_id=data.get("request_id", ""),
)
result: dict[str, Any] = {
"success": True,
"task_id": data.get("task_id", task_id),
"status": data.get("status", STATUS_RUNNING),
"result": data.get("result"),
"created_at": data.get("created_at"),
"finished_at": data.get("finished_at"),
"expires_at": data.get("expires_at"),
}
# 失败时提取错误信息
if data.get("status") == STATUS_FAILED:
error_obj = data.get("error", {})
result["error"] = {
"code": error_obj.get("code", "TaskFailed"),
"message": error_obj.get("message", "任务执行失败"),
}
return result
# ── 单例 ──────────────────────────────────────────────────────────────────
_client: Optional[MediaKitClient] = None
def get_mediakit_client() -> MediaKitClient:
"""获取 MediaKit 客户端单例."""
global _client
if _client is None:
_client = MediaKitClient()
return _client
def reset_mediakit_client() -> None:
"""重置客户端(测试用)."""
global _client
_client = None
@@ -13,7 +13,6 @@
from __future__ import annotations from __future__ import annotations
import logging import logging
import random
from typing import Any, List from typing import Any, List
from sqlalchemy.orm import Session from sqlalchemy.orm import Session
@@ -30,11 +29,9 @@ from packages.domain.editing_mode import EditingMode
from packages.domain.plan_generator_utils import ( from packages.domain.plan_generator_utils import (
create_clips_from_configs, create_clips_from_configs,
distribute_assets, distribute_assets,
extract_scene_points_from_metadata,
generate_default_clips, generate_default_clips,
map_clip_types_for_mode, map_clip_types_for_mode,
) )
from packages.domain.smart_match import SCORE_RANDOM_NOISE_MAX, score_asset
from packages.domain.template_clip_config import TemplateClipConfig from packages.domain.template_clip_config import TemplateClipConfig
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -131,7 +128,6 @@ class PlanGeneratorService:
editing_mode, editing_mode,
random_selection=random_preview, random_selection=random_preview,
asset_durations=asset_durations, asset_durations=asset_durations,
user_id=created_by_user_id,
) )
# 5. 持久化所有 clips 并计算总时长 # 5. 持久化所有 clips 并计算总时长
@@ -219,82 +215,19 @@ class PlanGeneratorService:
*, *,
random_selection: bool = False, random_selection: bool = False,
asset_durations: dict[str, float] | None = None, asset_durations: dict[str, float] | None = None,
user_id: str = "",
) -> None: ) -> None:
"""按 editing_mode 将素材分配到 clips(就地修改,未持久化). """按 editing_mode 将素材分配到 clips(就地修改,未持久化).
先用 smart_match 评分对素材排序(高分优先),再委托给 委托给 plan_generator_utils.distribute_assets 纯函数。
plan_generator_utils.distribute_assets 纯函数完成分配。
""" """
# 预览随机模式:素材顺序已 shuffle,纯随机起点即可,不读 DB 评分/缓存
asset_scene_points: dict[str, list[float]] = {}
if not random_selection:
# 正式生成:smart_match 评分排序(高分优先)+ 场景切换点缓存
if self._asset_repo:
asset_ids = self._sort_assets_by_smart_score(asset_ids)
# 读取素材 metadata 中的场景切换点缓存(后台 SceneChange 检测写入):
# 有缓存的素材片段起点从随机镜头段选取,无缓存走随机起点兜底
asset_scene_points = self._fetch_asset_scene_points(asset_ids)
# 正式生成也随机重排片段顺序(降重,默认开启无开关)
# smart_match 决定选哪些素材,shuffle 只改变分配到 clips 的顺序
asset_ids = list(asset_ids) # 复制避免修改调用方原列表
random.shuffle(asset_ids)
# 查询已有视频的已用区间(跨视频避让)
external_used_segments = None
if user_id and self._clip_repo:
try:
external_used_segments = self._clip_repo.list_used_segments_by_user(user_id, limit_recent=50)
except Exception:
logger.warning("跨视频避让查询失败,回退到纯随机", exc_info=True)
distribute_assets( distribute_assets(
clips, clips,
asset_ids, asset_ids,
editing_mode, editing_mode,
random_selection=random_selection, random_selection=random_selection,
asset_durations=asset_durations, asset_durations=asset_durations,
asset_scene_points=asset_scene_points,
external_used_segments=external_used_segments,
) )
def _fetch_asset_scene_points(self, asset_ids: List[str]) -> dict[str, list[float]]:
"""从素材 metadata 读取场景切换点缓存(无缓存的素材不包含在结果中)。"""
points_map: dict[str, list[float]] = {}
if not self._asset_repo:
return points_map
for asset_id in asset_ids:
asset = self._asset_repo.get(asset_id)
if asset:
points = extract_scene_points_from_metadata(getattr(asset, "metadata", None))
if points:
points_map[asset_id] = points
return points_map
def _sort_assets_by_smart_score(self, asset_ids: List[str]) -> List[str]:
"""按 smart_match 综合评分降序排列素材 ID(注入随机噪声)。
评分高的素材(质量好、时长合适、新鲜、使用次数少)倾向排在前面;
排序时给每个素材的得分注入 0~SCORE_RANDOM_NOISE_MAX 的随机噪声,
使得分接近的素材排名每次浮动,避免一键生成反复选出相同素材组合,
从素材组合层面降低成片查重率。分差大于噪声上限时排名保持稳定。
"""
scored: list[tuple[str, float]] = []
for asset_id in asset_ids:
asset = self._asset_repo.get(asset_id)
if asset:
score, _ = score_asset(asset)
scored.append((asset_id, score))
else:
scored.append((asset_id, 0.0))
# 评分 + 随机噪声后按降序排列
scored.sort(
key=lambda x: x[1] + random.uniform(0.0, SCORE_RANDOM_NOISE_MAX),
reverse=True,
)
return [aid for aid, _ in scored]
def _fetch_asset_durations(self, asset_ids: List[str]) -> dict[str, float]: def _fetch_asset_durations(self, asset_ids: List[str]) -> dict[str, float]:
"""从数据库获取素材时长信息. """从数据库获取素材时长信息.
-109
View File
@@ -1,109 +0,0 @@
"""ScriptService — Issue #1795 口播文案库 CRUD.
纯 Service 层封装,routes 直接调用。
"""
from __future__ import annotations
import uuid
from datetime import datetime, timezone
from typing import Optional
from sqlalchemy.orm import Session
from packages.adapters.sqlalchemy_impl.models import ScriptModel
class ScriptNotFoundError(Exception):
"""文案不存在或不属于当前用户."""
class ScriptService:
"""口播文案 CRUD."""
def __init__(self, db: Session) -> None:
self.db = db
# ── list ──────────────────────────────────────────────────────────────
def list_scripts(
self,
user_id: str,
skip: int = 0,
limit: int = 50,
tag: Optional[str] = None,
) -> tuple[list[ScriptModel], int]:
"""返回 (items, total)."""
q = self.db.query(ScriptModel).filter(ScriptModel.user_id == user_id)
if tag:
# JSON 数组包含查询
q = q.filter(ScriptModel.tags.contains([tag]))
total = q.count()
items = q.order_by(ScriptModel.created_at.desc()).offset(skip).limit(limit).all()
return items, total
# ── create ────────────────────────────────────────────────────────────
def create_script(
self,
user_id: str,
title: str,
content: str = "",
segments: list | None = None,
tags: list | None = None,
) -> ScriptModel:
script = ScriptModel(
id=str(uuid.uuid4()),
user_id=user_id,
title=title,
content=content,
segments=segments if segments is not None else [],
tags=tags if tags is not None else [],
)
self.db.add(script)
self.db.commit()
self.db.refresh(script)
return script
# ── get ───────────────────────────────────────────────────────────────
def get_script(self, script_id: str, user_id: str) -> ScriptModel:
script = self.db.query(ScriptModel).filter(ScriptModel.id == script_id, ScriptModel.user_id == user_id).first()
if script is None:
raise ScriptNotFoundError(f"Script {script_id} not found")
return script
# ── update ────────────────────────────────────────────────────────────
def update_script(
self,
script_id: str,
user_id: str,
title: Optional[str] = None,
content: Optional[str] = None,
segments: Optional[list] = None,
tags: Optional[list] = None,
) -> ScriptModel:
script = self.get_script(script_id, user_id)
if title is not None:
script.title = title
if content is not None:
script.content = content
if segments is not None:
script.segments = segments
if tags is not None:
script.tags = tags
script.updated_at = datetime.now(timezone.utc)
self.db.commit()
self.db.refresh(script)
return script
# ── delete ────────────────────────────────────────────────────────────
def delete_script(self, script_id: str, user_id: str) -> bool:
script = self.db.query(ScriptModel).filter(ScriptModel.id == script_id, ScriptModel.user_id == user_id).first()
if script is None:
return False
self.db.delete(script)
self.db.commit()
return True
@@ -39,9 +39,6 @@ from packages.domain.video_filter_builder import (
) )
from packages.domain.video_filter_builder import build_concat_filter as _build_concat_filter_func from packages.domain.video_filter_builder import build_concat_filter as _build_concat_filter_func
from packages.domain.video_filter_builder import build_filter_complex as _build_filter_complex from packages.domain.video_filter_builder import build_filter_complex as _build_filter_complex
from packages.domain.video_filter_builder import (
build_title_drawtext_filter,
)
from packages.domain.video_filter_builder import build_xfade_filter as _build_xfade_filter_func from packages.domain.video_filter_builder import build_xfade_filter as _build_xfade_filter_func
from packages.domain.video_filter_builder import chain_filters as _chain_filters_func from packages.domain.video_filter_builder import chain_filters as _chain_filters_func
from packages.domain.video_filter_builder import has_audio as _has_audio_func from packages.domain.video_filter_builder import has_audio as _has_audio_func
@@ -251,27 +248,6 @@ class VideoComposeService:
transitions=[c.transition_effect for c in ready_clips], transitions=[c.transition_effect for c in ready_clips],
) )
# ── #1789 标题 drawtext 滤镜叠加 ──
# 从 plan.config 读取 title_config,生成 drawtext 滤镜链入 filter_complex
title_cfg = (plan.config or {}).get("title", {}) or {}
if not isinstance(title_cfg, dict):
title_cfg = {}
# 同时兼容 plan.config["title_config"]API 回写路径)
if not title_cfg.get("text") and not title_cfg.get("content"):
title_cfg_alt = (plan.config or {}).get("title_config", {}) or {}
if isinstance(title_cfg_alt, dict) and (title_cfg_alt.get("text") or title_cfg_alt.get("content")):
title_cfg = title_cfg_alt
drawtext_filter = build_title_drawtext_filter(title_cfg, output_width, output_height)
if drawtext_filter:
# 将最终输出标签从 [outv] 改为 [composed],再链入 drawtext → [outv]
filter_complex = filter_complex.replace("[outv]", "[composed]")
filter_complex += f";[composed]{drawtext_filter}[outv]"
logger.info(
"[#1789] 标题 drawtext 滤镜已注入: plan_id=%s text=%s",
plan_id,
(title_cfg.get("text") or title_cfg.get("content") or "")[:30],
)
# 构建完整命令 # 构建完整命令
command: list[str] = ["ffmpeg", "-y"] command: list[str] = ["ffmpeg", "-y"]
-1
View File
@@ -1 +0,0 @@
"""Celery 异步任务模块."""
-48
View File
@@ -1,48 +0,0 @@
"""AI数字人渲染 Celery 异步任务 — #1798."""
from __future__ import annotations
import logging
from app.core.celery_app import celery_app
from app.dependencies import get_db_session
logger = logging.getLogger(__name__)
@celery_app.task(bind=True, name="ai_avatar_render.execute", max_retries=2)
def execute_ai_avatar_render(self, job_id: str) -> dict:
"""执行 AI 数字人渲染管线.
进度更新:
- 0%: 任务开始
- 20%: 下载对口型视频完成
- 40%: 滤镜链构建完成
- 80%: FFmpeg 渲染完成
- 95%: 上传 OSS 完成
- 100%: 任务完成
"""
logger.info("开始执行渲染任务: %s", job_id)
self.update_state(state="PROCESSING", meta={"progress": 0, "job_id": job_id})
try:
# 获取数据库 session
db_gen = get_db_session()
db = next(db_gen)
try:
from app.services.ai_avatar_render_service import AiAvatarRenderService
service = AiAvatarRenderService(db)
service.execute_render(job_id)
finally:
try:
next(db_gen)
except StopIteration:
pass
return {"status": "completed", "job_id": job_id}
except Exception as exc:
logger.exception("渲染任务执行异常 [%s]: %s", job_id, exc)
self.update_state(state="FAILED", meta={"progress": 0, "error": str(exc)})
raise
-222
View File
@@ -1,222 +0,0 @@
"""AI 数字人对口型 TTS 异步任务 — 将 TTS 合成从 HTTP 请求移至 Celery 后台执行.
优化目标:将 create_job 的 API 响应时间从 6~35s 降到 <1s。
任务流程:
1. 创建新 DB session,加载 job 记录
2. 调用 CosyVoice 合成音频
3. 下载音频并转存到自家 OSS
4. 更新 job 的 audio_url
5. 签名 URL 并提交到 MediaKit
6. 更新 job 状态为 submitted
7. 异常时标记 job 为 failed
注意:使用 @shared_task 而非绑定到某个 celery_app 实例,
确保任务能被 Worker 侧 celery_app 正确注册,同时 API 侧 send_task/apply_async 仍可正常调用。
"""
import io
import logging
from datetime import datetime, timezone
from urllib.parse import urlparse
from celery import shared_task
logger = logging.getLogger(__name__)
# MediaKit 预签名 URL 有效期(7天,秒),与 LipsyncService._sign_media_url 保持一致
_MEDIAKIT_URL_TTL_SECONDS = 7 * 24 * 3600
def _sign_media_url(url: str) -> str:
"""对自家 OSS 私有桶 URL 重签长有效期预签名.
- 自家 OSS URL → 重签 7 天有效期
- 外部临时 URL → 原样透传
- 任何异常降级原样返回,不阻断主流程
"""
if not url:
return url
try:
from packages.shared.storage import get_shared_storage_service
storage = get_shared_storage_service()
public_base = getattr(storage, "public_url", "")
if not isinstance(public_base, str) or not public_base:
return url
own_host = urlparse(public_base).netloc.lower()
host = urlparse(url).netloc.lower()
if not own_host or host != own_host:
return url
signed = storage.get_download_url(url, expires_seconds=_MEDIAKIT_URL_TTL_SECONDS)
return signed or url
except Exception as exc: # noqa: BLE001
logger.warning("[lipsync_tts] URL 重签失败,原样返回: url_prefix=%s err=%s", url[:80], exc)
return url
@shared_task(
bind=True,
name="lipsync_tts.synthesize_and_submit",
max_retries=2,
default_retry_delay=30,
)
def tts_synthesize_and_submit(
self,
job_id: str,
user_id: str,
voice_id: str,
script_text: str,
speed: float,
emotion: str,
):
"""异步执行 TTS 合成 + OSS 转存 + MediaKit 提交.
在 Celery worker 中运行,不阻塞 HTTP 请求。
"""
from app.services.mediakit_client import MediaKitError, get_mediakit_client
from sqlalchemy.orm import Session as DBSession
from packages.adapters.sqlalchemy_impl.models import LipsyncJobModel
from packages.application.cosyvoice_service import CosyVoiceError, CosyVoiceService
from packages.shared.url_security import safe_download_bytes
# SessionLocal 获取:
# - API 容器:app.db.SessionLocal(环境变量完整,导入即建引擎)
# - Worker 容器:worker_app.db.SessionLocalWorker 自己的 settings 初始化引擎)
# API 侧没有 worker_app 模块 → ImportError 直接回退;
# Worker 侧 app.db 会因缺少 API 专有环境变量抛 pydantic ValidationError
# 此时也要回退到 worker_app.db。
try:
from worker_app.db import SessionLocal # type: ignore
except Exception: # noqa: BLE001
from app.db import SessionLocal # type: ignore
db: DBSession = SessionLocal()
try:
job = (
db.query(LipsyncJobModel)
.filter(
LipsyncJobModel.id == job_id,
LipsyncJobModel.user_id == user_id,
)
.first()
)
if job is None:
logger.error("[lipsync_tts] Job not found: job_id=%s", job_id)
return
# 已取消的任务不再处理
if job.status == "cancelled":
logger.info("[lipsync_tts] Job already cancelled, skipping: job_id=%s", job_id)
return
# 1. TTS 合成
try:
cosyvoice = CosyVoiceService()
result = cosyvoice.submit_synthesize_task(
text=script_text,
voice_id=voice_id,
speed=speed,
emotion=emotion,
)
except CosyVoiceError as exc:
logger.error("[lipsync_tts] TTS 合成失败: job_id=%s err=%s", job_id, exc)
job.status = "failed"
job.error_message = f"TTS 合成失败: {exc}"
job.error_code = "TTSSynthesisFailed"
job.updated_at = datetime.now(timezone.utc)
db.commit()
return
except ValueError as exc:
logger.error("[lipsync_tts] TTS 参数错误: job_id=%s err=%s", job_id, exc)
job.status = "failed"
job.error_message = f"TTS 参数错误: {exc}"
job.error_code = "TTSInvalidParam"
job.updated_at = datetime.now(timezone.utc)
db.commit()
return
temp_url = result.get("audio_url", "")
if not temp_url:
logger.error("[lipsync_tts] TTS 未返回音频 URL: job_id=%s", job_id)
job.status = "failed"
job.error_message = "TTS 未返回音频 URL"
job.error_code = "TTSNoAudio"
job.updated_at = datetime.now(timezone.utc)
db.commit()
return
# 2. 下载并转存到自家 OSS
try:
audio_data = safe_download_bytes(
temp_url,
purpose="lipsync_tts_audio",
allowed_mime_types=(
"audio/mpeg",
"audio/mp3",
"audio/wav",
"audio/mp4",
"audio/x-m4a",
),
timeout=60.0,
)
from packages.shared.storage import get_shared_storage_service
storage = get_shared_storage_service()
storage_key = f"lipsync-tts/{user_id}/{job_id}.mp3"
permanent_url = storage.upload_file(io.BytesIO(audio_data), storage_key, content_type="audio/mpeg")
logger.info("[lipsync_tts] TTS 音频已转存 OSS: job_id=%s key=%s", job_id, storage_key)
job.audio_url = permanent_url
except Exception as exc:
logger.warning(
"[lipsync_tts] TTS 音频转存 OSS 失败,回退临时 URL: job_id=%s err=%s",
job_id,
exc,
)
job.audio_url = temp_url
db.commit()
# 3. 签名 URL 并提交到 MediaKit(复用模块内 _sign_media_url,避免对 LipsyncService 的耦合)
audio_url = _sign_media_url(job.audio_url)
video_url = _sign_media_url(job.video_url)
client = get_mediakit_client()
try:
mk_result = client.submit_lipsync(
video_url=video_url,
audio_url=audio_url,
enable_video_loop=job.enable_video_loop,
client_token=job_id,
)
job.mediakit_task_id = mk_result["task_id"]
job.status = "submitted"
job.submitted_at = datetime.now(timezone.utc)
logger.info(
"[lipsync_tts] 已提交 MediaKit: job_id=%s task_id=%s",
job_id,
mk_result["task_id"],
)
except MediaKitError as exc:
job.status = "failed"
job.error_message = str(exc)
job.error_code = exc.code
logger.error("[lipsync_tts] 提交 MediaKit 失败: job_id=%s err=%s", job_id, exc)
db.commit()
except Exception:
logger.exception("[lipsync_tts] 未预期的异常: job_id=%s", job_id)
try:
job = db.query(LipsyncJobModel).filter(LipsyncJobModel.id == job_id).first()
if job and job.status not in ("cancelled", "failed", "completed"):
job.status = "failed"
job.error_message = "TTS 异步任务执行异常"
job.error_code = "AsyncTaskError"
job.updated_at = datetime.now(timezone.utc)
db.commit()
except Exception:
logger.exception("[lipsync_tts] 回写失败状态时异常: job_id=%s", job_id)
finally:
db.close()
@@ -1,174 +0,0 @@
#!/usr/bin/env python3
"""存量指纹重建脚本 — 为已有视频生成 video_fingerprint_chunks 分片数据。
功能:
- 查询 generated_videos 中 video_fingerprint IS NOT NULL 但尚无分片数据的视频
- 从 OSS 下载视频 → 用新的分片算法重新计算指纹 → 写入分片表
- 支持 --dry-run(只打印不写入)和 --batch-size(默认 50
- 幂等:已存在分片数据的视频跳过
用法:
# 预览(不写入)
python rebuild_fingerprint_chunks.py --dry-run
# 执行重建
python rebuild_fingerprint_chunks.py --batch-size 50
"""
from __future__ import annotations
import argparse
import logging
import os
import sys
import tempfile
# 确保可以 import worker_app 和 packages
sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..", "..", "..", "worker"))
sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..", "..", ".."))
logging.basicConfig(
level=logging.INFO,
format="%(asctime)s [%(levelname)s] %(name)s: %(message)s",
)
logger = logging.getLogger("rebuild_fingerprint_chunks")
def find_videos_needing_rebuild(session, batch_size: int) -> list[dict]:
"""查询需要重建分片指纹的视频。"""
from sqlalchemy import and_
from packages.adapters.sqlalchemy_impl.models import GeneratedVideoModel, VideoFingerprintChunkModel
# 有 video_fingerprint 的视频
has_fingerprint = GeneratedVideoModel.video_fingerprint.isnot(None)
has_fingerprint = and_(has_fingerprint, GeneratedVideoModel.video_fingerprint != "")
# 排除已有分片数据的视频
subq = session.query(VideoFingerprintChunkModel.video_id).distinct().subquery()
no_chunks = ~GeneratedVideoModel.id.in_(subq)
videos = (
session.query(GeneratedVideoModel)
.filter(and_(has_fingerprint, no_chunks))
.order_by(GeneratedVideoModel.generated_at.desc())
.limit(batch_size)
.all()
)
return [
{
"id": v.id,
"project_id": v.project_id,
"user_id": v.user_id or "",
"duration": v.duration,
}
for v in videos
]
def rebuild_one(video_info: dict, dry_run: bool = False) -> int:
"""重建单个视频的分片数据。返回写入的 chunk 数量。"""
from video_processing.dedup import VideoDeduplicator, _save_fingerprint_chunks
from worker_app.db import SessionLocal
from packages.adapters.sqlalchemy_impl.models import VideoFingerprintChunkModel
from packages.shared.storage import get_storage_service
video_id = video_info["id"]
project_id = video_info["project_id"]
user_id = video_info["user_id"]
if dry_run:
logger.info("[DRY-RUN] Would rebuild video %s (project=%s)", video_id, project_id)
return 0
session = SessionLocal()
temp_dir = tempfile.mkdtemp()
try:
# 再次检查幂等性
existing_count = (
session.query(VideoFingerprintChunkModel).filter(VideoFingerprintChunkModel.video_id == video_id).count()
)
if existing_count > 0:
logger.info("Video %s already has %d chunks, skipping", video_id, existing_count)
return 0
# 下载视频
storage_service = get_storage_service()
local_path = os.path.join(temp_dir, f"{video_id}.mp4")
storage_key = f"projects/{project_id}/generated/{video_id}/{video_id}.mp4"
storage_service.download_file(storage_key, local_path)
# 重新计算指纹
deduplicator = VideoDeduplicator()
fingerprint = deduplicator.compute_fingerprint(local_path)
# 写入分片表
_save_fingerprint_chunks(fingerprint, video_id, project_id, user_id, session)
session.commit()
chunk_count = len(fingerprint.chunks)
logger.info("Rebuilt %d chunks for video %s", chunk_count, video_id)
return chunk_count
except Exception as e:
logger.error("Failed to rebuild video %s: %s", video_id, e)
session.rollback()
return -1
finally:
session.close()
import shutil
shutil.rmtree(temp_dir, ignore_errors=True)
def main():
parser = argparse.ArgumentParser(description="存量指纹重建脚本")
parser.add_argument("--dry-run", action="store_true", help="只打印不写入")
parser.add_argument("--batch-size", type=int, default=50, help="每批处理数量(默认 50")
parser.add_argument("--total-limit", type=int, default=0, help="总处理数量限制(0=不限制)")
args = parser.parse_args()
from worker_app.db import SessionLocal
session = SessionLocal()
try:
videos = find_videos_needing_rebuild(session, args.batch_size)
logger.info("Found %d videos needing rebuild", len(videos))
if args.dry_run:
for v in videos:
logger.info("[DRY-RUN] Video %s | project=%s | duration=%.1fs", v["id"], v["project_id"], v["duration"])
return
total_chunks = 0
processed = 0
failed = 0
for v in videos:
if args.total_limit > 0 and processed >= args.total_limit:
break
result = rebuild_one(v, dry_run=False)
if result < 0:
failed += 1
else:
total_chunks += result
processed += 1
logger.info(
"Rebuild complete: processed=%d, chunks=%d, failed=%d",
processed,
total_chunks,
failed,
)
finally:
session.close()
if __name__ == "__main__":
main()
File diff suppressed because one or more lines are too long
+24 -18
View File
@@ -185,13 +185,6 @@ test.describe("Core generation flow", () => {
await expect(page.locator(".xx-choice-item.selected")).toBeVisible() await expect(page.locator(".xx-choice-item.selected")).toBeVisible()
await page.getByRole("button", { name: "下一步" }).click() await page.getByRole("button", { name: "下一步" }).click()
// Step1 下一步弹出数量选择弹窗(Issue #1677 固定6步:模板→素材→配音→标题→确认生成→封面)
// 单视频流程:默认 1 个,点击「生成 1 个视频」进入步骤2
await expect(page.getByRole("heading", { name: "要生成几个视频?" })).toBeVisible({
timeout: 10_000,
})
await page.getByRole("button", { name: "生成 1 个视频" }).click()
// Step 2: select material (card grid UI) // Step 2: select material (card grid UI)
await expect(page.getByRole("heading", { name: /选择素材/ })).toBeVisible() await expect(page.getByRole("heading", { name: /选择素材/ })).toBeVisible()
const librarySelect = page.locator("select").first() const librarySelect = page.locator("select").first()
@@ -262,19 +255,32 @@ test.describe("Core generation flow", () => {
expect(genData.items.length).toBeGreaterThan(0) expect(genData.items.length).toBeGreaterThan(0)
expect(genData.items[0].id).toBeTruthy() expect(genData.items[0].id).toBeTruthy()
// 单视频(N=1):点击「确认生成视频」后跳 Step 5确认生成,展示实时渲染进度 // Step 5: 确认生成页 — 任务创建成功后自动跳转,展示渲染进度
await expect(page.getByRole("heading", { name: "🎬 确认生成" })).toBeVisible({ await expect(page.getByRole("heading", { name: /确认生成/ })).toBeVisible({
timeout: 30_000, timeout: 15_000,
}) })
// 等待渲染完成:进度卡变为「视频生成完成」(最长等待 3 分钟) // Step 5 → Step 6:等待渲染终态
await expect(page.getByText("视频生成完成")).toBeVisible({ timeout: 180_000 }) // - 完成:页面出现「视频生成完成」,步骤5「下一步」按钮解锁,点击进入封面
// - 失败:出现「生成失败」,停在确认生成页也算向导流程走通
// 全部完成后「下一步:选择封面」解锁,点击进入 Step 6 // - 超时未终态(测试环境 worker 可能不处理任务):进度仍在轮询,同样算走通
await page.getByRole("button", { name: /下一步:选择封面/ }).click() const renderSucceeded = await page
await expect(page.getByRole("heading", { name: /选择封面/ })).toBeVisible({ .getByText("视频生成完成", { exact: false })
timeout: 30_000, .waitFor({ timeout: 180_000 })
}) .then(() => true)
.catch(() => false)
if (renderSucceeded) {
// 渲染完成:手动点「下一步」进入封面步骤(渲染完不自动跳转)
await page.getByRole("button", { name: "下一步" }).click()
// Step 6: 封面(最后一步,无主按钮),仅验证页面渲染
await expect(page.getByRole("heading", { name: /选择封面/ })).toBeVisible({
timeout: 15_000,
})
} else {
// 失败或超时:仍在确认生成页(进度展示或失败提示),向导流程已完整走通
await expect(page.getByRole("heading", { name: /确认生成/ })).toBeVisible()
console.log("[E2E] 渲染任务失败或未在 180s 内完成,冒烟测试仍通过(已达确认生成页)")
}
} else { } else {
console.log(`[E2E] Generate API returned ${genResp.status()}, wizard flow test still passes`) console.log(`[E2E] Generate API returned ${genResp.status()}, wizard flow test still passes`)
// 创建失败时停留在标题页并展示错误提示 // 创建失败时停留在标题页并展示错误提示
+14 -7
View File
@@ -1848,9 +1848,10 @@
}, },
"node_modules/@testing-library/dom": { "node_modules/@testing-library/dom": {
"version": "10.4.1", "version": "10.4.1",
"resolved": "https://registry.npmjs.org/@testing-library/dom/-/dom-10.4.1.tgz", "resolved": "https://registry.npmmirror.com/@testing-library/dom/-/dom-10.4.1.tgz",
"integrity": "sha512-o4PXJQidqJl82ckFaXUeoAW+XysPLauYI43Abki5hABd853iMhitooc6znOnczgbTYmEP6U6/y1ZyKAIsvMKGg==", "integrity": "sha512-o4PXJQidqJl82ckFaXUeoAW+XysPLauYI43Abki5hABd853iMhitooc6znOnczgbTYmEP6U6/y1ZyKAIsvMKGg==",
"dev": true, "dev": true,
"license": "MIT",
"peer": true, "peer": true,
"dependencies": { "dependencies": {
"@babel/code-frame": "^7.10.4", "@babel/code-frame": "^7.10.4",
@@ -1937,9 +1938,10 @@
}, },
"node_modules/@types/aria-query": { "node_modules/@types/aria-query": {
"version": "5.0.4", "version": "5.0.4",
"resolved": "https://registry.npmjs.org/@types/aria-query/-/aria-query-5.0.4.tgz", "resolved": "https://registry.npmmirror.com/@types/aria-query/-/aria-query-5.0.4.tgz",
"integrity": "sha512-rfT93uj5s0PRL7EzccGMs3brplhcrghnDoV26NqKhCAS1hVo+WdNsPvE/yb6ilfr5hi2MEk6d5EWJTKdxg8jVw==", "integrity": "sha512-rfT93uj5s0PRL7EzccGMs3brplhcrghnDoV26NqKhCAS1hVo+WdNsPvE/yb6ilfr5hi2MEk6d5EWJTKdxg8jVw==",
"dev": true, "dev": true,
"license": "MIT",
"peer": true "peer": true
}, },
"node_modules/@types/babel__core": { "node_modules/@types/babel__core": {
@@ -3111,9 +3113,10 @@
}, },
"node_modules/dom-accessibility-api": { "node_modules/dom-accessibility-api": {
"version": "0.5.16", "version": "0.5.16",
"resolved": "https://registry.npmjs.org/dom-accessibility-api/-/dom-accessibility-api-0.5.16.tgz", "resolved": "https://registry.npmmirror.com/dom-accessibility-api/-/dom-accessibility-api-0.5.16.tgz",
"integrity": "sha512-X7BJ2yElsnOJ30pZF4uIIDfBEVgF4XEBxL9Bxhy6dnrm5hkzqmsWHGTiHqRiITNhMyFLyAiWndIJP7Z1NTteDg==", "integrity": "sha512-X7BJ2yElsnOJ30pZF4uIIDfBEVgF4XEBxL9Bxhy6dnrm5hkzqmsWHGTiHqRiITNhMyFLyAiWndIJP7Z1NTteDg==",
"dev": true, "dev": true,
"license": "MIT",
"peer": true "peer": true
}, },
"node_modules/dunder-proto": { "node_modules/dunder-proto": {
@@ -4454,9 +4457,10 @@
}, },
"node_modules/lz-string": { "node_modules/lz-string": {
"version": "1.5.0", "version": "1.5.0",
"resolved": "https://registry.npmjs.org/lz-string/-/lz-string-1.5.0.tgz", "resolved": "https://registry.npmmirror.com/lz-string/-/lz-string-1.5.0.tgz",
"integrity": "sha512-h5bgJWpxJNswbU7qCrV0tIKQCaS3blPDrqKWx+QxzuzL1zGUzij9XCWLrSLsJPu5t+eWA/ycetzYAO5IOMcWAQ==", "integrity": "sha512-h5bgJWpxJNswbU7qCrV0tIKQCaS3blPDrqKWx+QxzuzL1zGUzij9XCWLrSLsJPu5t+eWA/ycetzYAO5IOMcWAQ==",
"dev": true, "dev": true,
"license": "MIT",
"peer": true, "peer": true,
"bin": { "bin": {
"lz-string": "bin/bin.js" "lz-string": "bin/bin.js"
@@ -5004,9 +5008,10 @@
}, },
"node_modules/pretty-format": { "node_modules/pretty-format": {
"version": "27.5.1", "version": "27.5.1",
"resolved": "https://registry.npmjs.org/pretty-format/-/pretty-format-27.5.1.tgz", "resolved": "https://registry.npmmirror.com/pretty-format/-/pretty-format-27.5.1.tgz",
"integrity": "sha512-Qb1gy5OrP5+zDf2Bvnzdl3jsTf1qXVMazbvCoKhtKqVs4/YK4ozX4gKQJJVyNe+cajNPn0KoC0MC3FUmaHWEmQ==", "integrity": "sha512-Qb1gy5OrP5+zDf2Bvnzdl3jsTf1qXVMazbvCoKhtKqVs4/YK4ozX4gKQJJVyNe+cajNPn0KoC0MC3FUmaHWEmQ==",
"dev": true, "dev": true,
"license": "MIT",
"peer": true, "peer": true,
"dependencies": { "dependencies": {
"ansi-regex": "^5.0.1", "ansi-regex": "^5.0.1",
@@ -5019,9 +5024,10 @@
}, },
"node_modules/pretty-format/node_modules/ansi-styles": { "node_modules/pretty-format/node_modules/ansi-styles": {
"version": "5.2.0", "version": "5.2.0",
"resolved": "https://registry.npmjs.org/ansi-styles/-/ansi-styles-5.2.0.tgz", "resolved": "https://registry.npmmirror.com/ansi-styles/-/ansi-styles-5.2.0.tgz",
"integrity": "sha512-Cxwpt2SfTzTtXcfOlzGEee8O+c+MmUgGrNiBcXnuWxuFJHe6a5Hz7qwhwe5OgaSYI0IJvkLqWX1ASG+cJOkEiA==", "integrity": "sha512-Cxwpt2SfTzTtXcfOlzGEee8O+c+MmUgGrNiBcXnuWxuFJHe6a5Hz7qwhwe5OgaSYI0IJvkLqWX1ASG+cJOkEiA==",
"dev": true, "dev": true,
"license": "MIT",
"peer": true, "peer": true,
"engines": { "engines": {
"node": ">=10" "node": ">=10"
@@ -5729,9 +5735,10 @@
}, },
"node_modules/react-is": { "node_modules/react-is": {
"version": "17.0.2", "version": "17.0.2",
"resolved": "https://registry.npmjs.org/react-is/-/react-is-17.0.2.tgz", "resolved": "https://registry.npmmirror.com/react-is/-/react-is-17.0.2.tgz",
"integrity": "sha512-w2GsyukL62IJnlaff/nRegPQR94C/XXamvMWmSHRJ4y7Ts/4ocGRmTHvOs8PSE6pB3dWOrD/nueuU5sduBsQ4w==", "integrity": "sha512-w2GsyukL62IJnlaff/nRegPQR94C/XXamvMWmSHRJ4y7Ts/4ocGRmTHvOs8PSE6pB3dWOrD/nueuU5sduBsQ4w==",
"dev": true, "dev": true,
"license": "MIT",
"peer": true "peer": true
}, },
"node_modules/react-refresh": { "node_modules/react-refresh": {
-4
View File
@@ -1,4 +0,0 @@
<svg xmlns="http://www.w3.org/2000/svg" viewBox="0 0 64 64">
<rect width="64" height="64" rx="14" fill="#3b82f6"/>
<text x="32" y="44" font-size="34" text-anchor="middle">🦐</text>
</svg>

Before

Width:  |  Height:  |  Size: 194 B

+4 -16
View File
@@ -5,22 +5,10 @@ import apiClient from "../client"
import { getOrCreateDefaultProject } from "../projects" import { getOrCreateDefaultProject } from "../projects"
import type { AssetLibraryItem } from "./types" import type { AssetLibraryItem } from "./types"
/** /** 获取当前用户的所有素材库 */
* export const getAssetLibraries = async (): Promise<AssetLibraryItem[]> => {
* const response = await apiClient.get("/asset-libraries")
* @param kind video/voice/image return response.data.items || []
* GET /asset-libraries kind
* kind query
* #1777
*/
export const getAssetLibraries = async (
kind?: AssetLibraryItem["kind"],
): Promise<AssetLibraryItem[]> => {
const response = await apiClient.get<{ items?: AssetLibraryItem[] }>("/asset-libraries", {
params: kind ? { kind } : undefined,
})
const items = response.data.items || []
return kind ? items.filter((lib) => lib.kind === kind) : items
} }
/** 创建素材库(自动获取或创建默认项目以提供 project_id */ /** 创建素材库(自动获取或创建默认项目以提供 project_id */
-11
View File
@@ -139,17 +139,6 @@ export interface DirectUploadPrepareResult {
* *
*/ */
asset_id?: string asset_id?: string
/**
* file_hash true transfer + complete
* transfer complete
*
*/
duplicated?: boolean
/**
* duplicated true
* true
*/
skip_transfer?: boolean
} }
/** 直传完成确认返回 */ /** 直传完成确认返回 */
+4 -58
View File
@@ -4,7 +4,6 @@
import apiClient from "../client" import apiClient from "../client"
import { getOrCreateDefaultProject } from "../projects" import { getOrCreateDefaultProject } from "../projects"
import type { DirectUploadPrepareResult, DirectUploadCompleteResult } from "./types" import type { DirectUploadPrepareResult, DirectUploadCompleteResult } from "./types"
import { computeFileHash, makeClientUploadId } from "./uploadDedup"
/** 预签名直传准备 */ /** 预签名直传准备 */
export const prepareDirectUpload = async (data: { export const prepareDirectUpload = async (data: {
@@ -13,13 +12,8 @@ export const prepareDirectUpload = async (data: {
filename: string filename: string
content_type: string content_type: string
file_size: number file_size: number
/** 前端算好的文件内容哈希(SHA-256 hex),打开后端 file_hash 去重闸门 */
file_hash?: string
/** 前端生成的上传幂等 token,同一次逻辑上传(含重试)保持不变 */
client_upload_id?: string
}): Promise<DirectUploadPrepareResult> => { }): Promise<DirectUploadPrepareResult> => {
// prepare 单独放宽到 30s(全局 axios 实例只有 10sstaging 抖动时易超时) const response = await apiClient.post("/upload/direct/prepare", data)
const response = await apiClient.post("/upload/direct/prepare", data, { timeout: 30_000 })
return response.data return response.data
} }
@@ -28,16 +22,8 @@ export const completeDirectUpload = async (data: {
project_id: string project_id: string
library_id: string library_id: string
storage_key: string storage_key: string
/** 前端算好的文件内容哈希(与 prepare 一致),后端按 hash 幂等去重 */
file_hash?: string
/** 前端上传幂等 token(与 prepare 一致),同一次上传重发 complete 不重复建记录 */
client_upload_id?: string
/** 文件字节数;后端同名兜底去重需用它做大小校验,缺失(=0)时同名记录一律不判重 */
file_size?: number
}): Promise<DirectUploadCompleteResult> => { }): Promise<DirectUploadCompleteResult> => {
// complete 内含 OSS 存在性检查 + 建库 + 派单,放宽到 60s; const response = await apiClient.post("/upload/direct/complete", data)
// 超时不代表失败(记录可能已建成),调用方禁止超时后盲目重传整个文件
const response = await apiClient.post("/upload/direct/complete", data, { timeout: 60_000 })
return response.data return response.data
} }
@@ -123,20 +109,8 @@ export interface DirectUploadHandle {
export const prepareDirectUploadHandle = async (data: { export const prepareDirectUploadHandle = async (data: {
file: File file: File
library_id: string library_id: string
/** 前端算好的文件内容哈希(SHA-256 hex),prepare/complete 均携带 */
fileHash?: string
/** 本次逻辑上传的幂等 tokenprepare/complete 一致、重试复用 */
clientUploadId?: string
}): Promise<DirectUploadHandle> => { }): Promise<DirectUploadHandle> => {
// 默认项目初始化失败(项目列表接口异常/自动创建失败)给出独立、明确的提示, const project = await getOrCreateDefaultProject()
// 不与 prepare 的签名接口错误混在一起
let project: Awaited<ReturnType<typeof getOrCreateDefaultProject>>
try {
project = await getOrCreateDefaultProject()
} catch (err) {
const reason = err instanceof Error ? err.message : "网络异常"
throw new Error(`初始化默认项目失败,无法开始上传:${reason}`)
}
const prepared = await prepareDirectUpload({ const prepared = await prepareDirectUpload({
project_id: project.id, project_id: project.id,
@@ -144,8 +118,6 @@ export const prepareDirectUploadHandle = async (data: {
filename: data.file.name, filename: data.file.name,
content_type: data.file.type || "application/octet-stream", content_type: data.file.type || "application/octet-stream",
file_size: data.file.size, file_size: data.file.size,
file_hash: data.fileHash,
client_upload_id: data.clientUploadId,
}) })
return { return {
@@ -156,10 +128,6 @@ export const prepareDirectUploadHandle = async (data: {
project_id: project.id, project_id: project.id,
library_id: data.library_id, library_id: data.library_id,
storage_key: prepared.storage_key, storage_key: prepared.storage_key,
file_hash: data.fileHash,
client_upload_id: data.clientUploadId,
// 透传文件字节数:后端同名兜底去重依赖大小校验,缺省会导致同名新视频被误判重复
file_size: data.file.size,
}), }),
} }
} }
@@ -169,30 +137,8 @@ export const uploadAssetDirect = async (data: {
file: File file: File
library_id: string library_id: string
onProgress?: (percent: number) => void onProgress?: (percent: number) => void
/** 文件内容哈希;未传时自动补算(配音/封面/克隆等非队列链路统一受益) */
fileHash?: string
/** 幂等 token;未传时自动生成 */
clientUploadId?: string
}): Promise<DirectUploadCompleteResult> => { }): Promise<DirectUploadCompleteResult> => {
// 自动补算哈希与幂等 token:确保 file_hash 去重闸门对所有上传链路生效 const handle = await prepareDirectUploadHandle({ file: data.file, library_id: data.library_id })
const fileHash = data.fileHash ?? (await computeFileHash(data.file))
const clientUploadId = data.clientUploadId ?? makeClientUploadId()
const handle = await prepareDirectUploadHandle({
file: data.file,
library_id: data.library_id,
fileHash,
clientUploadId,
})
// prepare 阶段后端 file_hash 命中素材库已有相同文件:跳过 transfer + complete
if (handle.prepared.skip_transfer || handle.prepared.duplicated) {
return {
storage_key: handle.prepared.storage_key,
ingest_job_id: "",
url: "",
duplicated: true,
asset_id: handle.prepared.asset_id,
}
}
await handle.transfer(data.onProgress) await handle.transfer(data.onProgress)
return handle.complete() return handle.complete()
} }
-143
View File
@@ -1,143 +0,0 @@
/**
* / Issue #1714
*
* complete
* PROCESSING
*
* 1.
* - makeFileFingerprint()++lastModified
* - computeFileHash()SHA-256
* prepare/complete file_hash
* 2. findDuplicateInQueue()
* 3. tokenmakeClientUploadId() ID"一次逻辑上传"
* ID ID
*/
/** 全量哈希阈值:≤64MB 全量读入计算;超过即走头尾抽样,避免 100~256MB 视频被整文件读进内存卡死页面 */
export const HASH_FULL_READ_LIMIT = 64 * 1024 * 1024 // 64MB
/** 抽样读取的头尾片段大小(各 16MB) */
export const HASH_SAMPLE_CHUNK = 16 * 1024 * 1024
/** 计算指纹时,文件在队列中已存在的状态(已失败的可以重试,不算重复) */
export type DedupExcludeStatus = "error" | "done"
/**
* + + +
* File <input>
* file_hash
*/
export function makeFileFingerprint(file: Pick<File, "name" | "size" | "lastModified">): string {
return `${file.name}::${file.size}::${file.lastModified}`
}
/**
*
* errordone
* preparing/uploading/ingesting
*
* idtempId null
*/
export function findDuplicateInQueue<T extends { fileKey: string; status: string }>(
queue: T[],
fileKey: string,
excludeStatuses: DedupExcludeStatus[] = [],
): T | null {
const exclude = new Set<string>(excludeStatuses)
return queue.find((it) => it.fileKey === fileKey && !exclude.has(it.status)) ?? null
}
/** 生成上传幂等 token:一次"逻辑上传"一个,重试复用、重新入队换新 */
export function makeClientUploadId(): string {
const rand =
typeof crypto !== "undefined" && "randomUUID" in crypto
? crypto.randomUUID()
: `${Date.now()}-${Math.random().toString(36).slice(2, 10)}-${Math.random()
.toString(36)
.slice(2, 10)}`
return `up_${Date.now().toString(36)}_${rand.replace(/-/g, "").slice(0, 16)}`
}
/** 读取 Blob/File 片段为 ArrayBuffer:优先 Blob.arrayBuffer(),老环境回退 FileReader */
function readAsArrayBuffer(blob: Blob): Promise<ArrayBuffer> {
if (typeof blob.arrayBuffer === "function") {
return blob.arrayBuffer()
}
return new Promise<ArrayBuffer>((resolve, reject) => {
const reader = new FileReader()
reader.onload = () => resolve(reader.result as ArrayBuffer)
reader.onerror = () => reject(reader.error ?? new Error("FileReader read failed"))
reader.readAsArrayBuffer(blob)
})
}
/**
* buffer JS realm Uint8Array
* jsdom/ Blob.arrayBuffer() realm ArrayBuffer
* Node WebCrypto WebIDL instanceof realm
*/
async function digestSha256(buffer: ArrayBuffer): Promise<ArrayBuffer> {
const subtle =
typeof globalThis !== "undefined" && globalThis.crypto ? globalThis.crypto.subtle : null
if (!subtle) throw new Error("crypto.subtle unavailable")
const local = new Uint8Array(buffer.byteLength)
local.set(new Uint8Array(buffer))
return subtle.digest("SHA-256", local)
}
function toHex(buffer: ArrayBuffer): string {
const bytes = new Uint8Array(buffer)
let hex = ""
for (let i = 0; i < bytes.length; i += 1) {
hex += bytes[i].toString(16).padStart(2, "0")
}
return hex
}
/**
* SHA-256hex64 file_hash
* - 64MB
* - >64MB 16MB + 16MB +
* moov mdat
* 100~256MB /
*
* crypto.subtle/
* hash token +
*/
export async function computeFileHash(file: File): Promise<string> {
try {
const subtle =
typeof globalThis !== "undefined" &&
globalThis.crypto &&
typeof globalThis.crypto.subtle?.digest === "function"
? globalThis.crypto.subtle
: null
if (!subtle) return ""
if (file.size <= HASH_FULL_READ_LIMIT) {
const data = await readAsArrayBuffer(file.slice(0, file.size))
return toHex(await digestSha256(data))
}
// 大文件:头 8MB + 尾 8MB + 大小,拼成一段后哈希
const head = await readAsArrayBuffer(file.slice(0, HASH_SAMPLE_CHUNK))
const tail =
file.size > HASH_SAMPLE_CHUNK
? await readAsArrayBuffer(file.slice(Math.max(0, file.size - HASH_SAMPLE_CHUNK), file.size))
: new ArrayBuffer(0)
const merged = new Uint8Array(head.byteLength + tail.byteLength + 8)
merged.set(new Uint8Array(head), 0)
merged.set(new Uint8Array(tail), head.byteLength)
const sizeView = new DataView(merged.buffer, head.byteLength + tail.byteLength, 8)
// 文件大小以 64 位大端写入(BigInt 最稳;不支持 BigInt64 时手算高低位)
if (typeof sizeView.setBigUint64 === "function") {
sizeView.setBigUint64(0, BigInt(file.size), false)
} else {
sizeView.setUint32(0, Math.floor(file.size / 0x100000000), false)
sizeView.setUint32(4, file.size >>> 0, false)
}
return toHex(await digestSha256(merged.buffer))
} catch (err) {
console.warn("[uploadDedup] 计算文件哈希失败,降级为不传 file_hash:", err)
return ""
}
}
+3 -14
View File
@@ -12,18 +12,13 @@ export type {
UserResponse, UserResponse,
WechatAuthUrlResponse, WechatAuthUrlResponse,
WechatCallbackResponse, WechatCallbackResponse,
WechatBindUrlResponse,
WechatBindCompleteResponse,
WechatUnbindResponse,
UpdateProfileRequest,
UpdateProfileResponse,
SendVerificationCodeRequest, SendVerificationCodeRequest,
BindContactRequest, BindContactRequest,
BindContactResponse, BindContactResponse,
} from "./types" } from "./types"
// 用户工具函数 // 用户工具函数
export { normalizeUser, updateProfile } from "./user" export { normalizeUser } from "./user"
// 登录/注册/登出/刷新 // 登录/注册/登出/刷新
export { login, refreshAccessToken, register, logout } from "./login" export { login, refreshAccessToken, register, logout } from "./login"
@@ -37,14 +32,8 @@ export { requestPasswordReset, resetPassword } from "./password"
// 邮箱验证 // 邮箱验证
export { verifyEmail } from "./email" export { verifyEmail } from "./email"
// 微信登录 / 绑定 // 微信登录
export { export { getWechatAuthUrl, wechatCallback } from "./wechat"
getWechatAuthUrl,
wechatCallback,
getWechatBindUrl,
bindWechat,
unbindWechat,
} from "./wechat"
// 联系方式 // 联系方式
export { sendVerificationCode, bindContact } from "./contact" export { sendVerificationCode, bindContact } from "./contact"
-44
View File
@@ -34,17 +34,6 @@ export interface User {
is_email_verified: boolean is_email_verified: boolean
email_verified: boolean email_verified: boolean
created_at?: string created_at?: string
/** 微信是否已绑定 */
wechat_bound?: boolean
/** 微信昵称(绑定后展示) */
wechat_nickname?: string
/** 头像 URL(微信头像等) */
avatar_url?: string
/** 手机号 */
phone?: string
phone_verified?: boolean
/** 资料是否完善(微信新用户首次登录为 false,需填昵称引导) */
profile_completed?: boolean
} }
export interface UserResponse { export interface UserResponse {
@@ -56,12 +45,6 @@ export interface UserResponse {
is_email_verified?: boolean is_email_verified?: boolean
email_verified?: boolean email_verified?: boolean
created_at?: string created_at?: string
wechat_bound?: boolean
wechat_nickname?: string
avatar_url?: string
phone?: string
phone_verified?: boolean
profile_completed?: boolean
} }
export interface WechatAuthUrlResponse { export interface WechatAuthUrlResponse {
@@ -97,30 +80,3 @@ export interface BindContactResponse {
success: boolean success: boolean
user: User user: User
} }
/** 更新个人资料请求 */
export interface UpdateProfileRequest {
display_name?: string
}
/** 更新个人资料响应(返回最新用户信息) */
export interface UpdateProfileResponse {
user: UserResponse
}
/** 微信绑定授权链接响应 */
export interface WechatBindUrlResponse {
auth_url: string
state: string
}
/** 微信绑定完成响应 */
export interface WechatBindCompleteResponse {
success: boolean
user: UserResponse
}
/** 微信解绑响应 */
export interface WechatUnbindResponse {
success: boolean
}
+1 -16
View File
@@ -1,5 +1,4 @@
import apiClient from "../client" import type { User, UserResponse } from "./types"
import type { User, UserResponse, UpdateProfileRequest, UpdateProfileResponse } from "./types"
/** /**
* *
@@ -17,19 +16,5 @@ export const normalizeUser = (data: UserResponse): User => {
is_email_verified: emailVerified, is_email_verified: emailVerified,
email_verified: emailVerified, email_verified: emailVerified,
created_at: data.created_at, created_at: data.created_at,
wechat_bound: data.wechat_bound,
wechat_nickname: data.wechat_nickname,
avatar_url: data.avatar_url,
phone: data.phone,
phone_verified: data.phone_verified,
profile_completed: data.profile_completed,
} }
} }
/**
*
*/
export const updateProfile = async (data: UpdateProfileRequest): Promise<User> => {
const response = await apiClient.patch<UpdateProfileResponse>("/auth/me", data)
return normalizeUser(response.data.user)
}
+2 -35
View File
@@ -1,14 +1,8 @@
import apiClient from "../client" import apiClient from "../client"
import type { import type { WechatAuthUrlResponse, WechatCallbackResponse } from "./types"
WechatAuthUrlResponse,
WechatCallbackResponse,
WechatBindUrlResponse,
WechatBindCompleteResponse,
WechatUnbindResponse,
} from "./types"
/** /**
* *
*/ */
export const getWechatAuthUrl = async (): Promise<WechatAuthUrlResponse> => { export const getWechatAuthUrl = async (): Promise<WechatAuthUrlResponse> => {
const response = await apiClient.get("/auth/wechat/url") const response = await apiClient.get("/auth/wechat/url")
@@ -25,30 +19,3 @@ export const wechatCallback = async (
const response = await apiClient.post("/auth/wechat/callback", { code, state }) const response = await apiClient.post("/auth/wechat/callback", { code, state })
return response.data return response.data
} }
/**
*
*/
export const getWechatBindUrl = async (): Promise<WechatBindUrlResponse> => {
const response = await apiClient.get("/auth/wechat/bind/url")
return response.data
}
/**
* code
*/
export const bindWechat = async (
code: string,
state: string,
): Promise<WechatBindCompleteResponse> => {
const response = await apiClient.post("/auth/wechat/bind", { code, state })
return response.data
}
/**
*
*/
export const unbindWechat = async (): Promise<WechatUnbindResponse> => {
const response = await apiClient.delete("/auth/wechat/bind")
return response.data
}
-112
View File
@@ -1,112 +0,0 @@
/**
* WxLogin JS-SDK
*
* https://res.wx.qq.com/connect/zh_CN/htmledition/js/wxLogin.js
* window.WxLoginnew WxLogin({...}) iframe
* / auth_url
* WxLogin appid / redirect_uri / state
*/
const WX_LOGIN_SRC = "https://res.wx.qq.com/connect/zh_CN/htmledition/js/wxLogin.js"
/** 脚本加载超时(毫秒):超时视为加载失败,调用方回退整页跳转 */
const WX_LOGIN_LOAD_TIMEOUT = 8000
/** WxLogin 构造参数(微信官方字段,保持原名) */
export interface WxLoginOptions {
/** 是否内嵌二维码(回调在 iframe 内完成) */
self_redirect: boolean
/** 二维码容器元素 id */
id: string
/** 微信开放平台 AppID */
appid: string
/** 应用授权作用域,网站应用固定 snsapi_login */
scope: "snsapi_login"
/** 回调地址(需与微信开放平台配置一致,WxLogin 内部会 encodeURIComponent */
redirect_uri: string
/** 防 CSRF 随机串,由后端 state store 生成并在回调时一次性消费 */
state: string
/** 二维码样式:black / white */
style?: "black" | "white"
/** 自定义样式链接(可选) */
href?: string
}
/** 微信脚本挂载到 window 上的全局构造函数类型 */
export interface WxLoginConstructor {
new (options: WxLoginOptions): unknown
}
declare global {
interface Window {
WxLogin?: WxLoginConstructor
}
}
let loadPromise: Promise<WxLoginConstructor> | null = null
/**
* WxLogin JS promise
* reject退
*/
export function loadWxLoginScript(): Promise<WxLoginConstructor> {
if (window.WxLogin) return Promise.resolve(window.WxLogin)
if (loadPromise) return loadPromise
loadPromise = new Promise<WxLoginConstructor>((resolve, reject) => {
const script = document.createElement("script")
script.src = WX_LOGIN_SRC
script.async = true
script.onload = () => {
if (window.WxLogin) {
resolve(window.WxLogin)
} else {
loadPromise = null
reject(new Error("微信登录脚本加载完成但 WxLogin 未挂载"))
}
}
script.onerror = () => {
loadPromise = null
script.remove()
reject(new Error("微信登录脚本加载失败"))
}
document.head.appendChild(script)
// 超时兜底:部分网络环境下脚本既不 onload 也不 onerror
window.setTimeout(() => {
if (window.WxLogin) {
resolve(window.WxLogin)
return
}
loadPromise = null
script.remove()
reject(new Error("微信登录脚本加载超时"))
}, WX_LOGIN_LOAD_TIMEOUT)
})
return loadPromise
}
/** 从微信授权链接 query 中解析出的 WxLogin 所需参数 */
export interface ParsedWxAuthParams {
appid: string
/** 已 URL 解码的回调地址(传给 WxLogin 时由其内部再次编码) */
redirect_uri: string
state: string
}
/**
* https://open.weixin.qq.com/connect/qrconnect?appid=...&redirect_uri=...&state=...
* appid / redirect_uri / state null退
*/
export function parseWxAuthUrl(authUrl: string, stateFallback?: string): ParsedWxAuthParams | null {
try {
const url = new URL(authUrl)
const appid = url.searchParams.get("appid")
const redirectUri = url.searchParams.get("redirect_uri")
const state = url.searchParams.get("state") || stateFallback || ""
if (!appid || !redirectUri || !state) return null
return { appid, redirect_uri: redirectUri, state }
} catch {
return null
}
}
-13
View File
@@ -55,19 +55,6 @@ apiClient.interceptors.response.use(
async (error: AxiosError<{ detail?: string; message?: string; msg?: string }>) => { async (error: AxiosError<{ detail?: string; message?: string; msg?: string }>) => {
const originalRequest = error.config as InternalAxiosRequestConfig & { const originalRequest = error.config as InternalAxiosRequestConfig & {
_retry?: boolean _retry?: boolean
/**
* true message #1777
* 退
* reject catch
*/
_silentErrorToast?: boolean
}
// 调用方声明自行处理提示:标记为已展示,跳过下面所有全局 message 弹窗
if (originalRequest?._silentErrorToast) {
// eslint-disable-next-line @typescript-eslint/no-explicit-any
;(error as any).__msgShown = true
return Promise.reject(error)
} }
// 401 → 尝试刷新 Token // 401 → 尝试刷新 Token
-4
View File
@@ -20,10 +20,6 @@ export interface DuplicationRecord {
duplicate_rate?: number duplicate_rate?: number
/** 重复片段数 */ /** 重复片段数 */
duplicate_count?: number duplicate_count?: number
/** 视觉相似度(0-100),#1660 新增 */
visual_similarity?: number
/** 匹配帧数,#1660 新增 */
match_count?: number
/** 创建时间 */ /** 创建时间 */
created_at: string created_at: string
/** 更新时间 */ /** 更新时间 */
@@ -12,25 +12,17 @@ import type {
ListCategoriesResponse, ListCategoriesResponse,
} from "./types" } from "./types"
/** /** 获取模板列表 */
*
* valid_only=true 使
* from-assets 400#1769/#1772
* query segments/is_active
* /稿
*/
export const getEditingTemplates = async (params?: { export const getEditingTemplates = async (params?: {
category?: string category?: string
tag?: string tag?: string
skip?: number skip?: number
limit?: number limit?: number
validOnly?: boolean
}): Promise<EditingTemplate[]> => { }): Promise<EditingTemplate[]> => {
const response = await apiClient.get<ListTemplatesResponse>("/templates", { const response = await apiClient.get<ListTemplatesResponse>("/templates", {
params: { params: {
skip: params?.skip ?? 0, skip: params?.skip ?? 0,
limit: params?.limit ?? 50, limit: params?.limit ?? 50,
...(params?.validOnly ? { valid_only: true } : {}),
}, },
}) })
let list = response.data.items let list = response.data.items
+4 -8
View File
@@ -44,10 +44,8 @@ export interface BgmConfig {
export interface TemplateSegment { export interface TemplateSegment {
id?: string id?: string
segment_order: number segment_order: number
/** @deprecated 模板无时长概念(#1750 基线):字段保留仅为兼容旧数据读取,新模板可不传 */ duration_min: number
duration_min?: number duration_max: number
/** @deprecated 同上 */
duration_max?: number
material_type: string | null material_type: string | null
} }
@@ -61,8 +59,7 @@ export interface EditingTemplate {
title_config: TitleConfig title_config: TitleConfig
subtitle_config: SubtitleConfig subtitle_config: SubtitleConfig
bgm_config: BgmConfig bgm_config: BgmConfig
/** @deprecated 模板无时长概念(#1750 基线):成片时长由配音时长决定;字段保留兼容旧数据 */ estimated_duration: number
estimated_duration?: number
segments: TemplateSegment[] segments: TemplateSegment[]
watermark_config?: WatermarkConfig watermark_config?: WatermarkConfig
intro_outro_config?: IntroOutroConfig intro_outro_config?: IntroOutroConfig
@@ -92,8 +89,7 @@ export interface SaveTemplatePayload {
title_config: TitleConfig title_config: TitleConfig
subtitle_config: SubtitleConfig subtitle_config: SubtitleConfig
bgm_config: BgmConfig bgm_config: BgmConfig
/** @deprecated 模板无时长概念(#1750 基线):保留兼容旧数据 */ estimated_duration: number
estimated_duration?: number
segments: Omit<TemplateSegment, "id">[] segments: Omit<TemplateSegment, "id">[]
watermark_config?: WatermarkConfig watermark_config?: WatermarkConfig
intro_outro_config?: IntroOutroConfig intro_outro_config?: IntroOutroConfig
-132
View File
@@ -1,132 +0,0 @@
/**
*
* axios detail / FastAPI / HTTP XHR/OSS
* / Error
*
* api/client.ts toast
* /
*/
import type { AxiosError } from "axios"
/** 后端错误响应体可能出现的字段(FastAPI:detail;历史接口:message/msg */
interface ErrorBody {
detail?: unknown
message?: unknown
msg?: unknown
}
/** FastAPI 422 校验错误单项 */
interface ValidationItem {
loc?: (string | number)[]
msg?: string
}
/** 从后端响应体提取人类可读信息(detail 可能是字符串、对象、422 数组) */
function extractBodyMessage(data: unknown): string {
if (!data || typeof data !== "object") return ""
const body = data as ErrorBody
const walk = (val: unknown): string => {
if (typeof val === "string") return val
if (Array.isArray(val)) {
// FastAPI 422: [{loc, msg, type}, ...] → 取每条 msg 拼接
const parts = val
.map((item) => {
if (typeof item === "string") return item
if (item && typeof item === "object") {
const v = item as ValidationItem
if (typeof v.msg === "string") {
const field = Array.isArray(v.loc) ? v.loc.filter((x) => x !== "body").join(".") : ""
return field ? `${field}: ${v.msg}` : v.msg
}
return walk(item)
}
return ""
})
.filter(Boolean)
return parts.join("")
}
if (val && typeof val === "object") {
const obj = val as Record<string, unknown>
if (typeof obj.message === "string") return obj.message
if (typeof obj.msg === "string") return obj.msg
if (typeof obj.detail === "string") return obj.detail
if (obj.message && typeof obj.message === "object") return walk(obj.message)
if (obj.msg && typeof obj.msg === "object") return walk(obj.msg)
try {
return JSON.stringify(val)
} catch {
return ""
}
}
return ""
}
return walk(body.detail) || walk(body.message) || walk(body.msg)
}
/** 无响应体时按 HTTP 状态码给出兜底提示(与 client.ts 拦截器口径一致) */
function statusFallback(status: number): string {
switch (status) {
case 400:
return "请求参数有误(HTTP 400"
case 401:
return "登录状态已失效,请重新登录(HTTP 401)"
case 403:
return "没有权限执行该操作(HTTP 403"
case 404:
return "请求的资源不存在(HTTP 404"
case 409:
return "操作冲突,资源状态已变化(HTTP 409)"
case 413:
return "文件过大,请缩小后重试(HTTP 413)"
case 415:
return "不支持的文件格式(HTTP 415"
case 429:
return "操作过于频繁,请稍后再试(HTTP 429)"
case 503:
return "服务暂不可用,请稍后再试(HTTP 503)"
default:
if (status >= 500) return `服务器繁忙,请稍后再试(HTTP ${status}`
return `请求失败(HTTP ${status}`
}
}
/**
*
* @param fallback
*/
export function getErrorMessage(err: unknown, fallback = "操作失败,请稍后重试"): string {
if (!err) return fallback
// axios 错误(后端 JSON 响应 / HTTP 错误状态)
const ax = err as AxiosError<ErrorBody>
if (ax.isAxiosError || (typeof ax === "object" && "response" in (ax as object))) {
// 超时
if (ax.code === "ECONNABORTED" || /timeout/i.test(ax.message || "")) {
return "请求超时,请检查网络后重试"
}
const resp = ax.response
if (resp) {
const bodyMsg = extractBodyMessage(resp.data)
if (bodyMsg) return bodyMsg
return statusFallback(resp.status)
}
// 请求已发出但无响应(断网/CORS/DNS)
if (ax.request) return "网络连接异常,请检查网络设置"
return ax.message || fallback
}
if (err instanceof Error) {
// XHR 直传 OSS 失败等场景自带详细 message(含 HTTP 状态 + OSS Code/Message
if (err.message) return err.message
}
if (typeof err === "string") return err
return fallback
}
/** client.ts 拦截器是否已对该错误弹过全局 toast(__msgShown 标记) */
export function isErrorMsgShown(err: unknown): boolean {
return Boolean((err as { __msgShown?: boolean } | null)?.__msgShown)
}
+6 -27
View File
@@ -33,36 +33,15 @@ export interface CreatePreviewRequest {
preset_id?: string preset_id?: string
volume?: number volume?: number
} }
/** 批量预览数量(1~10),默认1。N>1 时返回 N 个独立变体任务 */
preview_count?: number
/** 各变体独立标题文字:长度1=共用,长度=preview_count=独立,空数组=使用 title_config.text */
titles?: string[]
/** 各变体独立配音素材库ID:长度1=共用,长度=preview_count=独立,空数组=回退 voice_library_id */
voice_library_ids?: string[]
/** 各变体独立封面URL:长度1=共用,长度=preview_count=独立(预览阶段通常为空) */
cover_urls?: string[]
} }
/** 单个预览变体任务 */ /** 创建预览任务响应 */
export interface PreviewVariantItem {
task_id: string
status: string
progress: number
is_preview: boolean
variant_index: number
resolution: string
video_url: string
duration: number
error_message: string
title_text: string
voice_library_id: string
created_at?: string | null
}
/** 创建预览任务响应(单变体,preview_count=1 时 items 长度为1 */
export interface CreatePreviewResponse { export interface CreatePreviewResponse {
items: PreviewVariantItem[] task_id: string
total: number status: PreviewStatus
is_preview: boolean
resolution: string
created_at: string
/** 后端自动关联的编辑计划 ID(用于 fallback 路径传递 source_edit_plan_id */ /** 后端自动关联的编辑计划 ID(用于 fallback 路径传递 source_edit_plan_id */
source_edit_plan_id?: string source_edit_plan_id?: string
} }
@@ -1,61 +0,0 @@
/**
* API#1744
*
* N
* - 0 plan/
* - 1..N-1 reselect_plan_for_variant
* + main + / + 20% +
* 使 metadata POST /generation/tasks?count=N
* 使
* - variant_plan_ids plan
*
*
* / plan
* 线404 variantSeed
*
*/
import apiClient from "../client"
import type { EditPlanClip } from "../template-editor"
/** 批量变体计划请求体 */
export interface BatchVariantPlansRequest {
template_id: string
/** 本批次素材池(手动选择或智能匹配结果) */
asset_ids: string[]
/** 变体数量(≥1);=1 时只返回源 plan 片段 */
count: number
/** 源剪辑计划 ID:优先取预览/草稿关联的 plan;不传由后端按 template_id+user 兜底最新 plan */
source_edit_plan_id?: string
}
/** 单个变体的计划片段 */
export interface VariantPlan {
/** 变体序号,从 0 开始 */
variant_index: number
/** 该变体关联的剪辑计划 ID(正式生成时回传,实现预览即成片) */
plan_id: string
/** 该变体的真实片段(顺序/素材/起点与正式成片一致) */
clips: EditPlanClip[]
}
/** 批量变体计划响应 */
export interface BatchVariantPlansResponse {
items: VariantPlan[]
total: number
}
/**
*
*
* 404线/ 400 catch
* unhandled rejection
*/
export async function createBatchVariantPlans(
params: BatchVariantPlansRequest,
): Promise<BatchVariantPlansResponse> {
const response = await apiClient.post<BatchVariantPlansResponse>(
"/generation/variant-plans",
params,
)
return response.data
}

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