Compare commits

..

5 Commits

Author SHA1 Message Date
CI Bot 6f222db4c7 fix(editor): 移除asset_id min_length限制+order默认None (#1468)
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
CI/CD Pipeline / Deploy Staging (Watchtower auto-deploy) (pull_request) Has been skipped
CI/CD Pipeline / Staging E2E Tests (pull_request) Has been skipped
CI/CD Pipeline / Staging API Integration Tests (pull_request) Has been skipped
CI/CD Pipeline / ACR Image Cleanup (pull_request) Has been skipped
CI/CD Pipeline / Check if frontend-only change (pull_request) Successful in 45s
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 Web Image (pull_request) Has been skipped
CI/CD Pipeline / PR Build Worker Image (pull_request) Successful in 44s
PR Automation / Auto Merge on CI Green + Approved (pull_request) Successful in 1m49s
CI/CD Pipeline / Validate - Type Check (mypy) (pull_request) Successful in 1m52s
CI/CD Pipeline / Validate - Migration (alembic) (pull_request) Successful in 1m53s
Preview Deploy / Deploy Preview Environment (pull_request) Successful in 2m23s
CI/CD Pipeline / PR Build API Image (pull_request) Successful in 1m44s
PR Automation / Auto Approve on CI Green (pull_request) Successful in 3m35s
CI/CD Pipeline / Unit Tests (pull_request) Successful in 3m9s
CI/CD Pipeline / Validate - Code Quality (pull_request) Successful in 5m54s
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 / Canary Release to Production (pull_request) Has been skipped
CI/CD Pipeline / Deploy Production (pull_request) Has been skipped
CI/CD Pipeline / Production Browser E2E (pull_request) Has been skipped
AI Code Review / AI Code Review (pull_request) Successful in 6m51s
CI/CD Pipeline / Integration Tests (pull_request) Successful in 2m15s
CI/CD Pipeline / CI Gate (pull_request) Successful in 6s
ACR Cleanup / ACR Image Cleanup (pull_request_target) Successful in 39s
Preview Cleanup / Cleanup Preview Environment (pull_request) Successful in 57s
AI Code Review 第三轮修复:
1. EditorClipBatchItem.asset_id 移除 min_length=1,允许空字符串(占位片段场景)
2. EditorClipBatchItem.order 默认值改为 None(非0),仅在显式传入时加入 dict
3. 服务层 order 逻辑:dict 无 order key 时回退到索引 i
4. 测试更新:9 passed
2026-08-23 15:54:34 +08:00
CI Bot ee561b0d64 fix(editor): B017 lint修复 pytest.raises(Exception)→ValidationError
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
CI/CD Pipeline / Deploy Staging (Watchtower auto-deploy) (pull_request) Has been skipped
CI/CD Pipeline / Staging API Integration Tests (pull_request) Has been skipped
CI/CD Pipeline / Staging E2E Tests (pull_request) Has been skipped
CI/CD Pipeline / ACR Image Cleanup (pull_request) Has been skipped
CI/CD Pipeline / Check if frontend-only change (pull_request) Successful in 32s
CI/CD Pipeline / Frontend Lint (pull_request) Has been skipped
CI/CD Pipeline / Frontend Unit Tests (pull_request) Has been skipped
CI/CD Pipeline / Validate - Migration (alembic) (pull_request) Successful in 1m20s
CI/CD Pipeline / Validate - Type Check (mypy) (pull_request) Successful in 1m40s
PR Automation / Auto Merge on CI Green + Approved (pull_request) Successful in 1m39s
CI/CD Pipeline / PR Build Web Image (pull_request) Has been skipped
Preview Deploy / Deploy Preview Environment (pull_request) Successful in 1m43s
AI Code Review / AI Code Review (pull_request) Failing after 2m9s
CI/CD Pipeline / PR Build Worker Image (pull_request) Successful in 31s
CI/CD Pipeline / PR Build API Image (pull_request) Successful in 1m42s
PR Automation / Auto Approve on CI Green (pull_request) Successful in 3m41s
CI/CD Pipeline / Unit Tests (pull_request) Successful in 3m30s
CI/CD Pipeline / Validate - Code Quality (pull_request) Successful in 5m17s
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 / 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
CI/CD Pipeline / Integration Tests (pull_request) Has been cancelled
CI/CD Pipeline / CI Gate (pull_request) Has been cancelled
2026-08-23 15:39:08 +08:00
CI Bot 0386a1f08f fix(editor): 批量更新clips加事务保护+asset_id校验 (#1468)
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
CI/CD Pipeline / Deploy Staging (Watchtower auto-deploy) (pull_request) Has been skipped
CI/CD Pipeline / Staging API Integration Tests (pull_request) Has been skipped
CI/CD Pipeline / Staging E2E Tests (pull_request) Has been skipped
CI/CD Pipeline / ACR Image Cleanup (pull_request) Has been skipped
CI/CD Pipeline / Check if frontend-only change (pull_request) Successful in 31s
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 Web Image (pull_request) Has been skipped
CI/CD Pipeline / PR Build Worker Image (pull_request) Successful in 40s
CI/CD Pipeline / Validate - Migration (alembic) (pull_request) Successful in 1m43s
PR Automation / Auto Merge on CI Green + Approved (pull_request) Successful in 1m42s
CI/CD Pipeline / Validate - Type Check (mypy) (pull_request) Successful in 1m53s
AI Code Review / AI Code Review (pull_request) Failing after 2m13s
Preview Deploy / Deploy Preview Environment (pull_request) Successful in 2m33s
CI/CD Pipeline / PR Build API Image (pull_request) Successful in 2m4s
CI/CD Pipeline / Validate - Code Quality (pull_request) Failing after 3m14s
PR Automation / Auto Approve on CI Green (pull_request) Successful in 3m33s
CI/CD Pipeline / Unit Tests (pull_request) Successful in 3m31s
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 / 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
CI/CD Pipeline / Integration Tests (pull_request) Successful in 1m46s
CI/CD Pipeline / CI Gate (pull_request) Failing after 8s
AI Code Review 修复:
1. batch_update_clips 改用 replace_all_clips_transactional:
   清空→创建→标记ready 在同一事务内完成,失败自动 rollback
2. EditorClipBatchItem.asset_id 加 min_length=1 校验,
   禁止空字符串避免脏数据

新增 8 个测试(5 endpoint + 3 service),全量相关 92 passed。
2026-08-23 15:32:12 +08:00
CI Bot dd7513a512 style: auto-format with black + isort + prettier [skip ci-format-check]
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
CI/CD Pipeline / Deploy Staging (Watchtower auto-deploy) (pull_request) Has been skipped
CI/CD Pipeline / Staging E2E Tests (pull_request) Has been skipped
CI/CD Pipeline / Staging API Integration Tests (pull_request) Has been skipped
CI/CD Pipeline / ACR Image Cleanup (pull_request) Has been skipped
CI/CD Pipeline / Check if frontend-only change (pull_request) Successful in 49s
CI/CD Pipeline / Frontend Lint (pull_request) Has been skipped
CI/CD Pipeline / Frontend Unit Tests (pull_request) Has been skipped
CI/CD Pipeline / Validate - Migration (alembic) (pull_request) Successful in 1m32s
CI/CD Pipeline / PR Build Web Image (pull_request) Has been skipped
AI Code Review / AI Code Review (pull_request) Failing after 1m50s
PR Automation / Auto Merge on CI Green + Approved (pull_request) Successful in 1m52s
CI/CD Pipeline / Validate - Type Check (mypy) (pull_request) Successful in 2m1s
CI/CD Pipeline / PR Build Worker Image (pull_request) Successful in 31s
Preview Deploy / Deploy Preview Environment (pull_request) Successful in 2m25s
CI/CD Pipeline / PR Build API Image (pull_request) Successful in 1m41s
PR Automation / Auto Approve on CI Green (pull_request) Successful in 3m44s
CI/CD Pipeline / Unit Tests (pull_request) Successful in 3m9s
CI/CD Pipeline / Validate - Code Quality (pull_request) Successful in 5m2s
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 / 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
CI/CD Pipeline / Integration Tests (pull_request) Successful in 1m43s
CI/CD Pipeline / CI Gate (pull_request) Successful in 7s
2026-08-23 07:17:14 +00:00
CI Bot 0cb67afd61 refactor(worker): generate_video直接读取edit_plan_clips渲染,删除内存重建逻辑
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
CI/CD Pipeline / Deploy Staging (Watchtower auto-deploy) (pull_request) Has been skipped
CI/CD Pipeline / Staging E2E Tests (pull_request) Has been skipped
CI/CD Pipeline / Staging API Integration Tests (pull_request) Has been skipped
CI/CD Pipeline / ACR Image Cleanup (pull_request) Has been skipped
CI/CD Pipeline / Check if frontend-only change (pull_request) Successful in 40s
CI/CD Pipeline / Frontend Unit Tests (pull_request) Has been skipped
CI/CD Pipeline / Frontend Lint (pull_request) Has been skipped
CI/CD Pipeline / Validate - Migration (alembic) (pull_request) Successful in 1m34s
CI/CD Pipeline / Validate - Type Check (mypy) (pull_request) Successful in 1m46s
CI/CD Pipeline / PR Build Web Image (pull_request) Has been skipped
PR Automation / Auto Merge on CI Green + Approved (pull_request) Successful in 1m54s
AI Code Review / AI Code Review (pull_request) Failing after 2m10s
Preview Deploy / Deploy Preview Environment (pull_request) Successful in 2m16s
CI/CD Pipeline / PR Build Worker Image (pull_request) Successful in 31s
CI/CD Pipeline / PR Build API Image (pull_request) Successful in 1m46s
CI/CD Pipeline / Validate - Code Quality (pull_request) Has been cancelled
CI/CD Pipeline / Unit Tests (pull_request) Has been cancelled
CI/CD Pipeline / Integration Tests (pull_request) Has been cancelled
CI/CD Pipeline / Build Production API Image (pull_request) Has been cancelled
CI/CD Pipeline / Build Production Web Image (pull_request) Has been cancelled
CI/CD Pipeline / Build Production Worker Image (pull_request) Has been cancelled
CI/CD Pipeline / Deploy Production (pull_request) Has been cancelled
CI/CD Pipeline / Production Browser E2E (pull_request) Has been cancelled
CI/CD Pipeline / Canary Release to Production (pull_request) Has been cancelled
CI/CD Pipeline / CI Gate (pull_request) Has been cancelled
PR Automation / Auto Approve on CI Green (pull_request) Has been cancelled
核心改造:
- 有source_edit_plan_id时,Worker直接调用RenderAdapter.render_plan(plan_id)
  从数据库加载edit_plan+clips渲染,不再_download_all_assets+_build_plan_and_clips
- 新增_sync_task_config_to_plan:将title_config/bgm_config/分辨率同步到plan.config,
  配音下载后通过voiceover_audio_path传给RenderAdapter
- 无source_edit_plan_id时保留旧路径(标记DEPRECATED)
- RenderAdapter.render_plan新增voiceover_audio_path参数透传给_do_render

新增API:
- PUT /templates/{id}/editor/clips 批量替换clips(delete_all+create+mark_ready)
- EditorClipBatchItem/EditorClipBatchUpdateRequest/EditorClipBatchUpdateResponse schema

删除死代码:
- worker_app/tasks/edit_plan_generation.py(worker.render_edit_plan,452行)
- worker_app/tasks/compose_video.py(worker.compose_video,198行)
- celery_app.py imports清理、tasks/__init__.py清理
- test_edit_plan_worker_failure.py、test_cover_url_finalize.py(测试已删除模块)
- templates_editor/generation.py的send_task改为worker.generate_video

新增3个单测,全量13790 passed
2026-08-23 15:14:01 +08:00
1096 changed files with 41061 additions and 107526 deletions
-1
View File
@@ -1 +0,0 @@
CI re-trigger after runner add-host/DNS fix. This file is harmless and not referenced.
-7
View File
@@ -196,10 +196,3 @@ DOUBAO_MODEL=doubao-seed-1-6-250615
DOUBAO_BASE_URL=https://ark.cn-beijing.volces.com/api/v3
DOUBAO_TIMEOUT=30
DOUBAO_MAX_RETRIES=2
# ==================== 积分/会员系统 (#1895) ====================
# 积分扣点总开关:默认 false(对现有用户零影响)。
# P2 阶段各业务路由逐个接入 @points_gate 时,用
# `if settings.points_enabled: ...`
# 包裹扣点逻辑;所有路由接入完成并验证通过后再在 staging/prod 打开。
POINTS_ENABLED=false
-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
File diff suppressed because it is too large Load Diff
+2 -2
View File
@@ -2,13 +2,13 @@ name: CI Trigger Monitor
on:
schedule:
- cron: '*/10 * * * *' # 每10分钟检查一次(与pr-auto-scan同步降频)
- cron: '*/5 * * * *' # 每5分钟检查一次
workflow_dispatch:
inputs:
stale_threshold:
description: 'CI未触发告警阈值(分钟)'
required: false
default: '10'
default: '5'
permissions:
contents: read
-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
@@ -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:
schedule:
# - cron: "*/15 * * * *" # DISABLED: temporarily to stop failure spam (2026-09-02) # 每10分钟扫描一次(脚本自带240s墙钟上限,降频减负)
- cron: "*/5 * * * *" # 每5分钟扫描一次
workflow_dispatch:
permissions:
+2 -3
View File
@@ -18,7 +18,7 @@ jobs:
name: Auto Approve on CI Green
runs-on: ci-check
if: github.event_name == 'pull_request' && !github.event.pull_request.draft
timeout-minutes: 10 # 等待CI全绿+审批,需要充足时间
timeout-minutes: 3 # 等待模式:等CI全绿后自动合并,不遗漏任何PR
steps:
- name: Checkout code
shell: sh
@@ -61,8 +61,7 @@ jobs:
name: Auto Merge on CI Green + Approved
runs-on: ci-check
if: github.event_name == 'pull_request' && !github.event.pull_request.draft && github.event.pull_request.base.ref == 'develop'
needs: [auto-approve] # 修复竞态:必须等审批完成后再尝试合并
timeout-minutes: 15 # 等待审批+CI就绪+合并,需要充足时间
timeout-minutes: 3 # 短作业模式:检查一次,不满足就退出,由pr-auto-scan每5分钟定时兜底
steps:
- name: Checkout code
shell: sh
-6
View File
@@ -24,11 +24,6 @@ ruff_cache/
.env.production
.env.staging
!.env.example
# 配置模板不受忽略规则限制
!deploy/configs/.env.staging
!deploy/configs/.env.production
# 渲染后的 env 文件包含真实密钥,绝不能提交
.env.rendered
# OS / editor
.DS_Store
@@ -59,4 +54,3 @@ frontend-v21-ui-prototype-final.html
!.vscode/settings.json
.vscode/extensions.json
.coverage
.env.current
-1
View File
@@ -1 +0,0 @@
retrigger3
-1
View File
@@ -263,4 +263,3 @@ pytest --cov=packages --cov-report=html
---
**License**: MIT
<!-- CI trigger: 1788229339 -->
@@ -1,49 +0,0 @@
"""Add unique index on asset_libraries(project_id, kind)
Revision ID: 058_uq_asset_lib_project_kind
Revises: 057_title_config
Create Date: 2026-08-30
同一项目下同 kind 的素材库业务上唯一(前端 getOrCreate 语义、TTS 保存自动建库)。
加唯一索引兜底并发创建竞态,避免重复素材库。
"""
import sqlalchemy as sa
from alembic import op
revision = "058_uq_asset_lib_project_kind"
down_revision = "057_title_config"
branch_labels = None
depends_on = None
def upgrade() -> None:
# 建唯一索引前清洗历史重复:同 (project_id, kind) 只保留 created_at 最新的一条。
# project_id 为 NULL 的系统级行不参与去重(NULL 在唯一索引中互不冲突)。
op.execute("""
DELETE FROM asset_libraries
WHERE id IN (
SELECT id FROM (
SELECT id,
ROW_NUMBER() OVER (
PARTITION BY project_id, kind
ORDER BY created_at DESC, id DESC
) AS rn
FROM asset_libraries
WHERE project_id IS NOT NULL
) t
WHERE t.rn > 1
)
""")
# 与 model 的 UniqueConstraint 定义保持一致(pg_constraint + pg_index 同时注册),
# 避免 Alembic autogenerate 检测到 schema drift
op.create_unique_constraint(
"uq_asset_libraries_project_kind",
"asset_libraries",
["project_id", "kind"],
)
def downgrade() -> None:
op.drop_constraint("uq_asset_libraries_project_kind", "asset_libraries", type_="unique")
@@ -1,23 +0,0 @@
"""add duplicate_rate to generated_videos
Revision ID: 059_duplicate_rate
Revises: 058_uq_asset_lib_project_kind
Create Date: 2026-08-31
"""
import sqlalchemy as sa
from alembic import op
revision = "059_duplicate_rate"
down_revision = "058_uq_asset_lib_project_kind"
branch_labels = None
depends_on = None
def upgrade() -> None:
op.add_column("generated_videos", sa.Column("duplicate_rate", sa.Float(), nullable=True))
def downgrade() -> None:
op.drop_column("generated_videos", "duplicate_rate")
@@ -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,
)
@@ -1,27 +0,0 @@
"""add sentence_timings to lipsync_jobs
Revision ID: 075_add_sentence_timings
Revises: 074_ai_avatar_render_script_id_optional
Create Date: 2026-09-12
"""
import sqlalchemy as sa
from alembic import op
revision = "075_add_sentence_timings"
down_revision = "074_render_script_id_optional"
branch_labels = None
depends_on = None
def upgrade() -> None:
with op.batch_alter_table("lipsync_jobs") as batch:
batch.add_column(
sa.Column("sentence_timings", sa.JSON(), nullable=True),
)
def downgrade() -> None:
with op.batch_alter_table("lipsync_jobs") as batch:
batch.drop_column("sentence_timings")
-133
View File
@@ -1,133 +0,0 @@
"""add membership & points system
Revision ID: 076_membership_points
Revises: 075_add_sentence_timings
Create Date: 2026-09-15
"""
import sqlalchemy as sa
from sqlalchemy import text
from alembic import op
revision = "076_membership_points"
down_revision = "075_add_sentence_timings"
branch_labels = None
depends_on = None
def upgrade() -> None:
# 1. users 表新增字段
with op.batch_alter_table("users") as batch:
batch.add_column(
sa.Column("is_member", sa.Boolean(), nullable=False, server_default=sa.text("false")),
)
batch.add_column(
sa.Column("member_type", sa.String(20), nullable=True),
)
batch.add_column(
sa.Column("member_expires_at", sa.DateTime(), nullable=True),
)
batch.add_column(
sa.Column("points_balance", sa.Integer(), nullable=False, server_default=sa.text("0")),
)
# 2. points_accounts 积分账户表
op.create_table(
"points_accounts",
sa.Column("id", sa.String(36), primary_key=True),
sa.Column("user_id", sa.String(36), nullable=False, unique=True, index=True),
sa.Column("balance", sa.Integer(), nullable=False, server_default=sa.text("0")),
sa.Column("total_earned", sa.Integer(), nullable=False, server_default=sa.text("0")),
sa.Column("total_spent", sa.Integer(), nullable=False, server_default=sa.text("0")),
sa.Column(
"created_at",
sa.DateTime(),
nullable=False,
server_default=sa.text("NOW()"),
),
sa.Column(
"updated_at",
sa.DateTime(),
nullable=False,
server_default=sa.text("NOW()"),
),
)
# 3. points_transactions 积分流水表
op.create_table(
"points_transactions",
sa.Column("id", sa.String(36), primary_key=True),
sa.Column("user_id", sa.String(36), nullable=False, index=True),
sa.Column("account_id", sa.String(36), nullable=False, index=True),
sa.Column("type", sa.String(20), nullable=False, index=True),
sa.Column("source", sa.String(50), nullable=False, index=True),
sa.Column("amount", sa.Integer(), nullable=False),
sa.Column("balance_after", sa.Integer(), nullable=False),
sa.Column("description", sa.String(255), nullable=False, server_default=""),
sa.Column("ref_id", sa.String(100), nullable=False, server_default=""),
sa.Column(
"created_at",
sa.DateTime(),
nullable=False,
server_default=sa.text("NOW()"),
),
)
# 4. points_orders 积分/会员订单表
op.create_table(
"points_orders",
sa.Column("id", sa.String(36), primary_key=True),
sa.Column("user_id", sa.String(36), nullable=False, index=True),
sa.Column("order_type", sa.String(20), nullable=False),
sa.Column("product_code", sa.String(50), nullable=False),
sa.Column("amount_cents", sa.Integer(), nullable=False),
sa.Column("original_amount_cents", sa.Integer(), nullable=False, server_default=sa.text("0")),
sa.Column("discount", sa.Float(), nullable=False, server_default=sa.text("1.0")),
sa.Column("points_amount", sa.Integer(), nullable=False, server_default=sa.text("0")),
sa.Column("status", sa.String(20), nullable=False, server_default="pending", index=True),
sa.Column("payment_method", sa.String(50), nullable=True),
sa.Column("payment_id", sa.String(100), nullable=True),
sa.Column("paid_at", sa.DateTime(), nullable=True),
sa.Column(
"created_at",
sa.DateTime(),
nullable=False,
server_default=sa.text("NOW()"),
),
)
# 5. daily_usage_records 每日使用记录表
op.create_table(
"daily_usage_records",
sa.Column("id", sa.String(36), primary_key=True),
sa.Column("user_id", sa.String(36), nullable=False, index=True),
sa.Column("usage_date", sa.DateTime(), nullable=False),
sa.Column("usage_type", sa.String(50), nullable=False, server_default="free_clip"),
sa.Column("count", sa.Integer(), nullable=False, server_default=sa.text("0")),
sa.Column(
"updated_at",
sa.DateTime(),
nullable=False,
server_default=sa.text("NOW()"),
),
sa.UniqueConstraint(
"user_id",
"usage_date",
"usage_type",
name="uq_daily_usage_user_date_type",
),
)
def downgrade() -> None:
op.drop_table("daily_usage_records")
op.drop_table("points_orders")
op.drop_table("points_transactions")
op.drop_table("points_accounts")
with op.batch_alter_table("users") as batch:
batch.drop_column("points_balance")
batch.drop_column("member_expires_at")
batch.drop_column("member_type")
batch.drop_column("is_member")
View File
-46
View File
@@ -1,27 +1,20 @@
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_libraries import router as asset_libraries_router
from app.api.routes.assets import router as assets_router
from app.api.routes.auth import router as auth_router
from app.api.routes.chunked_upload import router as chunked_upload_router
from app.api.routes.classification_jobs import router as classification_jobs_router
from app.api.routes.clips_standalone import router as clips_standalone_router
from app.api.routes.cover_templates import router as cover_templates_router
from app.api.routes.duplication import router as duplication_router
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_preview import router as generation_preview_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.ingest_jobs import router as ingest_jobs_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.points import points_router, usage_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.scripts_ai import router as scripts_ai_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.tags import router as tags_router
@@ -44,11 +37,6 @@ api_router.include_router(
auth_router,
tags=["Auth"],
)
api_router.include_router(
lipsync_router,
prefix="/lipsync",
tags=["Lipsync"],
)
api_router.include_router(
projects_router,
prefix="/projects",
@@ -111,11 +99,6 @@ api_router.include_router(
prefix="/generation",
tags=["Generation"],
)
api_router.include_router(
generation_variant_plans_router,
prefix="/generation",
tags=["Generation"],
)
api_router.include_router(
generation_cover_router,
prefix="/generation",
@@ -159,10 +142,6 @@ api_router.include_router(
prefix="/templates",
tags=["Template"],
)
api_router.include_router(
clips_standalone_router,
tags=["Clips"],
)
api_router.include_router(
templates_editor_router,
prefix="/templates/{template_id}/editor",
@@ -186,28 +165,3 @@ api_router.include_router(
internal_render_router,
tags=["Internal"],
)
api_router.include_router(
scripts_router,
prefix="/scripts",
tags=["ScriptLibrary"],
)
api_router.include_router(
scripts_ai_router,
prefix="/scripts",
tags=["ScriptLibrary AI"],
)
api_router.include_router(
ai_avatar_render_router,
prefix="/ai-avatar/render",
tags=["AI Avatar Render"],
)
api_router.include_router(
points_router,
prefix="/points",
tags=["Points"],
)
api_router.include_router(
usage_router,
prefix="/usage",
tags=["Usage"],
)
@@ -1,91 +0,0 @@
"""默认模板兜底共享逻辑(P0 #1922).
提供 get_or_create_default_template_id(db, user_id) 共享函数,
供 templates.py 列表查询、clips_standalone.py 独立端点、dependencies.py
resolve_draft_plan_id 三处复用,避免三处各写一套兜底逻辑产生分叉。
根因:PR#1918 清理模板管理 API 时误删了 GET /templates 自动创建默认模板
兜底,前端 PR#1913 去掉空 tid 拦截后首次进入生成页拼出
/templates//editor/clips/from-assets(双斜杠)→ FastAPI 404,阻断新用户首次
生成。
"""
from __future__ import annotations
import logging
from typing import Optional
from sqlalchemy.orm import Session
logger = logging.getLogger(__name__)
def get_or_create_default_template_id(db: Session, user_id: str) -> Optional[str]:
"""获取或自动创建默认配音模板的 id。
判定逻辑(不做异常降级,只有确实创建失败时才回滚重查):
1. 查用户名下 is_active=True 且有 TemplateClipConfig 的模板 → 返回其 id
2. 无则调用 CreateTemplateUseCase 创建一条默认 voice_over 模板;
3. 创建异常时 rollback 再重查一次(防并发唯一键冲突),重查仍无返回 None。
"""
from packages.adapters.sqlalchemy_impl.models import (
TemplateClipConfigModel,
TemplateModel,
)
from packages.adapters.sqlalchemy_impl.template_repository import (
SQLAlchemyTemplateRepository,
)
from packages.application.template.commands import (
CreateTemplateCommand,
SegmentCommand,
)
from packages.application.template.use_cases import CreateTemplateUseCase
existing = (
db.query(TemplateModel)
.filter(TemplateModel.user_id == user_id, TemplateModel.is_active.is_(True))
.order_by(TemplateModel.created_at.asc())
.first()
)
if existing is not None:
has_seg = (
db.query(TemplateClipConfigModel.id).filter(TemplateClipConfigModel.template_id == existing.id).first()
)
if has_seg:
return existing.id
try:
repo = SQLAlchemyTemplateRepository(db)
cmd = CreateTemplateCommand(
user_id=user_id,
name="默认配音模板",
mode="voice_over",
category="default",
tags=[],
title_config={},
subtitle_config={},
bgm_config={},
estimated_duration=0.0,
segments=[SegmentCommand(segment_order=0, duration_min=1.0, duration_max=30.0)],
)
tpl = CreateTemplateUseCase(repo).execute(cmd)
db.commit()
logger.info("auto-created default voice_over template: id=%s user=%s", tpl.id, user_id)
return tpl.id
except Exception:
db.rollback()
# 重查:可能并发请求已建好
existing = (
db.query(TemplateModel)
.filter(TemplateModel.user_id == user_id, TemplateModel.is_active.is_(True))
.order_by(TemplateModel.created_at.asc())
.first()
)
if existing is not None:
has_seg = (
db.query(TemplateClipConfigModel.id).filter(TemplateClipConfigModel.template_id == existing.id).first()
)
if has_seg:
return existing.id
logger.exception("failed to auto-create default template user=%s", user_id)
return None
+2 -2
View File
@@ -1,6 +1,6 @@
"""路由层共享辅助函数 — 消除跨文件重复定义。"""
from datetime import UTC, datetime
from datetime import datetime, timezone
from typing import Any
from fastapi import HTTPException, status
@@ -138,4 +138,4 @@ def format_utc_datetime(dt: datetime | None) -> str:
return dt
if dt.tzinfo is None:
return dt.isoformat() + "Z"
return dt.astimezone(UTC).isoformat().replace("+00:00", "Z")
return dt.astimezone(timezone.utc).isoformat().replace("+00:00", "Z")
+6 -6
View File
@@ -5,7 +5,7 @@
from __future__ import annotations
from typing import Literal
from typing import List, Literal
from app.services.ai_service import TITLE_STYLES, generate_smart_titles, semantic_match_assets
from fastapi import APIRouter
@@ -31,7 +31,7 @@ class GenerateTitlesRequest(BaseModel):
class GenerateTitlesResponse(BaseModel):
"""智能标题生成响应."""
titles: list[str] = Field(..., description="生成的标题列表")
titles: List[str] = Field(..., description="生成的标题列表")
style: str = Field(..., description="实际使用的风格")
source: str = Field(..., description="来源:doubao 或 fallback")
description: str = Field(..., description="原始描述")
@@ -53,7 +53,7 @@ class AssetMatchItem(BaseModel):
id: str = Field(..., description="素材ID")
name: str = Field(default="", description="素材名称")
tags: list[str] = Field(default_factory=list, description="标签列表")
tags: List[str] = Field(default_factory=list, description="标签列表")
description: str = Field(default="", description="素材描述")
@@ -61,7 +61,7 @@ class SemanticMatchRequest(BaseModel):
"""语义匹配请求."""
description: str = Field(..., min_length=1, max_length=500, description="目标视频内容描述")
assets: list[AssetMatchItem] = Field(..., min_length=1, max_length=100, description="待匹配素材列表")
assets: List[AssetMatchItem] = Field(..., min_length=1, max_length=100, description="待匹配素材列表")
top_k: int = Field(default=0, ge=0, le=100, description="返回前K个,0返回全部")
@@ -75,7 +75,7 @@ class SemanticMatchResultItem(AssetMatchItem):
class SemanticMatchResponse(BaseModel):
"""语义匹配响应."""
matches: list[SemanticMatchResultItem] = Field(..., description="按匹配度降序排列的素材列表")
matches: List[SemanticMatchResultItem] = Field(..., description="按匹配度降序排列的素材列表")
source: str = Field(..., description="来源:doubao / fallback")
description: str = Field(..., description="原始描述")
total: int = Field(..., description="输入素材总数")
@@ -99,7 +99,7 @@ def generate_titles(request: GenerateTitlesRequest):
return GenerateTitlesResponse(**result)
@router.get("/titles/styles", response_model=list[TitleStyleInfo])
@router.get("/titles/styles", response_model=List[TitleStyleInfo])
def list_title_styles():
"""获取支持的标题风格列表."""
return [
-327
View File
@@ -1,327 +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 datetime import UTC, datetime
from app.auth import AuthenticatedUser, get_current_user
from packages.middleware.points_gate import points_gate
from app.dependencies import get_db_session
from app.schemas.ai_avatar_render import (
AiAvatarRenderJobResponse,
CreateAiAvatarRenderRequest,
FinalizeRenderResponse,
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)
@points_gate("ai_digital_human", per_unit=15)
def create_render_job(
body: CreateAiAvatarRenderRequest,
current_user: AuthenticatedUser = Depends(get_current_user),
svc: AiAvatarRenderService = Depends(_get_service),
db: Session = Depends(get_db_session),
):
"""提交 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 as exc:
logger.exception("Celery 任务投递失败(创建): job_id=%s err=%s", job.id, exc)
job.status = "failed"
job.error_message = f"任务提交失败:{exc}"
job.updated_at = datetime.now(UTC)
svc.db.commit()
svc.db.refresh(job)
return AiAvatarRenderJobResponse.model_validate(job)
return AiAvatarRenderJobResponse.model_validate(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 as exc:
logger.exception("Celery 任务投递失败(重试): job_id=%s err=%s", job.id, exc)
job.status = "failed"
job.error_message = f"任务提交失败:{exc}"
job.updated_at = datetime.now(UTC)
svc.db.commit()
svc.db.refresh(job)
return AiAvatarRenderJobResponse.model_validate(job)
return AiAvatarRenderJobResponse.model_validate(job)
# ── POST /{job_id}/smart-cover — 从最终成片智能抽封面(步骤②)────────
@router.post("/{job_id}/smart-cover", response_model=SmartCoverResponse)
def generate_render_smart_cover(
job_id: str,
current_user: AuthenticatedUser = Depends(get_current_user),
db: Session = Depends(get_db_session),
):
"""从最终渲染成片智能抽帧生成封面(MediaKit 抽帧 + 评分选最佳帧 + 转存 OSS).
- 必须等渲染任务 completed 后才可调用(否则返回 400)
- 生成成功后自动更新 render_job 的 cover_config 与 output_cover_url
"""
from app.services.ai_avatar_render_service import AiAvatarRenderService
svc = AiAvatarRenderService(db)
job = svc.get_render_job(job_id, current_user.user.id)
if job is None:
raise HTTPException(status_code=404, detail="渲染任务不存在")
if job.status != "completed":
raise HTTPException(status_code=400, detail="请先完成视频生成")
video_url = (job.output_video_url or "").strip()
if not video_url:
raise HTTPException(status_code=400, detail="渲染成片视频 URL 为空")
try:
# 从最终成片抽帧,帧本身已含标题/B-roll,直接转存 OSS
cover_url = generate_smart_cover(video_url, job_id=job_id, max_frames=5)
except Exception as exc:
logger.error(
"渲染成片智能封面生成异常: user=%s render_id=%s video_url=%s err=%s",
current_user.user.id,
job_id,
video_url[:80],
exc,
exc_info=True,
)
cover_url = ""
if not cover_url:
return SmartCoverResponse(
cover_url="",
status="fallback_failed",
message="智能抽帧失败(MediaKit 不可用或抽帧异常),请稍后重试",
)
# 更新 render_job 的封面字段(异步写入 DB;失败不影响返回)
try:
job.cover_config = {
**(job.cover_config if isinstance(job.cover_config, dict) else {}),
"mode": "auto_frame",
"url": cover_url,
}
job.output_cover_url = cover_url
job.updated_at = datetime.now(UTC)
db.commit()
except Exception as exc:
logger.warning("更新 render_job 封面字段失败(不影响返回): job_id=%s err=%s", job_id, exc)
logger.info(
"渲染成片智能封面生成成功: user=%s render_id=%s cover_url=%s",
current_user.user.id,
job_id,
cover_url[:120],
)
return SmartCoverResponse(cover_url=cover_url, status="completed")
# ── POST /{job_id}/finalize — 封面选定后正式入库成片库 ────────────────────
@router.post("/{job_id}/finalize", response_model=FinalizeRenderResponse)
def finalize_render_job(
job_id: str,
current_user: AuthenticatedUser = Depends(get_current_user),
db: Session = Depends(get_db_session),
):
"""用户完成封面选择后,将视频正式保存到成片库.
- 必须等渲染任务 completed 后才可调用
- 如果已通过 smart-cover/custom-cover 设置了封面,会自动带上
- 返回成片库视频ID
- 幂等:已 finalize 的任务重复调用会返回 existing 记录
"""
from app.services.ai_avatar_render_service import AiAvatarRenderError, AiAvatarRenderService
svc = AiAvatarRenderService(db)
job = svc.get_render_job(job_id, current_user.user.id)
if job is None:
raise HTTPException(status_code=404, detail="渲染任务不存在")
if job.status != "completed":
raise HTTPException(status_code=400, detail="请先完成视频生成")
# 幂等检查(通过 generation_task_id=job_id 识别,finalize_job 内部也做了一次,这里提前返回简化)
from packages.adapters.sqlalchemy_impl.models import GeneratedVideoModel
existing = (
db.query(GeneratedVideoModel)
.filter(
GeneratedVideoModel.user_id == current_user.user.id,
GeneratedVideoModel.generation_task_id == job_id,
)
.first()
)
if existing is not None:
return FinalizeRenderResponse(
video_id=existing.id,
cover_url=existing.thumbnail_url or "",
status="already_finalized",
)
try:
video = svc.finalize_job(job_id, current_user.user.id)
return FinalizeRenderResponse(
video_id=video.id,
cover_url=video.thumbnail_url or job.output_cover_url or "",
status="success",
)
except AiAvatarRenderError as exc:
status_map = {
"RenderJobNotFound": 404,
"RenderNotCompleted": 400,
"OutputVideoMissing": 400,
}
raise HTTPException(
status_code=status_map.get(exc.code, 400),
detail=str(exc),
) from exc
except Exception as exc:
logger.error("渲染任务finalize失败: job_id=%s err=%s", job_id, exc, exc_info=True)
raise HTTPException(status_code=500, detail=f"保存到成片库失败: {str(exc)}") from exc
+24 -5
View File
@@ -20,7 +20,7 @@ from packages.application import (
GetProjectUseCase,
ListAssetLibrariesUseCase,
)
from packages.domain import AssetLibraryKind
from packages.domain import AssetLibrary, AssetLibraryKind
from ._helpers import check_project_access
@@ -120,11 +120,30 @@ def ensure_default_library(
kind = AssetLibraryKind(request.kind)
# Issue #1775: 幂等获取/创建——依赖唯一约束 uq_asset_libraries_project_kind
# 并发创建冲突时回滚重查返回已有记录,不再依赖应用层"先查后插",也不会 500。
# 查找该项目下同 kind 的素材库,返回第一个
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}素材库")
library = asset_library_repository.get_or_create_default_library(request.project_id, kind, name=default_name)
return _to_asset_library_response(library)
library = AssetLibrary(
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)
+17 -110
View File
@@ -1,5 +1,5 @@
import logging
from typing import Any, Optional
from typing import Any, List, Optional
from app.api.routes._helpers import check_project_access, format_utc_datetime
from app.auth import AuthenticatedUser, get_current_user
@@ -26,7 +26,6 @@ from app.schemas.asset import (
UpdateAssetReviewRequest,
)
from app.schemas.tag import TagAssetsRequest
from app.services.asset_segment_tracker import compute_asset_availability, get_asset_recent_use_counts
from fastapi import APIRouter, Depends, HTTPException, Query, Response
from packages.domain.smart_match import smart_select_assets
@@ -36,23 +35,6 @@ logger = logging.getLogger(__name__)
router = APIRouter()
def _asset_availability_fields(item) -> dict:
"""视频素材返回余量四字段;非视频/无时长/异常时返回 None + usable=True(零影响)。"""
try:
info = compute_asset_availability(item)
except Exception:
logger.warning("计算素材余量失败,按可用处理: asset_id=%s", getattr(item, "id", "?"), exc_info=True)
info = None
if info is None:
return {
"used_duration": None,
"available_duration": None,
"used_ratio": None,
"usable": True,
}
return info
def _to_asset_response(item, storage_service=None) -> AssetResponse:
# 生成签名文件 URL(用于视频播放 / 文件下载)
file_url = None
@@ -64,16 +46,10 @@ def _to_asset_response(item, storage_service=None) -> AssetResponse:
logger.warning("生成签名URL失败: storage_key=%s", item.storage_key, exc_info=True)
file_url = None
# 缩略图:存储的是 storage_key,需要生成签名 URL 供前端使用
# 不再降级使用视频文件 URL(浏览器 <img> 无法渲染 .mp4,会显示黑屏)
thumbnail_url = None
if item.thumbnail_url:
try:
svc = storage_service or get_storage_service()
thumbnail_url = svc.get_download_url(item.thumbnail_url)
except Exception:
logger.warning("生成缩略图签名URL失败: key=%s", item.thumbnail_url, exc_info=True)
thumbnail_url = None
# 缩略图:优先用已有 thumbnail_url,否则对视频素材复用文件签名 URL
thumbnail_url = item.thumbnail_url
if not thumbnail_url and item.mime_type and item.mime_type.startswith("video") and file_url:
thumbnail_url = file_url
return AssetResponse(
id=item.id,
@@ -97,7 +73,6 @@ def _to_asset_response(item, storage_service=None) -> AssetResponse:
created_at=format_utc_datetime(item.created_at),
uploaded_by_user_id=item.uploaded_by_user_id,
tag_ids=getattr(item, "tag_ids", []),
**_asset_availability_fields(item),
)
@@ -290,7 +265,7 @@ def list_assets(
else:
total = asset_repository.count_by_project_ids(project_ids, status=status_list)
# 跨项目分页:逐项目累积直到凑够一页
paged_items = []
paged_items: list = []
offset = skip
remaining = limit
for pid in project_ids:
@@ -390,7 +365,7 @@ def update_asset_review_status(
return _to_asset_response(updated)
@router.post("/batch", response_model=list[AssetResponse])
@router.post("/batch", response_model=List[AssetResponse])
def batch_get_assets(
request: BatchGetRequest,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
@@ -576,89 +551,21 @@ def smart_match_assets(
request.library_id, request.kind, status=["ready"], limit=10000
)
else:
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)
# ── 过滤前置:余量 + 高频使用,过滤在评分/截取 limit 之前完成 ──────────
# 旧实现先 smart_select_assets(limit=N) 再对这 N 条做过滤,过滤后不回补,
# 当排名靠前的素材恰好都被排除时返回空 items(前端回退全选,smart-match 名存实亡)。
# 现在先过滤全量候选,每级过滤后为空/不足则回退上一级,最后才评分截取。
# 调用统一智能选素材算法(kind 已在 DB 层过滤,无需重复过滤)
results = smart_select_assets(
filtered_assets,
limit=request.limit,
kind=None,
)
# 1) 余量过滤:usable=False(零重复可切区间耗尽且历史区间均达复用上限)的素材排除
usable_assets = []
exhausted_assets = []
for a in filtered_assets:
try:
avail = compute_asset_availability(a)
except Exception:
logger.warning(
"smart-match 余量计算失败,按可用处理: asset_id=%s",
getattr(a, "id", "?"),
exc_info=True,
)
avail = None
if avail is not None and not avail["usable"]:
exhausted_assets.append(a)
else:
usable_assets.append(a)
if exhausted_assets:
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
high_freq_assets = set()
if pool:
asset_ids = [getattr(a, "id", "") for a in pool if getattr(a, "id", "")]
if asset_ids:
try:
use_counts = get_asset_recent_use_counts(
db=asset_repository.session,
asset_ids=asset_ids,
recent_video_count=5,
)
for a in pool:
aid = getattr(a, "id", "")
count = use_counts.get(aid, 0)
if count > MAX_RECENT_USE_COUNT:
high_freq_assets.add(aid)
logger.info(
"smart-match 排除高频使用素材: asset_id=%s use_count=%d limit=%d",
aid, count, MAX_RECENT_USE_COUNT,
)
# 回退策略:排除后剩余素材不足(为空或不够 limit)时,
# 不再全部排除,保留全部可用素材
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:
logger.info(
"smart-match 高频排除后素材不足(%d<%s),保留全部 %d",
remaining_count,
request.limit if request.limit is not None else "不限",
len(pool),
)
except Exception:
logger.warning("smart-match 高频使用查询失败,跳过排除", exc_info=True)
# 3) 调用统一智能选素材算法(kind 已在 DB 层过滤,无需重复过滤)
results = smart_select_assets(pool, limit=request.limit, kind=None)
# 扁平结构:SmartMatchItem 继承 AssetResponse,素材字段直接在条目顶层,
# 前端无需解析 item.asset 包装层,item.id / item.usable / 余量字段直接可读
items = [
SmartMatchItem(
**_to_asset_response(r.asset).model_dump(),
asset=_to_asset_response(r.asset),
score=r.score,
breakdown=r.breakdown,
)
+4 -205
View File
@@ -1,5 +1,4 @@
"""
from __future__ import annotations
Canonical authentication API routes.
The route layer is intentionally thin: repository construction lives in
@@ -13,10 +12,10 @@ from typing import Optional
import jwt
from app.auth import AuthenticatedUser, blacklist_token, get_current_user
from app.config import settings
from app.dependencies import get_auth_email_service, get_auth_session_store, get_db_session, get_user_repository
from fastapi import APIRouter, Depends, Header, HTTPException, Request, status
from app.dependencies import get_auth_email_service, get_auth_session_store, get_user_repository
from fastapi import APIRouter, Depends, Header, HTTPException, status
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.smtp import NoopEmailService
@@ -85,23 +84,6 @@ class CurrentUserResponse(BaseModel):
phone: str = ""
phone_verified: 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):
@@ -126,7 +108,6 @@ async def register(
request: RegisterRequest,
user_repository: UserRepository = Depends(get_user_repository),
email_service=Depends(get_auth_email_service),
db=Depends(get_db_session),
) -> RegisterResponse:
use_case = RegisterUserUseCase(
user_repository=user_repository,
@@ -144,22 +125,6 @@ async def register(
if error or response is None:
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=_translate_auth_error(error))
# 新用户注册赠送 50 积分(失败不影响注册)
if settings.points_enabled:
try:
from packages.domain.points_service import PointsService
_svc = PointsService()
_svc.add_points(
user_id=response.user_id,
amount=50,
source="task_reward",
db=db,
description="新用户注册赠送",
)
except Exception as _bonus_err:
import logging
logging.getLogger(__name__).warning("注册送积分失败: user_id=%s err=%s", response.user_id, _bonus_err)
return RegisterResponse(
user_id=response.user_id,
email=response.email,
@@ -307,52 +272,9 @@ async def get_current_user_info(
phone=user.phone or "",
phone_verified=user.phone_verified,
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):
pass
@@ -504,7 +426,6 @@ async def get_wechat_auth_url() -> WechatAuthUrlResponse:
@router.post("/wechat/callback", response_model=WechatLoginResponse)
async def wechat_callback(
request: WechatCallbackRequest,
http_request: Request,
user_repository: UserRepository = Depends(get_user_repository),
) -> WechatLoginResponse:
"""微信登录回调处理"""
@@ -512,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 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 换微信用户信息
oauth_service = get_wechat_oauth_service()
wechat_user, err = oauth_service.handle_callback(request.code, request.state)
if err:
# state 校验失败 / 微信 errcode 等错误原文已在 service 内 log,这里带上 UA 上下文
logger.warning("[微信回调] 处理失败: err=%s 微信内置浏览器=%s", err, is_wechat_browser)
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 逻辑)
use_case = WechatSyncUseCase(user_repository=user_repository)
@@ -554,7 +456,7 @@ async def wechat_callback(
user = user_repository.find_by_id(response.user_id)
binding_complete = False
if user:
binding_complete = bool(
binding_complete = (
user.phone_verified and user.email_verified and user.email and "@wechat.local" not in user.email
)
@@ -570,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)
# ==================== 验证码 & 绑定 ====================
+9 -11
View File
@@ -8,13 +8,12 @@ import json
import logging
import shutil
import tempfile
from datetime import UTC, datetime, timedelta
from datetime import datetime, timedelta, timezone
from pathlib import Path
from typing import Any
from uuid import uuid4
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.core.celery_app import celery_app
from app.core.storage import OSSStorageService, get_storage_service
@@ -156,7 +155,7 @@ def _cleanup_expired_uploads() -> int:
if not CHUNK_STORAGE_ROOT.exists():
return 0
now = datetime.now(UTC)
now = datetime.now(timezone.utc)
cleaned = 0
for meta_file in CHUNK_STORAGE_ROOT.glob("*.meta.json"):
@@ -166,7 +165,7 @@ def _cleanup_expired_uploads() -> int:
expires_at = datetime.fromisoformat(meta["expires_at"])
if expires_at.tzinfo is None:
expires_at = expires_at.replace(tzinfo=UTC)
expires_at = expires_at.replace(tzinfo=timezone.utc)
# Only cleanup uploads that are not actively being uploaded
if expires_at < now and meta.get("status") != "uploading":
@@ -177,8 +176,8 @@ def _cleanup_expired_uploads() -> int:
meta_file.unlink()
cleaned += 1
logger.info(f"Cleaned up expired upload: {upload_id}")
except Exception:
logger.exception("Failed to cleanup upload metadata: %s", meta_file)
except Exception as e:
logger.warning(f"Failed to cleanup upload metadata {meta_file}: {e}")
return cleaned
@@ -226,7 +225,7 @@ async def init_chunked_upload(
# Generate upload ID
upload_id = uuid4().hex
now = datetime.now(UTC)
now = datetime.now(timezone.utc)
expires_at = now + timedelta(hours=CHUNK_EXPIRY_HOURS)
# Create chunk directory
@@ -382,8 +381,7 @@ async def complete_chunked_upload(
file_hash=request.file_hash,
)
)
celery_result = celery_app.send_task("worker.ingest_asset", args=[job.id])
_persist_celery_task_id(ingest_job_repository, job, getattr(celery_result, "id", ""))
celery_app.send_task("worker.ingest_asset", args=[job.id])
# Update metadata status
meta["status"] = "completed"
@@ -421,9 +419,9 @@ async def upload_chunk(
# Check expiry
expires_at = datetime.fromisoformat(meta["expires_at"])
if expires_at.tzinfo is None:
expires_at = expires_at.replace(tzinfo=UTC)
expires_at = expires_at.replace(tzinfo=timezone.utc)
if expires_at < datetime.now(UTC):
if expires_at < datetime.now(timezone.utc):
raise HTTPException(status_code=status.HTTP_410_GONE, detail="Upload has expired")
# Validate chunk index
@@ -1,90 +0,0 @@
"""独立的从素材创建片段端点(不依赖 template_id 路径参数).
POST /api/v1/clips/from-assets
- 与 /api/v1/templates/{template_id}/editor/clips/from-assets 功能一致
- 区别:template_id 从 body 传入(可选),为空时后端自动创建/查找默认模板
- 解决前端首次加载时 templateId 为空导致双斜杠 404 的问题(P0 #1922
- 内部复用 resolve_draft_plan_id 和 create_clips_from_assets_editor 的核心逻辑
"""
from __future__ import annotations
import logging
from app.auth import AuthenticatedUser, get_current_user
from app.dependencies import get_asset_repository, get_db_session
from app.services.edit_plan_service import EditPlanService
from app.services.edit_template_service import EditTemplateService
from fastapi import APIRouter, BackgroundTasks, Depends, HTTPException, status
from sqlalchemy.orm import Session
from packages.adapters.sqlalchemy_impl.asset_repository import SQLAlchemyAssetRepository
from ._default_template import get_or_create_default_template_id
from .templates_editor.clips import create_clips_from_assets_editor
from .templates_editor.dependencies import resolve_draft_plan_id
from .templates_editor.schemas import ClipsFromAssetsRequest, ClipsFromAssetsResponse
logger = logging.getLogger(__name__)
router = APIRouter(tags=["Clips"])
class StandaloneClipsRequest(ClipsFromAssetsRequest):
"""扩展请求:template_id 可选(不传则后端自动兜底默认模板)。"""
template_id: str | None = None
def _get_editor_services_direct(db: Session) -> tuple[EditTemplateService, EditPlanService]:
"""直接构造服务实例(非 Depends 版本,供独立端点内部调用)。"""
return EditTemplateService(db), EditPlanService(db)
@router.post("/clips/from-assets", response_model=ClipsFromAssetsResponse)
def create_clips_from_assets(
body: StandaloneClipsRequest,
background_tasks: BackgroundTasks,
db: Session = Depends(get_db_session),
current_user: AuthenticatedUser = Depends(get_current_user),
asset_repo: SQLAlchemyAssetRepository = Depends(get_asset_repository),
) -> ClipsFromAssetsResponse:
"""从素材批量创建片段(template_id 可选,为空自动兜底)。"""
user_id = str(current_user.user.id)
services = _get_editor_services_direct(db)
# 1. 解析/兜底 template_id,拿到 plan_id
template_id = (body.template_id or "").strip()
if not template_id:
template_id = get_or_create_default_template_id(db, user_id)
if not template_id:
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail="无法自动创建默认模板,请刷新页面重试",
)
plan_id = resolve_draft_plan_id(
template_id=template_id,
services=services,
current_user=current_user,
db=db,
auto_create_default=False, # 上面已兜底过
)
# 2. 构造标准化请求(去除独立端扩展字段),复用原端点核心逻辑
core_body = ClipsFromAssetsRequest(
asset_ids=body.asset_ids,
clip_type=body.clip_type,
clip_count=body.clip_count,
required_clips_count=body.required_clips_count,
)
# 3. 直接调用原端点函数(此时所有 Depends 依赖已手动传入)
return create_clips_from_assets_editor(
template_id=template_id,
body=core_body,
background_tasks=background_tasks,
plan_id=plan_id,
services=services,
asset_repo=asset_repo,
db=db,
current_user=current_user,
)
-9
View File
@@ -7,7 +7,6 @@ from typing import Any
from uuid import uuid4
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.dependencies import get_duplication_repository
from app.schemas.duplication import (
@@ -77,8 +76,6 @@ def _to_record_response(record: DuplicationRecord) -> DuplicationRecordResponse:
status=record.status,
duplicate_rate=record.duplicate_rate,
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(),
updated_at=record.updated_at.isoformat(),
)
@@ -93,8 +90,6 @@ def _to_detail_response(record: DuplicationRecord) -> DuplicationDetailResponse:
status=record.status,
duplicate_rate=record.duplicate_rate,
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(),
updated_at=record.updated_at.isoformat(),
segments=[
@@ -197,8 +192,6 @@ async def upload_for_duplication(
authenticated_user.user.id,
)
celery_app.send_task("worker.process_duplication_check", args=[record.id])
return DuplicationUploadResponse(
id=record.id,
status=record.status,
@@ -303,8 +296,6 @@ def retry_duplication(
detail=f"查重记录 {record_id} 不存在",
)
celery_app.send_task("worker.process_duplication_check", args=[updated.id])
return DuplicationUploadResponse(
id=updated.id,
status=updated.status,
+82 -419
View File
@@ -1,21 +1,17 @@
"""封面生成路由 — Generation 模块.
端点:
- POST /generate-cover AI 生成封面(从最终成片视频中抽帧,兼容预览片段回退
- POST /generate-cover AI 生成封面(从预览视频中抽帧)
挂载路径: /api/v1/generation/generate-cover
"""
from __future__ import annotations
import ipaddress
import logging
import re
from typing import Any, Optional
from urllib.parse import urlparse
from typing import Any, List, Optional
from app.auth import AuthenticatedUser, get_current_user
from packages.middleware.points_gate import points_gate
from app.dependencies import get_db_session, get_generated_video_repository
from app.services.edit_plan_service import EditPlanService
from app.services.edit_template_service import EditTemplateService
@@ -28,7 +24,6 @@ from packages.adapters.sqlalchemy_impl.generation_task_repository import (
)
from packages.application import ListGeneratedVideosByTaskUseCase
from packages.domain.config_schemas import normalize_plan_config
from packages.shared.storage import get_shared_storage_service
from .templates_editor.dependencies import get_draft_plan_id, get_editor_services
@@ -42,7 +37,7 @@ router = APIRouter(tags=["Generation"])
class GenerateCoverRequest(BaseModel):
"""AI 封面生成请求体"""
asset_ids: list[str] = Field(default_factory=list, description="素材 ID 列表(确定视频来源)")
asset_ids: List[str] = Field(default_factory=list, description="素材 ID 列表(确定视频来源)")
cover_type: str = Field(
default="ai_frame",
description="封面类型: ai_frame / manual / upload / ai_regenerate",
@@ -56,14 +51,6 @@ class GenerateCoverRequest(BaseModel):
default=None,
description="上传的封面图片 URL,仅 cover_type=upload 时有效",
)
generated_video_id: Optional[str] = Field(
default=None,
description="确认生成产出的最终视频 ID。传入后封面从该视频文件抽帧,而非预览片段。",
)
video_url: Optional[str] = Field(
default=None,
description="最终视频 URL(兜底)。当 generated_video_id 不可用时,直接从此 URL 对应的视频抽帧。",
)
class GenerateCoverResponse(BaseModel):
@@ -76,86 +63,6 @@ class GenerateCoverResponse(BaseModel):
# ── 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(
frame_url: str,
plan_id: str,
@@ -214,6 +121,8 @@ def _persist_cover_frame(
exc_info=True,
)
from packages.shared.storage import get_shared_storage_service
storage = get_shared_storage_service()
cover_key = f"covers/{plan_id}/cover_{uuid.uuid4().hex[:8]}.jpg"
storage.upload_file(
@@ -231,108 +140,7 @@ def _persist_cover_frame(
Path(tmp_path).unlink(missing_ok=True)
def _get_task_video_url(db: Session, task_id: str) -> Optional[str]:
"""从 GenerationTask 关联的 GeneratedVideo 中获取视频 storage_key / URL."""
try:
video_repo = get_generated_video_repository(db)
use_case = ListGeneratedVideosByTaskUseCase(video_repo)
videos = use_case.execute(task_id)
if videos:
return getattr(videos[0], "file_url", "") or ""
except Exception:
logger.warning("[封面生成] 获取任务视频失败: task_id=%s", task_id, exc_info=True)
return None
def _resolve_storage_key_to_url(storage_key: str) -> Optional[str]:
"""将 storage_key 或完整 URL 转换为可访问的裸 URL。"""
if not storage_key:
return None
try:
if storage_key.startswith("http"):
url = storage_key
else:
storage_svc = get_shared_storage_service()
url = storage_svc.get_url(storage_key)
if url:
url = re.sub(r"(?<!:)//", "/", url)
return url
except Exception as e:
logger.warning("[封面生成] storage_key 转 URL 失败: key=%s err=%s", storage_key, e)
return None
def _endpoint_host(value: str) -> str:
"""从 endpoint / URL 字符串中安全提取主机名(兼容有无 scheme 两种配置)。"""
v = (value or "").strip().lower()
if not v:
return ""
if "://" in v:
return (urlparse(v).hostname or "").lower()
# 无 scheme:去掉可能的端口(host:port),urlparse 补 // 以正确解析
return (urlparse("//" + v).hostname or "").lower()
def _is_private_or_reserved_host(host: str) -> bool:
"""判断主机名是否为内网/回环/链路本地/保留地址(IPv4 与 IPv6 统一处理)。
使用标准库 ipaddress 判定;非 IP 主机名(如 localhost)单独处理。
"""
h = host.strip().lower()
if h in {"localhost", "0.0.0.0", "::", "::1"}:
return True
try:
addr = ipaddress.ip_address(h)
# is_private 覆盖 10/8、172.16/12、192.168/16、127/8、169.254/16、
# ::1、fc00::/7、fe80::/10 等全部私有/保留段
return bool(addr.is_private or addr.is_loopback or addr.is_link_local or addr.is_reserved)
except ValueError:
return False
def _is_trusted_media_url(url: str) -> bool:
"""校验 URL 是否指向受信任的存储域名(OSS bucket / 本地存储),防止 SSRF。
用户可通过 video_url 传入视频地址,但服务端(MediaKit)会主动请求该 URL
因此必须限制为自家存储域名,拒绝内网地址、元数据地址等任意主机。
"""
if not url:
return False
try:
parsed = urlparse(url.strip())
if parsed.scheme not in ("http", "https"):
return False
host = (parsed.hostname or "").lower()
if not host:
return False
# 拒绝一切内网/回环/链路本地/保留地址(IPv4 + IPv6,标准库判定)
if _is_private_or_reserved_host(host):
return False
# 允许:自家 OSS bucket 域名(<bucket>.<endpoint>)或 endpoint 自身及其子域
try:
storage_svc = get_shared_storage_service()
trusted_hosts = set()
public_base = getattr(storage_svc, "public_url", "") or ""
h1 = _endpoint_host(public_base)
if h1:
trusted_hosts.add(h1)
h2 = _endpoint_host(getattr(storage_svc, "endpoint", "") or "")
if h2:
trusted_hosts.add(h2)
for trusted in trusted_hosts:
if host == trusted or host.endswith("." + trusted):
return True
except Exception:
logger.warning("[封面生成] 存储域名白名单初始化失败,URL 校验从严拒绝", exc_info=True)
return False
return False
except Exception:
logger.warning("[封面生成] video_url 白名单校验异常,从严拒绝: url=%s", url[:80], exc_info=True)
return False
@router.post("/generate-cover", response_model=GenerateCoverResponse)
@points_gate("ai_cover")
def generate_cover(
body: GenerateCoverRequest,
template_id: str = Query(..., description="模板 ID"),
@@ -341,16 +149,12 @@ def generate_cover(
db: Session = Depends(get_db_session),
current_user: AuthenticatedUser = Depends(get_current_user),
) -> GenerateCoverResponse:
"""AI 生成封面 — 优先从最终成片视频中抽帧,回退到预览片段.
"""AI 生成封面 — 从预览视频中抽帧.
流程(串行):
1. 优先使用前端传入的 generation_task_id 定位最终成片任务,
或自动查找 plan 关联的已完成最终成片任务(is_preview=False
2. 回退:从预览片段获取视频 URL(兼容旧流程)
3. 用裸 URL 让 MediaKit 下载视频并抽帧
4. 帧图下载后上传到 OSS covers/ 路径
MediaKit 的调用方式(strategy / max_frames / 轮询 / 重试 / 降级)不变。
1. 预览视频已渲染完成(通过 3 步查找获取 URL)
2. 用裸 URL 让 MediaKit 下载视频并抽帧
3. 帧图下载后上传到 OSS covers/ 路径
"""
_, plan_svc = services
plan = plan_svc.get_plan_or_raise(plan_id)
@@ -378,111 +182,27 @@ def generate_cover(
)
return GenerateCoverResponse(plan_id=plan_id, cover=cover_data)
# ── 查找用于抽帧的视频 URL ────────────────────────────────────────
# 优先级:
# 0. 请求体显式传入的 generation_task_id(最终成片任务)
# 1. plan.config.rendered_storage_key
# 2. plan.config.generation_task_id 对应的任务
# 3. source_edit_plan_id 关联的已完成「最终成片」任务(is_preview=False
# 4. source_edit_plan_id 关联的已完成预览任务(is_preview=True,兼容回退)
# 5. user + template 最近的已完成预览任务(兜底)
# ── 3 步查找预览视频 URL ──────────────────────────────────────────
# 第一步:从 plan.config 读取
logger.info("[封面生成] 步骤1: 从 plan.config 查找 rendered_storage_key: plan_id=%s", plan_id)
rendered_storage_key = (plan.config or {}).get("rendered_storage_key", "")
# 步骤 0:请求体传入最终视频标识(generated_video_id 或 video_url
if not rendered_storage_key:
# 0a:通过 generated_video_id 查找最终成片视频
if body.generated_video_id:
logger.info(
"[封面生成] 步骤0a: 使用 generated_video_id: plan_id=%s video_id=%s",
plan_id,
body.generated_video_id,
)
try:
gv_repo = get_generated_video_repository(db)
gv = gv_repo.get(body.generated_video_id)
if gv:
file_url = getattr(gv, "file_url", "") or ""
if file_url:
# 权限校验(双重,任何一层确认归属不符即拒绝):
# 1) GeneratedVideo.user_id 直接归属(老数据可能为空,为空时不据此放行)
gv_owner = (getattr(gv, "user_id", "") or "").strip()
if gv_owner and gv_owner != current_user.user.id:
raise HTTPException(status_code=403, detail="无权访问该视频")
# 2) 关联 generation_task 归属校验;关联任务缺失时不可静默放行:
# 若 video 自身无 owner 信息且关联任务也查不到,拒绝访问
gv_task_id = getattr(gv, "generation_task_id", "") or ""
task0 = None
if gv_task_id:
try:
task0 = SQLAlchemyGenerationTaskRepository(db).get(gv_task_id)
except Exception:
logger.warning(
"[封面生成] 步骤0a关联任务查询异常: plan_id=%s task_id=%s",
plan_id,
gv_task_id,
exc_info=True,
)
if task0 is not None:
task_owner = (getattr(task0, "created_by_user_id", "") or "").strip()
if task_owner and task_owner != current_user.user.id:
raise HTTPException(status_code=403, detail="无权访问该视频")
elif not gv_owner:
# video 无 owner 且关联任务不存在/无法确认归属 → 拒绝,防止越权
logger.warning(
"[封面生成] 步骤0a视频归属无法确认,拒绝访问: plan_id=%s video_id=%s",
plan_id,
body.generated_video_id,
)
raise HTTPException(status_code=403, detail="无权访问该视频")
rendered_storage_key = file_url
logger.info(
"[封面生成] ✅ 步骤0a找到最终成片: plan_id=%s video_id=%s url=%s",
plan_id,
body.generated_video_id,
file_url[:80],
)
except HTTPException:
raise
except Exception:
logger.warning(
"[封面生成] 步骤0a查找视频失败: plan_id=%s video_id=%s",
plan_id,
body.generated_video_id,
exc_info=True,
)
# 0b:直接使用 video_url(兜底)— 必须通过存储域名白名单校验,防止 SSRF
if not rendered_storage_key and body.video_url:
if _is_trusted_media_url(body.video_url):
logger.info(
"[封面生成] 步骤0b: 使用请求体传入的 video_url(白名单通过): plan_id=%s url=%s",
plan_id,
body.video_url[:80],
)
rendered_storage_key = body.video_url
else:
logger.warning(
"[封面生成] 步骤0b: video_url 不在受信任存储域名白名单内,已忽略: plan_id=%s url=%s",
plan_id,
body.video_url[:80],
)
# 步骤 2:通过 plan.config.generation_task_id 查找
# 第二步:如果还没有,通过 generation_task_id 查找预览任务的产物
if not rendered_storage_key:
generation_task_id = (plan.config or {}).get("generation_task_id", "")
logger.info(
"[封面生成] 步骤2: 通过 generation_task_id 查找: plan_id=%s task_id=%s", plan_id, generation_task_id
)
if generation_task_id:
logger.info(
"[封面生成] 步骤2: 通过 plan.config.generation_task_id 查找: plan_id=%s task_id=%s",
plan_id,
generation_task_id,
)
try:
_repo = SQLAlchemyGenerationTaskRepository(db)
task = _repo.get(generation_task_id)
gen_task_repo = SQLAlchemyGenerationTaskRepository(db)
task = gen_task_repo.get(generation_task_id)
if task:
rendered_storage_key = _get_task_video_url(db, task.id) or ""
if rendered_storage_key:
video_repo = get_generated_video_repository(db)
use_case = ListGeneratedVideosByTaskUseCase(video_repo)
videos = use_case.execute(task.id)
if videos:
rendered_storage_key = getattr(videos[0], "file_url", "") or ""
logger.info(
"[封面生成] ✅ 步骤2找到视频: plan_id=%s task_id=%s url=%s",
plan_id,
@@ -491,47 +211,26 @@ def generate_cover(
)
except Exception:
logger.warning(
"[封面生成] 步骤2查找失败: plan_id=%s",
"封面生成: 通过 generation_task_id 查找视频失败: plan_id=%s",
plan_id,
exc_info=True,
)
# 步骤 3:通过 source_edit_plan_id 查找已完成「最终成片」任务(is_preview=False
# 第 2.5 步:通过 plan_id 作为 source_edit_plan_id 查找关联的已完成预览任务
if not rendered_storage_key:
try:
_repo = SQLAlchemyGenerationTaskRepository(db)
logger.info("[封面生成] 步骤3: 查找最终成片任务(is_preview=False): plan_id=%s", plan_id)
all_tasks = _repo.list_by_source_edit_plan(plan_id)
for pt in all_tasks:
if getattr(pt, "status", "") == "completed" and not getattr(pt, "is_preview", False):
rendered_storage_key = _get_task_video_url(db, pt.id) or ""
if rendered_storage_key:
logger.info(
"[封面生成] ✅ 步骤3找到最终成片: plan_id=%s task_id=%s url=%s",
plan_id,
pt.id,
rendered_storage_key[:80],
)
break
except Exception:
logger.warning(
"[封面生成] 步骤3查找最终成片失败: plan_id=%s",
plan_id,
exc_info=True,
)
# 步骤 4:兼容回退 — 通过 source_edit_plan_id 查找已完成预览任务
if not rendered_storage_key:
try:
_repo = SQLAlchemyGenerationTaskRepository(db)
logger.info("[封面生成] 步骤4: 回退查找预览任务(is_preview=True): plan_id=%s", plan_id)
preview_tasks = _repo.list_by_source_edit_plan(plan_id)
gen_task_repo = SQLAlchemyGenerationTaskRepository(db)
logger.info("[封面生成] 步骤2.5: 通过 source_edit_plan_id 查找: plan_id=%s", plan_id)
preview_tasks = gen_task_repo.list_by_source_edit_plan(plan_id)
for pt in preview_tasks:
if getattr(pt, "status", "") == "completed" and getattr(pt, "is_preview", False):
rendered_storage_key = _get_task_video_url(db, pt.id) or ""
if rendered_storage_key:
video_repo = get_generated_video_repository(db)
use_case = ListGeneratedVideosByTaskUseCase(video_repo)
videos = use_case.execute(pt.id)
if videos:
rendered_storage_key = getattr(videos[0], "file_url", "") or ""
logger.info(
"[封面生成] ✅ 步骤4找到预览视频: plan_id=%s task_id=%s url=%s",
"[封面生成] ✅ 步骤2.5找到视频: plan_id=%s task_id=%s url=%s",
plan_id,
pt.id,
rendered_storage_key[:80],
@@ -539,50 +238,66 @@ def generate_cover(
break
except Exception:
logger.warning(
"[封面生成] 步骤4查找预览任务失败: plan_id=%s",
"封面生成: 通过 source_edit_plan_id 查找预览任务失败: plan_id=%s",
plan_id,
exc_info=True,
)
# 步骤 5:按 user + template 查找最近的已完成预览任务(兜底)
# 第三步:按 user + template 查找最近的已完成预览任务(兜底)
if not rendered_storage_key:
try:
_repo = SQLAlchemyGenerationTaskRepository(db)
logger.info(
"[封面生成] 步骤5: 通过 user+template 查找预览任务: plan_id=%s template_id=%s",
plan_id,
template_id,
)
preview_tasks = _repo.list_latest_completed_preview(
gen_task_repo = SQLAlchemyGenerationTaskRepository(db)
logger.info("[封面生成] 步骤3: 通过 user+template 查找: plan_id=%s template_id=%s", plan_id, template_id)
preview_tasks = gen_task_repo.list_latest_completed_preview(
user_id=str(current_user.user.id),
template_id=template_id,
)
if preview_tasks:
rendered_storage_key = _get_task_video_url(db, preview_tasks[0].id) or ""
if rendered_storage_key:
completed_preview = preview_tasks[0]
video_repo = get_generated_video_repository(db)
use_case = ListGeneratedVideosByTaskUseCase(video_repo)
videos = use_case.execute(completed_preview.id)
if videos:
rendered_storage_key = getattr(videos[0], "file_url", "") or ""
logger.info(
"[封面生成] ✅ 步骤5找到预览视频: plan_id=%s task_id=%s",
"封面视频: 通过 user+template 找到预览任务: plan_id=%s template_id=%s task_id=%s",
plan_id,
preview_tasks[0].id,
template_id,
completed_preview.id,
)
except Exception:
logger.warning(
"[封面生成] 步骤5 user+template 查找失败: plan_id=%s",
"封面警告: user+template 查找预览任务失败: plan_id=%s template_id=%s",
plan_id,
template_id,
exc_info=True,
)
# 将 storage_key 转换为可访问 URL;找不到视频时不立即报错,
# 因为步骤 E2 可以直接从源素材抽帧(历史数据或 Worker 抽帧失败时的兜底)
# 使用裸 URLrendered/* 已配置公开读);找不到渲染视频时不立即报错,
# 因为步骤 E 可以直接从源素材抽帧(历史数据或 Worker 抽帧失败时的兜底)
primary_video_url = None
if rendered_storage_key:
plan_svc.update_plan_config(plan_id, {"rendered_storage_key": rendered_storage_key})
primary_video_url = _resolve_storage_key_to_url(rendered_storage_key)
logger.info(
"[封面生成] 封面抽帧视频URL: plan_id=%s url=%s",
plan_id,
primary_video_url[:80] if primary_video_url else "",
)
try:
if rendered_storage_key.startswith("http"):
primary_video_url = rendered_storage_key
else:
from packages.shared.storage import get_shared_storage_service
storage_svc = get_shared_storage_service()
primary_video_url = storage_svc.get_url(rendered_storage_key)
if primary_video_url:
import re as _re
primary_video_url = _re.sub(r"(?<!:)//", "/", primary_video_url)
logger.info(
"获取预览视频URL用于封面生成: plan_id=%s url=%s",
plan_id,
primary_video_url[:80] if primary_video_url else "",
)
except Exception as e:
logger.warning("获取预览视频URL失败: plan_id=%s err=%s", plan_id, e)
primary_video_url = None
# 统一封面管道:优先从 GenerationTask.cover_url 读取渲染后视频抽帧的封面
# 多步查找 cover_url,和查找视频 URL 一样的 fallback 逻辑
@@ -595,7 +310,7 @@ def generate_cover(
if generation_task_id:
try:
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
logger.info(
"[封面生成] 统一管道封面(步骤A-direct): plan_id=%s task_id=%s url=%s",
@@ -611,67 +326,20 @@ def generate_cover(
exc_info=True,
)
# 步骤 A2:通过 generated_video_id 查找关联任务的 cover_url
if not cover_url_from_task and body.generated_video_id:
try:
gv_repo = get_generated_video_repository(db)
gv = gv_repo.get(body.generated_video_id)
if gv:
gv_task_id = getattr(gv, "generation_task_id", "") or ""
if 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]
cover_url_from_task = task_a2.cover_url
logger.info(
"[封面生成] 封面(步骤A2-video-task): plan_id=%s video_id=%s url=%s",
plan_id,
body.generated_video_id,
cover_url_from_task[:80],
)
except Exception:
logger.warning(
"[封面生成] 步骤A2读取 cover_url 失败: plan_id=%s video_id=%s",
plan_id,
body.generated_video_id,
exc_info=True,
)
# 步骤 B:通过 source_edit_plan_id 查找关联任务的 cover_url
# 优先最终成片任务(is_preview=False),其次预览任务
# 步骤 B:通过 source_edit_plan_id 查找关联预览任务的 cover_url
if not cover_url_from_task:
try:
all_tasks = gen_task_repo.list_by_source_edit_plan(plan_id)
# 先找最终成片
for pt in all_tasks:
if (
getattr(pt, "status", "") == "completed"
and not getattr(pt, "is_preview", False)
and getattr(pt, "cover_url", "")
):
preview_tasks = gen_task_repo.list_by_source_edit_plan(plan_id)
for pt in preview_tasks:
if getattr(pt, "status", "") == "completed" and getattr(pt, "cover_url", ""):
cover_url_from_task = pt.cover_url
logger.info(
"[封面生成] 封面(步骤B-final): plan_id=%s task_id=%s url=%s",
"[封面生成] 统一管道封面(步骤B-source_plan): plan_id=%s task_id=%s url=%s",
plan_id,
pt.id,
cover_url_from_task[:80],
)
break
# 再找预览
if not cover_url_from_task:
for pt in all_tasks:
if (
getattr(pt, "status", "") == "completed"
and getattr(pt, "is_preview", False)
and getattr(pt, "cover_url", "")
):
cover_url_from_task = pt.cover_url
logger.info(
"[封面生成] 封面(步骤B-preview): plan_id=%s task_id=%s url=%s",
plan_id,
pt.id,
cover_url_from_task[:80],
)
break
except Exception:
logger.warning(
"[封面生成] 步骤B查找 cover_url 失败: plan_id=%s",
@@ -734,13 +402,13 @@ def generate_cover(
snapshots = mk_client.extract_frames(
video_url=primary_video_url,
strategy="SpecifiedFrames",
max_frames=5, # 抽 5 帧,通过质量评分选最佳
max_frames=1,
poll_interval=2.0,
max_poll_attempts=5,
max_retries=0,
)
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:
cover_url_from_task = _persist_cover_frame(raw, plan_id)
logger.info(
@@ -768,12 +436,7 @@ def generate_cover(
storage_svc = get_shared_storage_service()
mk_client = get_mediakit_client()
# 从 plan.config 读取完整标题样式,E2 从源素材抽帧时叠加(源素材本身无标题)
# #1901 统一读 "title",兼容老数据 "title_config"
_e2_title_cfg = (plan.config or {}).get("title", {}) or {}
if not isinstance(_e2_title_cfg, dict) or not (_e2_title_cfg.get("text") or "").strip():
_alt = (plan.config or {}).get("title_config", {}) or {}
if isinstance(_alt, dict):
_e2_title_cfg = _alt
if not isinstance(_e2_title_cfg, dict):
_e2_title_cfg = {}
_e2_title_text = (_e2_title_cfg.get("text", "") or "").strip() if _e2_title_cfg.get("enabled", True) else ""
@@ -802,13 +465,13 @@ def generate_cover(
snapshots = mk_client.extract_frames(
video_url=src_url,
strategy="SpecifiedFrames",
max_frames=5, # 抽 5 帧,通过质量评分选最佳
max_frames=1,
poll_interval=2.0,
max_poll_attempts=5,
max_retries=0,
)
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:
cover_url_from_task = _persist_cover_frame(
raw,
@@ -834,7 +497,7 @@ def generate_cover(
if cover_url_from_task:
# 标题已在预览视频渲染时烧录(ASS字幕),封面帧自然包含标题
cover_data: dict[str, object] = { # type: ignore[no-redef]
cover_data = {
"type": "ai_frame",
"image_url": cover_url_from_task,
"frame_time": 0.0,
+98 -350
View File
@@ -5,17 +5,16 @@
from __future__ import annotations
import json
import logging
from app.auth import AuthenticatedUser, get_current_user
from packages.middleware.points_gate import points_gate
from app.core.storage import get_storage_service
from app.core.task_enqueue import (
GLOBAL_PENDING_LIMIT,
USER_PENDING_LIMIT,
GlobalQueueFull,
UserPendingLimitExceeded,
build_rate_limit_detail,
safe_enqueue_generation_task,
)
from app.dependencies import (
@@ -25,7 +24,6 @@ from app.dependencies import (
get_generation_task_repository,
)
from app.schemas.generation_task import (
BatchPreviewGenerationTaskResponse,
CreatePreviewGenerationTaskRequest,
PreviewGenerationTaskResponse,
)
@@ -101,7 +99,7 @@ def _resolve_strategy_id_from_template(template_id: str, db: Session, user_id: s
try:
new_repo = SQLAlchemyEditTemplateRepository(db)
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()
if mode:
logger.info(
@@ -196,19 +194,11 @@ def _to_preview_response(task, generated_videos: list | None = None) -> PreviewG
if started_at and completed_at:
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(
task_id=task.id,
status=task.status.value if hasattr(task.status, "value") else str(task.status),
progress=float(task.progress or 0.0),
is_preview=bool(getattr(task, "is_preview", True)),
variant_index=int(extra_meta.get("variant_index", 0) or 0),
resolution=getattr(task, "resolution", "") or "",
video_url=video_url,
duration=duration,
@@ -217,8 +207,6 @@ def _to_preview_response(task, generated_videos: list | None = None) -> PreviewG
transition_count=transition_count,
material_usage=material_usage,
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,
started_at=started_at,
finished_at=completed_at,
@@ -226,101 +214,50 @@ def _to_preview_response(task, generated_videos: list | None = None) -> PreviewG
)
def _resolve_preview_edit_plan_id(
*,
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)
@points_gate("ai_video", quantity_field="preview_count")
@router.post("/preview", response_model=PreviewGenerationTaskResponse, status_code=201)
def create_preview_generation_task(
request: CreatePreviewGenerationTaskRequest,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
generation_task_repository=Depends(get_generation_task_repository),
db: Session = Depends(get_db_session),
asset_repo=Depends(get_asset_repository),
) -> BatchPreviewGenerationTaskResponse:
"""创建预览生成任务(支持批量)
) -> PreviewGenerationTaskResponse:
"""创建预览生成任务。
preview_count=1 时行为与旧版完全一致(创建 1 个任务);
preview_count=N 时一次创建 N 个独立变体任务:
- 每个变体克隆独立编辑计划(独立 clips、独立随机素材起点),N 个预览内容互不相同
- 每个变体拥有独立 task_id / 状态 / 预览视频 URL,前端按 task_id 分别轮询
- 标题样式(font/color/position 等)全局共用;标题文字/配音/封面可按变体独立
titles[] / voice_library_ids[] / cover_urls[],长度1=共用,长度N=独立)
预览渲染品质与正式生成一致(1080p, CRF 23, medium preset),确认生成时可直接复用预览产物。
Args:
request: 预览任务创建请求(template_id + asset_ids 等)
Returns:
201 + 变体任务数组 {items: [...], total: N}
201 + 预览任务详情
"""
user_id = authenticated_user.user.id
count = max(1, request.preview_count)
logger.info(
"[预览生成] 接收请求: user_id=%s, template_id=%s, asset_count=%d, preview_count=%d",
user_id,
request.template_id,
len(request.asset_ids),
count,
request.preview_count,
)
# 预检查队列限流(按变体总数计)
# 预检查队列限流
try:
user_pending = generation_task_repository.count_pending_by_user(user_id)
global_pending = generation_task_repository.count_pending_total()
if user_pending + count > USER_PENDING_LIMIT:
raise UserPendingLimitExceeded(
user_id=user_id, pending_count=user_pending + count, limit=USER_PENDING_LIMIT
)
if global_pending + count > GLOBAL_PENDING_LIMIT:
raise GlobalQueueFull(pending_count=global_pending + count, limit=GLOBAL_PENDING_LIMIT)
if user_pending + 1 > USER_PENDING_LIMIT:
raise UserPendingLimitExceeded(user_id=user_id, pending_count=user_pending + 1, limit=USER_PENDING_LIMIT)
if global_pending + 1 > GLOBAL_PENDING_LIMIT:
raise GlobalQueueFull(pending_count=global_pending + 1, limit=GLOBAL_PENDING_LIMIT)
except UserPendingLimitExceeded as e:
raise HTTPException(
status_code=429,
detail=build_rate_limit_detail(e, generation_task_repository, scope="user"),
detail=f"您的待处理任务过多(当前 {e.pending_count - 1}/{e.limit}),请等待后再提交",
) from e
except GlobalQueueFull as e:
raise HTTPException(
status_code=503,
detail=build_rate_limit_detail(e, generation_task_repository, scope="global"),
detail="系统繁忙,请稍后再试",
) from e
# 确定视频比例:优先前端传入,否则从模板 mode 推断
@@ -328,92 +265,49 @@ def create_preview_generation_task(
if not video_ratio and request.template_id:
video_ratio = _infer_video_ratio_from_template(request.template_id, db, user_id)
# 根据 video_ratio 计算输出分辨率(默认竖屏 1080x1920
output_width, output_height = 1080, 1920
if video_ratio:
parts = video_ratio.split(":")
if len(parts) == 2:
try:
w, h = int(parts[0]), int(parts[1])
base = 1920
if w < h:
output_width = round(base * w / h)
output_height = base
else:
output_width = base
output_height = round(base * h / w)
output_width = output_width - output_width % 2
output_height = output_height - output_height % 2
except (ValueError, ZeroDivisionError):
output_width, output_height = 1080, 1920
resolution = f"{output_width}x{output_height}"
logger.info(
"[预览生成] 分辨率: video_ratio=%s%s (%dx%d)",
video_ratio,
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)
base_title_config = request.title_config or {}
# 处理标题配置:如果有标题文本,序列化到 custom_title 字段传递给 worker
title_config = request.title_config or {}
title_text = (title_config.get("text") or "").strip()
custom_title_value = ""
if title_text:
# 将标题文本和样式配置序列化为 JSON 存入 custom_title
# Worker 端会解析 JSON 获取完整标题配置
custom_title_value = json.dumps(title_config, ensure_ascii=False)
logger.info(
"[预览生成] 标题配置: text=%s, config_keys=%s",
title_text[:30],
list(title_config.keys()),
)
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:
for variant_index in range(count):
# 变体独立标题文字:titles[] 覆盖 title_config.text
variant_title_text = _variant_value(request.titles, variant_index, "")
variant_title_config = dict(base_title_config)
if variant_title_text.strip():
variant_title_config["text"] = variant_title_text.strip()
# 变体独立配音
variant_voice_library_id = _variant_value(
request.voice_library_ids, variant_index, request.voice_library_id
task = use_case.execute(
CreateGenerationTaskCommand(
project_id="",
asset_library_id="",
strategy_id=strategy_id,
voice_library_id=request.voice_library_id,
template_id=request.template_id,
asset_ids=list(request.asset_ids),
title_ids=list(request.title_ids),
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="",
bgm_config=request.bgm_config or {},
auto_retry_enabled=False,
auto_retry_max=0,
is_preview=True,
custom_title=custom_title_value,
)
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:
logger.warning("[预览生成] 创建失败: %s", e)
raise HTTPException(status_code=400, detail=str(e)) from e
@@ -421,204 +315,58 @@ def create_preview_generation_task(
logger.error("[预览生成] 创建失败: %s", e, exc_info=True)
raise HTTPException(status_code=500, detail="创建预览生成任务失败,请稍后再试") from e
# ── 独立变体 plan(#1743)──
# count=1:克隆源 plan(预览不污染源 plan,仅起点重算),行为与旧版一致;
# 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]:
# 关联编辑计划:如果前端未传 source_edit_plan_id,通过 template_id + user_id 查找
if not task.source_edit_plan_id and request.template_id:
try:
from packages.domain.variant_voice_resolver import resolve_variant_voice_ids
voices = resolve_variant_voice_ids(
count=count,
voice_library_id=request.voice_library_id,
voice_library_ids=request.voice_library_ids or None,
from packages.adapters.sqlalchemy_impl.edit_plan_repository import (
SQLAlchemyEditPlanRepository,
)
except Exception:
logger.warning("[预览生成] 配音解析失败(按无配音处理)", exc_info=True)
return [0.0] * count
try:
from app.api.routes.generation_tasks import _query_voice_durations
return _query_voice_durations(db, voices)
except Exception:
return [0.0] * count
voice_durations = _preview_voice_durations()
if source_plan_id and count == 1:
# 单预览:克隆一份(原逻辑)+ 配音时长分配
try:
from app.services.edit_plan_service import EditPlanService
_plan_svc = EditPlanService(db)
variant_plan = _plan_svc.clone_plan_for_variant(
source_plan_id,
created_by_user_id=user_id,
name_suffix="预览变体",
)
if voice_durations and voice_durations[0] > 0:
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,
_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,
)
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)
# 回写变体标题到 plan configworker 渲染时从 plan 读取 title 配置)
if task.source_edit_plan_id and (task.title_config or {}).get("text", "").strip():
try:
from app.api.routes.generation_tasks import _writeback_edit_plan_config
_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:
logger.warning("[预览生成] 任务入队失败: task_id=%s", task.id)
_mark_task_failed(generation_task_repository, task, "任务入队失败")
except UserPendingLimitExceeded as e:
_mark_task_failed(generation_task_repository, task, "待处理任务超限")
rate_limit_exc = rate_limit_exc or e
except GlobalQueueFull as e:
_mark_task_failed(generation_task_repository, task, "系统队列已满")
rate_limit_exc = rate_limit_exc or e
break
except Exception:
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(
status_code=429,
detail=build_rate_limit_detail(rate_limit_exc, generation_task_repository, scope="user"),
logger.warning(
"[预览生成] 查找关联编辑计划失败(不影响主流程): task_id=%s",
task.id,
exc_info=True,
)
# 入队执行;若入队失败则标记任务为 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(
status_code=503,
detail=build_rate_limit_detail(rate_limit_exc, generation_task_repository, scope="global"),
)
detail="系统繁忙,请稍后再试",
) from None
logger.info(
"[预览生成] 创建完成: %d 个变体任务, task_ids=%s",
len(responses),
[r.task_id for r in responses],
)
return BatchPreviewGenerationTaskResponse(items=responses, total=len(responses))
return _to_preview_response(task)
@router.get("/preview/{task_id}", response_model=PreviewGenerationTaskResponse)
+88 -418
View File
@@ -1,17 +1,16 @@
import logging
import random
import uuid
from typing import Any
from app.api.routes._helpers import check_project_access
from app.auth import AuthenticatedUser, get_current_user
from packages.middleware.points_gate import points_gate
from app.core.storage import OSSStorageService, get_storage_service
from app.core.task_enqueue import (
GLOBAL_PENDING_LIMIT,
USER_PENDING_LIMIT,
GlobalQueueFull,
UserPendingLimitExceeded,
build_rate_limit_detail,
safe_enqueue_generation_task,
)
from app.dependencies import (
@@ -49,22 +48,6 @@ logger = logging.getLogger(__name__)
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]:
"""[已下沉] 路由层兼容别名 → app.services.generation_common.query_voice_durations。"""
from app.services.generation_common import query_voice_durations
return query_voice_durations(db, voice_ids)
def _to_generation_task_response(task) -> GenerationTaskResponse:
return GenerationTaskResponse(
id=task.id,
@@ -87,6 +70,7 @@ def _to_generation_task_response(task) -> GenerationTaskResponse:
output_width=getattr(task, "output_width", 1280),
output_height=getattr(task, "output_height", 720),
cover_url=getattr(task, "cover_url", ""),
custom_title=getattr(task, "custom_title", ""),
title_config=getattr(task, "title_config", {}) or {},
logs=getattr(task, "logs", "[]"),
status=task.status,
@@ -110,9 +94,6 @@ def _to_generated_video_response(item, download_url: str | None = None) -> Gener
height=item.height,
fps=item.fps,
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),
)
@@ -131,16 +112,13 @@ def _select_assets_from_library(
assets: list,
mode: str,
count: int,
rng=None,
) -> list[str]:
"""根据选取模式从素材库中选取 ready 状态的视频素材 ID。
Args:
assets: 素材库中所有素材(Asset 实体列表)
mode: 选取模式 — all=全部, smart=智能匹配(多维度评分+多样性)
count: 选取数量,0 表示全部(仅 smart 模式有效)
rng: 可选随机源(smart 模式排序噪声用),生产环境不传则内部随机;
测试可注入固定种子或零噪声随机源获得确定性结果。
mode: 选取模式 — all=全部, random=随机, smart=智能匹配(多维度评分+多样性)
count: 选取数量,0 表示全部(仅 random/smart 模式有效)
Returns:
选中的素材 ID 列表
@@ -150,28 +128,69 @@ def _select_assets_from_library(
if not ready_video_assets:
return []
if mode == "random":
selected = (
ready_video_assets if count <= 0 else random.sample(ready_video_assets, min(count, len(ready_video_assets)))
)
return [a.id for a in selected]
if mode == "smart":
# 智能匹配:统一使用 packages/domain/smart_match.py 的多维评分+多样性选取
# 评分维度:质量分(40%) + 时长适配(30%) + 新鲜度(20%) + 未使用加分(10%)
# 排序注入随机噪声(#1743):同分素材每次选出不同组合,从素材组合层面降重
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]
# 默认 all 模式:返回全部 ready 视频素材
return [a.id for a in ready_video_assets]
def _writeback_edit_plan_config(
plan_id: str,
task_id: str,
title_config: dict | None,
db: Session,
) -> None:
"""[已下沉] 路由层兼容别名 → app.services.generation_common.writeback_edit_plan_config。"""
from app.services.generation_common import writeback_edit_plan_config
"""任务入队成功后,回写 EditPlan.configgeneration_task_id + title_config。
return writeback_edit_plan_config(plan_id, task_id, title_config, db)
用 merge 方式更新,不整体覆盖 config,避免丢失其他字段。
失败只记日志,不影响任务创建。
"""
if not plan_id:
return
try:
from packages.adapters.sqlalchemy_impl.models import EditPlanModel
plan_model = db.query(EditPlanModel).filter(EditPlanModel.id == plan_id).first()
if plan_model is None:
logger.warning("[生成任务] 回写plan.config失败: plan不存在 plan_id=%s", plan_id)
return
current_config = plan_model.config if isinstance(plan_model.config, dict) else {}
merged = dict(current_config)
merged["generation_task_id"] = task_id
if title_config:
merged["title_config"] = title_config
plan_model.config = merged
db.commit()
logger.info(
"[生成任务] 回写plan.config成功: plan_id=%s task_id=%s keys=%s",
plan_id,
task_id,
list(merged.keys()),
)
except Exception as e:
logger.warning(
"[生成任务] 回写plan.config异常(不影响任务创建): plan_id=%s error=%s",
plan_id,
e,
exc_info=True,
)
try:
db.rollback()
except Exception:
pass
def _resolve_project_and_library(
@@ -212,7 +231,6 @@ def _resolve_project_and_library(
@router.post("/tasks", response_model=BatchGenerationTaskResponse)
@points_gate("ai_video", quantity_field="count")
def create_generation_task(
request: CreateGenerationTaskRequest,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
@@ -261,8 +279,8 @@ def create_generation_task(
mode=request.asset_select_mode,
count=request.asset_select_count,
)
elif project_id and not resolved_asset_ids and request.asset_select_mode in ("smart",):
# 项目级模式:未指定 asset_ids 且选择了 smart 模式时,也自动选取
elif project_id and not resolved_asset_ids and request.asset_select_mode in ("random", "smart"):
# 项目级模式:未指定 asset_ids 且选择了 random/smart 模式时,也自动选取
assets = asset_repository.find_by_project(project_id)
if assets:
resolved_asset_ids = _select_assets_from_library(
@@ -276,92 +294,9 @@ def create_generation_task(
detail="当前项目没有符合条件的视频素材,请先上传并等待导入完成后再生成。",
)
# ── 兜底复用预览产物 ──
# 前端刷新后 previewTaskId 丢失,降级调 create 接口时,
# 如果同一 edit_plan 有已完成的预览任务,直接复用(秒出)。
if request.source_edit_plan_id and not request.is_preview:
try:
from packages.adapters.sqlalchemy_impl.models import (
GenerationTaskModel,
)
_preview_model = (
db.query(GenerationTaskModel)
.filter(
GenerationTaskModel.source_edit_plan_id == request.source_edit_plan_id,
GenerationTaskModel.is_preview.is_(True),
GenerationTaskModel.status == "completed",
GenerationTaskModel.created_by_user_id == authenticated_user.user.id,
)
.order_by(GenerationTaskModel.created_at.desc())
.first()
)
if _preview_model is not None:
# 校验分辨率一致性(与 confirm 端点逻辑相同)
req_w = request.output_width or 0
req_h = request.output_height or 0
src_w = getattr(_preview_model, "output_width", 0) or 0
src_h = getattr(_preview_model, "output_height", 0) or 0
resolution_match = (req_w == 0 or req_w == src_w) and (req_h == 0 or req_h == src_h)
if resolution_match:
from packages.adapters.sqlalchemy_impl.generation_task_repository import (
_to_domain,
)
preview_task = _to_domain(_preview_model)
# 如果传了标题,更新 title_config
fallback_title_config = None
if request.title_config and request.title_config.get("text", "").strip():
fallback_title_config = dict(preview_task.title_config or {})
fallback_title_config.update(request.title_config)
preview_task.mark_confirmed(
cover_url=request.cover_url or preview_task.cover_url,
output_width=request.output_width or preview_task.output_width,
output_height=request.output_height or preview_task.output_height,
title_config=fallback_title_config,
)
generation_task_repository.update(preview_task)
# 同步标题到 EditPlan.config
if fallback_title_config:
_writeback_edit_plan_config(
plan_id=request.source_edit_plan_id,
task_id=preview_task.id,
title_config=fallback_title_config,
db=db,
)
logger.info(
"[生成任务] 兜底复用预览产物: preview_task_id=%s, plan_id=%s",
preview_task.id,
request.source_edit_plan_id,
)
return BatchGenerationTaskResponse(
items=[_to_generation_task_response(preview_task)],
total=1,
)
else:
logger.info(
"[生成任务] 兜底复用跳过(分辨率不一致): plan_id=%s, src=%sx%s, req=%sx%s",
request.source_edit_plan_id,
src_w,
src_h,
req_w,
req_h,
)
except Exception:
logger.warning(
"[生成任务] 兜底复用预览产物异常(不影响主流程): plan_id=%s",
request.source_edit_plan_id,
exc_info=True,
)
use_case = CreateGenerationTaskUseCase(generation_task_repository)
count = request.count
created_tasks: list = []
created_tasks = []
failed_tasks = []
user_id = authenticated_user.user.id
# 同批次任务共享 batch_id,用于视频查重时批次内比对
@@ -380,12 +315,12 @@ def create_generation_task(
except UserPendingLimitExceeded as e:
raise HTTPException(
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
except GlobalQueueFull as e:
raise HTTPException(
status_code=503,
detail=build_rate_limit_detail(e, generation_task_repository, scope="global"),
detail="系统繁忙,请稍后再试",
) from e
# 画中画已下线:strategy_id 中的 pip/voice_pip 统一映射为 one_take
@@ -394,213 +329,20 @@ def create_generation_task(
logger.info("画中画已下线,strategy_id %s → one_take", effective_strategy_id)
effective_strategy_id = "one_take"
# 批量生成(count>1):每个变体必须走与单视频完全相同的独立选片流程(#1743/#1749)。
# - 变体 0clone 源 plan(不污染源 plan),变体 1..N-1 用 reselect_plan_for_variant
# 完整重跑选片(素材级去重:fresh 优先 → 受控复用 overlap≤20% → 短素材禁复用);
# - #1749:前端可回传 variant-plans 接口预生成的 plan_idvariant_plan_ids),直接复用;
# 回传 plan 仍按各变体配音幂等重分配段长(防 variant-plans 阶段未带配音/占位时长);
# - 配音时长:独立配音各自时长、统一配音同值,逐变体 apply_voice_duration_to_plan
# 成片总时长=配音时长(素材短→末帧冻结,禁慢放/禁截配音);
# - count>1 但没有源 plan 时,不允许 N 个任务兜底共用同一 plan,直接 4xx 中断。
# 在创建任何任务【之前】预生成/校验全部变体 plan:失败直接中断(此时无脏数据)。
variant_plan_ids: list[str] = []
if count > 1:
from app.services.edit_plan_service import EditPlanService
from packages.domain.variant_voice_resolver import VariantVoiceError, resolve_variant_voice_ids
_plan_svc = EditPlanService(db)
# 解析每变体配音(严格守卫:独立配音长度/缺值 → 400,禁静默 fallback
try:
variant_voices = resolve_variant_voice_ids(
count=count,
voice_library_id=request.voice_library_id,
voice_library_ids=request.voice_library_ids or None,
)
except VariantVoiceError as ve:
raise HTTPException(status_code=400, detail=str(ve)) from ve
# 各变体配音时长(查询硬化:异常 → 0.0 不阻断)
voice_durations = _query_voice_durations(db, variant_voices)
# 解析批量源 plan:优先前端传入;否则按 template_id + user 查最新(公共函数)
from app.services.generation_common import resolve_latest_plan_by_template
batch_source_plan_id = (
request.source_edit_plan_id
or resolve_latest_plan_by_template(db, template_id=request.template_id, user_id=user_id)
or ""
)
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(
status_code=400,
detail=f"variant_plan_ids 数量({len(request.variant_plan_ids)})与视频数量({count})不一致",
)
# 校验归属权
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)
# #1855 P0:批次区间避让表,从变体0实际clips构建初始值(公共函数)
from app.services.generation_common import collect_plan_segments as _collect_segments
_batch_segments = _collect_segments(_plan0.id, _plan_svc._clip_repo)
# 变体 1..N-1 独立选片(传入累积batch_segments做素材区间避让)
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,
batch_segments=_batch_segments,
)
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)
# #1855 P0:把新变体的clips区间追加到batch_segments,供下一变体避让
try:
_new_segs = _collect_segments(variant.id, _plan_svc._clip_repo)
for _aid, _ivs in _new_segs.items():
_batch_segments.setdefault(_aid, []).extend(_ivs)
except Exception:
logger.exception("[生成任务] 变体%d 区间收集失败(不阻断)", task_index)
# ③ 配音时长分配(回传 plan / clone 变体0 均需幂等分配;reselect 已在选片时分配,
# #1855apply_voice_duration_to_plan 已内置幂等判断,重复调用安全)
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
from app.services.generation_common import resolve_latest_plan_by_template
_single_plan = (
request.source_edit_plan_id
or resolve_latest_plan_by_template(db, template_id=request.template_id, user_id=user_id)
or ""
)
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:
for task_index in range(count):
# #1749count>1 时每个变体(含变体0)都关联各自独立 planclone/reselect/variant-plans)。
if count > 1 and variant_plan_ids:
effective_plan_id = variant_plan_ids[task_index]
else:
effective_plan_id = request.source_edit_plan_id
# 变体级独立配置: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)
for _ in range(count):
task = use_case.execute(
CreateGenerationTaskCommand(
project_id=project_id,
asset_library_id=asset_library_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,
asset_ids=resolved_asset_ids,
title_ids=request.title_ids,
voice_ids=request.voice_ids,
created_by_user_id=user_id,
source_edit_plan_id=effective_plan_id,
source_edit_plan_id=request.source_edit_plan_id,
asset_select_mode=request.asset_select_mode,
batch_id=batch_id,
video_title=request.video_title,
@@ -612,68 +354,12 @@ def create_generation_task(
source_task_id=request.source_task_id,
output_width=request.output_width,
output_height=request.output_height,
cover_url=variant_cover_url,
title_config=variant_title_config,
cover_url=request.cover_url,
custom_title=request.custom_title,
title_config=request.title_config or {},
)
)
# 变体序号写入 extra_meta(响应/排查时可辨识)
task.extra_meta["variant_index"] = task_index
try:
# 兜底关联编辑计划:前端未传 source_edit_plan_id 时,
# 通过 template_id + user_id 在 DB 层直接查找最新的 plan。
# 必须在 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:
try:
from packages.adapters.sqlalchemy_impl.models import EditPlanModel
_plan_model = (
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 _plan_model:
task.source_edit_plan_id = _plan_model.id
generation_task_repository.update(task)
logger.info(
"[生成任务] 自动关联编辑计划: task_id=%s plan_id=%s",
task.id,
_plan_model.id,
)
except Exception:
logger.warning(
"[生成任务] 查找关联编辑计划失败(不影响主流程): task_id=%s",
task.id,
exc_info=True,
)
# 回写 plan.config:必须在 enqueue 之前执行,
# 确保 worker 读取 plan 时 config 中已包含 generation_task_id。
# 批量场景下每个变体关联独立 plan,需各自回写自己的变体标题配置。
_effective_plan_id = task.source_edit_plan_id
if _effective_plan_id:
_writeback_edit_plan_config(
plan_id=_effective_plan_id,
task_id=task.id,
title_config=variant_title_config,
db=db,
)
if safe_enqueue_generation_task(
task,
generation_task_repository,
@@ -682,6 +368,15 @@ def create_generation_task(
log_task_status=True,
):
created_tasks.append(task)
# 只在首个成功任务时回写一次 plan.config
# 避免批量生成时循环覆盖 generation_task_id
if request.source_edit_plan_id and len(created_tasks) == 1:
_writeback_edit_plan_config(
plan_id=request.source_edit_plan_id,
task_id=task.id,
title_config=request.title_config,
db=db,
)
else:
failed_tasks.append(task)
except UserPendingLimitExceeded as _e:
@@ -690,7 +385,7 @@ def create_generation_task(
if not created_tasks:
raise HTTPException(
status_code=429,
detail=build_rate_limit_detail(_e, generation_task_repository, scope="user"),
detail="您的待处理任务过多,请等待完成后再提交",
) from _e
break
except GlobalQueueFull as _e:
@@ -698,7 +393,7 @@ def create_generation_task(
if not created_tasks:
raise HTTPException(
status_code=503,
detail=build_rate_limit_detail(_e, generation_task_repository, scope="global"),
detail="系统繁忙,请稍后再试",
) from _e
break
except HTTPException:
@@ -718,7 +413,6 @@ def confirm_generation(
authenticated_user: AuthenticatedUser = Depends(get_current_user),
generation_task_repository: Any = Depends(get_generation_task_repository),
project_repository: Any = Depends(get_project_repository),
db: Session = Depends(get_db_session),
) -> BatchGenerationTaskResponse:
"""确认生成 -- 复用预览渲染产物(预览与正式品质一致)。
@@ -747,29 +441,13 @@ def confirm_generation(
resolution_match = (req_w == 0 or req_w == src_w) and (req_h == 0 or req_h == src_h)
if resolution_match:
# 如果用户传了 custom_title,同步更新 title_config
confirmed_title_config = None
if request.custom_title and request.custom_title.strip():
confirmed_title_config = dict(getattr(source_task, "title_config", {}) or {})
confirmed_title_config["text"] = request.custom_title.strip()
source_task.mark_confirmed(
cover_url=request.cover_url,
custom_title=request.custom_title,
output_width=request.output_width,
output_height=request.output_height,
title_config=confirmed_title_config,
)
generation_task_repository.update(source_task)
# 同步标题到 EditPlan.config
if confirmed_title_config and source_task.source_edit_plan_id:
_writeback_edit_plan_config(
plan_id=source_task.source_edit_plan_id,
task_id=source_task.id,
title_config=confirmed_title_config,
db=db,
)
logger.info(
"[确认生成] 复用预览产物: task_id=%s, user_id=%s",
task_id,
@@ -800,6 +478,7 @@ def confirm_generation(
template_id=source_task.template_id,
asset_ids=source_task.asset_ids,
title_ids=source_task.title_ids,
voice_ids=source_task.voice_ids,
created_by_user_id=authenticated_user.user.id,
source_edit_plan_id=source_task.source_edit_plan_id or "",
asset_select_mode=source_task.asset_select_mode,
@@ -810,6 +489,7 @@ def confirm_generation(
output_width=request.output_width,
output_height=request.output_height,
cover_url=request.cover_url,
custom_title=request.custom_title,
)
)
@@ -823,15 +503,15 @@ def confirm_generation(
log_task_status=True,
):
logger.warning("[确认生成] 入队失败: task_id=%s", new_task.id)
except UserPendingLimitExceeded as _e:
except UserPendingLimitExceeded:
raise HTTPException(
status_code=429,
detail=build_rate_limit_detail(_e, generation_task_repository, scope="user"),
detail="您的待处理任务过多,请等待完成后再提交",
) from None
except GlobalQueueFull as _e:
except GlobalQueueFull:
raise HTTPException(
status_code=503,
detail=build_rate_limit_detail(_e, generation_task_repository, scope="global"),
detail="系统繁忙,请稍后再试",
) from None
return BatchGenerationTaskResponse(
@@ -913,24 +593,12 @@ def retry_generation_task(
if user_pending >= USER_PENDING_LIMIT:
raise HTTPException(
status_code=429,
detail=build_rate_limit_detail(
UserPendingLimitExceeded(
user_id=user_id,
pending_count=user_pending,
limit=USER_PENDING_LIMIT,
),
generation_task_repository,
scope="user",
),
detail=f"您的待处理任务过多(当前 {user_pending}/{USER_PENDING_LIMIT}),请等待完成后再提交",
)
if global_pending >= GLOBAL_PENDING_LIMIT:
raise HTTPException(
status_code=503,
detail=build_rate_limit_detail(
GlobalQueueFull(pending_count=global_pending, limit=GLOBAL_PENDING_LIMIT),
generation_task_repository,
scope="global",
),
detail="系统繁忙,请稍后再试",
)
use_case = CreateGenerationTaskUseCase(generation_task_repository)
@@ -943,6 +611,7 @@ def retry_generation_task(
template_id=task.template_id,
asset_ids=task.asset_ids,
title_ids=task.title_ids,
voice_ids=task.voice_ids,
created_by_user_id=user_id,
source_edit_plan_id=task.source_edit_plan_id or "",
asset_select_mode=getattr(task, "asset_select_mode", ""),
@@ -953,6 +622,7 @@ def retry_generation_task(
output_width=getattr(task, "output_width", 1280),
output_height=getattr(task, "output_height", 720),
cover_url=getattr(task, "cover_url", ""),
custom_title=getattr(task, "custom_title", ""),
)
)
try:
@@ -964,15 +634,15 @@ def retry_generation_task(
log_task_status=True,
):
logger.warning("[生成任务] 重试入队失败: task_id=%s", retried.id)
except UserPendingLimitExceeded as _e:
except UserPendingLimitExceeded:
raise HTTPException(
status_code=429,
detail=build_rate_limit_detail(_e, generation_task_repository, scope="user"),
detail="您的待处理任务过多,请等待完成后再提交",
) from None
except GlobalQueueFull as _e:
except GlobalQueueFull:
raise HTTPException(
status_code=503,
detail=build_rate_limit_detail(_e, generation_task_repository, scope="global"),
detail="系统繁忙,请稍后再试",
) from None
return _to_generation_task_response(retried)
@@ -1,240 +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()
def _get_or_create_default_template_id(db: Session, user_id: str) -> str | None:
"""为用户查找一个有效模板;若不存在则自动创建默认配音模板。
前端 #1911 删除了模板选择 UI,当调用方未传 template_id/source_edit_plan_id
时(如剪辑页首次进入直接选片),后端兜底查找/创建默认模板,避免 400。
Returns:
template_id(字符串);失败时返回 None。
"""
from packages.adapters.sqlalchemy_impl.models import TemplateClipConfigModel, TemplateModel
from packages.adapters.sqlalchemy_impl.template_repository import SQLAlchemyTemplateRepository
from packages.application.template.commands import CreateTemplateCommand, SegmentCommand
from packages.application.template.use_cases import CreateTemplateUseCase
# 1. 先查已有有效模板(is_active=True 且存在片段配置)
existing = (
db.query(TemplateModel)
.filter(
TemplateModel.user_id == user_id,
TemplateModel.is_active.is_(True),
)
.order_by(TemplateModel.created_at.asc())
.first()
)
if existing is not None:
# 验证该模板是否有片段配置;若没有继续尝试创建默认
has_seg = (
db.query(TemplateClipConfigModel.id).filter(TemplateClipConfigModel.template_id == existing.id).first()
)
if has_seg:
return existing.id
# 2. 无有效模板 → 自动创建默认配音模板
try:
repo = SQLAlchemyTemplateRepository(db)
cmd = CreateTemplateCommand(
user_id=user_id,
name="默认配音模板",
mode="voice_over",
category="default",
tags=[],
title_config={},
subtitle_config={},
bgm_config={},
estimated_duration=0.0,
segments=[
SegmentCommand(
segment_order=0,
duration_min=1.0,
duration_max=30.0,
material_type=None,
),
],
)
use_case = CreateTemplateUseCase(repo)
tpl = use_case.execute(cmd)
logger.info(
"[variant-plans] 自动创建默认模板: user=%s tpl=%s",
user_id,
tpl.id,
)
return tpl.id
except Exception:
logger.exception("[variant-plans] 自动创建默认模板失败: user=%s", user_id)
return None
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":
# 不再强制要求 template_id / source_edit_plan_id
# 后端在路由内会自动查找/创建默认模板兜底(#1911 后前端不再显式选模板)。
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 查最新(公共函数)
from app.services.generation_common import resolve_latest_plan_by_template
source_plan_id = request.source_edit_plan_id.strip()
template_id = request.template_id.strip()
# P0 兜底:前端 #1911 已删除模板选择 UI,调用方可能不传 template_id
# 此时自动为该用户查找/创建默认模板。
if not source_plan_id and not template_id:
template_id = _get_or_create_default_template_id(db, user_id) or ""
if not source_plan_id and template_id:
source_plan_id = resolve_latest_plan_by_template(db, template_id=template_id, user_id=user_id) or ""
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))
+3 -3
View File
@@ -1,4 +1,4 @@
from datetime import UTC, datetime
from datetime import datetime, timezone
import psycopg
import redis
@@ -13,7 +13,7 @@ router = APIRouter(tags=["Health"])
async def health_check():
return {
"status": "healthy",
"timestamp": datetime.now(UTC).isoformat(),
"timestamp": datetime.now(timezone.utc).isoformat(),
"version": settings.APP_VERSION,
}
@@ -33,7 +33,7 @@ async def startup_check():
all_ready = all(check["status"] == "healthy" for check in checks.values())
response = {
"status": "started" if all_ready else "starting",
"timestamp": datetime.now(UTC).isoformat(),
"timestamp": datetime.now(timezone.utc).isoformat(),
"checks": checks,
}
if not all_ready:
+3 -9
View File
@@ -3,7 +3,7 @@ from typing import Any
from app.core.celery_app import celery_app
from app.dependencies import get_ingest_job_repository
from app.schemas.ingest_job import IngestJobResponse, SubmitIngestJobRequest
from fastapi import APIRouter, Depends, HTTPException
from fastapi import APIRouter, Depends
from packages.application import SubmitIngestJobCommand, SubmitIngestJobUseCase
@@ -17,7 +17,7 @@ def get_ingest_job(
) -> IngestJobResponse:
job = ingest_job_repository.get(job_id)
if job is None:
raise HTTPException(status_code=404, detail=f"IngestJob {job_id} not found")
raise ValueError(f"IngestJob {job_id} not found")
return IngestJobResponse(
id=job.id,
project_id=job.project_id,
@@ -43,13 +43,7 @@ def submit_ingest_job(
)
)
celery_result = 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
celery_app.send_task("worker.ingest_asset", args=[job.id])
return IngestJobResponse(
id=job.id,
-347
View File
@@ -1,347 +0,0 @@
"""对口型 API 路由 — #1796 MediaKit 对口型, #1809 参数调整, #1845 配音前置.
接口:
POST /api/v1/lipsync/jobs 提交对口型任务(支持 TTS/直传/预合成 三种模式)
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 取消任务
POST /api/v1/lipsync/tts-preview #1845 步骤1 TTS 预合成(同步 HTTP~2-3s
"""
from __future__ import annotations
import logging
import math
from datetime import UTC
from app.auth import AuthenticatedUser, get_current_user
from app.config import settings
from app.dependencies import (
get_db_session,
get_voice_clone_profile_repository,
)
from app.schemas.lipsync import (
AiAvatarTtsPreviewRequest,
AiAvatarTtsPreviewResponse,
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
from packages.domain.points_rules import calculate_points_cost
from packages.domain.points_service import PointsService
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 解析
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),
db: Session = Depends(get_db_session),
svc: LipsyncService = Depends(_get_service),
):
user_id = current_user.user.id
# ── 积分扣点(#1895 P2) ──
_points_deducted = 0
_points_scene = "ai_digital_human"
_points_svc = PointsService() if settings.points_enabled else None
if _points_svc is not None:
# 口型同步:TTS 模式按 script_text 估时长(240字/分钟);音频直传按 audio_duration(秒→分钟)
if body.audio_url and body.audio_duration and body.audio_duration > 0:
est_minutes = max(1.0, math.ceil(body.audio_duration / 60.0))
elif body.script_text:
est_minutes = max(1.0, math.ceil(len(body.script_text) / 240))
else:
est_minutes = 1.0
_points_deducted = calculate_points_cost(
_points_scene,
is_member=getattr(current_user.user, "is_member", False),
duration_minutes=est_minutes,
member_type=getattr(current_user.user, "member_type", None),
)
_deduct_res = _points_svc.deduct_points(user_id, _points_deducted, _points_scene, db)
if not _deduct_res["success"]:
raise HTTPException(
status_code=402,
detail={
"code": "INSUFFICIENT_POINTS",
"message": f"积分不足,需要 {_points_deducted} 积分,当前余额 {_deduct_res['balance']}",
"required": _points_deducted,
"balance": _deduct_res["balance"],
},
)
"""提交对口型任务.
三种模式:
- TTS 直生(旧版/降级):传 {video_url, voice_id, script_text, speed?, emotion?}
后端 dispatch Celery 异步任务。
- 直接音频:传 {video_url, audio_url},后端同步下载+算timings+提交MediaKit。
- 预合成音频(#1845 新主路径):传 {video_url, audio_url, audio_duration, sentence_timings}
后端同步ffprobe+写入timings+直接提交MediaKit~2-3s)。
"""
try:
job = svc.create_job(
user_id=user_id,
video_url=body.video_url,
audio_url=body.audio_url,
audio_duration=body.audio_duration,
sentence_timings=body.sentence_timings,
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:
if _points_deducted > 0 and _points_svc is not None:
try:
_points_svc.refund_points(user_id, _points_deducted, _points_scene, db)
except Exception as refund_err:
logger.warning(f"对口型 ValueError 退积分异常: err={refund_err}")
raise HTTPException(status_code=400, detail=str(exc)) from exc
except MediaKitError as exc:
if _points_deducted > 0 and _points_svc is not None:
try:
_points_svc.refund_points(user_id, _points_deducted, _points_scene, db)
except Exception as refund_err:
logger.warning(f"对口型 MediaKitError 退积分异常: err={refund_err}")
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:
logger.error("创建对口型任务异常: %s", exc, exc_info=True)
if _points_deducted > 0 and _points_svc is not None:
try:
_points_svc.refund_points(user_id, _points_deducted, _points_scene, db)
except Exception as refund_err:
logger.warning(f"对口型异常退积分异常: err={refund_err}")
raise HTTPException(
status_code=400,
detail=f"创建对口型任务失败: {exc}",
) from exc
# 创建成功但状态为 failed(同步路径失败已抛异常到上面 except;此处处理 Celery 调度失败等)
# 若任务已创建且状态为 failed,退费
if _points_deducted > 0 and _points_svc is not None and getattr(job, "status", None) == "failed":
try:
_points_svc.refund_points(user_id, _points_deducted, _points_scene, db, ref_id=job.id)
except Exception as refund_err:
logger.warning(f"对口型任务失败退积分异常: job_id={job.id}, err={refund_err}")
return job
# ── POST /tts-preview — #1845 步骤1 TTS 预合成 ──────────────────────────
@router.post("/tts-preview", response_model=AiAvatarTtsPreviewResponse)
def preview_tts(
body: AiAvatarTtsPreviewRequest,
current_user: AuthenticatedUser = Depends(get_current_user),
db: Session = Depends(get_db_session),
svc: LipsyncService = Depends(_get_service),
):
user_id = current_user.user.id
# ── 积分扣点(#1895 P2) ──
_points_deducted = 0
_points_scene = "ai_digital_human"
_points_svc = PointsService() if settings.points_enabled else None
if _points_svc is not None:
est_minutes = max(1.0, math.ceil(len(body.script_text or "") / 240)) if body.script_text else 1.0
_points_deducted = calculate_points_cost(
_points_scene,
is_member=getattr(current_user.user, "is_member", False),
duration_minutes=est_minutes,
member_type=getattr(current_user.user, "member_type", None),
)
_deduct_res = _points_svc.deduct_points(user_id, _points_deducted, _points_scene, db)
if not _deduct_res["success"]:
raise HTTPException(
status_code=402,
detail={
"code": "INSUFFICIENT_POINTS",
"message": f"积分不足,需要 {_points_deducted} 积分,当前余额 {_deduct_res['balance']}",
"required": _points_deducted,
"balance": _deduct_res["balance"],
},
)
"""步骤1「生成配音」同步 TTS 预合成.
同步执行 TTS 合成 → 下载音频 → ffprobe 时长 → 句子时间戳计算,
不创建 LipsyncJob、不转存 OSS,直接返回 CosyVoice 临时 URL~24h 有效)。
耗时约 2-3 秒。
"""
try:
result = svc.preview_tts(
user_id=user_id,
voice_id=body.voice_id,
script_text=body.script_text,
speed=body.speed,
emotion=body.emotion,
)
except MediaKitError as exc:
if _points_deducted > 0 and _points_svc is not None:
try:
_points_svc.refund_points(user_id, _points_deducted, _points_scene, db)
except Exception as refund_err:
logger.warning(f"TTS 预合成 MediaKitError 退积分异常: err={refund_err}")
status_code = 400
if exc.code in ("VoiceForbidden",):
status_code = 403
elif exc.code in ("TTSNoAudio",):
status_code = 502
raise HTTPException(
status_code=status_code,
detail={
"code": exc.code,
"message": str(exc),
},
) from exc
except Exception as exc:
logger.error("TTS 预合成异常: %s", exc, exc_info=True)
if _points_deducted > 0 and _points_svc is not None:
try:
_points_svc.refund_points(user_id, _points_deducted, _points_scene, db)
except Exception as refund_err:
logger.warning(f"TTS 预合成异常退积分异常: err={refund_err}")
raise HTTPException(
status_code=400,
detail=f"TTS 合成失败: {exc}",
) from exc
return result
# ── 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),
):
"""获取对口型任务详情."""
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"):
# 三层防御 ①:如果距上次更新超过 30 秒,同步刷新一次(避免 background task
# 静默失败导致前端永远看到 running);否则挂后台异步刷新(避免阻塞轮询)。
from datetime import datetime as _dt
_now = _dt.now(UTC)
_stale = job.updated_at is None or (_now - job.updated_at).total_seconds() > 30
if _stale:
try:
refreshed = svc.refresh_job_status(job_id, current_user.user.id)
if refreshed is not None:
job = refreshed
except Exception as exc: # noqa: BLE001
logger.error("同步刷新对口型状态失败 job_id=%s err=%s", job_id, exc, exc_info=True)
background.add_task(svc.refresh_job_status, job_id, current_user.user.id)
else:
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
-321
View File
@@ -1,321 +0,0 @@
"""积分 & 会员 API 路由 (#1895)
导出两个 router
- points_router: 积分相关路由,前缀 /points
- usage_router: 每日额度路由,前缀 /usage
"""
from __future__ import annotations
import logging
from datetime import datetime
from typing import Optional
from app.auth import AuthenticatedUser, get_current_user
from app.dependencies import get_db_session
from app.schemas.points import (
DailyUsageResponse,
MembershipStatusResponse,
PointRuleItem,
PointsBalanceResponse,
PointsCheckRequest,
PointsCheckResponse,
PointsDeductRequest,
PointsOrderResponse,
PointsPackageItem,
PointsPackagesResponse,
PointsRechargeRequest,
PointsRefundRequest,
PointsRulesResponse,
PointsTransactionsResponse,
SimpleMessageResponse,
)
from fastapi import APIRouter, Depends, HTTPException, Query
from sqlalchemy.orm import Session
from packages.domain.points_rules import (
FREE_USER_MULTIPLIER,
MEMBER_DISCOUNT,
POINTS_PACKAGES,
POINTS_SCENES,
calculate_points_cost,
)
from packages.domain.points_service import PointsService
logger = logging.getLogger(__name__)
# ── 两个 router ──
points_router = APIRouter()
usage_router = APIRouter()
def _get_service() -> PointsService:
return PointsService()
def _is_member(user: AuthenticatedUser) -> bool:
"""判断用户是否为付费会员。"""
return getattr(user.user, "is_member", False)
def _member_type(user: AuthenticatedUser) -> str | None:
return getattr(user.user, "member_type", None)
# ════════════════════════════════════════════════════════════════
# 积分相关路由 (prefix=/points)
# ════════════════════════════════════════════════════════════════
@points_router.get("/balance", response_model=PointsBalanceResponse)
def get_balance(
current_user: AuthenticatedUser = Depends(get_current_user),
db: Session = Depends(get_db_session),
):
"""查询当前用户积分余额 + 会员状态。"""
svc = _get_service()
account = svc.get_or_create_account(current_user.user.id, db)
return PointsBalanceResponse(
balance=account["balance"],
total_earned=account["total_earned"],
total_spent=account["total_spent"],
is_member=_is_member(current_user),
member_type=_member_type(current_user),
member_expires_at=getattr(current_user.user, "member_expires_at", None),
)
@points_router.get("/transactions", response_model=PointsTransactionsResponse)
def get_transactions(
page: int = Query(1, ge=1),
page_size: int = Query(20, ge=1, le=100),
type: Optional[str] = Query(None, description="筛选类型: add/deduct"),
source: Optional[str] = Query(None, description="筛选来源场景"),
start_date: Optional[datetime] = Query(None),
end_date: Optional[datetime] = Query(None),
current_user: AuthenticatedUser = Depends(get_current_user),
db: Session = Depends(get_db_session),
):
"""查询积分流水(分页+筛选)。"""
svc = _get_service()
result = svc.get_transactions(
user_id=current_user.user.id,
db=db,
page=page,
page_size=page_size,
type_filter=type,
source_filter=source,
start_date=start_date,
end_date=end_date,
)
return PointsTransactionsResponse(**result)
@points_router.get("/rules", response_model=PointsRulesResponse)
def get_rules(
_current_user: AuthenticatedUser = Depends(get_current_user),
):
"""查询所有积分消耗规则。"""
rules = []
for scene_key, scene_data in POINTS_SCENES.items():
rules.append(
PointRuleItem(
scene_key=scene_key,
name=scene_data["name"],
base_points=scene_data["base_points"],
unit=scene_data["unit"],
extra_per_30s=scene_data.get("extra_per_30s"),
)
)
return PointsRulesResponse(
rules=rules,
free_user_multiplier=FREE_USER_MULTIPLIER,
)
@points_router.get("/packages", response_model=PointsPackagesResponse)
def get_packages(
current_user: AuthenticatedUser = Depends(get_current_user),
):
"""查询可购买的积分包列表。"""
packages = []
for code, pkg in POINTS_PACKAGES.items():
unit_price = f"¥{pkg['price_cents'] / 100 / pkg['points']:.3f}/积分"
packages.append(
PointsPackageItem(
code=code,
name=pkg["name"],
points=pkg["points"],
price_cents=pkg["price_cents"],
unit_price=unit_price,
)
)
mt = _member_type(current_user)
discount = MEMBER_DISCOUNT.get(mt) if mt else None
return PointsPackagesResponse(packages=packages, user_discount=discount)
@points_router.post("/check", response_model=PointsCheckResponse)
def check_points(
body: PointsCheckRequest,
current_user: AuthenticatedUser = Depends(get_current_user),
db: Session = Depends(get_db_session),
):
"""消费前检查余额是否足够。"""
is_mem = _is_member(current_user)
mt = _member_type(current_user)
# 混剪场景先检查免费额度
is_free_quota = False
if body.scene_key == "ai_video" and not is_mem:
svc = _get_service()
if svc.check_daily_free_clip(current_user.user.id, db):
is_free_quota = True
required = calculate_points_cost(
body.scene_key,
is_mem,
quantity=body.quantity or 1,
duration_minutes=body.duration_minutes or 0,
member_type=mt,
)
svc = _get_service()
account = svc.get_or_create_account(current_user.user.id, db)
balance = account["balance"]
return PointsCheckResponse(
allowed=is_free_quota or balance >= required,
required_points=required,
current_balance=balance,
remaining_after=balance - required,
is_free_quota=is_free_quota,
)
@points_router.post("/deduct", response_model=SimpleMessageResponse)
def deduct_points(
body: PointsDeductRequest,
current_user: AuthenticatedUser = Depends(get_current_user),
db: Session = Depends(get_db_session),
):
"""积分扣减(内部服务调用)。"""
svc = _get_service()
result = svc.deduct_points(
user_id=current_user.user.id,
amount=body.amount,
source=body.scene_key,
db=db,
description=body.description or "",
ref_id=body.ref_id or "",
)
if not result["success"]:
raise HTTPException(
status_code=402,
detail={
"code": "INSUFFICIENT_POINTS",
"message": f"积分不足,需要 {body.amount},余额 {result['balance']}",
},
)
return SimpleMessageResponse(
success=True,
message=f"扣减 {body.amount} 积分成功",
data={"transaction_id": result["transaction_id"], "balance": result["balance"]},
)
@points_router.post("/refund", response_model=SimpleMessageResponse)
def refund_points(
body: PointsRefundRequest,
current_user: AuthenticatedUser = Depends(get_current_user),
db: Session = Depends(get_db_session),
):
"""积分退还(内部服务调用)。"""
from packages.adapters.sqlalchemy_impl.models import PointsTransactionModel
txn = (
db.query(PointsTransactionModel)
.filter(PointsTransactionModel.id == body.transaction_id)
.first()
)
if txn is None:
raise HTTPException(status_code=404, detail="交易记录不存在")
if txn.user_id != current_user.user.id:
raise HTTPException(status_code=403, detail="无权退还他人积分")
svc = _get_service()
result = svc.refund_points(
user_id=current_user.user.id,
amount=txn.amount,
source=txn.source,
db=db,
ref_id=body.transaction_id,
description=body.reason or f"退还: {txn.description}",
)
if not result["success"]:
raise HTTPException(status_code=500, detail="退还失败")
return SimpleMessageResponse(
success=True,
message=f"退还 {txn.amount} 积分成功",
data={"transaction_id": result["transaction_id"], "balance": result["balance"]},
)
@points_router.post("/recharge", response_model=PointsOrderResponse)
def create_recharge_order(
body: PointsRechargeRequest,
current_user: AuthenticatedUser = Depends(get_current_user),
db: Session = Depends(get_db_session),
):
"""创建积分充值订单。"""
svc = _get_service()
try:
order = svc.create_order(
user_id=current_user.user.id,
order_type="points",
product_code=body.package_id,
db=db,
)
except ValueError as e:
raise HTTPException(status_code=400, detail=str(e)) from None
return PointsOrderResponse(**order)
@points_router.get("/subscription/membership", response_model=MembershipStatusResponse)
def get_membership_status(
current_user: AuthenticatedUser = Depends(get_current_user),
db: Session = Depends(get_db_session),
):
"""获取当前用户会员状态(聚合信息)。"""
svc = _get_service()
account = svc.get_or_create_account(current_user.user.id, db)
is_mem = _is_member(current_user)
max_resolution = "1080p" if is_mem else "720p"
return MembershipStatusResponse(
is_member=is_mem,
member_type=_member_type(current_user),
member_expires_at=getattr(current_user.user, "member_expires_at", None),
points_balance=account["balance"],
max_resolution=max_resolution,
)
# ════════════════════════════════════════════════════════════════
# 每日额度路由 (prefix=/usage)
# ════════════════════════════════════════════════════════════════
@usage_router.get("/daily", response_model=DailyUsageResponse)
def get_daily_usage(
current_user: AuthenticatedUser = Depends(get_current_user),
db: Session = Depends(get_db_session),
):
"""查询今日免费混剪额度使用情况。"""
svc = _get_service()
result = svc.get_daily_usage(current_user.user.id, db)
return DailyUsageResponse(**result)
# 为了向后兼容,也导出一个不带后缀的 router(方便旧引用)
router = points_router
+1 -41
View File
@@ -1,14 +1,13 @@
from typing import Any
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 (
CreateProjectRequest,
ListProjectsResponse,
ProjectResponse,
)
from fastapi import APIRouter, Depends, HTTPException, Response, status
from pydantic import BaseModel
from packages.application import (
CreateProjectCommand,
@@ -17,20 +16,10 @@ from packages.application import (
GetProjectUseCase,
ListProjectsUseCase,
)
from packages.domain import AssetLibraryKind
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:
return ProjectResponse(
id=item.id,
@@ -83,35 +72,6 @@ def create_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)
def delete_project(
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
-243
View File
@@ -1,243 +0,0 @@
"""Scripts AI 能力路由 — Issue #1893.
三个 AI 工具接口(均挂载在 /api/v1/scripts 前缀下):
- POST /extract-from-douyin 从抖音视频提取文案(yt-dlp 下载 + ASR 转写)
- POST /ai-rewrite AI 文案改写(复用豆包 LLM)
- POST /ai-generate-titles AI 标题生成(复用 generate_smart_titles
"""
from __future__ import annotations
import logging
import re
import tempfile
from app.auth import AuthenticatedUser, get_current_user
from app.dependencies import get_db_session
from app.schemas.scripts_ai import (
AiGenerateTitlesRequest,
AiGenerateTitlesResponse,
AiRewriteRequest,
AiRewriteResponse,
ExtractFromDouyinRequest,
ExtractFromDouyinResponse,
)
from app.services.script_asr_service import (
ASRNotConfiguredError,
ASRTranscriptionError,
transcribe_to_text,
)
from fastapi import APIRouter, Depends, HTTPException, status
from sqlalchemy.orm import Session
from packages.middleware.points_gate import points_gate
from packages.shared.ai_client import get_doubao_client
logger = logging.getLogger(__name__)
router = APIRouter()
# 抖音 URL 校验:支持短链 v.douyin.com 和长链 www.douyin.com/video/
_DOUYIN_URL_RE = re.compile(
r"^(https?://)?(v\.douyin\.com/\S+|www\.douyin\.com/video/\S+)$",
re.IGNORECASE,
)
def _validate_douyin_url(url: str) -> None:
"""校验抖音 URL 格式,不合法时抛 HTTPException(400)."""
if not url or not url.strip():
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="链接不能为空",
)
if not _DOUYIN_URL_RE.match(url.strip()):
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="无效的抖音链接,仅支持 v.douyin.com 短链或 www.douyin.com/video/ 长链",
)
# ── 1. 从抖音视频提取文案 ─────────────────────────────────────────────────────
@router.post(
"/extract-from-douyin",
response_model=ExtractFromDouyinResponse,
)
@points_gate("douyin_extract")
def extract_from_douyin(
request: ExtractFromDouyinRequest,
current_user: AuthenticatedUser = Depends(get_current_user),
db: Session = Depends(get_db_session),
) -> ExtractFromDouyinResponse:
"""从抖音视频下载无水印视频并通过 ASR 提取文案."""
source_url = request.url.strip()
_validate_douyin_url(source_url)
# 确保 URL 有 schemeyt-dlp 需要完整 URL
url_for_download = source_url
if not re.match(r"^https?://", url_for_download, re.IGNORECASE):
url_for_download = "https://" + url_for_download
# 使用临时目录下载视频,退出时自动清理
try:
with tempfile.TemporaryDirectory(prefix="douyin_extract_") as temp_dir:
import yt_dlp
ydl_opts = {
"format": "best[ext=mp4]/best",
"outtmpl": f"{temp_dir}/%(id)s.%(ext)s",
"quiet": True,
"no_warnings": True,
"noplaylist": True,
}
try:
ydl = yt_dlp.YoutubeDL(ydl_opts)
info = ydl.extract_info(url_for_download, download=True)
except Exception as exc:
logger.error("抖音视频下载失败: url=%s error=%s", source_url, exc)
raise HTTPException(
status_code=status.HTTP_502_BAD_GATEWAY,
detail=f"视频下载失败: {exc}",
) from exc
if info is None:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="无法解析该抖音链接",
)
video_path = ydl.prepare_filename(info)
duration = float(info.get("duration") or 0)
# ASR 转写
try:
text = transcribe_to_text(video_path)
except ASRNotConfiguredError as exc:
raise HTTPException(
status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
detail=str(exc),
) from exc
except ASRTranscriptionError as exc:
raise HTTPException(
status_code=status.HTTP_502_BAD_GATEWAY,
detail=str(exc),
) from exc
except HTTPException:
raise
return ExtractFromDouyinResponse(
text=text,
duration_seconds=duration,
source_url=source_url,
)
# ── 2. AI 文案改写 ───────────────────────────────────────────────────────────
@router.post(
"/ai-rewrite",
response_model=AiRewriteResponse,
)
@points_gate("ai_rewrite")
def ai_rewrite(
request: AiRewriteRequest,
current_user: AuthenticatedUser = Depends(get_current_user),
db: Session = Depends(get_db_session),
) -> AiRewriteResponse:
"""使用豆包大模型改写文案."""
content = (request.content or "").strip()
if not content:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="文案内容不能为空",
)
style = request.style or "口语化"
client = get_doubao_client()
if not client.is_available:
raise HTTPException(
status_code=status.HTTP_502_BAD_GATEWAY,
detail="AI 服务不可用,请联系管理员配置豆包大模型 API Key",
)
system_prompt = (
"你是一个专业的短视频文案改写专家。请对以下文案进行改写,"
"要求:保留原意、口语化、适合短视频口播、调整语序避免查重。"
)
if style:
system_prompt += f"\n风格要求:{style}"
user_prompt = f"请改写以下文案:\n\n{content}"
messages = [
{"role": "system", "content": system_prompt},
{"role": "user", "content": user_prompt},
]
try:
rewritten = client.chat_completion(
messages=messages,
temperature=0.8,
max_tokens=2048,
)
except Exception as exc:
logger.error("AI 改写调用失败: %s", exc)
raise HTTPException(
status_code=status.HTTP_502_BAD_GATEWAY,
detail=f"AI 改写失败: {exc}",
) from exc
if not rewritten:
raise HTTPException(
status_code=status.HTTP_502_BAD_GATEWAY,
detail="AI 改写未返回有效结果",
)
return AiRewriteResponse(
original=content,
rewritten=rewritten.strip(),
style=style,
)
# ── 3. AI 标题生成 ───────────────────────────────────────────────────────────
@router.post(
"/ai-generate-titles",
response_model=AiGenerateTitlesResponse,
)
@points_gate("ai_title")
def ai_generate_titles(
request: AiGenerateTitlesRequest,
current_user: AuthenticatedUser = Depends(get_current_user),
db: Session = Depends(get_db_session),
) -> AiGenerateTitlesResponse:
"""使用现有 generate_smart_titles 生成标题."""
content = (request.content or "").strip()
if not content:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="文案内容不能为空",
)
# count 限制在 1-5Pydantic ge=1 le=5 已校验),但为兼容直接调用场景截断
count = max(1, min(5, request.count))
from app.services.ai_service import generate_smart_titles
result = generate_smart_titles(
description=content,
style="viral",
count=count,
)
titles = result.get("titles", [])[:count]
return AiGenerateTitlesResponse(titles=titles)
+6 -5
View File
@@ -4,7 +4,8 @@ from __future__ import annotations
import logging
from dataclasses import replace
from datetime import UTC, datetime
from datetime import datetime, timezone
from typing import List
from app.auth import AuthenticatedUser, get_current_user
from app.dependencies import get_user_repository
@@ -66,7 +67,7 @@ def _get_plan_price(plan_id: str, billing_cycle: str) -> float:
def _build_subscription_info(user: AuthenticatedUser) -> SubscriptionInfo:
"""构建订阅信息响应"""
now = datetime.now(UTC)
now = datetime.now(timezone.utc)
if user.user.subscription_expires_at:
period_end = user.user.subscription_expires_at.isoformat()
period_start = now.isoformat()
@@ -99,10 +100,10 @@ async def get_current_subscription(
return _build_subscription_info(current_user)
@router.get("/billing-records", response_model=list[BillingRecord])
@router.get("/billing-records", response_model=List[BillingRecord])
async def get_billing_records(
current_user: AuthenticatedUser = Depends(get_current_user),
) -> list[BillingRecord]:
) -> List[BillingRecord]:
"""获取账单记录列表"""
from packages.adapters.sqlalchemy_impl.billing_repository import SQLAlchemyBillingRepository
from packages.adapters.sqlalchemy_impl.session import SessionLocal
@@ -250,7 +251,7 @@ async def payment_callback(
# 计算到期时间
days = 365 if billing_cycle == "yearly" else 30
expires_at = datetime.now(UTC) + timedelta(days=days)
expires_at = datetime.now(timezone.utc) + timedelta(days=days)
repo.update_subscription_on_payment(user_id, plan, expires_at)
return {"success": True, "message": "支付成功", "record_id": record_id}
+1 -7
View File
@@ -375,13 +375,7 @@ def retry_project_task(
storage_key=job.storage_key,
)
)
celery_result = 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
celery_app.send_task("worker.ingest_asset", args=[retried.id])
return ProjectTaskResponse(
id=f"ingest:{retried.id}",
task_type="ingest",
+381 -39
View File
@@ -1,13 +1,4 @@
"""Template 列表路由(供生成页自动选模板).
保留:
- GET /templates:列表查询(生成页使用)
- 默认模板自动创建兜底逻辑(复用 _default_template.get_or_create_default_template_id
其他模板 CRUD / 分类 / 标签 / 收藏 / 复制 / 校验 / 使用统计等 HTTP 端点
已在 PR#1918 中删除(前端 PR#1911 已删除 my-templates / editing-planner /
templates 管理页面)。
"""
"""Template CRUD + generate + category routes."""
from __future__ import annotations
@@ -16,19 +7,53 @@ import logging
from app.auth import AuthenticatedUser, get_current_user
from app.dependencies import get_db_session
from app.schemas.template import (
CategoryResponse,
CopyTemplateRequest,
CreateCategoryRequest,
CreateTemplateRequest,
GenerateWarningResponse,
ListCategoriesResponse,
ListTagsResponse,
ListTemplatesResponse,
SegmentResponse,
TemplateResponse,
TemplateUsageResponse,
ToggleFavoriteResponse,
UpdateTemplateRequest,
ValidateTemplateRequest,
ValidateTemplateResponse,
)
from fastapi import APIRouter, Depends, Query
from fastapi import APIRouter, Depends, HTTPException, Query, Response, status
from sqlalchemy.orm import Session
logger = logging.getLogger(__name__)
from packages.adapters.sqlalchemy_impl.template_repository import SQLAlchemyTemplateRepository
from packages.application.template.commands import ListTemplatesFilter
from packages.application.template.use_cases import CountTemplatesUseCase, ListTemplatesUseCase
from ._default_template import get_or_create_default_template_id
from packages.application.template.commands import (
CopyTemplateCommand,
CreateCategoryCommand,
CreateTemplateCommand,
ListTemplatesFilter,
SegmentCommand,
UpdateTemplateCommand,
ValidateTemplateCommand,
)
from packages.application.template.use_cases import (
CopyTemplateUseCase,
CountTemplatesUseCase,
CreateCategoryUseCase,
CreateTemplateUseCase,
DeleteCategoryUseCase,
DeleteTemplateUseCase,
GetTemplateUseCase,
ListCategoriesUseCase,
ListTagsUseCase,
ListTemplatesUseCase,
NotFoundError,
UpdateTemplateUseCase,
ValidateTemplateUseCase,
ValidationError,
)
router = APIRouter()
@@ -37,32 +62,349 @@ def _get_template_repository(session: Session = Depends(get_db_session)) -> SQLA
return SQLAlchemyTemplateRepository(session)
@router.get("", response_model=ListTemplatesResponse, summary="获取模板列表")
def _segment_to_response(seg) -> SegmentResponse:
return SegmentResponse(
id=seg.id,
template_id=seg.template_id,
segment_order=seg.segment_order,
duration_min=seg.duration_min,
duration_max=seg.duration_max,
material_type=seg.material_type,
created_at=seg.created_at,
updated_at=seg.updated_at,
)
def _to_response(template, usage_count: int = 0) -> TemplateResponse:
return TemplateResponse(
id=template.id,
user_id=template.user_id,
name=template.name,
mode=template.mode,
category=template.category,
tags=template.tags,
title_config=template.title_config,
subtitle_config=template.subtitle_config,
bgm_config=template.bgm_config,
estimated_duration=template.estimated_duration,
segments=[_segment_to_response(s) for s in getattr(template, "segments", [])],
is_active=template.is_active,
usage_count=usage_count,
created_at=template.created_at,
updated_at=template.updated_at,
)
# ── Template CRUD ──
@router.get("", response_model=ListTemplatesResponse)
def list_templates(
mode: str | None = Query(None, description="编辑模式:generic/vlog/storyboard,不传返回全部"),
skip: int = Query(0, ge=0),
limit: int = Query(50, ge=1, le=200),
category: str | None = Query(None, description="按分类筛选"),
tag: str | None = Query(None, description="按标签筛选"),
page: int = Query(1, ge=1, description="页码,从 1 开始"),
page_size: int = Query(20, ge=1, le=100, description="每页条数,默认 20"),
current_user: AuthenticatedUser = Depends(get_current_user),
repo: SQLAlchemyTemplateRepository = Depends(_get_template_repository),
db: Session = Depends(get_db_session),
):
"""获取用户可用的模板列表(仅返回 active 状态)。"""
user_id = str(current_user.user.id)
# P0 兜底:无有效模板时自动创建默认配音模板(解决新用户首次进入生成页 404)
get_or_create_default_template_id(db, user_id)
keyword: str | None = Query(None, description="按名称关键词搜索"),
mode: str | None = Query(None, description="按剪辑模式筛选"),
authenticated_user: AuthenticatedUser = Depends(get_current_user),
template_repository: SQLAlchemyTemplateRepository = Depends(_get_template_repository),
) -> ListTemplatesResponse:
user_id = authenticated_user.user.id
try:
tpl_filter = ListTemplatesFilter(
category=category,
tag=tag,
keyword=keyword,
mode=mode,
)
use_case = ListTemplatesUseCase(template_repository)
templates = use_case.execute(user_id, skip=skip, limit=limit, filter=tpl_filter)
count_use_case = CountTemplatesUseCase(template_repository)
total = count_use_case.execute(user_id, filter=tpl_filter)
list_uc = ListTemplatesUseCase(repo)
count_uc = CountTemplatesUseCase(repo)
filters = ListTemplatesFilter(
category=category,
tag=tag,
mode=mode,
valid_only=True, # 仅返回 active + 有片段配置
# 批量查询使用次数
items = []
for t in templates:
usage = template_repository.get_usage_count(t.id)
items.append(_to_response(t, usage_count=usage))
except Exception:
logger.exception("list_templates 查询失败: user_id=%s", user_id)
return ListTemplatesResponse(items=[], total=0)
return ListTemplatesResponse(
items=items,
total=total,
)
skip = (page - 1) * page_size
templates = list_uc.execute(user_id, skip=skip, limit=page_size, filter=filters)
total = count_uc.execute(user_id, filter=filters)
items = [TemplateResponse.model_validate(tpl, from_attributes=True) for tpl in templates]
return ListTemplatesResponse(items=items, total=total)
@router.get("/{template_id}", response_model=TemplateResponse)
def get_template(
template_id: str,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
template_repository: SQLAlchemyTemplateRepository = Depends(_get_template_repository),
) -> TemplateResponse:
user_id = authenticated_user.user.id
try:
use_case = GetTemplateUseCase(template_repository)
template = use_case.execute(template_id, user_id)
usage = template_repository.get_usage_count(template_id)
except Exception as _e:
logger.exception("get_template 查询失败: template_id=%s", template_id)
raise HTTPException(status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail="模板查询失败") from _e
if template is None:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Template not found")
return _to_response(template, usage_count=usage)
@router.post("", response_model=TemplateResponse, status_code=status.HTTP_201_CREATED)
def create_template(
request: CreateTemplateRequest,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
template_repository: SQLAlchemyTemplateRepository = Depends(_get_template_repository),
) -> TemplateResponse:
user_id = authenticated_user.user.id
command = CreateTemplateCommand(
user_id=user_id,
name=request.name,
mode=request.mode,
category=request.category,
tags=request.tags,
title_config=request.title_config,
subtitle_config=request.subtitle_config,
bgm_config=request.bgm_config,
estimated_duration=request.estimated_duration,
segments=[
SegmentCommand(
segment_order=s.segment_order,
duration_min=s.duration_min,
duration_max=s.duration_max,
material_type=s.material_type,
)
for s in request.segments
],
)
use_case = CreateTemplateUseCase(template_repository)
try:
template = use_case.execute(command)
except ValidationError as exc:
raise HTTPException(status_code=status.HTTP_422_UNPROCESSABLE_ENTITY, detail=str(exc)) from exc
return _to_response(template)
@router.patch("/{template_id}", response_model=TemplateResponse)
def update_template(
template_id: str,
request: UpdateTemplateRequest,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
template_repository: SQLAlchemyTemplateRepository = Depends(_get_template_repository),
) -> TemplateResponse:
user_id = authenticated_user.user.id
command = UpdateTemplateCommand(
template_id=template_id,
user_id=user_id,
name=request.name,
mode=request.mode,
category=request.category,
tags=request.tags,
title_config=request.title_config,
subtitle_config=request.subtitle_config,
bgm_config=request.bgm_config,
estimated_duration=request.estimated_duration,
segments=(
[
SegmentCommand(
segment_order=s.segment_order,
duration_min=s.duration_min,
duration_max=s.duration_max,
material_type=s.material_type,
)
for s in request.segments
]
if request.segments is not None
else None
),
)
use_case = UpdateTemplateUseCase(template_repository)
try:
template = use_case.execute(command)
except NotFoundError as _e:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Template not found") from _e
except ValidationError as exc:
raise HTTPException(status_code=status.HTTP_422_UNPROCESSABLE_ENTITY, detail=str(exc)) from exc
return _to_response(template)
@router.delete("/{template_id}", status_code=status.HTTP_204_NO_CONTENT, response_model=None, response_class=Response)
def delete_template(
template_id: str,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
template_repository: SQLAlchemyTemplateRepository = Depends(_get_template_repository),
) -> Response:
user_id = authenticated_user.user.id
use_case = DeleteTemplateUseCase(template_repository)
deleted = use_case.execute(template_id, user_id)
if not deleted:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Template not found")
return
@router.post("/{template_id}/copy", response_model=TemplateResponse, status_code=status.HTTP_201_CREATED)
def copy_template(
template_id: str,
request: CopyTemplateRequest,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
template_repository: SQLAlchemyTemplateRepository = Depends(_get_template_repository),
) -> TemplateResponse:
"""复制模板(含所有片段配置)"""
user_id = authenticated_user.user.id
command = CopyTemplateCommand(
template_id=template_id,
user_id=user_id,
new_name=request.new_name,
)
use_case = CopyTemplateUseCase(template_repository)
try:
template = use_case.execute(command)
except NotFoundError as _e:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Template not found") from _e
except ValidationError as exc:
raise HTTPException(status_code=status.HTTP_422_UNPROCESSABLE_ENTITY, detail=str(exc)) from exc
return _to_response(template)
@router.get("/{template_id}/usage", response_model=TemplateUsageResponse)
def get_template_usage(
template_id: str,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
template_repository: SQLAlchemyTemplateRepository = Depends(_get_template_repository),
) -> TemplateUsageResponse:
"""获取模板使用次数(关联的剪辑计划数量)"""
user_id = authenticated_user.user.id
# 鉴权:确保模板存在且属于当前用户
use_case = GetTemplateUseCase(template_repository)
template = use_case.execute(template_id, user_id)
if template is None:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Template not found")
usage = template_repository.get_usage_count(template_id)
return TemplateUsageResponse(template_id=template_id, usage_count=usage)
@router.post("/{template_id}/toggle-favorite", response_model=ToggleFavoriteResponse)
def toggle_favorite(
template_id: str,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
template_repository: SQLAlchemyTemplateRepository = Depends(_get_template_repository),
) -> ToggleFavoriteResponse:
"""切换模板收藏状态(当前为兼容端点,始终返回 false)"""
user_id = authenticated_user.user.id
use_case = GetTemplateUseCase(template_repository)
try:
template = use_case.execute(template_id, user_id)
except Exception as _e:
logger.exception("toggle_favorite 查询失败: template_id=%s", template_id)
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Template not found") from _e
if template is None:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Template not found")
return ToggleFavoriteResponse(id=template_id, is_favorite=False)
# ── Validate template ──
@router.post("/{template_id}/validate", response_model=ValidateTemplateResponse)
def validate_template(
template_id: str,
request: ValidateTemplateRequest = ValidateTemplateRequest(),
authenticated_user: AuthenticatedUser = Depends(get_current_user),
template_repository: SQLAlchemyTemplateRepository = Depends(_get_template_repository),
) -> ValidateTemplateResponse:
user_id = authenticated_user.user.id
command = ValidateTemplateCommand(
template_id=template_id,
user_id=user_id,
voiceover_duration=request.voiceover_duration,
)
use_case = ValidateTemplateUseCase(template_repository)
try:
result = use_case.execute(command)
except NotFoundError as _e:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Template not found") from _e
except ValidationError as exc:
raise HTTPException(status_code=status.HTTP_422_UNPROCESSABLE_ENTITY, detail=str(exc)) from exc
return ValidateTemplateResponse(
template=_to_response(result.template),
warnings=[GenerateWarningResponse(code=w.code, message=w.message, details=w.details) for w in result.warnings],
)
# ── Category CRUD ──
@router.get("/categories/list", response_model=ListCategoriesResponse)
def list_categories(
authenticated_user: AuthenticatedUser = Depends(get_current_user),
template_repository: SQLAlchemyTemplateRepository = Depends(_get_template_repository),
) -> ListCategoriesResponse:
user_id = authenticated_user.user.id
try:
use_case = ListCategoriesUseCase(template_repository)
categories = use_case.execute(user_id)
except Exception:
logger.exception("list_categories 查询失败: user_id=%s", user_id)
return ListCategoriesResponse(items=[])
return ListCategoriesResponse(
items=[CategoryResponse(id=c.id, user_id=c.user_id, name=c.name, created_at=c.created_at) for c in categories],
)
@router.post("/categories", response_model=CategoryResponse, status_code=status.HTTP_201_CREATED)
def create_category(
request: CreateCategoryRequest,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
template_repository: SQLAlchemyTemplateRepository = Depends(_get_template_repository),
) -> CategoryResponse:
user_id = authenticated_user.user.id
command = CreateCategoryCommand(user_id=user_id, name=request.name)
use_case = CreateCategoryUseCase(template_repository)
category = use_case.execute(command)
return CategoryResponse(
id=category.id,
user_id=category.user_id,
name=category.name,
created_at=category.created_at,
)
@router.delete(
"/categories/{category_id}", status_code=status.HTTP_204_NO_CONTENT, response_model=None, response_class=Response
)
def delete_category(
category_id: str,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
template_repository: SQLAlchemyTemplateRepository = Depends(_get_template_repository),
) -> Response:
user_id = authenticated_user.user.id
use_case = DeleteCategoryUseCase(template_repository)
deleted = use_case.execute(category_id, user_id)
if not deleted:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Category not found")
return Response(status_code=204)
# ── Tags ──
@router.get("/tags/list", response_model=ListTagsResponse)
def list_tags(
authenticated_user: AuthenticatedUser = Depends(get_current_user),
template_repository: SQLAlchemyTemplateRepository = Depends(_get_template_repository),
) -> ListTagsResponse:
"""获取用户所有模板标签(去重排序)"""
user_id = authenticated_user.user.id
try:
use_case = ListTagsUseCase(template_repository)
tags = use_case.execute(user_id)
except Exception:
logger.exception("list_tags 查询失败: user_id=%s", user_id)
return ListTagsResponse(items=[])
return ListTagsResponse(items=tags)
@@ -1,9 +1,10 @@
"""模板编辑器 API 路由包.
模块拆分
将原来 2560 行的 templates_editor.py 巨无霸拆分为 12 个模块:
- schemas.py: 所有 Pydantic model
- dependencies.py: 依赖注入
- _utils.py: 工具函数
- _fallback.py: 自动兜底逻辑
- draft.py: 草稿管理(详情/更新/发布/版本/回滚)
- clips.py: 片段管理(CRUD/分割/合并/重排/批量删除/从素材创建)
- adjustments.py: 片段调整(速度/音量/裁剪/批量调速)
@@ -12,6 +13,7 @@
- export.py: 导出配置
- subtitles.py: 字幕管理
- ai_features.py: AI 推荐
- generation.py: 生成(触发/进度/记录)
- timeline.py: 时间线
挂载路径: /api/v1/templates/{template_id}/editor/
@@ -28,10 +30,11 @@ from .adjustments import router as adjustments_router
from .ai_features import router as ai_features_router
from .bgm import router as bgm_router
from .clips import router as clips_router
from .dependencies import get_draft_plan_id, get_editor_services, resolve_draft_plan_id # noqa: F401
from .dependencies import get_draft_plan_id, get_editor_services # noqa: F401
from .draft import router as draft_router
from .effects import router as effects_router
from .export import router as export_router
from .generation import router as generation_router
from .subtitles import router as subtitles_router
from .timeline import router as timeline_router
@@ -48,6 +51,7 @@ _sub_routers = [
export_router,
subtitles_router,
ai_features_router,
generation_router,
timeline_router,
]
+227
View File
@@ -0,0 +1,227 @@
"""模板编辑器自动兜底逻辑.
generate_editor_draft 触发生成前的自动修复流程:
1. draft → editing 状态迁移
2. 无片段时从模板复制片段配置
3. 为无素材片段分配指定素材
4. 项目有素材库时自动选素材
"""
from __future__ import annotations
import logging
import random
from typing import Any
from app.services.edit_plan_service import EditPlanService
from sqlalchemy.orm import Session
from packages.adapters.sqlalchemy_impl.template_clip_config_repository import (
SQLAlchemyTemplateClipConfigRepository,
)
from packages.adapters.sqlalchemy_impl.template_repository import (
SQLAlchemyTemplateRepository,
)
from packages.domain.edit_plan import EditPlanStatus
logger = logging.getLogger(__name__)
def _auto_fallback_draft_to_editing(svc: EditPlanService, plan_id: str, plan_check) -> None:
"""自动兜底 1: draft → editing"""
if plan_check.status == EditPlanStatus.DRAFT:
logger.info("模板编辑器自动兜底: plan=%s draft→editing", plan_id)
svc.transition_status(plan_id, EditPlanStatus.EDITING)
def _auto_fallback_copy_template_clips(svc: EditPlanService, plan_id: str, plan_check, db: Session) -> None:
"""自动兜底 2: 无片段 + 有 template_id → 从模板复制片段配置"""
existing_clips = svc.count_clips(plan_id)
if existing_clips == 0 and plan_check.template_id:
logger.info(
"模板编辑器自动兜底: plan=%s 无片段,从模板 %s 复制片段配置",
plan_id,
plan_check.template_id,
)
clip_config_repo = SQLAlchemyTemplateClipConfigRepository(db)
configs = clip_config_repo.list_by_template(plan_check.template_id)
if configs:
for cfg in configs:
svc.create_clip(
plan_id=plan_id,
clip_type=cfg.clip_type.value if hasattr(cfg.clip_type, "value") else cfg.clip_type,
order=cfg.order,
template_clip_config_id=cfg.id,
duration=cfg.default_duration,
transition_effect=(
cfg.transition_effect.value
if hasattr(cfg.transition_effect, "value")
else cfg.transition_effect
),
)
logger.info(
"模板编辑器自动兜底: plan=%s 从 template_clip_configs 复制了 %d 个片段",
plan_id,
len(configs),
)
else:
tpl_repo = SQLAlchemyTemplateRepository(db)
segments = tpl_repo.list_segments(plan_check.template_id)
for seg in segments:
avg_duration = (seg.duration_min + seg.duration_max) / 2
svc.create_clip(
plan_id=plan_id,
clip_type="main",
order=seg.segment_order,
duration=avg_duration,
config={
"material_type": seg.material_type or "",
"template_segment_id": seg.id,
},
)
logger.info(
"模板编辑器自动兜底: plan=%s 从旧模板 segments 复制了 %d 个片段",
plan_id,
len(segments),
)
def _auto_fallback_assign_assets(svc: EditPlanService, plan_id: str, plan_check) -> list:
"""自动兜底 3: 为没有素材的片段分配素材。返回剩余无素材片段列表。"""
all_clips = svc.list_clips(plan_id)
clips_without_asset = [c for c in all_clips if not c.asset_id]
config_asset_ids = (plan_check.config or {}).get("asset_ids", [])
logger.info(
"模板编辑器自动兜底3 诊断: plan=%s total_clips=%d " "clips_without_asset=%d config_asset_ids=%r",
plan_id,
len(all_clips),
len(clips_without_asset),
config_asset_ids[:5] if config_asset_ids else [],
)
if clips_without_asset and config_asset_ids:
logger.info(
"模板编辑器自动兜底3: plan=%s%d 个无素材片段分配 %d 个指定素材",
plan_id,
len(clips_without_asset),
len(config_asset_ids),
)
assigned = 0
for i, clip in enumerate(clips_without_asset):
asset_idx = i % len(config_asset_ids)
try:
svc.assign_asset(clip.id, config_asset_ids[asset_idx])
assigned += 1
except Exception as exc:
logger.error(
"模板编辑器自动兜底3: plan=%s clip=%s 分配素材 %s 失败: %s",
plan_id,
clip.id,
config_asset_ids[asset_idx],
exc,
)
logger.info(
"模板编辑器自动兜底3: plan=%s 素材分配完成 assigned=%d/%d",
plan_id,
assigned,
len(clips_without_asset),
)
# 重新检查剩余无素材片段
all_clips_after = svc.list_clips(plan_id)
clips_without_asset = [c for c in all_clips_after if not c.asset_id]
if clips_without_asset:
logger.warning(
"模板编辑器自动兜底3: plan=%s 仍有 %d 个片段无素材",
plan_id,
len(clips_without_asset),
)
elif not clips_without_asset:
logger.info("模板编辑器自动兜底3: plan=%s 所有片段已有素材,跳过", plan_id)
elif not config_asset_ids:
logger.info(
"模板编辑器自动兜底3: plan=%s config.asset_ids 为空,跳过分配",
plan_id,
)
return clips_without_asset
def _auto_fallback_auto_material_mode(
svc: EditPlanService,
plan_id: str,
plan_check,
clips_without_asset: list,
asset_library_repo: Any,
asset_repo: Any,
user_id: str = "",
) -> None:
"""自动兜底 4: 自动选素材分配给无素材片段
查找策略(按优先级):
1. plan 有 project_id → 从项目素材库查找
2. plan 无 project_id 但有 user_id → 从用户上传的素材中查找
"""
if not clips_without_asset:
return
ready_videos: list = []
source_desc = ""
# 策略 1: 通过 project_id 查找项目素材库
if plan_check.project_id:
libs = asset_library_repo.find_by_project(plan_check.project_id)
video_lib = None
for lib in libs:
lib_kind = lib.kind.value if hasattr(lib.kind, "value") else lib.kind
if lib_kind == "video":
video_lib = lib
break
if video_lib:
assets = asset_repo.find_by_library(video_lib.id)
ready_videos = [
a
for a in assets
if (a.status.value if hasattr(a.status, "value") else a.status) == "ready"
and a.mime_type
and a.mime_type.startswith("video")
]
source_desc = f"素材库 {video_lib.name}"
# 策略 2: 通过 user_id 查找用户上传的素材
if not ready_videos and user_id and hasattr(asset_repo, "find_ready_videos_by_user"):
logger.info(
"模板编辑器自动兜底4: plan=%s project_id 为空,尝试通过 user_id=%s 查找素材",
plan_id,
user_id,
)
ready_videos = asset_repo.find_ready_videos_by_user(user_id)
source_desc = f"用户上传 (user_id={user_id[:8]}...)"
if not ready_videos:
logger.warning(
"模板编辑器自动兜底4: plan=%s 未找到可用素材 (project_id=%s, user_id=%s)",
plan_id,
plan_check.project_id or "(empty)",
user_id[:8] + "..." if user_id else "(empty)",
)
return
logger.info(
"模板编辑器自动兜底4: plan=%s 自动选素材分配给 %d 个无素材片段 (来源: %s, 共 %d 个)",
plan_id,
len(clips_without_asset),
source_desc,
len(ready_videos),
)
random.shuffle(ready_videos)
for i, clip in enumerate(clips_without_asset):
asset = ready_videos[i % len(ready_videos)]
svc.assign_asset(clip.id, asset.id)
logger.info(
"模板编辑器自动兜底4: plan=%s%s 分配了 %d 个素材给 %d 个片段",
plan_id,
source_desc,
len(ready_videos),
len(clips_without_asset),
)
@@ -65,8 +65,8 @@ def _build_asset_analyses(
if url:
video_urls.append(url)
valid_asset_ids.append(aid)
except Exception:
logger.exception("获取素材URL失败: asset_id=%s", aid)
except Exception as e:
logger.warning("获取素材URL失败: asset_id=%s error=%s", aid, str(e))
if not video_urls:
logger.info("无可用视频素材,跳过视频理解分析")
@@ -108,7 +108,7 @@ def _build_asset_analyses(
return analyses
except Exception as e:
logger.exception("MediaKit 视频理解异常,将降级到无分析模式: %s", e)
logger.warning("MediaKit 视频理解异常,将降级到无分析模式: %s", str(e))
return {}
@@ -177,7 +177,7 @@ def editor_ai_recommend(
try:
db.rollback()
except Exception:
logger.exception("db rollback failed in ai_recommend")
pass
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail="AI推荐结果保存失败,请稍后重试",
+35 -828
View File
@@ -15,41 +15,16 @@
from __future__ import annotations
import json
import logging
import random
import re
from app.auth import AuthenticatedUser, get_current_user
from app.core.storage import get_storage_service
from app.dependencies import get_asset_repository, get_db_session
# 默认转场时长(与 worker 端保持一致)
_DEFAULT_TRANSITION_DURATION = 0.5
from app.services.asset_segment_tracker import (
REUSE_RATIO_LIMIT,
SEGMENT_EDGE_GAP,
get_used_segments,
make_reuse_callback,
record_used_segments,
remove_used_segment,
)
from app.dependencies import get_asset_repository
from app.services.edit_plan_service import EditPlanService
from app.services.edit_template_service import EditTemplateService, TemplateNotFoundError
from fastapi import APIRouter, BackgroundTasks, Depends, HTTPException, Query, status
from sqlalchemy.orm import Session
from app.services.edit_template_service import EditTemplateService
from fastapi import APIRouter, Depends, HTTPException, Query, status
from packages.adapters.sqlalchemy_impl.asset_repository import SQLAlchemyAssetRepository
from packages.domain.plan_generator_utils import (
_calc_random_start_time,
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.shared.mediakit_client import get_mediakit_client
from .dependencies import get_draft_plan_id, get_editor_services
from .schemas import (
@@ -70,9 +45,6 @@ from .schemas import (
logger = logging.getLogger(__name__)
router = APIRouter(tags=["Template Editor"])
# 编辑器默认片段时长(秒)
_DEFAULT_EDITOR_CLIP_DURATION = 5.0
def _clip_to_response(clip, asset_url: str | None = None) -> EditorClipResponse:
"""统一构造片段响应 — 与 edit_plan_clips 表字段完全对齐"""
@@ -132,8 +104,8 @@ def _build_asset_url_map(
result: dict[str, str | None] = {}
try:
storage = get_storage_service()
except Exception as e:
logger.exception("获取存储服务失败,跳过asset_url生成: %s", e)
except Exception:
logger.warning("获取存储服务失败,跳过asset_url生成")
return {aid: None for aid in asset_ids}
# 批量查询所有 Asset(单次 SQL IN 查询,避免 N+1)
@@ -141,7 +113,7 @@ def _build_asset_url_map(
assets = asset_repo.find_by_ids(unique_ids)
asset_map = {a.id: a for a in assets}
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}
for aid in unique_ids:
@@ -156,7 +128,7 @@ def _build_asset_url_map(
continue
result[aid] = storage.get_download_url(storage_key, expires_seconds=3600)
except Exception:
logger.exception("生成素材签名URL失败: asset_id=%s", aid)
logger.warning("生成素材签名URL失败: asset_id=%s", aid, exc_info=True)
result[aid] = None
return result
@@ -183,7 +155,10 @@ def list_draft_clips(
url_map = _build_asset_url_map(asset_ids, asset_repo)
return EditorClipListResponse(
items=[_clip_to_response(c, asset_url=url_map.get(getattr(c, "asset_id", "") or "")) for c in clips],
items=[
_clip_to_response(c, asset_url=url_map.get(getattr(c, "asset_id", "") or ""))
for c in clips
],
total=total,
)
@@ -297,7 +272,9 @@ def split_draft_clip(
try:
result = plan_svc.split_clip(clip_id, body.split_time)
except ValueError as exc:
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=str(exc)) from exc
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST, detail=str(exc)
) from exc
left = result["left_clip"]
right = result["right_clip"]
asset_ids = [getattr(left, "asset_id", "") or "", getattr(right, "asset_id", "") or ""]
@@ -327,7 +304,9 @@ def merge_draft_clips(
try:
merged = plan_svc.merge_clips(body.clip_ids)
except ValueError as exc:
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=str(exc)) from exc
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST, detail=str(exc)
) from exc
asset_id = getattr(merged, "asset_id", "") or ""
url_map = _build_asset_url_map([asset_id], asset_repo) if asset_id else {}
return {
@@ -375,812 +354,40 @@ def batch_delete_editor_clips(
return ClipBatchDeleteResponse(deleted_count=deleted, plan_id=plan_id)
def _safe_segment_duration(value, default: float) -> float:
"""安全地将数据库中的时长值转换为正浮点数.
处理 None、无效类型、负数、NaN 等异常情况。
"""
if value is None:
return default
try:
result = float(value)
except (ValueError, TypeError):
return default
if result != result or result <= 0: # NaN check or non-positive
return default
return result
def _get_template_segments(
template_id: str,
user_id: str,
tpl_svc: EditTemplateService,
) -> list[tuple[int, float, float]]:
"""获取模板的片段配置(顺序、最短时长、最长时长).
单一数据源:模板主表为 ``templates``(用户自建,归属 user_id/
``edit_templates``(全局模板库),片段配置主表为 ``template_clip_configs``
(由 ``EditTemplateService.list_clip_configs_for_editor`` 统一读取)。
不再使用"新表抛异常 → 降级直查配置表 → 再降级查 segments"的异常控制流,
也不在正常请求中打印 ``ValueError: 模板不存在`` 堆栈。
Args:
template_id: 模板 ID
user_id: 当前登录用户 ID(用于归属校验)
tpl_svc: 模板编辑器服务
Returns:
[(segment_order, duration_min, duration_max), ...] 按 order 排序;
模板存在但未配置片段时返回空列表。
Raises:
TemplateNotFoundError: 模板不存在、已删除或不归属于当前用户。
"""
clip_configs = tpl_svc.list_clip_configs_for_editor(template_id, user_id)
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])
def _recommended_time_conflicts(
start: float,
duration: float,
used: list[tuple[float, float]],
edge_gap: float = SEGMENT_EDGE_GAP,
) -> bool:
"""检查推荐起始时间是否与已使用时间段冲突.
冲突检测统一加 ``edge_gap`` 秒边缘间隙:已用区间按 [s-gap, e+gap] 扩边后判定,
避免推荐片段与已用片段首尾紧贴导致画面观感重复。
"""
end = start + duration
for used_start, used_end in used:
if start < used_end + edge_gap and end > used_start - edge_gap:
return True
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(
asset_ids: list[str],
asset_repo,
) -> dict[str, float]:
"""调用 MediaKit 视频理解,获取智能选片推荐起始时间.
尝试让 MediaKit 分析视频内容,返回每个素材的推荐起始时间。
任何异常都优雅降级,返回空字典(调用方降级到随机选择)。
"""
try:
client = get_mediakit_client()
if not client.is_available:
logger.info("MediaKit 未配置,使用随机起始时间")
return {}
storage = get_storage_service()
video_urls: list[str] = []
valid_asset_ids: list[str] = []
for asset_id in asset_ids[:10]:
asset = asset_repo.get(asset_id)
if not asset or not getattr(asset, "storage_key", None):
continue
mime = getattr(asset, "mime_type", "")
if not mime.startswith("video/"):
continue
try:
url = storage.get_download_url(asset.storage_key)
if url:
video_urls.append(url)
valid_asset_ids.append(asset_id)
except Exception:
logger.exception("获取素材URL失败: asset_id=%s", asset_id)
if not video_urls:
return {}
prompt = (
"请分析每段视频,找出最精彩的5秒片段应该从哪个时间点开始。"
"考虑因素:画面清晰度、主体是否明确、是否有明显的动作或场景变化。"
"请严格以JSON数组格式返回,不要包含其他文字:"
'[{"asset_id": "素材ID", "recommended_start_time": 12.5, "reason": "原因"}]'
)
contents = client.analyze_videos(
video_urls=video_urls,
prompt=prompt,
level="Economy",
poll_interval=2.0,
max_poll_attempts=15,
)
if not contents:
logger.info("MediaKit 分析无结果,降级为随机选择")
return {}
# 按索引映射结果:contents[i] 对应 valid_asset_ids[i]
recommendations: dict[str, float] = {}
for idx, content_text in enumerate(contents):
if idx >= len(valid_asset_ids):
break
asset_id = valid_asset_ids[idx]
if not content_text:
continue
# 尝试从文本中提取 JSON
parsed = False
# 尝试直接解析
try:
data = json.loads(content_text.strip())
if isinstance(data, list) and data:
for item in data:
if isinstance(item, dict) and "recommended_start_time" in item:
recommendations[asset_id] = float(item["recommended_start_time"])
parsed = True
break
except (json.JSONDecodeError, ValueError, TypeError):
pass
# 尝试从 markdown 代码块中提取 JSON
if not parsed:
json_match = re.search(r"\[\s*(\{.*?\})\s*\]", content_text, re.DOTALL)
if json_match:
try:
item = json.loads(json_match.group(1))
if isinstance(item, dict) and "recommended_start_time" in item:
recommendations[asset_id] = float(item["recommended_start_time"])
parsed = True
except (json.JSONDecodeError, ValueError, TypeError):
pass
# 尝试正则提取
if not parsed:
time_match = re.search(r'recommended_start_time["\s:]+([\d.]+)', content_text)
if time_match:
try:
recommendations[asset_id] = float(time_match.group(1))
except (ValueError, TypeError):
pass
if recommendations:
logger.info("MediaKit 智能选片推荐: %s", recommendations)
else:
logger.info("MediaKit 结果解析失败,降级为随机选择")
return recommendations
except Exception as e:
logger.exception("MediaKit 智能选片异常,降级为随机选择: %s", e)
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)
def create_clips_from_assets_editor(
template_id: str,
body: ClipsFromAssetsRequest,
background_tasks: BackgroundTasks,
plan_id: str = Depends(get_draft_plan_id),
services: tuple[EditTemplateService, EditPlanService] = Depends(get_editor_services),
asset_repo: SQLAlchemyAssetRepository = Depends(get_asset_repository),
db: Session = Depends(get_db_session),
current_user: AuthenticatedUser = Depends(get_current_user),
) -> ClipsFromAssetsResponse:
"""从素材批量创建片段(按模板segment配置创建,MediaKit异步更新).
逻辑:
1. 从模板读取 segments,片段数量优先级:显式 clip_count(1-10)→ 旧字段
required_clips_count(兼容,超10截断)→ 默认 3(产品默认 3 段)。
片段数大于模板 segment 数时按顺序循环复用 segment 配置。
2. 每个片段时长在对应 segment 的 duration_min ~ duration_max 之间随机取值(保留一位小数)
3. 素材按片段顺序轮询分配,素材不够时同一素材切多个片段
4. 使用 replace_all_clips_transactional 原子性地清空旧片段并创建新的(随机起始时间)
5. 立即返回响应(目标 <1秒)
6. 后台异步任务:调用 MediaKit 智能选片并更新片段的 start_time
7. 素材时长为 0 或缺失时报 400,不创建无效片段
"""
tpl_svc, plan_svc = services
user_id = str(current_user.user.id)
# 1. 查询模板片段配置。模板不存在/已删除/无权限 → 404;
# 模板存在但确实未配置片段 → 422(配置错误,与 404 区分)。
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:
raise HTTPException(
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
detail="模板未配置片段",
)
# 1.5 归一化片段数量:
# 优先级:显式 clip_count → 旧字段 required_clips_count(由 schema 归一化到 clip_count
# → 默认 3(产品默认 3 段)。按 N 循环复用 segment 配置;N <= len(segments) 时截取前 N 个
# (保持向后兼容:原模板有 N 个 segment、前端不传 clip_count 且 N<=10 时按模板段数创建;
# 默认模板仅有 1 个通用 segment 时按 clip_count=3 循环生成 3 段)。
requested_clip_count = getattr(body, "clip_count", None)
if requested_clip_count is None:
# schema 未显式传 clip_count 且无 legacy:使用模板 segments 数量,若超出 10 则截断
requested_clip_count = len(segments) if 1 <= len(segments) <= 10 else 3
requested_clip_count = max(1, min(int(requested_clip_count), 10))
effective_segments: list[tuple[int, float, float]] = []
for i in range(requested_clip_count):
src = segments[i % len(segments)]
effective_segments.append((i, float(src[1]), float(src[2])))
segments = effective_segments
# 防御:schema validator 已过滤 null/空串,这里再归一化一次,
# 避免异常入参(undefined → null)导致后续 /assets/{id} 404 / 422
asset_ids = [str(aid).strip() for aid in (body.asset_ids or []) if isinstance(aid, str) and aid.strip()]
if not asset_ids:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="素材列表为空,无法创建片段",
)
# 2. 获取素材实际时长(去重查询)
unique_asset_ids = list(dict.fromkeys(asset_ids))
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:
asset = asset_repo.get(asset_id)
if asset and hasattr(asset, "duration"):
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)
# 读取素材 metadata 中持久化的历史已用区间(跨任务/跨调用去重),
# 格式与 _calc_random_start_time 的 used_segments 参数一致
used_segments: dict[str, list[tuple[float, float]]] = get_used_segments(db, unique_asset_ids)
# 受控复用回调:可用区间耗尽时复用最久未用且未达复用上限(3次)的历史区间,
# 复用片段时长累加到 reused_durations 供 15% 占比控制
reused_durations: dict[str, float] = {}
# 本条成片中每个素材被分配的片段总时长(复用占比分母)
asset_assigned_durations: dict[str, float] = {}
# 受控复用回调:区间耗尽时复用最久未用且 use_count<3 的历史区间;
# 回调内部预判复用后占比是否超 15%,超限拒绝复用(返回 None)
reuse_cb = make_reuse_callback(
db,
asset_durations,
reused_durations,
assigned_tracker=asset_assigned_durations,
)
clips_data: list[dict] = []
def _reuse_ratio_exceeded(aid: str, extra: float = 0.0) -> bool:
"""该素材在本条成片中「已复用片段时长 / 已分配片段总时长」是否已超 15%
在为下一片段选素材时调用:本片段尚未分配,复用状态只在分配后的回调里
更新,因此直接检查当前占比——一旦已超 15%,该素材不再参与后续分配。
assigned=0(首个片段)放行;reused=0(尚未发生复用)时不误拦正常分配。
"""
assigned = asset_assigned_durations.get(aid, 0.0)
if assigned <= 0:
return False
return reused_durations.get(aid, 0.0) / assigned > REUSE_RATIO_LIMIT
# 素材耗尽标志:某轮循环中所有素材均被跳过时为 True
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 之间随机取值(保留一位小数)
raw_duration = random.uniform(dur_min, dur_max)
# 加上转场补偿,确保最终输出时长 = 模板设定总时长
raw_duration += transition_compensation
# 贪心分配素材:按"已使用次数"升序排列候选素材(使用最少的优先),
# 同次数随机打散,避免"A-B-C-D"的固定组合反复出现。
# 跳过时长缺失、复用占比已超 10% 阈值的素材;
# 选中后计算起点,若该素材可用区间耗尽且复用被闸门拒绝(calc 返回 None),
# 继续尝试下一个素材
asset_id = ""
clip_duration = 0.0
start_time: float | None = None
# 动态按使用次数排序:优先选使用最少的素材,同次数随机打散
asset_use_counts = {aid: len(used_segments.get(aid, [])) for aid in asset_ids}
# 排序键:smart_match 评分(注入随机噪声)→ 使用次数 → 纯随机。
# 噪声让得分接近的素材排名每次浮动,避免同一批素材反复选出相同组合,
# 从素材组合层面降低成片查重率;分差 > 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)
if candidate_total <= 0:
continue
candidate_duration = min(round(raw_duration, 1), candidate_total)
if candidate_duration <= 0:
continue
if _reuse_ratio_exceeded(candidate, candidate_duration):
logger.info(
"from-assets 素材复用占比超 %.0f%% 阈值,跳过分配: asset_id=%s",
REUSE_RATIO_LIMIT * 100,
candidate,
)
continue
# 起始时间选取(不调用 MediaKit,保证接口快速返回):
# 1) 素材有场景切换点缓存时,优先从随机镜头段中选起点(不同片段来自不同镜头,
# 画面内容本质不同),与 used_segments 做冲突避让(含 1.5s 边缘间隙)
# 2) 无缓存 / 镜头段全冲突 → _calc_random_start_time 随机起点兜底;
# 100 次避不开历史区间时走受控复用回调(复用片段累加 reused_durations
# 回调内部预判复用后占比超 10% 则拒绝并返回 None)
candidate_start = None
if candidate in asset_scene_points:
candidate_start = pick_scene_aware_start(
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:
# 该素材可用区间耗尽且复用被闸门/use_count 上限拒绝 → 尝试下一素材
logger.info(
"from-assets 素材无可用可切区间(复用被拒),轮询下一素材: asset_id=%s",
candidate,
)
continue
asset_id = candidate
clip_duration = candidate_duration
start_time = candidate_start
break
if not asset_id or start_time is None:
# 所有素材时长缺失、复用占比超阈值,或区间耗尽且复用被拒 → 素材可切区间不足
all_assets_exhausted = True
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="素材可切区间不足,请补充新素材",
"""从素材批量创建片段"""
_, plan_svc = services
clips = []
for i, asset_id in enumerate(body.asset_ids):
try:
clip = plan_svc.create_clip(
plan_id,
clip_type="main",
order=body.start_order + i if hasattr(body, "start_order") else i,
duration=5.0,
asset_id=asset_id,
)
# 记录已使用时间段(内存,供本次后续片段避开)
used_segments.setdefault(asset_id, []).append((start_time, start_time + clip_duration))
asset_assigned_durations[asset_id] = asset_assigned_durations.get(asset_id, 0.0) + clip_duration
# 同步写入素材 metadata(不 commit,与下方 replace_all_clips_transactional
# 处于同一事务,任一步失败整体回滚,不留脏数据);
# 复用区间与历史记录高度重叠时 record 内部自动累加 use_count
record_used_segments(db, asset_id, start_time, start_time + clip_duration, plan_id)
clips_data.append(
{
"order": _seg_order,
"asset_id": asset_id,
"start_time": start_time,
"duration": clip_duration,
"clip_type": body.clip_type or "main",
}
)
# 按原始 segment order 排序,确保 clips_data 的 order 字段有序(0,1,2,3...
clips_data.sort(key=lambda c: c["order"])
# 4. 事务性替换:清空旧片段 → 创建新片段 → 标记ready(单事务,失败自动回滚)
created_count = plan_svc.replace_all_clips_transactional(plan_id, clips_data)
clips.append(clip)
except ValueError:
pass
logger.info(
"from-assets按模板创建片段(异步): template_id=%s plan_id=%s segments=%d created=%d by user=%s",
"模板编辑器从素材创建片段: template_id=%s plan_id=%s count=%d by user=%s",
template_id,
plan_id,
len(segments),
created_count,
len(clips),
current_user.user.id,
)
# 5. 触发后台任务:异步调用 MediaKit 并更新片段起始时间
background_tasks.add_task(
_update_mediakit_recommendations_async,
plan_id,
unique_asset_ids,
)
# 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(
created_count=created_count,
created_count=len(clips),
plan_id=plan_id,
clip_ids=[],
duplicate_warning=duplicate_warning,
exhaustion_warning=exhaustion_warning,
clip_ids=[c.id for c in clips],
)
def _update_mediakit_recommendations_async( # pragma: no cover
plan_id: str,
asset_ids: list[str],
) -> None:
"""后台任务:使用 SceneChange 智能选帧并更新片段的起始时间.
优先使用 SceneChange 策略检测视频镜头切换点,将每个素材按镜头段拆分,
各片段优先从不同镜头段中选取起始时间,实现「不同片段展示不同场景」的效果。
降级策略:
1. SceneChange 优先 → detect_scene_changes 内部已含 TimeInterval 降级
2. 若 detect_scene_changes 仍返回 None → 回退到旧的 analyze_videos 方式
3. 所有方式都失败 → 保持现有随机 start_time,不影响视频生成
此函数在后台异步执行,不影响接口响应时间。
失败时静默处理,不影响已创建的片段。
"""
from collections import defaultdict
from packages.adapters.sqlalchemy_impl.asset_repository import SQLAlchemyAssetRepository
from packages.adapters.sqlalchemy_impl.session import SessionLocal
db = None
try:
# 复用应用全局 Session(避免每次创建新连接池导致资源泄漏)
if SessionLocal is None:
logger.warning("后台任务: SessionLocal 未初始化,跳过 MediaKit 更新")
return
db = SessionLocal()
# 初始化服务
asset_repo = SQLAlchemyAssetRepository(db)
plan_svc = EditPlanService(db)
# 查询该 plan 的所有片段(分批获取,避免硬编码 limit 截断)
batch_size = 500
all_clips = []
offset = 0
while True:
batch = plan_svc.list_clips(plan_id, skip=offset, limit=batch_size)
if not batch:
break
all_clips.extend(batch)
if len(batch) < batch_size:
break
offset += batch_size
clips = all_clips
if not clips:
logger.info("后台任务: plan_id=%s 无片段,跳过更新", plan_id)
return
# 批量预加载所有涉及的素材(消除 N+1 查询)
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)}
# 按 asset_id 预分组片段对象(按 order 排序,保证按模板顺序分配镜头段)
clips_by_asset: dict[str, list] = defaultdict(list)
for clip in clips:
aid = getattr(clip, "asset_id", "") or ""
if aid:
clips_by_asset[aid].append(clip)
for aid in clips_by_asset:
clips_by_asset[aid].sort(key=lambda c: c.order)
# 读取素材全部历史已用区间(跨任务/跨 plan 持久化记录)
historical_segments = get_used_segments(db, unique_asset_ids)
# 已更新的片段ID(用于排除已移动的旧时间段)
updated_clip_ids: set[str] = set()
# 已更新的时间段
updated_segments: dict[str, list[tuple[float, float]]] = {}
updated_count = 0
# 尝试获取存储服务(用于生成视频 URL)
try:
storage = get_storage_service()
except Exception as e:
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
asset = assets_map.get(asset_id)
if not asset:
continue
asset_total = float(getattr(asset, "duration", 0.0) or 0.0)
if asset_total <= 0:
continue
# 获取素材视频 URL
video_url: str | None = None
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(
"后台任务: 命中场景点缓存: asset_id=%s scenes=%d",
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,
)
continue
# 为每个片段分配不同的镜头段
scene_segments_pool = list(scene_segments) # 可消费的镜头段池
for clip in asset_clips:
clip_duration = clip.duration
recommended_start: float | None = None
# 从镜头段池中依次尝试,选一个不冲突的
for seg_idx, (seg_start, seg_end) in enumerate(scene_segments_pool):
candidate_start = pick_start_in_scene_segment(seg_start, seg_end, clip_duration)
if candidate_start is None:
continue # 镜头段太短,跳过
# 检查越界
if candidate_start + clip_duration > asset_total:
continue
# 检查与已用区间冲突
other_segs = _get_other_segments(asset_id, clip.id)
if _recommended_time_conflicts(candidate_start, clip_duration, other_segs):
continue
recommended_start = candidate_start
# 消费该镜头段(从池中移除,下一个片段用不同镜头段)
scene_segments_pool.pop(seg_idx)
break
if recommended_start is None:
# 镜头段用完或都冲突 → 尝试 _calc_random_start_time 兜底
used_segs_for_calc: dict[str, list[tuple[float, float]]] = {
asset_id: _get_other_segments(asset_id, clip.id)
}
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:
old_start = clip.start_time
old_end = old_start + clip_duration
plan_svc.update_clip(clip.id, start_time=recommended_start)
try:
if remove_used_segment(db, asset_id, old_start, old_end, plan_id=plan_id):
record_used_segments(
db,
asset_id,
recommended_start,
recommended_start + clip_duration,
plan_id,
)
except Exception:
logger.exception(
"后台任务: 同步素材区间记录失败,回滚本次片段更新: clip_id=%s",
clip.id,
)
db.rollback()
continue
db.commit()
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,
asset_id,
recommended_start,
)
except Exception:
logger.exception("后台任务: 单个片段更新失败: clip_id=%s", clip.id)
try:
db.rollback()
except Exception:
pass
continue
logger.info("后台任务完成: plan_id=%s 成功更新 %d 个片段", plan_id, updated_count)
except Exception:
# 后台任务失败不影响已创建的片段,静默处理
logger.exception("后台任务异常: plan_id=%s", plan_id)
if db:
try:
db.rollback()
except Exception:
pass
finally:
if db:
try:
db.close()
except Exception:
pass
@@ -2,17 +2,16 @@
核心依赖:
- get_editor_services: 获取模板+计划服务
- get_draft_plan_id: Depends 形式的路径依赖(template_id 路径参数必填)
- resolve_draft_plan_id: 纯函数版本,供 clips_standalone 等非路径参数场景复用
(支持空 tid 时自动兜底创建默认模板)
- get_draft_plan_id: 根据 template_id 获取或创建草稿,返回 plan_id
- _check_queue_limits: 生成队列限流检查
"""
from __future__ import annotations
import logging
from app.api.routes._default_template import get_or_create_default_template_id
from app.auth import AuthenticatedUser, get_current_user
from app.core.task_enqueue import GLOBAL_PENDING_LIMIT, USER_PENDING_LIMIT
from app.dependencies import get_db_session
from app.services.edit_plan_service import EditPlanService
from app.services.edit_template_service import EditTemplateService
@@ -33,61 +32,46 @@ def get_editor_services(
return EditTemplateService(db), EditPlanService(db)
def resolve_draft_plan_id(
def get_draft_plan_id(
template_id: str,
services: tuple[EditTemplateService, EditPlanService],
current_user: AuthenticatedUser,
db: Session,
auto_create_default: bool = True,
services: tuple[EditTemplateService, EditPlanService] = Depends(get_editor_services),
current_user: AuthenticatedUser = Depends(get_current_user),
db: Session = Depends(get_db_session),
) -> str:
"""根据 template_id 获取或创建草稿,返回 plan_id(纯函数,不带 Depends)。
"""路径依赖:根据 template_id 获取或创建草稿,返回 plan_id.
当 auto_create_default=True 且 template_id 为空时,自动调用
get_or_create_default_template_id 创建默认模板(用于 clips_standalone
等非路径参数场景)。
这是模板编辑器路由的核心依赖——所有编辑器端点都先经过这里,
确保 template_id → plan_id 的映射始终存在。
兼容策略:优先从新模板系统(edit_templates 表)查找,
若不存在则回退到旧模板系统(templates 表),确保用户自建模板可用。
"""
tpl_svc, plan_svc = services
user_id = str(current_user.user.id)
# 0. 空 tid 兜底
if not template_id:
if auto_create_default:
tid = get_or_create_default_template_id(db, user_id)
if not tid:
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail="无法自动创建默认模板,请刷新页面重试",
)
template_id = tid
else:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="template_id 不能为空",
)
# 1. 门禁:校验模板存在且可访问
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="模板不存在")
# 2. 草稿已存在 → 直接返回
# 1. 草稿已存在 → 直接返回
draft = tpl_svc.get_template_draft(template_id)
if draft is not None:
return draft.id
# 3. 全局模板(新系统)→ 用新服务创建草稿
if is_global_template:
# 2. 新系统有模板 → 用新服务创建草稿
if tpl_svc.get_template(template_id) is not None:
draft = tpl_svc.create_template_draft(template_id, user_id=user_id)
return draft.id
# 4. 旧模板(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 packages.domain.edit_template import EditTemplate, EditTemplateStatus
from packages.domain.template_clip_config import ClipType, TemplateClipConfig
# 构造伪 EditTemplate 对象(只填 generate_from_template 需要的字段)
pseudo_template = EditTemplate(
id=old_template.id,
name=old_template.name,
@@ -95,6 +79,7 @@ def resolve_draft_plan_id(
status=EditTemplateStatus.ACTIVE,
)
# 将旧模板 segments 转换为 clip_configs
clip_configs: list[TemplateClipConfig] = []
for seg in old_template.segments or []:
clip_configs.append(
@@ -118,6 +103,7 @@ def resolve_draft_plan_id(
)
plan = result["plan"]
# 标记为模板草稿(后续可复用 tpl_svc.get_template_draft 的查找逻辑)
plan_svc.update_plan_config(plan.id, {"is_template_draft": True})
logger.info(
@@ -129,21 +115,27 @@ def resolve_draft_plan_id(
return plan.id
def get_draft_plan_id(
template_id: str,
services: tuple[EditTemplateService, EditPlanService] = Depends(get_editor_services),
current_user: AuthenticatedUser = Depends(get_current_user),
db: Session = Depends(get_db_session),
) -> str:
"""路径依赖:根据 template_id 获取或创建草稿,返回 plan_id.
Depends 版本:路径参数 template_id 由 FastAPI 保证非空,不自动兜底。
兜底逻辑走 resolve_draft_plan_id(auto_create_default=False)。
"""
return resolve_draft_plan_id(
template_id=template_id,
services=services,
current_user=current_user,
db=db,
auto_create_default=False,
)
def _check_queue_limits(gen_task_repo, user_id: str) -> None:
"""队列限流预检查"""
try:
has_count = (
hasattr(gen_task_repo, "count_pending_by_user")
and hasattr(gen_task_repo, "count_pending_total")
)
if has_count:
user_pending = gen_task_repo.count_pending_by_user(user_id)
global_pending = gen_task_repo.count_pending_total()
if user_pending >= USER_PENDING_LIMIT:
raise HTTPException(
status_code=429,
detail=f"您的待处理任务过多(当前 {user_pending}/{USER_PENDING_LIMIT}),请等待完成后再提交",
)
if global_pending >= GLOBAL_PENDING_LIMIT:
raise HTTPException(
status_code=503,
detail="系统繁忙,请稍后再试",
)
except HTTPException:
raise
except Exception as e:
logger.warning("[模板编辑器队列限流] 检查失败,跳过: %s", e)
@@ -150,7 +150,7 @@ def rollback_template(
try:
tpl = tpl_svc.rollback_to_version(template_id, request.version)
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)
return EditorRollbackResponse(
@@ -41,17 +41,17 @@ def list_editor_transition_presets(
_: AuthenticatedUser = Depends(get_current_user),
) -> TransitionPresetListResponse:
"""获取转场预设列表"""
from packages.domain.transition_presets import TRANSITION_PRESET_LIBRARY
from packages.domain.transition_presets import TRANSITION_PRESETS
items = [
{
"id": p.id,
"name": p.name,
"category": p.category,
"duration": p.default_duration,
"description": p.description,
"id": p["id"],
"name": p["name"],
"category": p.get("category", "通用"),
"duration": p.get("default_duration", 0.5),
"description": p.get("description", ""),
}
for p in TRANSITION_PRESET_LIBRARY
for p in TRANSITION_PRESETS
]
return TransitionPresetListResponse(items=items, total=len(items))
@@ -123,17 +123,17 @@ def list_editor_filter_presets(
_: AuthenticatedUser = Depends(get_current_user),
) -> FilterPresetListResponse:
"""获取滤镜预设列表"""
from packages.domain.filter_presets import FILTER_PRESET_LIBRARY
from packages.domain.filter_presets import FILTER_PRESETS
items = [
{
"id": p.id,
"name": p.name,
"category": p.category,
"thumbnail": p.lut_url,
"description": p.description,
"id": p["id"],
"name": p["name"],
"category": p.get("category", "通用"),
"thumbnail": p.get("thumbnail", ""),
"description": p.get("description", ""),
}
for p in FILTER_PRESET_LIBRARY
for p in FILTER_PRESETS
]
return FilterPresetListResponse(items=items, total=len(items))
+385
View File
@@ -0,0 +1,385 @@
"""草稿生成路由.
端点:
- POST /generate 触发生成
- GET /generation-status 生成进度
- GET /generations 生成记录列表
"""
from __future__ import annotations
import json
import logging
from typing import Any, Optional
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.dependencies import (
get_asset_library_repository,
get_asset_repository,
get_db_session,
get_generated_video_repository,
)
from app.schemas.generation_task import GenerationTaskResponse
from app.services.edit_plan_service import EditPlanService
from app.services.edit_template_service import EditTemplateService
from fastapi import APIRouter, Depends, HTTPException, status
from sqlalchemy.orm import Session
from packages.adapters.sqlalchemy_impl.generation_task_repository import (
SQLAlchemyGenerationTaskRepository,
)
from packages.application.generated_videos import ListGeneratedVideosByTaskUseCase
from packages.application.generation_tasks import (
CreateGenerationTaskCommand,
CreateGenerationTaskUseCase,
)
from packages.domain.edit_plan import EditPlanStatus
from ._fallback import (
_auto_fallback_assign_assets,
_auto_fallback_auto_material_mode,
_auto_fallback_copy_template_clips,
_auto_fallback_draft_to_editing,
)
from .dependencies import _check_queue_limits, get_draft_plan_id, get_editor_services
from .schemas import (
ClipStatusItem,
EditPlanGenerateRequest,
EditPlanGenerateResponse,
EditPlanGenerationsResponse,
EditPlanGenerationStatusResponse,
)
logger = logging.getLogger(__name__)
router = APIRouter(tags=["Template Editor"])
# DEPRECATED: 前端已改用 /generation/tasks 体系,此路由保留仅供旧版兼容,计划下线
@router.post("/generate", response_model=EditPlanGenerateResponse)
def generate_editor_draft(
template_id: str,
request: Optional[EditPlanGenerateRequest] = None,
plan_id: str = Depends(get_draft_plan_id),
services: tuple[EditTemplateService, EditPlanService] = Depends(get_editor_services),
db: Session = Depends(get_db_session),
current_user: AuthenticatedUser = Depends(get_current_user),
asset_library_repo: Any = Depends(get_asset_library_repository),
asset_repo: Any = Depends(get_asset_repository),
) -> EditPlanGenerateResponse:
"""触发模板草稿渲染生成"""
req = request or EditPlanGenerateRequest()
_, plan_svc = services
plan_check = plan_svc.get_plan_or_raise(plan_id)
# 自动兜底流程
_auto_fallback_draft_to_editing(plan_svc, plan_id, plan_check)
_auto_fallback_copy_template_clips(plan_svc, plan_id, plan_check, db)
clips_without_asset = _auto_fallback_assign_assets(plan_svc, plan_id, plan_check)
_auto_fallback_auto_material_mode(
plan_svc,
plan_id,
plan_check,
clips_without_asset,
asset_library_repo,
asset_repo,
user_id=str(current_user.user.id),
)
# 检查是否可复用已完成的预览产物(预览品质已与正式一致)
gen_task_repo = SQLAlchemyGenerationTaskRepository(db)
reusable_task = _find_reusable_preview_task(gen_task_repo, plan_id, plan_check)
if reusable_task:
# 复用预览产物:标记为正式产出,跳过渲染
# 如果前端传了 title_config,需要创建新任务(因为预览任务的 custom_title 可能不同)
title_config_reuse = req.title_config or {}
title_text_reuse = (title_config_reuse.get("text") or "").strip()
existing_custom_title = getattr(reusable_task, "custom_title", "") or ""
if title_text_reuse and existing_custom_title:
# 如果新标题和已有标题不同,不能复用,走新建任务流程
new_title_json = json.dumps(title_config_reuse, ensure_ascii=False)
if new_title_json != existing_custom_title:
logger.info(
"[模板生成] 标题已变更,跳过复用: task_id=%s",
reusable_task.id,
)
reusable_task = None
elif title_text_reuse and not existing_custom_title:
# 原来没标题,现在有标题,不能复用
logger.info(
"[模板生成] 新增标题,跳过复用: task_id=%s",
reusable_task.id,
)
reusable_task = None
elif not title_text_reuse and existing_custom_title:
# 原来有标题,现在移除了,不能复用
logger.info(
"[模板生成] 移除标题,跳过复用: task_id=%s",
reusable_task.id,
)
reusable_task = None
if reusable_task:
# 复用预览产物:标记为正式产出,跳过渲染
reusable_task.mark_confirmed()
gen_task_repo.update(reusable_task)
# 将产物 URL 写入 plan config
rendered_url = _get_task_output_url(reusable_task, gen_task_repo, db)
plan_svc.update_plan_config(
plan_id,
{
"generation_task_id": reusable_task.id,
"rendered_storage_key": rendered_url, # 统一用 rendered_storage_key
},
)
plan_svc.transition_status(plan_id, EditPlanStatus.COMPLETED)
updated_plan = plan_svc.get_plan_or_raise(plan_id)
logger.info(
"模板编辑器复用预览产物: template_id=%s plan_id=%s task_id=%s by user=%s",
template_id,
plan_id,
reusable_task.id,
current_user.user.id,
)
return EditPlanGenerateResponse(
plan_id=plan_id,
plan_status=updated_plan.status.value if hasattr(updated_plan.status, "value") else updated_plan.status,
generation_task_id=reusable_task.id,
clip_count=len((plan_check.config or {}).get("clips", [])),
)
# 检查是否可生成(含最后防线自动修复 + 诊断日志)
try:
can_gen, reason = plan_svc.can_generate(plan_id)
except ValueError as exc:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=str(exc)) from exc
if not can_gen:
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=reason)
try:
clip_count = plan_svc.mark_clips_ready(plan_id)
user_id = current_user.user.id
_check_queue_limits(gen_task_repo, user_id)
gen_task_use_case = CreateGenerationTaskUseCase(gen_task_repo)
plan = plan_svc.get_plan_or_raise(plan_id)
config_asset_ids = (plan.config or {}).get("asset_ids", [])
# 从 plan config 读取封面 URL(由 generate-cover 保存)
cover_url_from_config = (plan.config or {}).get("cover", {}).get("image_url", "")
# 处理标题配置:序列化 title_config 为 JSON 存入 custom_title
title_config = req.title_config or {}
title_text = (title_config.get("text") or "").strip()
custom_title_value = ""
if title_text:
custom_title_value = json.dumps(title_config, ensure_ascii=False)
logger.info(
"[模板生成] 标题配置: text=%s, config_keys=%s",
title_text[:30],
list(title_config.keys()),
)
gen_task = gen_task_use_case.execute(
CreateGenerationTaskCommand(
project_id=plan.project_id or "",
template_id=plan.template_id,
created_by_user_id=current_user.user.id,
source_edit_plan_id=plan_id,
asset_ids=list(config_asset_ids) if config_asset_ids else [],
cover_url=cover_url_from_config,
custom_title=custom_title_value,
),
)
plan_svc.update_plan_config(plan_id, {"generation_task_id": gen_task.id})
plan_svc.transition_status(plan_id, EditPlanStatus.RENDERING)
celery_app.send_task("worker.generate_video", args=[gen_task.id])
updated_plan = plan_svc.get_plan_or_raise(plan_id)
logger.info(
"模板编辑器触发生成: template_id=%s plan_id=%s gen_task_id=%s clips=%d by user=%s",
template_id,
plan_id,
gen_task.id,
clip_count,
current_user.user.id,
)
return EditPlanGenerateResponse(
plan_id=plan_id,
plan_status=updated_plan.status.value if hasattr(updated_plan.status, "value") else updated_plan.status,
generation_task_id=gen_task.id,
clip_count=clip_count,
)
except HTTPException:
raise
except Exception as _e:
logger.exception(
"模板编辑器触发生成失败: template_id=%s plan_id=%s",
template_id,
plan_id,
)
try:
plan_svc.transition_status(plan_id, EditPlanStatus.FAILED)
except Exception:
pass
raise HTTPException(
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
detail="生成失败,请稍后重试",
) from _e
def _find_reusable_preview_task(gen_task_repo, plan_id: str, plan) -> "object | None":
"""查找该 plan 关联的已完成预览任务,判断是否可复用。
复用条件:
1. 存在 source_edit_plan_id == plan_id 的已完成预览任务
2. plan 在预览完成后未被修改(updated_at <= 预览完成时间)
Returns:
可复用的 GenerationTask,或 None
"""
try:
tasks = gen_task_repo.list_by_source_edit_plan(plan_id)
except Exception:
return None
for task in tasks:
if not getattr(task, "is_preview", False):
continue
if not task.is_completed:
continue
# 检查 plan 是否在预览完成后被修改
completed_at = getattr(task, "completed_at", None)
if completed_at and hasattr(plan, "updated_at"):
plan_updated = plan.updated_at
# 如果 plan.updated_at 为空,无法判断是否修改过,跳过
if plan_updated is None:
continue
# 如果 plan 在预览完成后又被修改了,不能复用
if plan_updated > completed_at:
continue
return task
return None
def _get_task_output_url(task, gen_task_repo, db) -> str:
"""获取任务的输出视频 URL。"""
try:
video_repo = get_generated_video_repository(db)
use_case = ListGeneratedVideosByTaskUseCase(video_repo)
videos = use_case.execute(task.id)
if videos:
url = getattr(videos[0], "file_url", "") or ""
# 规范化:合并路径中的双斜杠(保留协议头 ://)
if url:
import re as _re
url = _re.sub(r"(?<!:)//", "/", url)
return url
except Exception:
pass
return ""
# DEPRECATED: 前端已改用 /generation/tasks 体系,此路由保留仅供旧版兼容,计划下线
@router.get("/generation-status", response_model=EditPlanGenerationStatusResponse)
def get_editor_generation_status(
template_id: str,
plan_id: str = Depends(get_draft_plan_id),
services: tuple[EditTemplateService, EditPlanService] = Depends(get_editor_services),
storage_service: OSSStorageService = Depends(get_storage_service),
_: AuthenticatedUser = Depends(get_current_user),
) -> EditPlanGenerationStatusResponse:
"""查询草稿生成进度"""
_, plan_svc = services
try:
gen_status = plan_svc.get_generation_status(plan_id)
except ValueError as exc:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=str(exc)) from exc
plan = gen_status["plan"]
clips = gen_status["clips"]
clip_items = [
ClipStatusItem(
clip_id=c.id,
clip_type=c.clip_type,
order=c.order,
status=c.status.value if hasattr(c.status, "value") else c.status,
asset_id=c.asset_id or "",
text_content=c.text_content or "",
duration=c.duration,
)
for c in clips
]
raw_video_url = (plan.config or {}).get("rendered_storage_key", "") or (plan.config or {}).get("rendered_url", "")
video_url = ""
if raw_video_url:
if raw_video_url.startswith("http"):
video_url = raw_video_url # 已经是完整 URL
else:
try:
video_url = storage_service.get_url(raw_video_url) # storage_key -> 完整 URL
except Exception as e:
logger.warning("生成视频URL获取失败: template_id=%s error=%s", template_id, e)
video_url = raw_video_url
progress = gen_status.get("progress", 0.0)
error_message = gen_status.get("error_message", "")
gen_task_status = gen_status.get("generation_task_status")
plan_status_val = plan.status.value if hasattr(plan.status, "value") else plan.status
if plan_status_val == "completed" and progress < 100:
progress = 100.0
return EditPlanGenerationStatusResponse(
plan_id=plan_id,
plan_status=plan_status_val,
generation_task_id=gen_status["generation_task_id"],
generation_task_status=gen_task_status,
progress=progress,
video_url=video_url,
error_message=error_message,
clips=clip_items,
)
@router.get("/generations", response_model=EditPlanGenerationsResponse)
def list_editor_generations(
template_id: str,
plan_id: str = Depends(get_draft_plan_id),
services: tuple[EditTemplateService, EditPlanService] = Depends(get_editor_services),
db: Session = Depends(get_db_session),
_: AuthenticatedUser = Depends(get_current_user),
) -> EditPlanGenerationsResponse:
"""查询草稿关联的生成记录列表"""
_, plan_svc = services
plan_svc.get_plan_or_raise(plan_id)
gen_task_repo = SQLAlchemyGenerationTaskRepository(db)
tasks = gen_task_repo.list_by_source_edit_plan(plan_id)
items = [
GenerationTaskResponse(
id=t.id,
project_id=t.project_id,
asset_library_id=t.asset_library_id,
strategy_id=t.strategy_id,
voice_library_id=t.voice_library_id,
template_id=t.template_id,
asset_ids=t.asset_ids,
title_ids=t.title_ids,
voice_ids=t.voice_ids,
source_edit_plan_id=t.source_edit_plan_id or "",
status=t.status.value if hasattr(t.status, "value") else t.status,
progress=t.progress,
result_count=t.result_count,
error_message=t.error_message,
)
for t in tasks
]
return EditPlanGenerationsResponse(items=items, total=len(items))
@@ -6,22 +6,76 @@
from __future__ import annotations
import re as _re
from typing import Any, Optional
from typing import Any, Dict, List, Optional
from pydantic import BaseModel, Field, model_validator, validator
from app.schemas.generation_task import GenerationTaskResponse
from pydantic import BaseModel, Field, validator
_EXPORT_RESOLUTION_PATTERN = _re.compile(r"^\d+x\d+$")
_EXPORT_VALID_QUALITY_PRESETS = {"ultra_fast", "fast", "balanced", "high", "best"}
_EXPORT_VALID_FORMATS = {"mp4", "mov"}
# ── 生成状态相关 ────────────────────────────────────────────────────────────
class ClipStatusItem(BaseModel):
"""片段生成状态"""
clip_id: str
clip_type: str
order: int
status: str
asset_id: str
text_content: str
duration: float
class EditPlanGenerationStatusResponse(BaseModel):
"""剪辑计划生成进度响应体"""
plan_id: str
plan_status: str
generation_task_id: Optional[str] = None
generation_task_status: Optional[str] = None
progress: float = 0.0
video_url: str = ""
error_message: str = ""
clips: List[ClipStatusItem]
class EditPlanGenerateRequest(BaseModel):
"""模板编辑器触发生成请求体"""
title_config: Optional[Dict[str, Any]] = Field(
default_factory=dict,
description="标题配置(可选),渲染时烧录到视频中。支持字段: text/font/font_size/font_color/position/bold/stroke/shadow",
)
class EditPlanGenerateResponse(BaseModel):
"""剪辑计划触发生成响应体"""
plan_id: str
plan_status: str
generation_task_id: str
clip_count: int
class EditPlanGenerationsResponse(BaseModel):
"""剪辑计划关联的生成记录列表响应体"""
items: List[GenerationTaskResponse]
total: int
# ── AI 推荐 ────────────────────────────────────────────────────────────────
class AIRecommendRequest(BaseModel):
"""AI 推荐片段方案请求体"""
asset_ids: list[str] = Field(default_factory=list, description="素材 ID 列表")
asset_ids: List[str] = Field(default_factory=list, description="素材 ID 列表")
editing_mode: str = Field(default="one_take", description="剪辑模式: one_take / pip / voice_over / voice_pip")
target_duration: float = Field(default=30.0, ge=1.0, le=600.0, description="目标时长(秒)")
@@ -44,7 +98,7 @@ class AIRecommendResponse(BaseModel):
"""AI 推荐片段方案响应体"""
plan_id: str = Field(..., description="剪辑计划 ID")
clips: list[AIRecommendClipItem] = Field(..., description="推荐的片段列表")
clips: List[AIRecommendClipItem] = Field(..., description="推荐的片段列表")
config: dict[str, Any] = Field(..., description="推荐的 plan configcover/title/subtitle/bgm")
total_duration: float = Field(..., ge=0.0, description="推荐方案总时长(秒)")
confidence: float = Field(..., ge=0.0, le=1.0, description="AI 推荐置信度 (0~1)")
@@ -137,7 +191,7 @@ class ClipReorderItem(BaseModel):
class ClipReorderRequest(BaseModel):
"""片段重排序请求"""
items: list[ClipReorderItem] = Field(..., min_length=1, max_length=500, description="重排序条目列表")
items: List[ClipReorderItem] = Field(..., min_length=1, max_length=500, description="重排序条目列表")
class ClipReorderResponse(BaseModel):
@@ -151,7 +205,7 @@ class ClipReorderResponse(BaseModel):
class ClipBatchDeleteRequest(BaseModel):
"""批量删除片段请求"""
clip_ids: list[str] = Field(..., min_length=1, max_length=500, description="要删除的片段ID列表")
clip_ids: List[str] = Field(..., min_length=1, max_length=500, description="要删除的片段ID列表")
class ClipBatchDeleteResponse(BaseModel):
@@ -162,53 +216,11 @@ class ClipBatchDeleteResponse(BaseModel):
message: str = ""
# sentinel:区分「前端未传 clip_count」和「显式传 0/None」
_UNSET = object()
class ClipsFromAssetsRequest(BaseModel):
"""从素材批量创建片段请求"""
asset_ids: list[str] = Field(..., min_length=1, max_length=200, description="素材 ID 列表,按顺序追加到时间线末尾")
asset_ids: List[str] = Field(..., min_length=1, max_length=200, description="素材 ID 列表,按顺序追加到时间线末尾")
clip_type: str = Field(default="main", description="片段类型,默认 main")
clip_count: Optional[int] = Field(
default=None,
ge=1,
le=10,
description="片段数量(1-10);不传时使用旧字段 required_clips_count;两者都不传时回退为模板 segments 数量(默认 3 段)。",
)
required_clips_count: Optional[int] = Field(
default=None,
ge=1,
le=200,
description="[已废弃] 旧字段,请使用 clip_count;仅作向后兼容——clip_count 未显式传入时才回退本字段(超10截断到10)。",
)
@validator("asset_ids", pre=True)
def _drop_invalid_asset_ids(cls, v): # noqa: N805
"""容错过滤:前端异常情况下可能把 undefined 序列化成 null 或空串混入
asset_ids(会直接 422 或导致后续 /assets/{id} 404),这里统一剔除。
过滤后为空时由 Field(min_length=1) / 路由层 400 兜底。"""
if not isinstance(v, list):
return v
return [x for x in v if isinstance(x, str) and x.strip()]
@model_validator(mode="before")
@classmethod
def _backfill_clip_count(cls, data: Any) -> Any:
"""兼容旧字段 required_clips_count:仅当新字段 clip_count 未显式传入时才回退旧字段;
两者都没传时保持 clip_count=None,路由层按模板 segments 数量兜底。旧字段超 10 截断到 10。"""
if not isinstance(data, dict):
return data
has_new = "clip_count" in data and data["clip_count"] is not None
if not has_new:
legacy = data.get("required_clips_count")
if legacy is not None:
try:
data["clip_count"] = max(1, min(int(legacy), 10))
except (TypeError, ValueError):
pass
return data
class ClipsFromAssetsResponse(BaseModel):
@@ -218,9 +230,7 @@ class ClipsFromAssetsResponse(BaseModel):
created_count: int
plan_id: str = ""
message: str = ""
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="素材耗尽警告")
clip_ids: List[str] = Field(default_factory=list, description="创建的片段ID列表")
# ── 封面配置 ────────────────────────────────────────────────────────────────
@@ -302,7 +312,7 @@ class ExportPresetItem(BaseModel):
class ExportPresetListResponse(BaseModel):
"""导出预设列表响应"""
items: list[ExportPresetItem]
items: List[ExportPresetItem]
total: int
@@ -316,7 +326,7 @@ class FilterPresetResponse(BaseModel):
name: str
category: str
description: str
tags: list[str] = Field(default_factory=list)
tags: List[str] = Field(default_factory=list)
class FilterConfigResponse(BaseModel):
@@ -346,7 +356,7 @@ class FilterUpdateRequest(BaseModel):
class FilterPresetListResponse(BaseModel):
"""滤镜预设列表响应"""
items: list[FilterPresetResponse]
items: List[FilterPresetResponse]
total: int
@@ -360,7 +370,7 @@ class TransitionPresetResponse(BaseModel):
name: str
category: str
description: str
tags: list[str] = Field(default_factory=list)
tags: List[str] = Field(default_factory=list)
default_duration: float
min_duration: float
max_duration: float
@@ -402,7 +412,7 @@ class BatchTransitionResponse(BaseModel):
class TransitionPresetListResponse(BaseModel):
"""转场预设列表响应"""
items: list[TransitionPresetResponse]
items: List[TransitionPresetResponse]
total: int
@@ -458,7 +468,7 @@ class EditorClipResponse(BaseModel):
class EditorClipListResponse(BaseModel):
"""片段列表响应"""
items: list[EditorClipResponse]
items: List[EditorClipResponse]
total: int
@@ -496,7 +506,7 @@ class EditorClipBatchItem(BaseModel):
class EditorClipBatchUpdateRequest(BaseModel):
"""批量替换clips请求(全量覆盖)"""
clips: list[EditorClipBatchItem] = Field(default_factory=list)
clips: List[EditorClipBatchItem] = Field(default_factory=list)
class EditorClipBatchUpdateResponse(BaseModel):
@@ -584,4 +594,4 @@ class EditorTimelineResponse(BaseModel):
plan_id: str
total_duration: float
scenes: list[EditorTimelineSceneResponse]
scenes: List[EditorTimelineSceneResponse]
Regular → Executable
+58 -359
View File
@@ -2,46 +2,36 @@
from __future__ import annotations
import json
import logging
import math
import subprocess
import tempfile
from pathlib import Path
from typing import Any, Optional
from typing import Optional
from app.auth import AuthenticatedUser, get_current_user
from app.config import settings
from app.core.celery_app import celery_app
from app.core.storage import get_storage_service
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,
get_voice_clone_profile_repository,
get_voice_library_repository,
)
from app.schemas.tts import (
ListTTSJobResponse,
SaveToLibraryRequest,
SaveToLibraryResponse,
TTSJobResponse,
TTSPreviewRequest,
TTSPreviewResponse,
TTSStatusResponse,
TTSSynthesizeRequest,
TTSSynthesizeResponse,
)
from fastapi import APIRouter, Depends, HTTPException, Query, Response, WebSocket, WebSocketDisconnect, status
from sqlalchemy.exc import IntegrityError
from sqlalchemy.orm import Session
from packages.adapters.sqlalchemy_impl.tts_job_repository import (
SQLAlchemyTTSJobRepository,
)
from packages.application.cosyvoice_service import CosyVoiceError, CosyVoiceService
from packages.adapters.sqlalchemy_impl.voice_library_repository import SQLAlchemyVoiceLibraryRepository
from packages.application.cosyvoice_service import CosyVoiceService
from packages.application.tts_job.streaming_service import TTSStreamingService
from packages.application.tts_job.use_cases import (
CreateTTSJobUseCase,
@@ -52,14 +42,13 @@ from packages.application.tts_job.use_cases import (
TTSJobNotFoundError,
)
from packages.application.tts_job.workflow import TTSWorkflowService
from packages.domain import Asset, AssetLibrary, AssetLibraryKind, AssetStatus, ClassificationStatus
from packages.domain.points_rules import calculate_points_cost
from packages.domain.points_service import PointsService
from packages.application.voice_library.commands import CreateVoiceLibraryCommand
from packages.application.voice_library.use_cases import (
CreateVoiceLibraryUseCase,
QuotaExceededError,
)
from packages.domain.voice_presets import list_voices
from packages.ports.asset_library_repository import AssetLibraryRepository
from packages.ports.asset_repository import AssetRepository
from packages.ports.project_repository import ProjectRepository
from packages.shared.storage import SharedStorageService
from packages.ports.user_repository import UserRepository
logger = logging.getLogger(__name__)
@@ -132,7 +121,6 @@ def _to_response(job, sign_url=None) -> TTSJobResponse:
def synthesize(
request: TTSSynthesizeRequest,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
db: Session = Depends(get_db_session),
repository: SQLAlchemyTTSJobRepository = Depends(_get_repository),
cosyvoice_service: CosyVoiceService = Depends(get_cosyvoice_service),
voice_clone_repo=Depends(get_voice_clone_profile_repository),
@@ -144,82 +132,28 @@ def synthesize(
"""
user_id = authenticated_user.user.id
# ── 积分扣点(#1895 P2) ──
_points_deducted = 0
_points_scene = "ai_voice"
_points_svc = PointsService() if settings.points_enabled else None
if _points_svc is not None:
# 中文按 ~240 字/分钟粗估时长,至少按 1 分钟扣 1 分
est_minutes = max(1.0, math.ceil(len(request.text) / 240))
_points_deducted = calculate_points_cost(
_points_scene,
is_member=getattr(authenticated_user.user, "is_member", False),
duration_minutes=est_minutes,
member_type=getattr(authenticated_user.user, "member_type", None),
)
_deduct_res = _points_svc.deduct_points(user_id, _points_deducted, _points_scene, db)
if not _deduct_res["success"]:
raise HTTPException(
status_code=402,
detail={
"code": "INSUFFICIENT_POINTS",
"message": f"积分不足,需要 {_points_deducted} 积分,当前余额 {_deduct_res['balance']}",
"required": _points_deducted,
"balance": _deduct_res["balance"],
},
)
# 解析 voice_id:前端可能传克隆音色 profile UUID(而非 CosyVoice voice_id),
# 与 /tts/preview 保持一致:命中 profile → 校验归属 → 取 CosyVoice voice_id
actual_voice_id = request.voice_id
voice_clone_profile_id = request.voice_clone_profile_id
resolved_profile = None
if actual_voice_id:
resolved_profile = voice_clone_repo.get(actual_voice_id)
if resolved_profile is not None:
voice_clone_profile_id = actual_voice_id
# 显式传了 voice_clone_profile_id(且与 voice_id 不同)时再查一次归属
if voice_clone_profile_id and (resolved_profile is None or resolved_profile.id != voice_clone_profile_id):
resolved_profile = voice_clone_repo.get(voice_clone_profile_id)
if resolved_profile is None:
# 校验 voice_clone_profile_id 归属(防止越权使用他人克隆音色)
if request.voice_clone_profile_id:
profile = voice_clone_repo.get(request.voice_clone_profile_id)
if profile is None:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail="Voice clone profile not found",
)
if resolved_profile is not None:
if resolved_profile.user_id != user_id:
if profile.user_id != user_id:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail="无权访问该音色",
detail="Access denied to voice clone profile",
)
if not resolved_profile.voice_id:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="音色克隆尚未完成,请稍后再试",
)
# 命中克隆音色:无论 voice_id 直接传 profile UUID 还是显式传 voice_clone_profile_id
# job.voice_id 统一存解析后的 CosyVoice voice_id
actual_voice_id = resolved_profile.voice_id
# 语速/情绪等合成参数随 metadata 落库,workflow 提交 CosyVoice 时读取透传
synthesis_meta = {
"speed": request.speed,
"emotion": request.emotion or "",
"language": request.language or "zh-CN",
}
if request.metadata_:
synthesis_meta.update(request.metadata_)
use_case = CreateTTSJobUseCase(repository)
job = use_case.execute(
user_id=user_id,
input_text=request.text,
voice_id=actual_voice_id,
voice_id=request.voice_id,
voice_model=request.voice_model,
voice_clone_profile_id=voice_clone_profile_id,
metadata=synthesis_meta,
voice_clone_profile_id=request.voice_clone_profile_id,
metadata=request.metadata_,
)
# 提交 CosyVoice 合成任务
@@ -228,7 +162,6 @@ def synthesize(
cosyvoice_service=cosyvoice_service,
)
synthesis_error: Exception | None = None
try:
job = workflow.start_synthesis(job.id)
except Exception as e:
@@ -236,17 +169,10 @@ def synthesize(
# 但 DB 异常、网络异常等意外错误可能逃逸。
# 与音色克隆接口保持一致:标记 failed,返回 201,不抛 500。
logger.error(f"TTS 合成异常: job_id={job.id}, error={e}", exc_info=True)
synthesis_error = e
try:
job = workflow.process_synthesis_failure(job.id, str(e))
except Exception as inner_e:
logger.error(f"标记 TTS job 失败时出错: job_id={job.id}, error={inner_e}")
# 合成失败且已扣积分 → 退费
if synthesis_error is not None and _points_deducted > 0 and _points_svc is not None:
try:
_points_svc.refund_points(user_id, _points_deducted, _points_scene, db, ref_id=job.id)
except Exception as refund_err:
logger.warning(f"TTS 合失败退积分异常: job_id={job.id}, err={refund_err}")
# 若任务处于 processing 状态(异步模式),触发 Celery 后台轮询
if job.status.value == "processing":
@@ -261,17 +187,10 @@ def synthesize(
celery_app.send_task("worker.process_tts_synthesis", args=[job.id])
except Exception as e:
# Celery 调度失败,标记 job 为 failed
# e used below for refund context
try:
workflow.process_synthesis_failure(job.id, f"Celery 任务调度失败: {e}")
except Exception as inner_e:
logger.error(f"Celery 调度后标记失败时出错: job_id={job.id}, error={inner_e}")
# 调度失败退费
if _points_deducted > 0 and _points_svc is not None:
try:
_points_svc.refund_points(user_id, _points_deducted, _points_scene, db, ref_id=job.id)
except Exception as refund_err:
logger.warning(f"Celery 调度失败退积分异常: job_id={job.id}, err={refund_err}")
return TTSSynthesizeResponse(
job_id=job.id,
@@ -363,62 +282,6 @@ def delete_tts_job(
return
def _find_or_create_voice_library(
*,
user_id: str,
project_repository: ProjectRepository,
asset_library_repository: Any, # port Protocol 声明为 asyncSQLAlchemy 实现为同步,与 upload/asset_libraries 路由惯例一致用 Any
) -> AssetLibrary:
"""在用户可访问的项目中找到(或自动创建)voice 素材库。
与前端配音素材页逻辑一致:素材库挂在项目下,配音素材读取
getAssetsByKind("voice") → 用户所有可访问项目中的 voice 库。
优先使用已有 voice 库;没有则在第一个可访问项目中自动创建。
"""
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
# 所有项目都没有 voice 库 → 在第一个可访问项目中自动创建默认配音素材库。
# asset_libraries 有 (project_id, kind) 唯一索引兜底并发:若两个请求同时创建,
# 落败方捕获 IntegrityError 回滚后重新查询,返回抢先创建成功的库。
project = projects[0]
library = AssetLibrary.create(
project_id=project.id,
name="配音素材库",
kind=AssetLibraryKind.VOICE,
)
try:
return asset_library_repository.create(library)
except IntegrityError:
# 并发下另一个请求已抢先创建:回滚当前事务(立即 commit 模式下 session 已
# 自动回滚,rollback 为幂等 no-opUoW/flush 模式下必须显式回滚才能继续查询),
# 再重查返回抢先创建成功的库。
session = getattr(asset_library_repository, "session", None)
if session is not None:
try:
session.rollback()
except Exception:
logger.warning("IntegrityError 后回滚 session 失败(可能已关闭)", exc_info=True)
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 # IntegrityError 已处理,不保留异常链
@router.post(
"/jobs/{job_id}/save-to-library",
response_model=SaveToLibraryResponse,
@@ -429,17 +292,13 @@ def save_tts_job_to_library(
request: SaveToLibraryRequest = SaveToLibraryRequest(),
authenticated_user: AuthenticatedUser = Depends(get_current_user),
tts_repository: SQLAlchemyTTSJobRepository = Depends(_get_repository),
asset_repository: AssetRepository = Depends(get_asset_repository),
asset_library_repository: AssetLibraryRepository = Depends(get_asset_library_repository),
project_repository: ProjectRepository = Depends(get_project_repository),
storage_service: SharedStorageService = Depends(get_storage_service),
voice_library_repository: SQLAlchemyVoiceLibraryRepository = Depends(get_voice_library_repository),
user_repository: UserRepository = Depends(get_user_repository),
sign_url=Depends(get_audio_url_signer),
) -> SaveToLibraryResponse:
"""将已完成的 TTS 合成结果保存到配音素材库(assets 表新素材体系)
"""将已完成的 TTS 合成结果保存到配音
流程:把 TTS 输出音频转存到用户素材 OSS 路径 → 创建 file_type=audio、
status=ready 的 asset(挂用户 voice 素材库)→ 返回前端可用结构。
配额策略与素材上传一致(上传/ingest 链路无额外配额拦截)。
自动携带音色名、时长、语速等元信息。
"""
user_id = authenticated_user.user.id
@@ -457,220 +316,60 @@ def save_tts_job_to_library(
detail="TTS job is not completed yet",
)
if not job.output_audio_url and not job.output_audio_key:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="TTS job 缺少输出音频,无法保存",
)
# 素材名称
# 构建配音素材名称
name = request.name or f"TTS-{job.id[:8]}"
# 找到(或自动创建)用户 voice 素材库
library = _find_or_create_voice_library(
user_id=user_id,
project_repository=project_repository,
asset_library_repository=asset_library_repository,
)
# 转存音频到素材 OSS 路径(tts-outputs/ 下的产物归 TTS 任务所有,
# 素材独立持有副本,删除 TTS 任务不影响配音库素材)
audio_format = (job.format or "mp3").strip() or "mp3"
content_type_map = {
"mp3": "audio/mpeg",
"wav": "audio/wav",
"pcm": "audio/pcm",
"opus": "audio/opus",
}
content_type = content_type_map.get(audio_format, "audio/mpeg")
storage_key = f"uploads/voice/tts/{job.id}.{audio_format}"
tmp_path: Path | None = None
audio_duration: float | None = None
file_size = 0
try:
with tempfile.NamedTemporaryFile(suffix=f".{audio_format}", delete=False) as tmp:
tmp_path = Path(tmp.name)
# 优先用 OSS storage_key(走 oss2 SDK,私有 bucket 也可下载);
# 兜底用 output_audio_url(旧任务可能没有 key)。
# download_asset 自动识别输入:http(s):// 开头走 HTTP 下载,否则按 OSS key 走 SDK。
download_source = job.output_audio_key or job.output_audio_url
downloaded = storage_service.download_asset(download_source, tmp_path)
if not downloaded or not tmp_path.exists() or tmp_path.stat().st_size == 0:
raise HTTPException(
status_code=status.HTTP_502_BAD_GATEWAY,
detail="TTS 音频下载失败,无法保存到配音库",
)
file_size = tmp_path.stat().st_size
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:
raise
except Exception as e:
logger.error("TTS 音频转存素材失败: job_id=%s, error=%s", job.id, e, exc_info=True)
raise HTTPException(
status_code=status.HTTP_502_BAD_GATEWAY,
detail="TTS 音频转存失败,无法保存到配音库",
) from e
finally:
if tmp_path and tmp_path.exists():
try:
tmp_path.unlink()
except OSError:
pass
# 构建素材元信息
metadata_: dict[str, object] = {
# 构建元信息
metadata_ = {
"source": "tts_job",
"tts_job_id": job.id,
"format": job.format,
"sample_rate": job.sample_rate,
"voice_id": job.voice_id,
"voice_name": job.voice_model or "",
}
if job.metadata:
# 保留原始 job 的有用元信息
for key in ("speed", "language"):
if key in job.metadata:
metadata_[key] = job.metadata[key]
asset = Asset.create(
project_id=library.project_id,
library_id=library.id,
name=name,
storage_key=storage_key,
mime_type=content_type,
metadata=metadata_,
file_size=file_size,
duration=job.duration or audio_duration or None,
status=AssetStatus.READY,
classification_status=ClassificationStatus.PENDING, # 音频不参与内容分类,保持 pending 与 ingest 链路一致
uploaded_by_user_id=user_id,
)
try:
asset = asset_repository.create(asset)
except Exception as e:
# DB 写入失败:清理已上传到 OSS 的素材文件,避免产生无法索引的孤儿文件
logger.error("素材记录创建失败,清理 OSS 文件: %s, error=%s", storage_key, e, exc_info=True)
try:
storage_service.delete_file(storage_key)
except Exception:
logger.warning("清理孤儿 OSS 文件失败: %s", storage_key, exc_info=True)
raise HTTPException(
status_code=status.HTTP_502_BAD_GATEWAY,
detail="素材保存失败,请重试",
) from e
# 获取用户套餐(用于配额检查)
user = user_repository.find_by_id(user_id)
plan_name = getattr(user, "subscription_plan", "free") if user else "free"
return SaveToLibraryResponse(
id=asset.id,
name=asset.name,
audio_url=sign_url(storage_key),
duration=asset.duration or 0.0,
# 构建命令并执行
command = CreateVoiceLibraryCommand(
user_id=user_id,
name=name,
text=job.input_text,
voice_provider="cosyvoice",
voice_id=job.voice_id,
voice_name=job.voice_model or "",
audio_url=job.output_audio_url,
duration=job.duration,
file_size=job.file_size,
status="completed",
project_id=job.project_id or "",
tags=[],
metadata_=metadata_,
)
@router.post("/preview", response_model=TTSPreviewResponse)
def preview_tts(
request: TTSPreviewRequest,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
db: Session = Depends(get_db_session),
cosyvoice_service: CosyVoiceService = Depends(get_cosyvoice_service),
voice_clone_repo=Depends(get_voice_clone_profile_repository),
) -> TTSPreviewResponse:
"""TTS 预览(试听)——同步合成,立即返回音频 URL。
用于前端预览配音效果,限制文本长度 200 字以内。
支持预设音色和克隆音色:克隆音色传的是 profile UUID,需解析为 CosyVoice voice_id。
"""
user_id = authenticated_user.user.id
# ── 积分扣点(#1895 P2) ──
_points_deducted = 0
_points_scene = "ai_voice"
_points_svc = PointsService() if settings.points_enabled else None
if _points_svc is not None:
est_minutes = max(1.0, math.ceil(len(request.text) / 240))
_points_deducted = calculate_points_cost(
_points_scene,
is_member=getattr(authenticated_user.user, "is_member", False),
duration_minutes=est_minutes,
member_type=getattr(authenticated_user.user, "member_type", None),
)
_deduct_res = _points_svc.deduct_points(user_id, _points_deducted, _points_scene, db)
if not _deduct_res["success"]:
raise HTTPException(
status_code=402,
detail={
"code": "INSUFFICIENT_POINTS",
"message": f"积分不足,需要 {_points_deducted} 积分,当前余额 {_deduct_res['balance']}",
"required": _points_deducted,
"balance": _deduct_res["balance"],
},
)
# 解析 voice_id:前端可能传 VoiceCloneProfile UUID 或预设音色 ID
actual_voice_id = request.voice_id
profile = voice_clone_repo.get(request.voice_id)
if profile is not None:
# 命中克隆音色 profile — 校验归属权限
if profile.user_id != authenticated_user.user.id:
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail="无权访问该音色",
)
if not profile.voice_id:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="音色克隆尚未完成,请稍后再试",
)
actual_voice_id = profile.voice_id
use_case = CreateVoiceLibraryUseCase(voice_library_repository)
try:
result = cosyvoice_service.synthesize_speech(
text=request.text,
voice_id=actual_voice_id,
speed=request.speed,
emotion=request.emotion,
language=getattr(request, "language", "zh-CN"),
)
except (CosyVoiceError, ValueError) as e:
# 合成失败退费
if _points_deducted > 0 and _points_svc is not None:
try:
_points_svc.refund_points(user_id, _points_deducted, _points_scene, db)
except Exception as refund_err:
logger.warning(f"TTS 预览失败退积分异常: {refund_err}")
if isinstance(e, CosyVoiceError):
raise HTTPException(status_code=status.HTTP_502_BAD_GATEWAY, detail=f"TTS 合成失败: {e}") from e
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=str(e)) from e
item = use_case.execute(command, plan_name=plan_name or "free")
except QuotaExceededError as exc:
raise HTTPException(
status_code=status.HTTP_429_TOO_MANY_REQUESTS,
detail=f"配音库配额已满({exc.used}/{exc.limit}),请升级套餐",
) from exc
return TTSPreviewResponse(
audio_url=result.audio_url,
duration=result.duration if result.duration and result.duration > 0 else None,
return SaveToLibraryResponse(
id=item.id,
name=item.name,
audio_url=sign_url(item.audio_url) if item.audio_url else "",
duration=item.duration,
voice_id=item.voice_id,
voice_name=item.voice_name,
status=item.status,
)
+49 -341
View File
@@ -23,7 +23,6 @@ from app.schemas.upload import (
from fastapi import APIRouter, Depends, File, Form, HTTPException, UploadFile, status
from packages.application import SubmitIngestJobCommand, SubmitIngestJobUseCase
from packages.domain import Asset, AssetStatus
logger = logging.getLogger(__name__)
@@ -81,199 +80,12 @@ def _validate_mime_type(content_type: str | None) -> str:
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(
project_id: str,
library_id: str,
storage_key: str,
ingest_job_repository: Any,
file_hash: str = "",
asset_id: str = "",
) -> Any:
use_case = SubmitIngestJobUseCase(ingest_job_repository)
job = use_case.execute(
@@ -282,11 +94,9 @@ def _submit_ingest_job(
library_id=library_id,
storage_key=storage_key,
file_hash=file_hash,
asset_id=asset_id,
)
)
celery_result = celery_app.send_task("worker.ingest_asset", args=[job.id])
_persist_celery_task_id(ingest_job_repository, job, getattr(celery_result, "id", ""))
celery_app.send_task("worker.ingest_asset", args=[job.id])
return job
@@ -296,15 +106,9 @@ async def prepare_direct_upload(
authenticated_user: AuthenticatedUser = Depends(get_current_user),
project_repository: Any = Depends(get_project_repository),
asset_library_repository: Any = Depends(get_asset_library_repository),
asset_repository: Any = Depends(get_asset_repository),
storage_service: OSSStorageService = Depends(get_storage_service),
) -> DirectUploadPrepareResponse:
"""创建浏览器直传 OSS 的短期表单签名,并在签名前按 file_hash/client_upload_id 去重。
命中去重:直接返回 duplicated=True + skip_transfer=True(前端跳过 OSS 直传),
未命中:正常签名 OSS 并立即预建一条 PROCESSING 状态的 asset 记录占住
file_hash 闸门,响应带 asset_id 供前端/后续 complete 关联。
"""
"""创建浏览器直传 OSS 的短期表单签名"""
settings = get_settings()
max_size_bytes = settings.OSS_DIRECT_UPLOAD_MAX_MB * 1024 * 1024
if request.file_size > max_size_bytes:
@@ -323,39 +127,8 @@ async def prepare_direct_upload(
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]
safe_filename = request.filename.replace("/", "_").replace("\\", "_")
storage_key = f"uploads/{file_id}/{safe_filename}"
try:
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__}",
) 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(
upload_url=str(payload["url"]),
method=str(payload["method"]),
@@ -403,9 +154,6 @@ async def prepare_direct_upload(
expires_at=str(payload["expires_at"]),
fields={str(key): str(value) for key, value in dict(payload["fields"]).items()},
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),
storage_service: OSSStorageService = Depends(get_storage_service),
) -> DirectUploadCompleteResponse:
"""确认浏览器直传完成并创建导入任务(幂等:重复 complete 返回同一素材)"""
"""确认浏览器直传完成并创建导入任务。"""
require_project_and_library(
request.project_id,
request.library_id,
@@ -429,29 +177,6 @@ async def complete_direct_upload(
normalized_key = storage_service._normalize_storage_key(request.storage_key)
if not normalized_key.startswith("uploads/"):
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:
file_exists = storage_service.file_exists(normalized_key)
except Exception as error:
@@ -463,21 +188,26 @@ async def complete_direct_upload(
if not file_exists:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Uploaded file not found")
# 立即创建 Asset 记录(PROCESSING 状态),使前端刷新后即可看到新素材
mime_type = _infer_mime_type_from_storage_key(normalized_key)
pending_asset = _create_pending_asset(
asset_repository=asset_repository,
project_id=request.project_id,
library_id=request.library_id,
storage_key=normalized_key,
filename=filename,
mime_type=mime_type,
user_id=authenticated_user.user.id,
file_hash=request.file_hash,
client_upload_id=request.client_upload_id,
file_size=request.file_size,
)
# Issue #1776: 计数由 asset_repository.create() 自动维护
# ── 素材去重检测:同素材库 + 同 file_hash 视为重复 ──
if request.file_hash:
existing = asset_repository.find_by_library_and_file_hash(
library_id=request.library_id,
file_hash=request.file_hash,
)
if existing is not None:
logger.info(
"素材去重命中: library=%s hash=%s existing_asset=%s",
request.library_id,
request.file_hash,
existing.id,
)
return DirectUploadCompleteResponse(
storage_key=normalized_key,
ingest_job_id="",
duplicated=True,
asset_id=existing.id,
url=storage_service.get_url(normalized_key),
)
job = _submit_ingest_job(
project_id=request.project_id,
@@ -485,14 +215,8 @@ async def complete_direct_upload(
storage_key=normalized_key,
ingest_job_repository=ingest_job_repository,
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(
@@ -505,8 +229,7 @@ async def upload_asset(
project_id: str = Form(..., min_length=1, description="项目 ID"),
library_id: str = Form(..., min_length=1, description="素材库 ID"),
file: UploadFile = File(..., description="要上传的文件(视频、音频、图片等)"),
file_hash: str = Form(default="", description="文件哈希,用于去重检测"),
client_upload_id: str = Form(default="", description="客户端幂等 token(同一次上传的重试保持一致)"),
file_hash: str = Form(default="", description="文件 MD5 哈希,用于去重检测"),
authenticated_user: AuthenticatedUser = Depends(get_current_user),
ingest_job_repository: Any = Depends(get_ingest_job_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)
# 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)
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]
safe_filename = file.filename.replace("/", "_").replace("\\", "_") if file.filename else "unknown"
storage_key = f"uploads/{file_id}/{safe_filename}"
try:
@@ -560,32 +284,16 @@ async def upload_asset(
detail=f"Failed to upload file: {type(error).__name__}",
) 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(
project_id=project_id,
library_id=library_id,
storage_key=storage_key,
ingest_job_repository=ingest_job_repository,
file_hash=file_hash,
asset_id=pending_asset.id,
)
return UploadAssetResponse(
storage_key=storage_key,
ingest_job_id=job.id,
asset_id=pending_asset.id,
url=file_url,
)
-75
View File
@@ -15,7 +15,6 @@ from app.schemas.video_center import (
VideoItemResponse,
)
from fastapi import APIRouter, Depends, HTTPException, Query, Response
from pydantic import BaseModel, Field
from packages.application import (
GetGeneratedVideoUseCase,
@@ -53,9 +52,6 @@ def _to_video_response(item, storage: OSSStorageService | None = None) -> VideoI
generation_params=item.generation_params,
download_url=download_url,
generated_at=format_utc_datetime(item.generated_at) if hasattr(item, "generated_at") else "",
duplicate_rate=getattr(item, "duplicate_rate", None),
visual_similarity=getattr(item, "visual_similarity", None),
match_count=getattr(item, "match_count", None),
)
@@ -240,74 +236,3 @@ def get_batch_download_status(
status=api_status,
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 "所有视频查重数据已完整",
)
+18 -168
View File
@@ -3,20 +3,11 @@
from __future__ import annotations
import logging
import math
from typing import Optional
from app.auth import AuthenticatedUser, get_current_user
from app.config import settings
from app.core.celery_app import celery_app
from app.core.storage import get_storage_service
from app.dependencies import (
get_asset_repository,
get_cosyvoice_service,
get_db_session,
get_project_repository,
get_voice_clone_profile_repository,
)
from app.dependencies import get_cosyvoice_service, get_voice_clone_profile_repository
from app.schemas.voice_clone import (
CreateVoiceCloneRequest,
ListVoiceCloneResponse,
@@ -25,7 +16,6 @@ from app.schemas.voice_clone import (
VoiceCloneStatusResponse,
)
from fastapi import APIRouter, Depends, HTTPException, Query, Response, status
from sqlalchemy.orm import Session
from packages.adapters.sqlalchemy_impl.voice_clone_profile_repository import (
SQLAlchemyVoiceCloneProfileRepository,
@@ -42,14 +32,6 @@ from packages.application.voice_clone.use_cases import (
from packages.application.voice_clone.workflow import (
VoiceCloneWorkflowService,
)
from packages.domain.points_rules import calculate_points_cost
from packages.domain.points_service import PointsService
# remove duplicate
_DUMMY_DELETED = ()
from packages.ports.asset_repository import AssetRepository
from packages.ports.project_repository import ProjectRepository
from packages.shared.storage import SharedStorageService
logger = logging.getLogger(__name__)
@@ -101,68 +83,23 @@ def create_voice_clone(
request: CreateVoiceCloneRequest,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
workflow: VoiceCloneWorkflowService = Depends(_get_workflow_service),
asset_repository: AssetRepository = Depends(get_asset_repository),
project_repository: ProjectRepository = Depends(get_project_repository),
storage_service: SharedStorageService = Depends(get_storage_service),
) -> VoiceCloneProfileResponse:
"""创建音色克隆任务。
创建 VoiceCloneProfile → 提交 CosyVoice 克隆任务 → 触发 Celery 异步轮询。
参考音频两种来源(二选一):
- source_audio_url:前端直传后的音频 URL(兼容旧流程)
- asset_id:配音素材库中的音频素材,服务端用其 OSS storage_key 生成
预签名下载 URL(不依赖前端签名,避免签名过期导致克隆失败)
如果有参考音频,状态会变为 processing;否则保持 pending。
如果有 source_audio_url,状态会变为 processing;否则保持 pending。
"""
user_id = authenticated_user.user.id
source_audio_url = request.source_audio_url
clone_metadata = dict(request.metadata_ or {})
if request.asset_id:
if source_audio_url:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="asset_id 与 source_audio_url 只能传一个",
)
asset = asset_repository.find_by_id(request.asset_id)
if asset is None:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail="素材不存在",
)
# 归属校验:素材挂在项目素材库下,用户必须能访问该项目
project = project_repository.find_by_id(asset.project_id)
if project is None or not project.can_access(user_id):
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail="无权使用该素材",
)
# 类型校验:仅支持音频素材
if asset.file_type != "audio":
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="仅支持音频素材进行音色克隆",
)
if not asset.storage_key:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="该素材缺少音频文件,无法用于克隆",
)
# 用 OSS storage_key 生成服务端预签名 URL(7 天有效,覆盖克隆重试周期)
source_audio_url = storage_service.get_download_url(asset.storage_key, expires_seconds=7 * 24 * 3600)
clone_metadata["source_asset_id"] = asset.id
profile = workflow.start_clone(
user_id=user_id,
name=request.name,
description=request.description,
source_audio_url=source_audio_url,
source_audio_url=request.source_audio_url,
voice_model=request.voice_model,
language=request.language,
gender=request.gender,
max_retries=request.max_retries,
metadata=clone_metadata,
metadata=request.metadata_,
)
# 如果 profile 处于 processing 且有 task_id,触发 Celery 异步轮询
@@ -172,12 +109,12 @@ def create_voice_clone(
celery_app.send_task("worker.process_voice_clone", args=[profile.id])
logger.info(f"Celery task dispatched for voice clone {profile.id}")
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
try:
workflow.process_clone_failure(profile.id, f"Celery 任务调度失败: {e}")
except Exception:
logger.exception("Failed to mark profile as failed after dispatch error")
except Exception as inner_e:
logger.error(f"Failed to mark profile as failed after dispatch error: {inner_e}")
return _to_response(profile)
@@ -286,111 +223,32 @@ def retry_voice_clone(
celery_app.send_task("worker.process_voice_clone", args=[profile.id])
logger.info(f"Celery task dispatched for voice clone retry {profile.id}")
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
try:
workflow.process_clone_failure(profile.id, f"Celery 任务调度失败: {e}")
except Exception:
logger.exception("Failed to mark profile as failed after dispatch error")
except Exception as inner_e:
logger.error(f"Failed to mark profile as failed after dispatch error: {inner_e}")
return _to_response(profile)
_ALLOWED_PREVIEW_EMOTIONS = {
"",
# 7 种标准英文枚举(CosyVoice v3 官方值)
"neutral",
"happy",
"sad",
"angry",
"surprised",
"fearful",
"disgusted",
# 前端中文 7 标签
"中立",
"开心",
"难过",
"生气",
"惊讶",
"恐惧",
"厌恶",
# 旧英文 4 枚举 + 常见中文别名兼容
"natural",
"excited",
"calm",
"friendly",
"自然",
"愉快",
"高兴",
"快乐",
"兴奋",
"悲伤",
"愤怒",
"惊奇",
"吃惊",
"害怕",
"讨厌",
# 灵应 P1 指定别名
"中性",
"伤心",
"沉稳",
"亲切",
}
@router.get("/{clone_id}/preview", response_model=VoiceClonePreviewResponse)
def get_voice_clone_preview(
clone_id: str,
text: str = Query("", description="自定义试听文本,为空则使用默认示例"),
speed: float = Query(1.0, ge=0.5, le=2.0, description="语速,0.5-2.0,默认 1.0"),
emotion: str = Query(
"",
description="情绪:neutral/happy/sad/angry/surprised/fearful/disgusted,兼容旧值 natural/excited/calm/friendly,空为默认自然",
),
authenticated_user: AuthenticatedUser = Depends(get_current_user),
db: Session = Depends(get_db_session),
repository: SQLAlchemyVoiceCloneProfileRepository = Depends(get_voice_clone_profile_repository),
cosyvoice: CosyVoiceService = Depends(get_cosyvoice_service),
) -> VoiceClonePreviewResponse:
"""获取克隆音色试听音频(实时 TTS 合成)。
- 克隆音色必须处于 ready 状态
- 使用默认试听文本时,结果缓存 7 天(仅默认 text+speed=1.0+emotion=空 组合缓存)
- 可传入自定义 text/speed/emotion 试听不同效果
- 使用默认试听文本时,结果缓存 7 天
- 可传入自定义 text 参数试听不同文本
"""
import time
user_id = authenticated_user.user.id
_points_deducted = 0
_points_scene = "voice_clone_synth"
_points_svc = PointsService() if settings.points_enabled else None
_preview_text_for_points = text.strip() or CLONE_PREVIEW_TEMPLATE
if _points_svc is not None:
est_minutes = max(1.0, math.ceil(len(_preview_text_for_points) / 240))
_points_deducted = calculate_points_cost(
_points_scene,
is_member=getattr(authenticated_user.user, "is_member", False),
duration_minutes=est_minutes,
member_type=getattr(authenticated_user.user, "member_type", None),
)
_deduct_res = _points_svc.deduct_points(user_id, _points_deducted, _points_scene, db)
if not _deduct_res["success"]:
raise HTTPException(
status_code=402,
detail={
"code": "INSUFFICIENT_POINTS",
"message": f"积分不足,需要 {_points_deducted} 积分,当前余额 {_deduct_res['balance']}",
"required": _points_deducted,
"balance": _deduct_res["balance"],
},
)
if emotion not in _ALLOWED_PREVIEW_EMOTIONS:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail=f"不支持的 emotion 值: {emotion},可选: neutral/happy/sad/angry/surprised/fearful/disgusted 或中文 中立/中性/开心/难过/伤心/生气/愤怒/惊讶/吃惊/恐惧/害怕/厌恶/讨厌 或留空",
)
use_case = GetVoiceCloneUseCase(repository)
try:
profile = use_case.execute(clone_id, authenticated_user.user.id)
@@ -403,8 +261,8 @@ def get_voice_clone_preview(
detail=f"Voice clone is not ready (current status: {profile.status})",
)
# 仅默认试听文本 + 默认 speed + 默认 emotion 时使用缓存
use_cache = (not text.strip()) and abs(speed - 1.0) < 1e-6 and (not emotion)
# 有自定义文本时不缓存
use_cache = not text.strip()
if use_cache and clone_id in _clone_preview_cache:
audio_url, duration, file_size, cached_text, cached_at = _clone_preview_cache[clone_id]
@@ -425,20 +283,12 @@ def get_voice_clone_preview(
text=preview_text,
voice_id=profile.voice_id,
format="mp3",
speed=speed,
emotion=emotion,
speed=1.0,
)
except (CosyVoiceError, ValueError) as e:
if _points_deducted > 0 and _points_svc is not None:
try:
_points_svc.refund_points(user_id, _points_deducted, _points_scene, db)
except Exception as refund_err:
logger.warning(f"克隆音色试听失败退积分异常: clone_id={clone_id}, err={refund_err}")
if isinstance(e, CosyVoiceError):
raise HTTPException(status_code=502, detail=f"TTS 合成失败: {e}") from e
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=str(e)) from e
except CosyVoiceError as e:
raise HTTPException(status_code=502, detail=f"TTS 合成失败: {e}") from e
# 缓存(仅默认参数组合
# 缓存(仅默认试听文本
if use_cache:
_clone_preview_cache[clone_id] = (
result.audio_url,
+4 -263
View File
@@ -6,26 +6,12 @@
from __future__ import annotations
import logging
import shutil
import subprocess
import tempfile
import time
from pathlib import Path
from typing import Literal, Optional
from uuid import uuid4
from app.api.routes._helpers import get_user_plan
from app.auth import AuthenticatedUser, get_current_user
from app.core.storage import get_storage_service
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.dependencies import get_audio_url_signer, get_cosyvoice_service, get_db_session, get_user_repository
from app.schemas.voice import (
PresetVoiceItemResponse,
PresetVoiceListResponse,
@@ -38,7 +24,7 @@ from app.schemas.voice_library import (
UpdateVoiceLibraryRequest,
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 packages.adapters.sqlalchemy_impl.voice_clone_profile_repository import SQLAlchemyVoiceCloneProfileRepository
@@ -54,12 +40,8 @@ from packages.application.voice_library.use_cases import (
QuotaExceededError,
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.ports.user_repository import UserRepository
from packages.shared.storage import SharedStorageService
router = APIRouter()
logger = logging.getLogger(__name__)
@@ -105,8 +87,8 @@ def _resolve_preset_preview_url(
_preset_preview_cache[voice_id] = (audio_url, time.time())
logger.info("Preset voice preview generated: %s", voice_id)
return audio_url
except Exception:
logger.exception("Failed to generate preset voice preview: voice_id=%s", voice_id)
except Exception as e:
logger.warning("Failed to generate preview for %s, using fallback: %s", voice_id, e)
return fallback_url
@@ -127,7 +109,6 @@ def _resolve_all_preset_preview_urls(
try:
result_map[p.voice_id] = _resolve_preset_preview_url(p.voice_id, p.preview_url, cosyvoice)
except Exception:
logger.exception("Failed to resolve preset preview URL: voice_id=%s", p.voice_id)
result_map[p.voice_id] = p.preview_url
return result_map
@@ -526,243 +507,3 @@ def delete_voice(
if not deleted:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Voice not found")
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.conf.broker_url = settings.CELERY_BROKER_URL
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 上限
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):
"""用户 pending 任务数超限,返回 429。"""
def __init__(
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,
):
def __init__(self, user_id: str, pending_count: int, limit: int):
self.user_id = user_id
self.pending_count = pending_count
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}")
class GlobalQueueFull(Exception):
"""全局限流,返回 503。"""
def __init__(
self,
pending_count: int,
limit: int,
*,
running_count: int = 0,
queue_ahead: int = 0,
estimated_wait_seconds: int = 0,
):
def __init__(self, pending_count: int, limit: int):
self.pending_count = pending_count
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}")
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(
user_id: str,
generation_task_repository: Any,
@@ -279,17 +161,7 @@ def safe_enqueue_generation_task(
# ── 发送 Celery 任务 ──
try:
celery_result = 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
)
celery_app.send_task("worker.generate_video", args=[task.id])
except Exception as e:
logger.error(
"%s 入队失败,标记为失败: task_id=%s error=%s",
+1 -1
View File
@@ -6,7 +6,7 @@ ensuring proper lifecycle management and testability.
from __future__ import annotations
from collections.abc import Generator
from typing import Generator
import redis
from app.config import settings
+1 -1
View File
@@ -4,7 +4,7 @@
import logging
import time
from collections.abc import Callable
from typing import Callable
from fastapi import Request
from starlette.middleware.base import BaseHTTPMiddleware
@@ -10,7 +10,7 @@ Exposes:
import re
import time
from collections.abc import Callable
from typing import Callable
from fastapi import Request, Response
from prometheus_client import (
-126
View File
@@ -1,126 +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="标题配置(可含 title_image_dataurl:前端 Canvas 渲染的标题 PNG dataURL"
)
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:
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 SmartCoverResponse(BaseModel):
"""智能封面响应(封面从最终成片抽帧,不再叠加标题)."""
cover_url: str = Field("", description="封面图公网 URL(OSS,非临时);失败为空")
status: str = Field("completed", description="completed / fallback_failed")
message: str = Field("", description="失败原因(如有)")
class FinalizeRenderResponse(BaseModel):
"""封面选好后点「完成」,正式入库成片库的响应."""
video_id: str = Field(..., description="成片库视频ID")
cover_url: str = Field("", description="封面URL")
status: str = Field("success", description="success/already_finalized")
+3 -14
View File
@@ -53,14 +53,6 @@ class AssetResponse(BaseModel):
created_at: str
uploaded_by_user_id: str
tag_ids: list[str] = Field(default_factory=list)
# 片段级余量信息(仅视频素材返回,非视频/无时长记录为 None,前端按可用处理)
used_duration: float | None = Field(default=None, description="已使用片段时长(秒,历史区间合并去重后)")
available_duration: float | None = Field(default=None, description="剩余可用时长(秒)= 素材总时长 - 已用时长")
used_ratio: float | None = Field(default=None, description="已用时长占比(0~1")
usable: bool = Field(
default=True,
description="是否仍可用于新片段:零重复可切区间耗尽且所有历史区间复用次数" "use_count)均达上限时为 false",
)
MAX_BATCH_SIZE = 200
@@ -129,13 +121,10 @@ class SmartMatchRequest(BaseModel):
)
class SmartMatchItem(AssetResponse):
"""智能选素材结果条目(扁平结构)。
素材字段(id/usable/余量等)直接挂在条目顶层,前端拿到 item 即可读 item.id
与 AssetResponse 字段完全一致;score/breakdown 为智能匹配附加的评分字段。
"""
class SmartMatchItem(BaseModel):
"""智能选素材结果条目"""
asset: AssetResponse
score: float = Field(..., ge=0, le=100, description="综合得分 0-100")
breakdown: dict[str, float] = Field(default_factory=dict, description="各维度得分明细")
-3
View File
@@ -28,9 +28,6 @@ class DuplicationRecordResponse(BaseModel):
status: str = "pending"
duplicate_rate: float | None = None
duplicate_count: int = 0
# #1661 视觉相似度(归一化 0~1)/ 匹配视频数
visual_similarity: float | None = None
match_count: int | None = None
created_at: str
updated_at: str
-4
View File
@@ -25,10 +25,6 @@ class GeneratedVideoResponse(BaseModel):
review_status: str = "pending_review"
generation_params: dict = Field(default_factory=dict)
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):
+16 -113
View File
@@ -10,7 +10,7 @@ class ConfirmGenerationRequest(BaseModel):
output_width: int = Field(default=1080, ge=100, description="输出视频宽度")
output_height: int = Field(default=1920, ge=100, description="输出视频高度")
cover_url: str = Field(default="", description="自定义封面图片 URL")
custom_title: str = Field(default="", description="用户自定义标题文本,非空时同步到任务和编辑计划")
custom_title: str = Field(default="", description="自定义视频标题")
class CreateGenerationTaskRequest(BaseModel):
@@ -25,25 +25,15 @@ class CreateGenerationTaskRequest(BaseModel):
asset_library_id: str = ""
strategy_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 = ""
# ── 模板模式新增字段 ──
template_id: str = ""
asset_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 = ""
# ── variant-plans 轻量选片回传(#1749):正式生成直接复用,不再重选 ──
variant_plan_ids: list[str] = Field(
default_factory=list,
description="POST /generation/variant-plans 返回的各变体 plan_id(长度须=count);为空则走服务端选片",
)
# ── 标题配置(结构化)──
# ── 标题配置(结构化,优先于 custom_title 纯文本)──
title_config: dict | None = Field(
default=None,
description="标题样式对象,包含 text/font/font_size/font_color/position/bold/stroke/shadow 等。为空时不影响现有行为。",
@@ -55,9 +45,11 @@ class CreateGenerationTaskRequest(BaseModel):
# ── 素材库自动匹配 ──
asset_select_mode: str = Field(
default="all",
description="素材选取模式:all=全部ready视频, smart=智能匹配(按质量/时长评分)",
description="素材选取模式:all=全部ready视频, random=随机选取, smart=智能匹配(按质量/时长评分)",
)
asset_select_count: int = Field(
default=0, ge=0, le=100, description="选取数量,0表示全部(仅 random/smart 模式有效)"
)
asset_select_count: int = Field(default=0, ge=0, le=100, description="选取数量,0表示全部(仅 smart 模式有效)")
# ── 自动重试 ──
auto_retry_enabled: bool = Field(
default=False,
@@ -85,47 +77,7 @@ class CreateGenerationTaskRequest(BaseModel):
output_width: int = Field(default=1280, description="输出视频宽度")
output_height: int = Field(default=720, description="输出视频高度")
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
custom_title: str = Field(default="", description="自定义视频标题")
@model_validator(mode="after")
def _check_at_least_one_mode(self) -> "CreateGenerationTaskRequest":
@@ -134,9 +86,9 @@ class CreateGenerationTaskRequest(BaseModel):
if not has_project and not has_template:
raise ValueError("project_id 或 template_id 至少需要提供一个")
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:
raise ValueError("asset_library_id 或 asset_ids/title_ids 至少需要提供一个")
raise ValueError("asset_library_id 或 asset_ids/title_ids/voice_ids 至少需要提供一个")
return self
@@ -161,6 +113,7 @@ class GenerationTaskResponse(BaseModel):
output_width: int = 1280
output_height: int = 720
cover_url: str = ""
custom_title: str = ""
title_config: dict = Field(default_factory=dict)
status: str
progress: float
@@ -213,6 +166,7 @@ class CreatePreviewGenerationTaskRequest(BaseModel):
template_id: str
asset_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(
default="", description="配音素材库ID(用户上传的音频或AI配音),对应配音选择页面选择的配音素材"
)
@@ -235,44 +189,8 @@ class CreatePreviewGenerationTaskRequest(BaseModel):
)
title_config: dict = Field(
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")
def _check_template_id(self) -> "CreatePreviewGenerationTaskRequest":
@@ -282,13 +200,13 @@ class CreatePreviewGenerationTaskRequest(BaseModel):
@model_validator(mode="after")
def _check_asset_ids(self) -> "CreatePreviewGenerationTaskRequest":
if not self.asset_ids and not self.title_ids:
raise ValueError("asset_ids/title_ids 至少需要提供一个")
if not self.asset_ids and not self.title_ids and not self.voice_ids:
raise ValueError("asset_ids/title_ids/voice_ids 至少需要提供一个")
return self
class PreviewGenerationTaskResponse(BaseModel):
"""单个预览变体任务响应。
"""预览生成任务响应。
包含任务状态、进度、分辨率、生成结果 URL 等关键字段。
"""
@@ -297,7 +215,6 @@ class PreviewGenerationTaskResponse(BaseModel):
status: str
progress: float
is_preview: bool = True
variant_index: int = 0
resolution: str = ""
video_url: str = ""
duration: float = 0.0
@@ -306,21 +223,7 @@ class PreviewGenerationTaskResponse(BaseModel):
transition_count: int = 0
material_usage: dict = Field(default_factory=dict)
error_message: str = ""
title_text: str = ""
voice_library_id: str = ""
created_at: datetime | None = None
started_at: datetime | None = None
finished_at: datetime | None = None
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
-138
View File
@@ -1,138 +0,0 @@
"""对口型 API Schema 定义 — #1796 / #1809 / #1822 / #1845(配音前置).
支持三种输入模式:
1. TTS 直生模式(兼容旧版前端):传 voice_id + script_text+ speed/emotion),
后端 Celery 异步做 TTS 合成 + MediaKit 提交。
2. 直接音频模式:传 video_url + audio_url(音频已由调用方准备好)。
3. 预合成音频模式(#1845 配音前置新主路径):前端先调 POST /lipsync/tts-preview
拿到 audio_url + sentence_timings,再在 create_job 时传 audio_url + audio_duration
+ sentence_timings,后端跳过 TTS 和时间戳计算,直接 ffprobe 校验后提交 MediaKit。
"""
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
sentence_timings: Optional[list] = None
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 必填;audio_url 留空。
- 直接音频:video_url + audio_url 必填。
- 预合成音频(#1845 新主路径):audio_url 必填 + 可选 audio_duration/sentence_timings
后端同步 ffprobe 校验时长、写入 timings,直接提交 MediaKit。
"""
video_url: str = Field(..., description="人物视频 URL(MP4,≤30min,单人真人)")
# 模式 2/3:直接/预合成音频
audio_url: str = Field("", description="驱动音频 URLmp3/aac/wav/m4a/flac);直生模式留空")
audio_duration: Optional[float] = Field(None, ge=0, description="预合成音频时长(秒),可选;后端会 ffprobe 校验")
sentence_timings: Optional[list] = Field(None, description="预合成接口返回的句子时间戳,可选;若传入则直接写入 job")
# 模式 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="情绪(英文枚举 neutral/happy/sad/angry/surprised/fearful/disgusted,或中文 中立/开心/难过/生气/惊讶/恐惧/厌恶;空为默认自然)",
)
enable_video_loop: bool = Field(
True, description="音频长于视频时是否循环画面(AI数字人默认开启,防止音频长于视频被截断)"
)
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]
allowed_video_exts = (".mp4", ".mov", ".m4v", ".webm", ".avi", ".mkv", ".3gp")
if not any(lower.endswith(ext) for ext in allowed_video_exts):
raise ValueError("video_url 格式不支持,仅支持: " + ", ".join(allowed_video_exts))
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
# ── #1845 TTS 预合成接口 ────────────────────────────────────────────────
class AiAvatarTtsPreviewRequest(BaseModel):
"""步骤1「生成配音」预合成请求(同步 HTTP,~2-3s)."""
voice_id: str = Field(..., min_length=1, max_length=128, description="音色 ID")
script_text: str = Field(..., min_length=1, max_length=5000, description="要合成的文案")
speed: float = Field(1.0, ge=0.5, le=2.0, description="语速(0.5-2.0),默认 1.0")
emotion: str = Field(
"neutral",
max_length=32,
description="情绪(英文枚举 neutral/happy/sad/angry/surprised/fearful/disgusted,或中文 中立/开心/难过/生气/惊讶/恐惧/厌恶;默认 neutral)",
)
class AiAvatarTtsPreviewResponse(BaseModel):
"""TTS 预合成响应(临时 URL,24h 内有效,足够当前会话使用)."""
audio_url: str = Field(..., description="CosyVoice 临时音频 URL")
duration: float = Field(..., ge=0, description="音频总时长(秒),ffprobe 测得")
sentence_timings: list[dict] = Field(..., description="句子级精确时间戳")
-182
View File
@@ -1,182 +0,0 @@
"""积分 & 会员相关 Pydantic Schema (#1895)"""
from __future__ import annotations
from datetime import datetime
from typing import Any, Optional
from pydantic import BaseModel, Field
# ============ 余额 & 账户 ============
class PointsBalanceResponse(BaseModel):
"""积分余额 + 会员状态"""
balance: int = Field(..., description="当前积分余额")
total_earned: int = Field(..., description="累计获得积分")
total_spent: int = Field(..., description="累计消耗积分")
is_member: bool = Field(default=False, description="是否付费会员")
member_type: Optional[str] = Field(None, description="会员类型: monthly/quarterly/yearly")
member_expires_at: Optional[datetime] = Field(None, description="会员到期时间")
# ============ 流水 ============
class PointsTransactionItem(BaseModel):
"""单条积分流水"""
id: str
type: str = Field(..., description="类型: add/deduct")
source: str = Field(..., description="来源场景")
amount: int
balance_after: int
description: str = ""
ref_id: str = ""
created_at: Optional[str] = None
class PointsTransactionsResponse(BaseModel):
"""积分流水分页响应"""
items: list[PointsTransactionItem]
total: int
page: int
page_size: int
# ============ 规则 & 积分包 ============
class PointRuleItem(BaseModel):
"""单条积分规则"""
scene_key: str
name: str
base_points: int
unit: str
extra_per_30s: Optional[int] = None
class PointsRulesResponse(BaseModel):
"""所有积分消耗规则"""
rules: list[PointRuleItem]
free_user_multiplier: float = Field(..., description="免费用户积分上浮系数")
class PointsPackageItem(BaseModel):
"""积分包信息"""
code: str
name: str
points: int
price_cents: int
unit_price: str = Field("", description="单价描述,如 ¥0.099/积分")
class PointsPackagesResponse(BaseModel):
"""可购买的积分包列表"""
packages: list[PointsPackageItem]
user_discount: Optional[float] = Field(None, description="当前用户折扣(会员)")
# ============ 消费前检查 ============
class PointsCheckRequest(BaseModel):
"""消费前余额检查请求"""
scene_key: str
duration_minutes: Optional[float] = None
quantity: Optional[int] = 1
class PointsCheckResponse(BaseModel):
"""消费前余额检查响应"""
allowed: bool
required_points: int
current_balance: int
remaining_after: int
is_free_quota: bool = False
# ============ 手动扣减 / 退还(内部接口) ============
class PointsDeductRequest(BaseModel):
"""积分扣减请求"""
scene_key: str
amount: int
description: Optional[str] = ""
ref_id: Optional[str] = ""
class PointsRefundRequest(BaseModel):
"""积分退还请求"""
transaction_id: str
reason: Optional[str] = ""
class PointsRechargeRequest(BaseModel):
"""积分充值请求"""
package_id: str = Field(..., description="积分包 code,如 starter_pack")
# ============ 订单 ============
class PointsOrderResponse(BaseModel):
"""订单信息"""
id: str
order_type: str
product_code: str
amount_cents: int
status: str
created_at: Optional[str] = None
# ============ 每日额度 ============
class DailyUsageResponse(BaseModel):
"""今日免费额度使用情况"""
free_clips_used: int
free_clips_limit: int
free_clips_remaining: int
reset_at: str
# ============ 会员状态(聚合) ============
class MembershipStatusResponse(BaseModel):
"""当前用户会员状态(聚合信息)"""
is_member: bool
member_type: Optional[str] = None
member_expires_at: Optional[datetime] = None
points_balance: int
max_resolution: str = Field(
default="1080p",
description="可用最高分辨率: 720p(free) / 1080p(paid)",
)
# ============ 通用响应 ============
class SimpleMessageResponse(BaseModel):
"""简单消息响应"""
success: bool
message: str
data: Optional[dict[str, Any]] = None
-45
View File
@@ -1,45 +0,0 @@
"""Script (口播文案库) Pydantic schemas — Issue #1795."""
from __future__ import annotations
from datetime import datetime
from typing import 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
-60
View File
@@ -1,60 +0,0 @@
"""Scripts AI 能力 Pydantic schemas — Issue #1893.
抖音文案提取、AI 改写、AI 标题生成的请求/响应模型。
"""
from __future__ import annotations
from typing import List, Optional
from pydantic import BaseModel, Field
# ── 抖音文案提取 ─────────────────────────────────────────────────────────────
class ExtractFromDouyinRequest(BaseModel):
"""从抖音视频提取文案请求."""
url: str = Field(..., description="抖音视频链接(短链或长链)")
class ExtractFromDouyinResponse(BaseModel):
"""从抖音视频提取文案响应."""
text: str = Field(..., description="ASR 识别出的文案文本")
duration_seconds: float = Field(..., description="视频时长(秒)")
source_url: str = Field(..., description="原始视频链接")
# ── AI 改写 ─────────────────────────────────────────────────────────────────
class AiRewriteRequest(BaseModel):
"""AI 文案改写请求."""
content: str = Field(..., description="原文内容")
style: Optional[str] = Field("口语化", description="改写风格,如 口语化/正式/活泼")
class AiRewriteResponse(BaseModel):
"""AI 文案改写响应."""
original: str = Field(..., description="原文")
rewritten: str = Field(..., description="改写后的文案")
style: str = Field(..., description="使用的改写风格")
# ── AI 标题生成 ──────────────────────────────────────────────────────────────
class AiGenerateTitlesRequest(BaseModel):
"""AI 标题生成请求."""
content: str = Field(..., description="文案内容")
count: int = Field(3, ge=1, le=5, description="生成标题数量(1-5,默认3")
class AiGenerateTitlesResponse(BaseModel):
"""AI 标题生成响应."""
titles: List[str] = Field(..., description="生成的标题列表")
+84 -22
View File
@@ -1,14 +1,9 @@
"""Template API schemas(精简版:仅保留列表接口 + 默认模板自动兜底所需字段).
前端 PR#1911 删除 my-templates / editing-planner / templates 管理页后,
模板 CRUD / 分类 / 标签 / 收藏 / 复制 / 校验 / 使用统计等端点全部下线,
对应 Request/Response 模型也一并清理。
"""
"""Template API schemas."""
from __future__ import annotations
from datetime import datetime
from typing import Any, Optional
from typing import Any, Dict, List, Optional
from pydantic import BaseModel, Field
@@ -42,12 +37,12 @@ class TemplateResponse(BaseModel):
name: str
mode: str
category: str = ""
tags: list[str] = Field(default_factory=list)
title_config: dict[str, Any] = Field(default_factory=dict)
subtitle_config: dict[str, Any] = Field(default_factory=dict)
bgm_config: dict[str, Any] = Field(default_factory=dict)
tags: List[str] = Field(default_factory=list)
title_config: Dict[str, Any] = Field(default_factory=dict)
subtitle_config: Dict[str, Any] = Field(default_factory=dict)
bgm_config: Dict[str, Any] = Field(default_factory=dict)
estimated_duration: float = 0.0
segments: list[SegmentResponse] = Field(default_factory=list)
segments: List[SegmentResponse] = Field(default_factory=list)
is_active: bool = True
is_favorite: bool = False
usage_count: int = 0
@@ -55,29 +50,96 @@ class TemplateResponse(BaseModel):
updated_at: datetime
class ToggleFavoriteResponse(BaseModel):
id: str
is_favorite: bool
class ListTemplatesResponse(BaseModel):
items: list[TemplateResponse]
items: List[TemplateResponse]
total: int = 0
# ── Template Request(保留给内部 _get_or_create_default_template_id 兜底创建默认模板使用)──
# ── Template Request ──
class CreateTemplateRequest(BaseModel):
name: str
mode: str
category: str = ""
tags: list[str] = Field(default_factory=list)
title_config: dict[str, Any] = Field(default_factory=dict)
subtitle_config: dict[str, Any] = Field(default_factory=dict)
bgm_config: dict[str, Any] = Field(default_factory=dict)
tags: List[str] = Field(default_factory=list)
title_config: Dict[str, Any] = Field(default_factory=dict)
subtitle_config: Dict[str, Any] = Field(default_factory=dict)
bgm_config: Dict[str, Any] = Field(default_factory=dict)
estimated_duration: float = 0.0
segments: list[SegmentRequest] = Field(default_factory=list)
segments: List[SegmentRequest] = Field(default_factory=list)
class UpdateTemplateRequest(BaseModel):
name: Optional[str] = None
mode: Optional[str] = None
category: Optional[str] = None
tags: Optional[List[str]] = None
title_config: Optional[Dict[str, Any]] = None
subtitle_config: Optional[Dict[str, Any]] = None
bgm_config: Optional[Dict[str, Any]] = None
estimated_duration: Optional[float] = None
segments: Optional[List[SegmentRequest]] = None
# ── Validate ──
class ValidateTemplateRequest(BaseModel):
voiceover_duration: Optional[float] = None # 配音实际时长(秒)
class GenerateWarningResponse(BaseModel):
"""兼容老 import(如校验逻辑内部复用);模板管理页已下线,可按需进一步清理。"""
code: str
message: str
details: dict[str, Any] = Field(default_factory=dict)
details: Dict[str, Any] = Field(default_factory=dict)
class ValidateTemplateResponse(BaseModel):
template: TemplateResponse
warnings: List[GenerateWarningResponse] = Field(default_factory=list)
# ── Category ──
class CategoryResponse(BaseModel):
id: str
user_id: str
name: str
created_at: datetime
class CreateCategoryRequest(BaseModel):
name: str
class ListCategoriesResponse(BaseModel):
items: List[CategoryResponse]
# ── Copy Template ──
class CopyTemplateRequest(BaseModel):
new_name: str
# ── Tags ──
class ListTagsResponse(BaseModel):
items: List[str]
# ── Usage Stats ──
class TemplateUsageResponse(BaseModel):
template_id: str
usage_count: int
+4 -4
View File
@@ -3,7 +3,7 @@
from __future__ import annotations
from datetime import datetime
from typing import Optional
from typing import List, Optional
from pydantic import BaseModel, Field
@@ -15,7 +15,7 @@ class TitleLibraryItemResponse(BaseModel):
text: str
category: str = "default"
description: str = ""
tags: list[str] = Field(default_factory=list)
tags: List[str] = Field(default_factory=list)
usage_count: int = 0
is_active: bool = True
created_at: datetime
@@ -32,7 +32,7 @@ class CreateTitleLibraryRequest(BaseModel):
text: str = Field(..., min_length=1, max_length=500)
category: str = "default"
description: str = ""
tags: list[str] = Field(default_factory=list)
tags: List[str] = Field(default_factory=list)
class UpdateTitleLibraryRequest(BaseModel):
@@ -40,4 +40,4 @@ class UpdateTitleLibraryRequest(BaseModel):
text: Optional[str] = Field(None, min_length=1, max_length=500)
category: Optional[str] = None
description: Optional[str] = None
tags: Optional[list[str]] = None
tags: Optional[List[str]] = None
+4 -26
View File
@@ -3,7 +3,7 @@
from __future__ import annotations
from datetime import datetime
from typing import Any, Optional
from typing import Any, Dict, List, Optional
from pydantic import BaseModel, Field
@@ -16,14 +16,10 @@ class TTSSynthesizeRequest(BaseModel):
output_name: str = Field("", description="输出文件名")
language: str = Field("zh-CN", description="语言")
speed: float = Field(1.0, ge=0.5, le=2.0, description="语速")
emotion: str = Field(
"",
description="情绪(中文/英文:自然/兴奋/沉稳/亲切/开心/悲伤/愤怒/惊讶/恐惧/厌恶 等;通过 instruction 自然语言指令控制)",
)
voice_model: str = Field("", description="语音模型名称")
voice_clone_profile_id: str = Field("", description="关联的音色克隆档案 ID")
format: str = Field("mp3", description="输出格式(mp3/wav/pcm")
metadata_: Optional[dict[str, Any]] = Field(default=None, alias="metadata", description="额外元数据")
metadata_: Optional[Dict[str, Any]] = Field(default=None, alias="metadata", description="额外元数据")
class Config:
populate_by_name = True
@@ -49,7 +45,7 @@ class TTSJobResponse(BaseModel):
error_message: str = ""
retry_count: int = 0
max_retries: int = 3
metadata_: Optional[dict[str, Any]] = Field(default=None, alias="metadata", description="额外元数据")
metadata_: Optional[Dict[str, Any]] = Field(default=None, alias="metadata", description="额外元数据")
started_at: Optional[datetime] = None
completed_at: Optional[datetime] = None
created_at: datetime
@@ -83,7 +79,7 @@ class TTSSynthesizeResponse(BaseModel):
class ListTTSJobResponse(BaseModel):
"""TTS 任务列表响应。"""
items: list[TTSJobResponse]
items: List[TTSJobResponse]
total: int
page: int
page_size: int
@@ -105,21 +101,3 @@ class SaveToLibraryResponse(BaseModel):
voice_id: str
voice_name: str
status: str
class TTSPreviewRequest(BaseModel):
"""TTS 预览(试听)请求。"""
text: str = Field(..., min_length=1, max_length=200, description="合成文本,限制 200 字")
voice_id: str = Field(..., min_length=1, description="音色 ID")
speed: float = Field(1.0, ge=0.5, le=2.0, description="语速")
emotion: str = Field("", description="情绪(中文/英文:自然/兴奋/沉稳/亲切/开心/悲伤/愤怒/惊讶/恐惧/厌恶 等)")
language: str = Field("zh-CN", description="语言(zh-CN/en-US 等)")
pitch: float = Field(1.0, ge=0.5, le=2.0, description="音调(预留,当前未使用)")
class TTSPreviewResponse(BaseModel):
"""TTS 预览(试听)响应。"""
audio_url: str = Field(..., description="合成音频 URL")
duration: Optional[float] = Field(default=None, 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)
file_size: int = Field(..., gt=0)
file_hash: str = Field(default="", max_length=64, description="文件 MD5 哈希,用于去重检测")
client_upload_id: str = Field(default="", max_length=64, description="客户端幂等 token(同一次上传的重试保持一致)")
class DirectUploadPrepareResponse(BaseModel):
@@ -26,25 +25,20 @@ class DirectUploadPrepareResponse(BaseModel):
expires_at: str
fields: dict[str, str]
max_size_bytes: int
duplicated: bool = False
skip_transfer: bool = False
asset_id: str = ""
class DirectUploadCompleteRequest(BaseModel):
project_id: str = Field(..., min_length=1)
library_id: str = Field(..., min_length=1)
storage_key: str = Field(..., min_length=1, max_length=255)
file_hash: str = Field(default="", max_length=64, description="文件哈希,用于去重检测")
client_upload_id: str = Field(default="", max_length=64, description="客户端幂等 token(同一次上传的重试保持一致)")
file_size: int = Field(default=0, ge=0, description="文件大小(字节),用于无 hash 时的兜底去重")
file_hash: str = Field(default="", max_length=64, description="文件 MD5 哈希,用于去重检测")
class DirectUploadCompleteResponse(BaseModel):
storage_key: str
ingest_job_id: str
duplicated: bool = Field(default=False, description="是否为重复素材/重复 complete(命中幂等去重)")
asset_id: str = Field(default="", description="素材 asset_id重复 complete 时返回已存在记录")
duplicated: bool = Field(default=False, description="是否为重复素材(命中去重)")
asset_id: str = Field(default="", description="重复素材 asset_idduplicated=true 时返回")
url: str = Field(default="", description="Public URL of uploaded file")
@@ -52,5 +46,5 @@ class UploadAssetResponse(BaseModel):
storage_key: str
ingest_job_id: str
url: str = Field(..., description="Public URL of uploaded file")
duplicated: bool = Field(default=False, description="是否为重复素材/重复提交(命中幂等去重)")
asset_id: str = Field(default="", description="素材 asset_id重复提交时返回已存在记录")
duplicated: bool = Field(default=False, description="是否为重复素材(命中去重)")
asset_id: str = Field(default="", description="重复素材 asset_idduplicated=true 时返回")
-4
View File
@@ -22,10 +22,6 @@ class VideoItemResponse(BaseModel):
generation_params: dict = Field(default_factory=dict)
download_url: str | None = None
generated_at: str = ""
# #1660 查重率(百分比 0~100)/ 视觉相似度(0~1)/ 匹配帧数
duplicate_rate: float | None = None
visual_similarity: float | None = None
match_count: int | None = None
class ListVideosResponse(BaseModel):
+2 -2
View File
@@ -3,7 +3,7 @@
from __future__ import annotations
from datetime import datetime
from typing import Optional
from typing import List, Optional
from pydantic import BaseModel, Field
@@ -61,7 +61,7 @@ class ShareResponse(BaseModel):
class ShareListResponse(BaseModel):
"""分享列表响应."""
items: list[ShareResponse]
items: List[ShareResponse]
total: int = 0
skip: int = 0
limit: int = 20
+3 -3
View File
@@ -6,7 +6,7 @@
from __future__ import annotations
from datetime import datetime
from typing import Literal, Optional
from typing import List, Literal, Optional
from pydantic import BaseModel, Field
@@ -56,7 +56,7 @@ class UnifiedVoiceItemResponse(BaseModel):
status: str = "completed"
"""状态"""
tags: list[str] = Field(default_factory=list)
tags: List[str] = Field(default_factory=list)
"""标签列表"""
# 克隆音色特有字段
@@ -113,7 +113,7 @@ class PresetVoiceItemResponse(BaseModel):
preview_url: str = ""
"""预览音频 URL"""
tags: list[str] = Field(default_factory=list)
tags: List[str] = Field(default_factory=list)
"""标签列表"""
+5 -6
View File
@@ -3,7 +3,7 @@
from __future__ import annotations
from datetime import datetime
from typing import Any, Optional
from typing import Any, Dict, List, Optional
from pydantic import BaseModel, Field
@@ -13,13 +13,12 @@ class CreateVoiceCloneRequest(BaseModel):
name: str = Field(..., min_length=1, max_length=100, description="音色名称")
description: str = Field("", description="音色描述")
source_audio_url: str = Field("", description="参考音频 URL(与 asset_id 二选一)")
asset_id: str = Field("", description="参考音频素材 ID(配音素材库中的音频 asset,与 source_audio_url 二选一)")
source_audio_url: str = Field("", description="参考音频 URL")
voice_model: str = Field("", description="语音模型名称")
language: str = Field("zh-CN", description="语言")
gender: str = Field("unknown", description="性别")
max_retries: int = Field(3, ge=1, le=10, description="最大重试次数")
metadata_: Optional[dict[str, Any]] = Field(default=None, alias="metadata", description="额外元数据")
metadata_: Optional[Dict[str, Any]] = Field(default=None, alias="metadata", description="额外元数据")
class Config:
populate_by_name = True
@@ -41,7 +40,7 @@ class VoiceCloneProfileResponse(BaseModel):
error_message: str = ""
retry_count: int = 0
max_retries: int = 3
metadata_: Optional[dict[str, Any]] = Field(default=None, alias="metadata", description="额外元数据")
metadata_: Optional[Dict[str, Any]] = Field(default=None, alias="metadata", description="额外元数据")
created_at: datetime
updated_at: datetime
@@ -62,7 +61,7 @@ class VoiceCloneStatusResponse(BaseModel):
class ListVoiceCloneResponse(BaseModel):
"""音色克隆列表响应。"""
items: list[VoiceCloneProfileResponse]
items: List[VoiceCloneProfileResponse]
total: int
+4 -4
View File
@@ -3,7 +3,7 @@
from __future__ import annotations
from datetime import datetime
from typing import Optional
from typing import List, Optional
from pydantic import BaseModel, Field
@@ -21,7 +21,7 @@ class VoiceLibraryItemResponse(BaseModel):
file_size: int = 0
status: str = "completed"
project_id: Optional[str] = None
tags: list[str] = Field(default_factory=list)
tags: List[str] = Field(default_factory=list)
created_at: datetime
updated_at: datetime
@@ -42,7 +42,7 @@ class CreateVoiceLibraryRequest(BaseModel):
file_size: int = 0
status: str = "completed"
project_id: Optional[str] = None
tags: list[str] = Field(default_factory=list)
tags: List[str] = Field(default_factory=list)
class UpdateVoiceLibraryRequest(BaseModel):
@@ -55,4 +55,4 @@ class UpdateVoiceLibraryRequest(BaseModel):
duration: Optional[float] = None
file_size: Optional[int] = None
status: Optional[str] = None
tags: Optional[list[str]] = None
tags: Optional[List[str]] = None
@@ -1,215 +0,0 @@
"""AI 数字人封面服务 — MediaKit 抽帧 + 质量评分选最佳帧 + 转存 OSS.
与 generation_cover.py 的智能选帧能力对齐(不再用 FFmpeg 简单截帧):
1. MediaKit extract_frames 抽取多帧(默认 5 帧,SpecifiedFrames 策略)
2. cover_frame_scorer.score_frames 按清晰度/亮度/色彩评分选最佳
3. 下载最佳帧并转存 OSS,返回公网封面 URL
设计原则:封面一律从最终成片(已叠加标题/B-roll)抽帧,帧本身已含标题,
本服务**不再叠加标题**。对口型阶段的裸视频封面入口已删除(废弃)。
降级: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 抽帧轮询参数:poll_interval=2s × max_poll=30 → 最长 60s(与 mediakit_client 默认值/lipsync 轮询保持一致,防止合成视频下载+抽帧超时)
COVER_POLL_INTERVAL = 2.0
COVER_MAX_POLL_ATTEMPTS = 30
# 帧图片下载超时(秒)
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:
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."""
if not video_url:
return ""
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",
video_url[:80],
max_frames,
)
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 ""
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(预签名).
封面来自最终成片抽帧,帧本身已含标题,本函数不再做任何文字/图片叠加。
"""
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)
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。失败返回空字符串。
封面从最终成片抽帧,不再叠加任何标题(帧本身已含)。
"""
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,658 +0,0 @@
"""AI数字人渲染合成 Service — #1798.
职责:
- 创建/查询/取消渲染任务
- 调用 Celery 异步任务执行渲染
- B-roll 合成 + 标题叠加 + 封面提取
- 用户隔离
"""
from __future__ import annotations
import base64
import binascii
import logging
import os
import subprocess
import tempfile
import uuid
from datetime import UTC, datetime
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_broll_overlay_filter,
build_title_drawtext_filter,
build_title_overlay_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(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(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(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. 上传到 OSS (95%) — 封面不再自动生成,改由前端主动抽帧
5. 更新任务状态 (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(UTC)
job.progress = 5
job.updated_at = datetime.now(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%)
# 用 ffprobe 探测输入视频分辨率,确保 B-roll 缩放与标题位置与实际输出一致。
# AI 数字人对口型输出为 9:16 竖屏,默认兜底 720x1280;探测失败时使用默认值不阻断渲染。
output_width, output_height = self._probe_video_resolution(input_video_path)
if output_width <= 0 or output_height <= 0:
output_width, output_height = 720, 1280
logger.info(
"[数字人渲染] ffprobe 探测分辨率失败或无效,使用默认竖屏尺寸 %sx%s",
output_width,
output_height,
)
else:
logger.info("[数字人渲染] 探测输入视频分辨率: %sx%s", output_width, output_height)
broll_filter, broll_label = build_broll_overlay_filter(
b_roll_segments=job.b_roll_segments,
video_duration=lipsync_job.output_duration,
output_width=output_width,
output_height=output_height,
)
# 标题叠加路径:优先前端 Canvas 渲染的 PNG 图层(所见即所得),
# 无 title_image_dataurl 时降级到 drawtext 重画文字。
title_cfg = job.title_config if isinstance(job.title_config, dict) else {}
title_dataurl = (title_cfg or {}).get("title_image_dataurl") if title_cfg else None
use_title_png = isinstance(title_dataurl, str) and title_dataurl.startswith("data:image/")
title_input_index = 1 + len(job.b_roll_segments or []) if use_title_png 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")
# 在临时目录里解码保存标题 PNG(with 退出自动清理)
title_png_path: Optional[str] = None
extra_inputs: list[str] = []
title_filter = None
if use_title_png:
try:
title_png_path = os.path.join(tmpdir, f"title_{job.id}.png")
self._save_title_dataurl_to_file(title_dataurl, dst_path=title_png_path)
extra_inputs.append(title_png_path)
logger.info(
"[数字人渲染] 标题 PNG 已保存: %s (input index %d)", title_png_path, title_input_index
)
except Exception as exc:
logger.warning("[数字人渲染] 标题 PNG 解码/保存失败,降级 drawtext: %s", exc)
title_png_path = None
extra_inputs = []
# 构建标题滤镜
final_label = None
if title_png_path and title_input_index is not None:
title_input_label = f"[{title_input_index}:v]"
base_label = f"[{broll_label}]" if broll_label else "[0:v]"
title_filter = build_title_overlay_filter(
title_cfg,
output_width=output_width,
output_height=output_height,
title_png_path=title_png_path,
title_input_label=title_input_label,
base_label=base_label,
output_label="vout_titled",
)
if not title_filter:
# build 返回 None → 文件不存在(极端并发情况),降级 drawtext
title_png_path = None
extra_inputs = []
if title_png_path:
# overlay 路径
if broll_filter and title_filter:
filter_complex = broll_filter + f";{title_filter}"
elif broll_filter:
filter_complex = broll_filter
final_label = broll_label
elif title_filter:
filter_complex = title_filter
else:
filter_complex = ""
if title_filter:
final_label = "vout_titled"
elif not final_label:
final_label = None
else:
# 降级:drawtext 重画文字
title_filter = build_title_drawtext_filter(
title_cfg,
output_width=output_width,
output_height=output_height,
)
if broll_filter and title_filter:
filter_complex = broll_filter + f";[{broll_label}]{title_filter}[vout_titled]"
final_label = "vout_titled"
elif broll_filter:
filter_complex = broll_filter
final_label = broll_label
elif title_filter:
filter_complex = f"[0:v]{title_filter}[vout_titled]"
final_label = "vout_titled"
else:
filter_complex = ""
final_label = None
cmd_list = self._build_ffmpeg_command(
input_video=input_video_path,
b_roll_segments=job.b_roll_segments,
extra_inputs=extra_inputs,
filter_complex=filter_complex,
final_label=final_label,
output_path=output_video_path,
)
try:
render_result = subprocess.run(
cmd_list,
capture_output=True,
text=True,
timeout=600,
)
except subprocess.TimeoutExpired as exc:
raise AiAvatarRenderError(
"FFmpeg 渲染超时(600s",
code="FFmpegTimeout",
) from exc
if render_result.returncode != 0:
stderr_tail = (render_result.stderr or "").strip()[-800:]
raise AiAvatarRenderError(
f"FFmpeg 渲染失败,退出码: {render_result.returncode}, stderr: {stderr_tail}",
code="FFmpegFailed",
)
job.progress = 80
self.db.commit()
# 4/5. 上传成片到 OSS (95%) —— 已砍掉自动抽封面逻辑(步骤⑤);
# 封面由前端在渲染完成后通过 /smart-cover 接口主动从成片抽帧,不阻塞渲染链路。
output_video_url = self._upload_to_oss(output_video_path, f"ai-avatar/{job_id}/output.mp4")
job.output_video_url = output_video_url
# 封面透传:如果用户已在 cover_config 中选定封面 URLmode=upload 的自定义上传 或
# mode=auto_frame 已有的智能封面结果),直接透传到 output_cover_url,不再重新截帧。
if isinstance(job.cover_config, dict):
_pre_cover_url = (
job.cover_config.get("url")
or job.cover_config.get("imageUrl")
or job.cover_config.get("cover_url")
or ""
)
if _pre_cover_url:
job.output_cover_url = _pre_cover_url
logger.info("[数字人渲染] 使用用户已选定封面 URL: job_id=%s", job_id)
# 获取输出视频时长
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(UTC)
job.updated_at = datetime.now(UTC)
self.db.commit()
logger.info("渲染任务完成: %s", job_id)
# 7. 渲染完成,停留在「待选封面」状态:不自动入库。
# 用户在前端选好封面、点「完成」后,由 /{job_id}/finalize 接口显式入库。
logger.info("渲染任务完成,等待用户选择封面后入库: job_id=%s", job_id)
except AiAvatarRenderError as exc:
job.status = "failed"
job.error_message = str(exc)
job.updated_at = datetime.now(UTC)
self.db.commit()
logger.error("渲染任务失败 [%s]: %s", job_id, exc)
raise
except Exception as exc:
job.status = "failed"
job.error_message = f"渲染异常: {str(exc)}"
job.updated_at = datetime.now(UTC)
self.db.commit()
logger.exception("渲染任务异常 [%s]", job_id)
raise
def _persist_to_library(self, job: AiAvatarRenderJob, cover_url: Optional[str] = None):
"""将渲染结果写入成片库,返回 GeneratedVideo 领域对象.
Args:
job: 渲染任务(必须 status=completed 且 output_video_url 非空)
cover_url: 可选的封面 URL 覆盖(finalize 时传入即优先使用,否则取 job.output_cover_url
"""
from packages.adapters.sqlalchemy_impl.generated_video_repository import (
SQLAlchemyGeneratedVideoRepository,
)
from packages.domain.generated_video import GeneratedVideo
clip_name = f"AI数字人_{job.id[:8]}"
# AI数字人入口是独立页面,前端可能不传 project_id(无项目概念),
# 兜底为 "ai_avatar" 避免 DB 非空约束/查询问题;generation_task_id 用 render_job_id 便于反查。
clip_project_id = (job.project_id or "").strip() or "ai_avatar"
clip_generation_task_id = job.id
effective_cover = (cover_url or "").strip() if cover_url else (job.output_cover_url or "").strip()
clip = GeneratedVideo.create(
project_id=clip_project_id,
generation_task_id=clip_generation_task_id,
name=clip_name,
file_url=job.output_video_url,
user_id=job.user_id,
duration=job.output_duration or 0.0,
thumbnail_url=effective_cover 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)
return clip
def finalize_job(self, job_id: str, user_id: str, cover_url: Optional[str] = None):
"""用户在前端点「完成」后调用:将已 completed 的渲染任务正式入库到成片库.
- 必须 status=completed 才可调用
- cover_url 若传入则优先使用并回写 job.output_cover_url;否则使用 job.output_cover_urlsmart-cover/custom-cover 已写入)
- 幂等:已入库则返回已存在的 GeneratedVideo
"""
from packages.adapters.sqlalchemy_impl.generated_video_repository import (
SQLAlchemyGeneratedVideoRepository,
)
from packages.adapters.sqlalchemy_impl.models import GeneratedVideoModel
job = self.get_render_job(job_id, user_id)
if job is None:
raise AiAvatarRenderError("渲染任务不存在", code="RenderJobNotFound")
if job.status != "completed":
raise AiAvatarRenderError(f"渲染任务未完成(当前状态: {job.status}),无法入库", code="RenderNotCompleted")
if not (job.output_video_url or "").strip():
raise AiAvatarRenderError("渲染成片视频 URL 为空,无法入库", code="OutputVideoMissing")
# 幂等检查:已入库直接返回现有记录(通过 generation_task_id=job_id 识别,
# 因为入库时 generation_task_id 被设置为 render_job_id 自身)
existing = (
self.db.query(GeneratedVideoModel)
.filter(
GeneratedVideoModel.user_id == user_id,
GeneratedVideoModel.generation_task_id == job_id,
)
.first()
)
if existing is not None:
logger.info("[数字人渲染] finalize 幂等命中,返回已存在记录: clip_id=%s job_id=%s", existing.id, job_id)
return SQLAlchemyGeneratedVideoRepository(self.db).get(existing.id)
# 传入 cover_url 时回写到 job
if cover_url and cover_url.strip():
job.output_cover_url = cover_url.strip()
# 同步更新 cover_config,保持 smart-cover 路径一致
if isinstance(job.cover_config, dict):
job.cover_config = {**job.cover_config, "mode": "auto_frame", "url": cover_url.strip()}
job.updated_at = datetime.now(UTC)
self.db.commit()
return self._persist_to_library(job, cover_url=cover_url)
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
@staticmethod
def _save_title_dataurl_to_file(dataurl: str, *, dst_path: str | None = None, job_id: str = "") -> str:
"""解码前端传来的 data:image/png;base64,... 并保存为本地 PNG 文件。
Args:
dataurl: 完整 dataURL 字符串
dst_path: 指定输出路径;为 None 时创建临时文件并返回路径
job_id: 仅在 dst_path 为空时用于临时文件命名
Returns:
保存后的本地文件路径
"""
if not isinstance(dataurl, str) or not dataurl.startswith("data:image/"):
raise ValueError("title_image_dataurl 不是合法的 data:image URL")
# 拆分 data:image/png;base64,<payload>
try:
header, b64 = dataurl.split(",", 1)
except ValueError as exc:
raise ValueError("title_image_dataurl 缺少 base64 payload") from exc
if "base64" not in header:
raise ValueError("title_image_dataurl 不是 base64 编码")
try:
png_bytes = base64.b64decode(b64, validate=True)
except (binascii.Error, ValueError) as exc:
raise ValueError(f"title_image_dataurl base64 解码失败: {exc}") from exc
if not png_bytes:
raise ValueError("title_image_dataurl 解码后为空")
if dst_path:
out_path = dst_path
with open(out_path, "wb") as f:
f.write(png_bytes)
return out_path
suffix = f"_title_{job_id}.png" if job_id else "_title.png"
with tempfile.NamedTemporaryFile(suffix=suffix, delete=False) as tmp:
tmp.write(png_bytes)
return tmp.name
@staticmethod
def _probe_video_resolution(video_path: str) -> tuple[int, int]:
"""用 ffprobe 探测视频分辨率,返回 (width, height);失败返回 (0, 0)。"""
try:
result = subprocess.run(
[
"ffprobe",
"-v",
"error",
"-select_streams",
"v:0",
"-show_entries",
"stream=width,height",
"-of",
"csv=p=0:s=x",
video_path,
],
capture_output=True,
text=True,
timeout=15,
)
if result.returncode == 0 and result.stdout.strip():
parts = result.stdout.strip().split("x")
if len(parts) == 2:
w, h = int(parts[0]), int(parts[1])
if w > 0 and h > 0:
return w, h
except Exception as exc:
logger.warning("[数字人渲染] ffprobe 探测分辨率失败: %s", exc)
return 0, 0
def _build_ffmpeg_command(
self,
*,
input_video: str,
b_roll_segments: list[dict[str, Any]],
extra_inputs: list[str] | None = None,
filter_complex: str,
final_label: Optional[str],
output_path: str,
) -> list[str]:
"""构建 FFmpeg 命令(list 形式,shell=False.
根因修复 #1798 P0OSS 预签名 URL 含 `&Expires=...&Signature=...` 特殊字符,
os.system(shell=True) 会把 `&` 解释为后台命令分隔符,导致 -filter_complex 被
当成独立命令报 sh: -filter_complex: not foundexit 127 → Python 32512)。
list + shell=False 彻底规避 shell 转义问题。
"""
cmd: list[str] = ["ffmpeg", "-i", input_video]
for seg in b_roll_segments:
asset_url = seg.get("asset_url", "")
if asset_url:
cmd.extend(["-i", asset_url])
# 额外输入(例如前端 Canvas 渲染的标题 PNG)
for extra in extra_inputs or []:
cmd.extend(["-i", extra])
if filter_complex and final_label:
cmd.extend(
[
"-filter_complex",
filter_complex,
"-map",
f"[{final_label}]",
"-map",
"0:a?",
]
)
elif filter_complex:
cmd.extend(["-filter_complex", filter_complex])
cmd.extend(
[
"-c:v",
"libx264",
"-preset",
"veryfast",
"-crf",
"23",
"-c:a",
"aac",
"-b:a",
"128k",
"-y",
output_path,
]
)
return cmd
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
+12 -12
View File
@@ -13,7 +13,7 @@
from __future__ import annotations
import logging
from typing import Any, Optional
from typing import Any, Dict, List, Optional
from packages.domain.ai_parsing import generate_titles_fallback as _generate_titles_fallback_base
from packages.domain.ai_parsing import keyword_match_fallback as _semantic_match_fallback_base
@@ -64,7 +64,7 @@ def _generate_titles_fallback(
description: str,
style: str = "viral",
count: int = 5,
) -> list[str]:
) -> List[str]:
"""本地降级:基于模板规则生成标题(薄包装,转发到 ai_parsing 模块)."""
style_info = TITLE_STYLES.get(style, TITLE_STYLES["viral"])
return _generate_titles_fallback_base(description, style_info, count)
@@ -74,7 +74,7 @@ def generate_smart_titles(
description: str,
style: str = "viral",
count: int = 5,
) -> dict[str, Any]:
) -> Dict[str, Any]:
"""生成智能标题.
Args:
@@ -164,16 +164,16 @@ def generate_smart_titles(
def _semantic_match_fallback(
description: str,
assets: list[dict[str, Any]],
) -> list[dict[str, Any]]:
assets: List[Dict[str, Any]],
) -> List[Dict[str, Any]]:
"""本地降级:基于关键词的简单匹配(薄包装,转发到 ai_parsing 模块)."""
return _semantic_match_fallback_base(description, assets)
def _parse_semantic_match_response(
content: str,
asset_ids: list[str],
) -> Optional[dict[str, float]]:
asset_ids: List[str],
) -> Optional[Dict[str, float]]:
"""从模型返回中解析素材匹配度(薄包装,转发到 ai_parsing 模块)."""
result = _parse_semantic_match_base(content, asset_ids)
if result is None:
@@ -183,9 +183,9 @@ def _parse_semantic_match_response(
def semantic_match_assets(
description: str,
assets: list[dict[str, Any]],
assets: List[Dict[str, Any]],
top_k: int = 0,
) -> dict[str, Any]:
) -> Dict[str, Any]:
"""智能素材语义匹配.
根据用户描述,评估每个素材的语义匹配度并排序。
@@ -336,13 +336,13 @@ class AIService:
description: str,
style: str = "viral",
count: int = 5,
) -> dict[str, Any]:
) -> Dict[str, Any]:
return generate_smart_titles(description, style, count)
def semantic_match(
self,
description: str,
assets: list[dict[str, Any]],
assets: List[Dict[str, Any]],
top_k: int = 0,
) -> dict[str, Any]:
) -> Dict[str, Any]:
return semantic_match_assets(description, assets, top_k)
@@ -1,494 +0,0 @@
"""素材片段级使用记录追踪与受控复用.
在素材 metadataassets.classification_result JSON)中持久化已使用的片段时间区间,
供 from-assets 创建片段时避开历史区间,实现跨任务/跨调用的片段去重;
素材可用区间耗尽后进入受控复用:允许有限次数(MAX_RANGE_USE_COUNT)复用最久未用
的历史区间,配合调用方的成片复用占比控制(MAX_REUSE_RATIO = 10%),把任意两条
成片的画面重复率控制在阈值内。
metadata 中的记录字段 ``used_time_ranges``::
"used_time_ranges": [
{
"start": 12.5, "end": 20.3,
"plan_id": "plan-xxx",
"created_at": "2026-08-29T12:00:00+00:00",
"use_count": 1, # 该区间累计被使用次数(复用一次 +1)
"last_used_at": "2026-08-29T12:00:00+00:00" # 最近一次使用时间
},
...
]
注意:本模块所有函数都不自行 commit,由调用方控制事务边界
from-assets 与 replace_all_clips_transactional 同事务;异步任务各自 commit)。
历史记录永不自动清空(自动轮回重置已下线,reset_used_segments 仅保留给运维/测试)。
"""
from __future__ import annotations
import json
import logging
from collections.abc import Callable
from datetime import UTC, datetime
from sqlalchemy.orm import Session
from packages.adapters.sqlalchemy_impl.models import AssetModel
logger = logging.getLogger(__name__)
USED_RANGES_KEY = "used_time_ranges"
# ── 受控复用配置常量 ─────────────────────────────────────────────────────────
MAX_RANGE_USE_COUNT = 2
"""单条历史区间最多被使用次数(含首次),达到后不再参与复用。"""
REUSE_RATIO_LIMIT = 0.10
"""单条成片中,单个素材的复用片段累计时长 / 该素材在成片中的总时长上限(10%)。
超过则该素材不再分配新片段(调用方在轮询分配时跳过)。"""
SEGMENT_EDGE_GAP = 1.5
"""冲突判定边缘间隙(秒):历史区间按 [start-gap, end+gap] 扩边后参与冲突检测,
避免两条片段首尾紧贴导致画面观感重复;记录仍存实际值。"""
# 判定"新片段与历史区间为同一次使用(复用)"的重叠率阈值:
# 重叠时长 / 新区间时长超过该比例视为复用该历史区间(累加 use_count)而非新增记录。
_REUSE_OVERLAP_RATIO = 0.6
def _now_iso() -> str:
return datetime.now(UTC).isoformat()
def _read_meta(model) -> dict:
"""读取素材 metadata dict。
兼容两种对象:
- ORM ``AssetModel``metadata 以 JSON 字符串存在 ``classification_result`` 列;
- 领域实体 ``Asset``(路由层 repository 返回):metadata 直接是 dict 属性
repository 与 classification_result 互转,见 asset_repository.py)。
"""
# 领域实体:metadata 已是 dict
meta = getattr(model, "metadata", None)
if isinstance(meta, dict):
return meta
raw = getattr(model, "classification_result", None)
if not raw:
return {}
try:
data = json.loads(raw) if isinstance(raw, str) else raw
return data if isinstance(data, dict) else {}
except Exception:
return {}
def _get_model(db: Session, asset_id: str, for_update: bool = False) -> AssetModel | None:
query = db.query(AssetModel).filter(AssetModel.id == asset_id)
if for_update:
# 行级锁(PostgreSQL SELECT ... FOR UPDATE):序列化同一素材的
# classification_result 读-改-写,避免并发事务丢失使用记录。
# SQLite 不支持时 SQLAlchemy 会忽略该子句(no-op)。
query = query.with_for_update()
return query.first()
def get_used_segments(db: Session, asset_ids: list[str]) -> dict[str, list[tuple[float, float]]]:
"""聚合多个素材的历史已用片段区间。
Returns:
``{asset_id: [(start, end), ...]}`` 格式,与 ``_calc_random_start_time`` 的
``used_segments`` 参数格式一致,可直接传入。
"""
if not asset_ids:
return {}
result: dict[str, list[tuple[float, float]]] = {}
models = db.query(AssetModel).filter(AssetModel.id.in_(list(set(asset_ids)))).all()
for model in models:
meta = _read_meta(model)
ranges = meta.get(USED_RANGES_KEY) or []
segments: list[tuple[float, float]] = []
for r in ranges:
try:
segments.append((float(r["start"]), float(r["end"])))
except (KeyError, TypeError, ValueError):
continue
if segments:
result[model.id] = segments
return result
def record_used_segments(
db: Session,
asset_id: str,
start: float,
end: float,
plan_id: str,
) -> None:
"""记录一次片段使用(不 commit.
若新区间与某条历史区间高度重叠(复用场景,如受控复用回调返回的区间、
MediaKit 挪到历史区间),则累加该记录的 ``use_count`` 并刷新 ``last_used_at``
不新增记录;否则追加一条新记录(use_count=1)。
"""
# 行级锁读取:与并发生成任务互斥,保证区间记录读-改-写一致
model = _get_model(db, asset_id, for_update=True)
if model is None:
logger.warning("[片段追踪] 素材不存在,跳过记录: asset_id=%s", asset_id)
return
meta = _read_meta(model)
ranges = list(meta.get(USED_RANGES_KEY) or [])
new_start = round(float(start), 3)
new_end = round(float(end), 3)
new_dur = max(new_end - new_start, 1e-6)
now = _now_iso()
for r in ranges:
try:
rs, re_ = float(r["start"]), float(r["end"])
except (KeyError, TypeError, ValueError):
continue
overlap = max(0.0, min(new_end, re_) - max(new_start, rs))
if overlap / new_dur >= _REUSE_OVERLAP_RATIO:
# 复用同一条历史区间:累加次数、刷新时间
r["use_count"] = int(r.get("use_count", 1)) + 1
r["last_used_at"] = now
r["plan_id"] = plan_id
meta[USED_RANGES_KEY] = ranges
model.classification_result = json.dumps(meta, ensure_ascii=False)
model.updated_at = datetime.now(UTC)
return
ranges.append(
{
"start": new_start,
"end": new_end,
"plan_id": plan_id,
"created_at": now,
"use_count": 1,
"last_used_at": now,
}
)
meta[USED_RANGES_KEY] = ranges
model.classification_result = json.dumps(meta, ensure_ascii=False)
model.updated_at = datetime.now(UTC)
def remove_used_segment(
db: Session,
asset_id: str,
start: float,
end: float,
plan_id: str | None = None,
tolerance: float = 0.5,
) -> bool:
"""删除素材 metadata 中匹配的一条使用记录(不 commit).
匹配规则:start/end 与记录值相差不超过 tolerance 秒;plan_id 非空时,
记录有 plan_id 则需相等,记录缺 plan_id(本功能上线前的旧数据)时按时间匹配。
Returns:
是否找到并删除了记录。
"""
model = _get_model(db, asset_id)
if model is None:
return False
meta = _read_meta(model)
ranges = list(meta.get(USED_RANGES_KEY) or [])
remaining: list[dict] = []
removed = False
for r in ranges:
try:
match = (
abs(float(r["start"]) - float(start)) <= tolerance and abs(float(r["end"]) - float(end)) <= tolerance
)
except (KeyError, TypeError, ValueError):
remaining.append(r)
continue
# plan_id 校验:传入 plan_id 时,记录有 plan_id 则必须相等;
# 记录本身缺 plan_id(旧数据)时退化为按时间匹配,避免旧区间永远删不掉
if plan_id is not None and r.get("plan_id") is not None and r.get("plan_id") != plan_id:
match = False
if match and not removed:
removed = True
continue
remaining.append(r)
if removed:
meta[USED_RANGES_KEY] = remaining
model.classification_result = json.dumps(meta, ensure_ascii=False)
model.updated_at = datetime.now(UTC)
return removed
def reset_used_segments(db: Session, asset_id: str) -> None:
"""清空单个素材的历史片段使用记录(不 commit).
仅供运维/测试使用;正常生成流程中历史记录永不自动清空(受控复用取代自动轮回)。
"""
model = _get_model(db, asset_id)
if model is None:
return
meta = _read_meta(model)
if meta.get(USED_RANGES_KEY):
meta[USED_RANGES_KEY] = []
model.classification_result = json.dumps(meta, ensure_ascii=False)
model.updated_at = datetime.now(UTC)
logger.info("[片段追踪] 素材区间记录手动清空: asset_id=%s", asset_id)
# ── 素材余量/可用性计算(Task H:素材库角标 + smart-match 过滤)──────────────
# 判定「是否还有空闲可切区间」时使用的最小片段时长(秒):空闲段长于此值才视为可切
_MIN_FREE_CLIP_DURATION = 3.0
def _merge_intervals(intervals: list[tuple[float, float]]) -> list[tuple[float, float]]:
"""合并重叠/相接的时间区间,返回升序不重叠区间列表。"""
if not intervals:
return []
ordered = sorted((float(a), float(b)) for a, b in intervals if b > a)
merged: list[tuple[float, float]] = [ordered[0]]
for start, end in ordered[1:]:
last_start, last_end = merged[-1]
if start <= last_end:
merged[-1] = (last_start, max(last_end, end))
else:
merged.append((start, end))
return merged
def _has_free_gap(used: list[tuple[float, float]], total: float, min_free: float = _MIN_FREE_CLIP_DURATION) -> bool:
"""素材 [0, total] 中是否存在长度 ≥ min_free 的空闲段(考虑边缘间隙)。"""
if total <= 0:
return False
# 历史区间按边缘间隙扩边后判定空闲(与选片冲突检测同一口径)
expanded = [(max(0.0, s - SEGMENT_EDGE_GAP), min(total, e + SEGMENT_EDGE_GAP)) for s, e in used]
merged = _merge_intervals(expanded)
cursor = 0.0
for start, end in merged:
if start - cursor >= min_free:
return True
cursor = max(cursor, end)
return total - cursor >= min_free
def compute_asset_availability(
model: "AssetModel | None",
min_free_clip_duration: float = _MIN_FREE_CLIP_DURATION,
) -> dict | None:
"""计算单个素材的余量与可用性(纯函数,不读写 DB)。
Returns:
视频素材返回 ``{"used_duration", "available_duration", "used_ratio", "usable"}``
非视频 / 无 model / 无时长信息返回 None(调用方按可用处理,零影响)。
usable=False 条件(与受控复用机制一致):
零重复可切区间已耗尽(不存在 ≥ min_free 的空闲段)且
所有历史区间 use_count 均达 MAX_RANGE_USE_COUNT 上限(无区间可复用)。
"""
if model is None:
return None
file_type = getattr(model, "file_type", None) or getattr(model, "mime_type", "") or ""
if file_type != "video" and not str(file_type).startswith("video/"):
return None
total = float(getattr(model, "duration", 0.0) or 0.0)
if total <= 0:
return None
meta = _read_meta(model)
raw_ranges = meta.get(USED_RANGES_KEY) or []
intervals: list[tuple[float, float]] = []
use_counts: list[int] = []
for r in raw_ranges:
try:
start = float(r["start"])
end = float(r["end"])
except (KeyError, TypeError, ValueError):
continue
if end <= start:
continue
intervals.append((start, end))
try:
use_counts.append(int(r.get("use_count", 1)))
except (TypeError, ValueError):
use_counts.append(1)
merged = _merge_intervals(intervals)
used_duration = round(sum(e - s for s, e in merged), 3)
used_duration = min(used_duration, total)
available_duration = round(max(total - used_duration, 0.0), 3)
used_ratio = round(min(used_duration / total, 1.0), 4)
has_free = _has_free_gap(intervals, total, min_free_clip_duration)
if has_free:
usable = True
else:
# 空闲段耗尽:仅当存在历史区间且全部达复用上限时才判定不可用;
# 无历史区间(理论上不会走到,因为 has_free=True)按可用处理
if not use_counts:
usable = True
else:
usable = any(uc < MAX_RANGE_USE_COUNT for uc in use_counts)
return {
"used_duration": used_duration,
"available_duration": available_duration,
"used_ratio": used_ratio,
"usable": usable,
}
def find_reusable_range(
db: Session,
asset_id: str,
clip_duration: float,
asset_total: float,
*,
max_use_count: int = MAX_RANGE_USE_COUNT,
) -> tuple[float, float] | None:
"""受控复用:在素材历史区间中选一条可复用区间返回 (start, end)。
选择规则:
1. 仅选 ``use_count < max_use_count`` 的历史区间;
2. 优先返回能完整容纳当前 clip_duration(起点后不越素材边界)的最久未用区间;
3. 没有能容纳的,则返回 last_used_at 最老(或缺失 last_used_at 的旧数据优先)
且 use_count 最低的区间起点(可能与其他历史区间重叠,属降级复用);
4. 无任何可复用区间(记录为空或全部达上限)返回 None。
本函数只读不写;复用次数的累加由后续 record_used_segments 完成。
"""
model = _get_model(db, asset_id)
if model is None:
return None
meta = _read_meta(model)
ranges = [r for r in (meta.get(USED_RANGES_KEY) or []) if int(r.get("use_count", 1)) < max_use_count]
if not ranges:
return None
def _last_used(r: dict) -> str:
return str(r.get("last_used_at") or r.get("created_at") or "")
max_start = max(0.0, asset_total - clip_duration)
# 2. 能完整容纳当前片段的候选:按 last_used_at 升序(最久未用优先)
fit = sorted(
[r for r in ranges if float(r["start"]) <= max_start + 1e-6],
key=_last_used,
)
if fit:
start = min(float(fit[0]["start"]), max_start)
return (start, start + clip_duration)
# 3. 降级:最久未用 + use_count 最低的区间起点
fallback = sorted(ranges, key=lambda r: (_last_used(r), int(r.get("use_count", 1))))[0]
start = min(float(fallback["start"]), max_start)
return (start, start + clip_duration)
def make_reuse_callback(
db: Session,
asset_durations: dict[str, float],
reused_tracker: dict[str, float] | None = None,
assigned_tracker: dict[str, float] | None = None,
ratio_limit: float = REUSE_RATIO_LIMIT,
) -> Callable[[str, float], tuple[float, float] | None]:
"""构造给 ``_calc_random_start_time`` 用的受控复用回调.
Args:
db: SQLAlchemy session
asset_durations: 素材 ID -> 总时长(回调需要素材总时长做边界约束)
reused_tracker: 可选的 ``{asset_id: 累计复用时长}``,回调成功返回复用区间时
会把本次片段时长累加进去,供调用方统计成片复用占比(10% 阈值)。
assigned_tracker: 可选的 ``{asset_id: 已分配片段总时长}``,配合 ratio_limit
在复用前预判:若复用本片段后占比 (reused + clip_duration) /
(assigned + clip_duration) 超过 ratio_limit,则拒绝复用、返回 None
(保证成片复用占比不超阈值)。
ratio_limit: 单条成片复用时长占比上限,默认 10%
Returns:
回调函数 ``(asset_id, clip_duration) -> (start, end) | None``。
回调内吞掉 DB 异常返回 None,不影响主生成流程。
"""
def _reuse(asset_id: str, clip_duration: float) -> tuple[float, float] | None:
try:
total = float(asset_durations.get(asset_id, 0.0) or 0.0)
if total <= 0:
return None
# 占比闸门:预判复用本片段后是否超限(仅当调用方提供了 assigned tracker
if assigned_tracker is not None:
assigned = float(assigned_tracker.get(asset_id, 0.0) or 0.0)
reused_amt = float((reused_tracker or {}).get(asset_id, 0.0) or 0.0)
if assigned > 0 and (reused_amt + clip_duration) / (assigned + clip_duration) > ratio_limit:
logger.info(
"[片段追踪] 复用占比预判超 %.0f%% 阈值,拒绝复用: asset_id=%s "
"reused=%.1f assigned=%.1f clip=%.1f",
ratio_limit * 100,
asset_id,
reused_amt,
assigned,
clip_duration,
)
return None
result = find_reusable_range(db, asset_id, clip_duration, total)
except Exception:
logger.warning("[片段追踪] 受控复用查询异常: asset_id=%s", asset_id, exc_info=True)
return None
if result is not None and reused_tracker is not None:
reused_tracker[asset_id] = reused_tracker.get(asset_id, 0.0) + clip_duration
return result
return _reuse
def get_asset_recent_use_counts(
db: Session,
asset_ids: list[str],
recent_video_count: int = 5,
) -> dict[str, int]:
"""统计每个素材在最近 N 个不同 plan_id 中的使用次数。
遍历素材 metadata 中的 used_time_ranges,统计有多少个不同的 plan_id(去重),
返回 {asset_id: count}。只统计最近 recent_video_count 个不同 plan_id 的使用次数。
Args:
db: 数据库会话
asset_ids: 素材 ID 列表
recent_video_count: 统计最近多少个不同 plan_id
Returns:
{asset_id: 在最近 recent_video_count 个 plan 中的使用次数}
"""
if not asset_ids:
return {}
result: dict[str, int] = {}
models = db.query(AssetModel).filter(AssetModel.id.in_(asset_ids)).all()
for model in models:
meta = _read_meta(model)
ranges = meta.get(USED_RANGES_KEY) or []
if not ranges:
result[model.id] = 0
continue
# 按 created_at 倒序收集不同 plan_id
sorted_ranges = sorted(
ranges,
key=lambda r: r.get("created_at") or "",
reverse=True,
)
recent_plan_ids: set[str] = set()
for r in sorted_ranges:
plan_id = r.get("plan_id")
if plan_id:
recent_plan_ids.add(plan_id)
if len(recent_plan_ids) >= recent_video_count:
break
result[model.id] = len(recent_plan_ids)
# 未找到的素材计为 0
for aid in asset_ids:
if aid not in result:
result[aid] = 0
return result

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