爆款/拆镜生成简化3个步骤 | 项目生成可携带附件控制

This commit is contained in:
2026-07-21 14:01:08 +08:00
parent 40efcf55cf
commit 79c09151ba
60 changed files with 4250 additions and 924 deletions
+32 -258
View File
@@ -1,7 +1,7 @@
from datetime import datetime, timezone, timedelta
import json
from fastapi import APIRouter, Depends, HTTPException, Query, status
from fastapi import APIRouter, Depends, HTTPException, Query
from sqlalchemy import delete, func, select, update
from sqlalchemy.ext.asyncio import AsyncSession
@@ -24,7 +24,6 @@ from app.models.credit_ratio import CreditRatio
from app.models.operation_log import OperationLog
from app.enums.user import FrontendUserKind, UserType
from app.enums.team import TEAM_UNASSIGNED_VALUE
from app.enums.generation_status import GenerationRecordPipelineStage
from app.schemas.admin import (
CreditAdjustRequest,
ModelConfigCreate,
@@ -38,17 +37,12 @@ from app.schemas.admin import (
UpdateMenusRequest,
ResetPasswordRequest,
UpdateFrontendUserKindRequest,
OperationLogOut,
)
from app.schemas.team import UpdateUserTeamRequest
from app.schemas.industry import IndustryConfigCreate, IndustryConfigOut
from app.schemas.industry import IndustryConfigCreate
from app.schemas.video_engine import VideoEngineCreate, VideoEngineOut
from app.schemas.image_engine import ImageEngineCreate, ImageEngineOut
from app.schemas.credit_ratio import CreditRatioCreate, CreditRatioOut
from app.services.generation.pipeline.db_lock_service import (
DatabaseRowLockBusy,
execute_with_lock_timeout,
)
from app.services.credits import add_credits, deduct_credits
from app.services.credit_record_meta_service import build_admin_adjust_meta
from app.services.admin_credit_record_service import list_admin_credit_records
@@ -57,23 +51,26 @@ from app.services.auth import hash_password, verify_password
from app.services.operation_log import log_operation
from app.services.private_portrait.reference_resolver import batch_resolve_private_portrait_reference_display_urls
from app.services.resource_signed_url_service import build_resource_signed_url
from app.services.payment import sync_pending_orders, process_refund
from app.services.payment import process_refund
from app.services.resource_capacity_service import batch_get_user_resource_capacity_usage, get_user_resource_capacity_usage
from app.services.team_service import batch_get_team_name_map, set_frontend_user_team
from app.services.generation.billing_service import (
OWNER_GENERATION_RECORD,
charge_generation_media_for_record,
get_next_credit_attempt_no,
)
from app.services.generation.refund_service import mark_generation_record_failed_and_refund_once
from app.utils.id_gen import generate_id
from app.schemas.generation import GenerationType, ASPECT_RATIOS, RESOLUTIONS
CST = timezone(timedelta(hours=8))
def _safe_json_object(value: str | None) -> dict | None:
if not value:
return None
try:
parsed = json.loads(value)
except (TypeError, json.JSONDecodeError):
return None
return parsed if isinstance(parsed, dict) else None
def _iso(dt):
"""Serialize datetime as naive ISO string (UTC→CST, strip tzinfo)."""
if dt is None:
@@ -1917,6 +1914,8 @@ async def list_token_usage(
async def admin_list_generation_records(
user_id: str | None = Query(None),
status: str | None = Query(None),
engine_id: str | None = Query(None),
include_media_references: bool | None = Query(None),
page: int = Query(1, ge=1),
page_size: int = Query(20, ge=1, le=500),
admin: User = Depends(get_admin_user),
@@ -1935,13 +1934,25 @@ async def admin_list_generation_records(
query = query.where(GenerationRecord.user_id == user_id)
if status:
query = query.where(GenerationRecord.status == status)
if engine_id:
query = query.where(GenerationRecord.engine_id == engine_id)
if include_media_references is not None:
query = query.where(GenerationRecord.include_media_references.is_(include_media_references))
# Count total
count_query = select(func.count(GenerationRecord.id)).where(GenerationRecord.deleted_at.is_(None))
count_query = (
select(func.count(GenerationRecord.id))
.join(Project, GenerationRecord.project_id == Project.id)
.where(GenerationRecord.deleted_at.is_(None), Project.deleted_at.is_(None))
)
if user_id:
count_query = count_query.where(GenerationRecord.user_id == user_id)
if status:
count_query = count_query.where(GenerationRecord.status == status)
if engine_id:
count_query = count_query.where(GenerationRecord.engine_id == engine_id)
if include_media_references is not None:
count_query = count_query.where(GenerationRecord.include_media_references.is_(include_media_references))
total_result = await db.execute(count_query)
total = total_result.scalar() or 0
@@ -1977,6 +1988,10 @@ async def admin_list_generation_records(
"video_url": build_resource_signed_url(record.video_url) if record.video_url else '',
"video_cover_url": build_resource_signed_url(record.video_cover_url) if record.video_cover_url else '',
"references": refs,
"engine_id": record.engine_id,
"engine_name": (_safe_json_object(record.engine_snapshot_json) or {}).get("name"),
"engine_snapshot": _safe_json_object(record.engine_snapshot_json),
"include_media_references": bool(record.include_media_references),
"credits_cost": record.credits_cost or 0,
"text_credits_cost": record.text_credits_cost or 0,
"text_tokens_used": record.text_tokens_used or 0,
@@ -1997,245 +2012,6 @@ async def admin_list_generation_records(
return {"total": total, "items": items}
@router.put("/generation-records/{record_id}/status")
async def admin_update_generation_status(
record_id: str,
body: dict,
admin: User = Depends(get_admin_user),
db: AsyncSession = Depends(get_db),
):
"""管理员只能终止正在执行或待生成的记录,禁止绕过流水线裸改生成/完成状态。"""
try:
result = await execute_with_lock_timeout(
db,
select(GenerationRecord).where(
GenerationRecord.id == record_id,
GenerationRecord.deleted_at.is_(None),
)
.with_for_update()
.limit(1),
)
except DatabaseRowLockBusy as exc:
raise HTTPException(status_code=409, detail=exc.detail) from exc
record = result.scalar_one_or_none()
if not record:
raise HTTPException(status_code=404, detail="记录不存在")
new_status = str(body.get("status") or "").strip()
if new_status in {"generating", "completed", "prompt_optimized"}:
raise HTTPException(
status_code=409,
detail="禁止直接修改为该状态;生成请调用生成接口,完成必须由下载/超分流水线落库",
)
if new_status != "failed":
raise HTTPException(status_code=400, detail="该接口仅允许管理员终止任务")
if record.status == "completed":
raise HTTPException(status_code=409, detail="已完成记录不能直接改为失败")
error_message = body.get("error_message") or record.error_message or "管理员终止生成任务"
await mark_generation_record_failed_and_refund_once(
db,
record=record,
error_message=error_message,
generation_attempt_no=int(record.generation_attempt_no or 1),
)
record.pipeline_stage = GenerationRecordPipelineStage.FAILED.value
record.provider_create_claim_token = None
record.provider_create_lease_until = None
record.poll_claim_token = None
record.poll_lease_until = None
record.next_poll_at = None
record.download_claim_token = None
record.download_lease_until = None
record.download_next_retry_at = None
# 若任务已进入超分,必须同时撤销超分数据库租约;执行中的超分 Worker
# 在回填前校验 lease_token,发现 token 被清除后会中止,不得覆盖管理员终止状态。
from app.enums.video_upscale import VideoUpscaleStage, VideoUpscaleTaskStatus
from app.models.video_upscale_task import VideoUpscaleTask
try:
upscale_result = await execute_with_lock_timeout(
db,
select(VideoUpscaleTask)
.where(VideoUpscaleTask.generation_record_id == record.id)
.with_for_update()
.limit(1),
)
except DatabaseRowLockBusy as exc:
raise HTTPException(status_code=409, detail=exc.detail) from exc
upscale = upscale_result.scalar_one_or_none()
if upscale and upscale.status not in {
VideoUpscaleTaskStatus.COMPLETED.value,
VideoUpscaleTaskStatus.FAILED.value,
}:
upscale.status = VideoUpscaleTaskStatus.FAILED.value
upscale.stage = VideoUpscaleStage.FAILED.value
upscale.last_error = error_message
upscale.failed_at = datetime.now(CST)
upscale.next_retry_at = None
upscale.lease_token = None
upscale.lease_until = None
await db.flush()
await log_operation(
db,
admin.id,
admin.username,
"管理员终止生成记录",
"PUT",
f"/admin/generation-records/{record_id}/status",
detail=json.dumps(
{
"record_id": record_id,
"new_status": new_status,
"generation_attempt_no": int(record.generation_attempt_no or 1),
"error_message": error_message,
},
ensure_ascii=False,
),
)
await db.commit()
# Redis 注册表只做调度加速;删除失败不回滚已提交的业务终止状态。
try:
from app.services.celery_download_recovery_service import remove_download_active
from app.services.generation.pipeline.owner_service import redis_owner_item_id
from app.services.redis_registry_service import redis_remove_registry_item
from app.config import settings
registry_id = redis_owner_item_id(
"generation_record",
record_id,
int(record.generation_attempt_no or 1),
)
await remove_download_active(registry_id)
await redis_remove_registry_item(
hash_key=settings.POLL_ACTIVE_REDIS_HASH_KEY,
zset_key=settings.POLL_ACTIVE_REDIS_ZSET_KEY,
item_id=registry_id,
log_context="admin_generation_record_terminate",
)
except Exception:
pass
return {"message": "ok"}
@router.post("/generation-records/{record_id}/generate")
async def admin_generate_record_resource(
record_id: str,
body: dict,
admin: User = Depends(get_admin_user),
db: AsyncSession = Depends(get_db),
):
"""管理员触发 GenerationRecord 图片或视频资源生成。"""
from app.services.generation.pipeline.generation_record_service import (
commit_and_enqueue_generation_record,
prepare_generation_record_execution,
)
try:
result = await execute_with_lock_timeout(
db,
select(GenerationRecord, Project.name)
.join(Project, GenerationRecord.project_id == Project.id)
.where(
GenerationRecord.id == record_id,
GenerationRecord.deleted_at.is_(None),
Project.deleted_at.is_(None),
)
.with_for_update(),
)
except DatabaseRowLockBusy as exc:
raise HTTPException(status_code=409, detail=exc.detail) from exc
row = result.first()
if not row:
raise HTTPException(status_code=404, detail="记录不存在")
record, project_name = row
type_str = "视频" if record.gen_type == GenerationType.video else "图片"
if record.status not in ("prompt_optimized", "failed"):
raise HTTPException(status_code=400, detail=f"当前状态不允许生成{type_str}")
if record.pipeline_stage == GenerationRecordPipelineStage.UPSCALE_FAILED.value:
raise HTTPException(status_code=409, detail="该任务为画质增强失败,请使用超分恢复命令处理")
attempt_no = await get_next_credit_attempt_no(
db,
owner_type=OWNER_GENERATION_RECORD,
owner_id=record.id,
)
if record.gen_type == GenerationType.video:
aspect_ratio = body.get("aspect_ratio", "16:9")
resolution = body.get("resolution", "720p")
if aspect_ratio not in ASPECT_RATIOS:
raise HTTPException(status_code=400, detail="不支持的画面比例")
if resolution not in RESOLUTIONS:
raise HTTPException(status_code=400, detail="不支持的分辨率")
from app.services.video_gen import get_active_engine
from app.services.video_upscale.snapshot_service import build_video_upscale_snapshot
engine = await get_active_engine(db)
try:
supported_provider_resolutions = json.loads(engine.supported_resolutions or "[]")
except (TypeError, json.JSONDecodeError):
supported_provider_resolutions = []
provider_resolution, upscale_enabled, upscale_snapshot_json = await build_video_upscale_snapshot(
db,
target_resolution=resolution,
aspect_ratio=aspect_ratio,
supported_provider_resolutions=supported_provider_resolutions,
)
record.aspect_ratio = aspect_ratio
record.resolution = resolution
record.provider_generation_resolution = provider_resolution
record.video_upscale_enabled_snapshot = upscale_enabled
record.video_upscale_snapshot_json = upscale_snapshot_json
elif record.gen_type == GenerationType.image:
from app.services.image_gen import get_active_image_engine
engine = await get_active_image_engine(db)
record.image_size = body.get("image_size") or record.image_size or "2K"
record.provider_generation_resolution = None
record.video_upscale_enabled_snapshot = False
record.video_upscale_snapshot_json = None
else:
raise HTTPException(status_code=400, detail="不支持的生成类型")
media_billing = await charge_generation_media_for_record(
db,
record=record,
project_name=project_name,
description_prefix=f"{type_str}生成(管理后台)-",
attempt_no=attempt_no,
engine_id=engine.id,
)
record.credits_cost = round(float(record.credits_cost or 0) + float(media_billing.total_charged or 0), 2)
prepare_generation_record_execution(record, engine=engine, attempt_no=attempt_no)
await db.flush()
await commit_and_enqueue_generation_record(db, record, reason="generation_record_admin_generate")
await log_operation(
db,
admin.id,
admin.username,
f"管理员触发生成{type_str}: {record_id}",
"POST",
f"/admin/generation-records/{record_id}/generate",
detail=json.dumps(
{
"record_id": record_id,
"gen_type": record.gen_type,
"project_name": project_name,
"generation_attempt_no": record.generation_attempt_no,
},
ensure_ascii=False,
),
)
return {"message": "ok", "record_id": record_id}
# ── File Uploads ─────────────────────────────────────────
import os
@@ -2251,7 +2027,6 @@ async def upload_pdf(
):
"""Upload a PDF file and save URL to system config."""
from app.config import settings
from app.utils.id_gen import generate_id
if not file.filename:
raise HTTPException(status_code=400, detail="请选择文件")
@@ -2313,7 +2088,6 @@ async def upload_logo(
):
"""Upload a Logo image file and save URL to system config."""
from app.config import settings
from app.utils.id_gen import generate_id
if not file.filename:
raise HTTPException(status_code=400, detail="请选择文件")
+158 -44
View File
@@ -55,6 +55,16 @@ from app.services.generation.billing_service import (
get_next_credit_attempt_no,
)
from app.services.generation.refund_service import mark_generation_record_failed_and_refund_once
from app.services.generation.ai.engine_service import (
get_image_engine,
get_video_engine,
image_supported_sizes,
parse_json_list,
)
from app.services.generation.media_reference_service import (
calculate_media_reference_usage,
validate_media_reference_usage_for_engine,
)
from app.services.media_token_usage_snapshot_service import sync_generation_record_media_token_snapshot
from app.services.credit_record_meta_service import build_generation_record_prompt_meta
from app.services.video_cover_service import async_create_video_cover_for_local_video
@@ -70,6 +80,36 @@ router = APIRouter(prefix="/generation-records", tags=["generation"])
logger = logging.getLogger("videogen")
def _engine_snapshot(record: GenerationRecord) -> dict | None:
if not record.engine_snapshot_json:
return None
try:
value = json.loads(record.engine_snapshot_json)
except (TypeError, json.JSONDecodeError):
return None
return value if isinstance(value, dict) else None
def _validate_video_engine_selection(engine, *, aspect_ratio: str, resolution: str, duration: int) -> None:
ratios = [str(item) for item in parse_json_list(engine.supported_ratios, [])]
resolutions = [str(item) for item in parse_json_list(engine.supported_resolutions, [])]
durations = [int(item) for item in parse_json_list(engine.supported_durations, []) if str(item).isdigit()]
if ratios and aspect_ratio not in ratios:
raise HTTPException(status_code=400, detail="当前视频引擎不支持所选画面比例")
if resolutions and resolution not in resolutions:
raise HTTPException(status_code=400, detail="当前视频引擎不支持所选分辨率")
if durations and duration not in durations:
raise HTTPException(status_code=400, detail="当前视频引擎不支持所选时长")
if int(engine.max_duration or 0) > 0 and duration > int(engine.max_duration):
raise HTTPException(status_code=400, detail="生成时长超过当前视频引擎上限")
def _validate_image_engine_selection(engine, *, image_size: str) -> None:
sizes = image_supported_sizes(engine)
if sizes and image_size not in sizes:
raise HTTPException(status_code=400, detail="当前图片引擎不支持所选画面分辨率")
def _record_to_out(record: GenerationRecord, project_name: str, refs_override: list[dict] | None = None) -> GenerationRecordOut:
refs = refs_override
if refs is None and record.media_references:
@@ -112,6 +152,10 @@ def _record_to_out(record: GenerationRecord, project_name: str, refs_override: l
video_cover_url=build_resource_signed_url(record.video_cover_url) if record.video_cover_url else '',
image_url=build_resource_signed_url(record.image_url) if record.image_url else '',
references=refs,
engine_id=record.engine_id,
engine_name=(_engine_snapshot(record) or {}).get("name"),
engine_snapshot=_engine_snapshot(record),
include_media_references=bool(record.include_media_references),
text_credits_cost=round(record.text_credits_cost or 0.00, 2),
# text_tokens_used=record.text_tokens_used or 0,
credits_cost=round(record.credits_cost or 0.00, 2),
@@ -436,7 +480,11 @@ async def generate_record_resource(
raise InvalidStatusError("该任务生成失败,请联系客服进行修复")
await assert_user_resource_capacity_available(db, current_user.id)
attempt_no = await get_next_credit_attempt_no(db, owner_type=OWNER_GENERATION_RECORD, owner_id=record.id)
attempt_no = await get_next_credit_attempt_no(
db, owner_type=OWNER_GENERATION_RECORD, owner_id=record.id
)
selected_engine_id = req.engine_id or record.engine_id
record.include_media_references = bool(req.include_media_references)
from app.services.generation.pipeline.generation_record_service import (
commit_and_enqueue_generation_record,
@@ -444,56 +492,83 @@ async def generate_record_resource(
)
if record.gen_type == GenerationType.video:
if req.aspect_ratio not in ASPECT_RATIOS:
aspect_ratio = req.aspect_ratio or record.aspect_ratio
resolution = req.resolution or record.resolution
if aspect_ratio not in ASPECT_RATIOS:
raise HTTPException(status_code=400, detail="不支持的画面比例")
if req.resolution not in RESOLUTIONS:
if resolution not in RESOLUTIONS:
raise HTTPException(status_code=400, detail="不支持的分辨率")
from app.services.video_gen import get_active_engine
engine = await get_video_engine(db, selected_engine_id)
_validate_video_engine_selection(
engine,
aspect_ratio=aspect_ratio,
resolution=resolution,
duration=int(record.duration or 5),
)
from app.services.video_upscale.snapshot_service import build_video_upscale_snapshot
engine = await get_active_engine(db)
try:
supported_provider_resolutions = json.loads(engine.supported_resolutions or "[]")
except (TypeError, json.JSONDecodeError):
supported_provider_resolutions = []
supported_provider_resolutions = parse_json_list(engine.supported_resolutions, [])
provider_resolution, upscale_enabled, upscale_snapshot_json = await build_video_upscale_snapshot(
db,
target_resolution=req.resolution,
aspect_ratio=req.aspect_ratio,
target_resolution=resolution,
aspect_ratio=aspect_ratio,
supported_provider_resolutions=supported_provider_resolutions,
)
billing = await charge_generation_media_by_params(
db, user_id=current_user.id, record_id=record.id, gen_type="video",
duration=record.duration or 5, resolution=req.resolution, engine_id=engine.id,
project_name=project_name, description_prefix=project_name + "-",
owner_type=OWNER_GENERATION_RECORD, attempt_no=attempt_no,
)
record.aspect_ratio = req.aspect_ratio
record.resolution = req.resolution
record.aspect_ratio = aspect_ratio
record.resolution = resolution
record.provider_generation_resolution = provider_resolution
record.video_upscale_enabled_snapshot = upscale_enabled
record.video_upscale_snapshot_json = upscale_snapshot_json
else:
from app.services.image_gen import get_active_image_engine
engine = await get_active_image_engine(db)
engine = await get_image_engine(db, selected_engine_id)
image_size = req.image_size or record.image_size or engine.default_size or "2K"
billing = await charge_generation_media_by_params(
db, user_id=current_user.id, record_id=record.id, gen_type="image",
image_size=image_size, engine_id=engine.id, project_name=project_name,
description_prefix=project_name + "-", owner_type=OWNER_GENERATION_RECORD, attempt_no=attempt_no,
)
_validate_image_engine_selection(engine, image_size=image_size)
record.image_size = image_size
record.provider_generation_resolution = None
record.video_upscale_enabled_snapshot = False
record.video_upscale_snapshot_json = None
record.credits_cost = round(float(record.credits_cost or 0) + float(billing.total_charged or 0), 2)
reference_usage = calculate_media_reference_usage(
record.media_references,
include=bool(record.include_media_references),
)
validate_media_reference_usage_for_engine(
reference_usage,
gen_type=record.gen_type,
engine=engine,
)
billing = await charge_generation_media_for_record(
db,
record=record,
project_name=project_name,
description_prefix=project_name + "-",
attempt_no=attempt_no,
engine_id=engine.id,
)
record.credits_cost = round(
float(record.credits_cost or 0) + float(billing.total_charged or 0), 2
)
prepare_generation_record_execution(record, engine=engine, attempt_no=attempt_no)
await db.flush()
await commit_and_enqueue_generation_record(db, record, reason="generation_record_api_generate")
record_id_snapshot = str(record.id)
await commit_and_enqueue_generation_record(
db, record, reason="generation_record_api_generate"
)
refreshed = await db.execute(
select(GenerationRecord, Project.name)
.join(Project, GenerationRecord.project_id == Project.id)
.where(GenerationRecord.id == record_id_snapshot)
.limit(1)
)
refreshed_row = refreshed.first()
if not refreshed_row:
raise RecordNotFoundError()
record, project_name = refreshed_row
refs = await resolve_private_portrait_reference_display_urls(
db, json.loads(record.media_references) if record.media_references else None, user_id=current_user.id
db,
json.loads(record.media_references) if record.media_references else None,
user_id=current_user.id,
)
return _record_to_out(record, project_name, refs_override=refs)
@@ -529,43 +604,82 @@ async def retry_generation(
raise InvalidStatusError("该任务生成失败,请联系客服进行修复")
await assert_user_resource_capacity_available(db, current_user.id)
attempt_no = await get_next_credit_attempt_no(db, owner_type=OWNER_GENERATION_RECORD, owner_id=record.id)
attempt_no = await get_next_credit_attempt_no(
db, owner_type=OWNER_GENERATION_RECORD, owner_id=record.id
)
from app.services.generation.pipeline.generation_record_service import (
commit_and_enqueue_generation_record,
prepare_generation_record_execution,
)
if record.gen_type == GenerationType.video:
from app.services.video_gen import get_active_engine
engine = await get_video_engine(db, record.engine_id)
_validate_video_engine_selection(
engine,
aspect_ratio=record.aspect_ratio or "16:9",
resolution=record.resolution or "480p",
duration=int(record.duration or 5),
)
from app.services.video_upscale.snapshot_service import build_video_upscale_snapshot
engine = await get_active_engine(db)
try:
supported_provider_resolutions = json.loads(engine.supported_resolutions or "[]")
except (TypeError, json.JSONDecodeError):
supported_provider_resolutions = []
provider_resolution, upscale_enabled, upscale_snapshot_json = await build_video_upscale_snapshot(
db, target_resolution=record.resolution or "480p", aspect_ratio=record.aspect_ratio or "16:9",
supported_provider_resolutions=supported_provider_resolutions,
db,
target_resolution=record.resolution or "480p",
aspect_ratio=record.aspect_ratio or "16:9",
supported_provider_resolutions=parse_json_list(engine.supported_resolutions, []),
)
record.provider_generation_resolution = provider_resolution
record.video_upscale_enabled_snapshot = upscale_enabled
record.video_upscale_snapshot_json = upscale_snapshot_json
else:
from app.services.image_gen import get_active_image_engine
engine = await get_active_image_engine(db)
engine = await get_image_engine(db, record.engine_id)
_validate_image_engine_selection(
engine, image_size=record.image_size or engine.default_size or "2K"
)
reference_usage = calculate_media_reference_usage(
record.media_references,
include=bool(record.include_media_references),
)
validate_media_reference_usage_for_engine(
reference_usage,
gen_type=record.gen_type,
engine=engine,
)
billing = await charge_generation_media_for_record(
db, record=record, project_name=project_name, description_prefix="资源生成重试-", attempt_no=attempt_no, engine_id=engine.id
db,
record=record,
project_name=project_name,
description_prefix="资源生成重试-",
attempt_no=attempt_no,
engine_id=engine.id,
)
record.credits_cost = round(
float(record.credits_cost or 0) + float(billing.total_charged or 0), 2
)
record.credits_cost = round(float(record.credits_cost or 0) + float(billing.total_charged or 0), 2)
record.manual_retry_count = int(record.manual_retry_count or 0) + 1
record.retry_count = int(record.manual_retry_count or 0)
prepare_generation_record_execution(record, engine=engine, attempt_no=attempt_no)
await db.flush()
await commit_and_enqueue_generation_record(db, record, reason="generation_record_api_retry")
record_id_snapshot = str(record.id)
await commit_and_enqueue_generation_record(
db, record, reason="generation_record_api_retry"
)
refreshed = await db.execute(
select(GenerationRecord, Project.name)
.join(Project, GenerationRecord.project_id == Project.id)
.where(GenerationRecord.id == record_id_snapshot)
.limit(1)
)
refreshed_row = refreshed.first()
if not refreshed_row:
raise RecordNotFoundError()
record, project_name = refreshed_row
refs = await resolve_private_portrait_reference_display_urls(
db, json.loads(record.media_references) if record.media_references else None, user_id=current_user.id
db,
json.loads(record.media_references) if record.media_references else None,
user_id=current_user.id,
)
return _record_to_out(record, project_name, refs_override=refs)
@@ -9,7 +9,7 @@ from sqlalchemy.ext.asyncio import AsyncSession
from app.dependencies import get_current_user, get_db
from app.models.user import User
from app.enums.common import ModuleProjectStatusEnum
from app.enums.common import ModuleProjectStatusEnum, ModuleEventTypeEnum
from app.enums.generation_task import GenerationOwnerType
from app.enums.hot_opening_replicate import HotOpeningLogEventEnum, HotOpeningStepCodeEnum, ModuleCodeEnum
from app.schemas.hot_opening_replicate import (
@@ -42,7 +42,7 @@ from app.services.hot_opening_replicate_service import (
update_hot_opening_material_input,
update_hot_opening_video_prompt_schema,
)
from app.services.module_generation_log_service import log_module_error
from app.services.module_generation_log_service import log_module_error, log_module_event_file
from app.services.module_async_recovery_service import (
TASK_HOT_IMAGE_PROMPT,
TASK_HOT_VIDEO_PROMPT,
@@ -283,29 +283,17 @@ async def create_task(
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
try:
project = await create_hot_opening_project(db, current_user, req)
project_id_value = str(project.id)
await bind_upload_resources(
db,
user_id=current_user.id,
module=UploadResourceModuleEnum.HOT_OPENING_REPLICATE.value,
source_model=UploadResourceSourceModelEnum.MODULE_GENERATION_PROJECT.value,
source_id=project_id_value,
resource_ids=[req.material_video_resource_id, req.material_image_resource_id],
urls=[req.material_video_url, req.material_image_url],
allow_common_migrate=True,
)
await db.commit()
except HTTPException:
await db.rollback()
raise
except Exception as exc:
await db.rollback()
_log_api_exception_from_locals(exc, locals(), f"创建爆款开头复刻项目失败: {exc}")
raise HTTPException(status_code=500, detail=f"创建爆款开头复刻项目失败: {exc}")
return await _reload_project_detail(db, current_user, project_id_value)
log_module_event_file(
module=MODULE,
event_type=ModuleEventTypeEnum.V1_CREATE_BLOCKED.value,
user_id=current_user.id,
message="拦截爆款开头复刻 V1 创建请求",
detail={"api_version": "v1", "flow_version": "v1"},
)
raise HTTPException(
status_code=410,
detail="V1 创建流程已停止,请使用 V2 API",
)
@router.get(
+12 -18
View File
@@ -10,6 +10,7 @@ from sqlalchemy.ext.asyncio import AsyncSession
from app.config import settings
from app.dependencies import get_current_user, get_db
from app.models.user import User
from app.enums.common import ModuleEventTypeEnum
from app.enums.generation_task import GenerationOwnerType
from app.enums.shot_replicate import (
ModuleCodeEnum,
@@ -862,24 +863,17 @@ async def create_replication_project_from_segment(
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
try:
segment = await get_segment_for_user(db, segment_id=segment_id, user=current_user, for_update=True)
project = await create_shot_replicate_project_from_segment(db, current_user=current_user, segment=segment, req=req)
project_id = project.id
await db.commit()
except HTTPException:
await db.rollback()
raise
except Exception as exc:
await db.rollback()
_log_api_exception_from_locals(exc, locals(), f"创建拆镜复刻项目失败: {exc}")
raise HTTPException(status_code=500, detail=f"创建拆镜复刻项目失败: {exc}")
return ShotReplicateActionOut(
message="已从拆镜片段创建复刻项目,素材视频已锁定",
project_id=project_id,
step_id=None,
detail=await _reload_project_detail(db, current_user, project_id),
log_module_event_file(
module=MODULE,
event_type=ModuleEventTypeEnum.V1_CREATE_BLOCKED.value,
user_id=current_user.id,
step_id=segment_id,
message="拦截拆镜复刻 V1 创建请求",
detail={"api_version": "v1", "flow_version": "v1", "segment_id": segment_id},
)
raise HTTPException(
status_code=410,
detail="V1 创建流程已停止,请使用 V2 API",
)
+8
View File
@@ -0,0 +1,8 @@
from fastapi import APIRouter
from app.api.v2.hot_opening_replicate import router as hot_opening_router
from app.api.v2.shot_replicate import router as shot_replicate_router
api_router_v2 = APIRouter()
api_router_v2.include_router(hot_opening_router)
api_router_v2.include_router(shot_replicate_router)
@@ -0,0 +1,255 @@
from __future__ import annotations
from fastapi import APIRouter, Body, Depends, HTTPException, Path
from sqlalchemy.exc import IntegrityError
from sqlalchemy.ext.asyncio import AsyncSession
from app.dependencies import get_current_user, get_db
from app.models.chat_generation_task import ChatGenerationTask
from app.models.user import User
from app.schemas.hot_opening_replicate import HotOpeningActionOut, HotOpeningDeleteOut, HotOpeningTaskDetailOut
from app.schemas.module_generation_v2 import (
HotOpeningTaskCreateV2,
ModuleVideoPromptRetryV2,
ModuleVideoPromptSchemaUpdateV2,
)
from app.services.generation.pipeline.enqueue_service import enqueue_generation_create
from app.services.hot_opening_replicate_service import project_to_detail_out
from app.services.module_generation_v2.config import HOT_OPENING_V2
from app.services.module_generation_v2.dispatch_service import (
dispatch_video_prompt_v2,
ensure_v2_celery_enabled,
)
from app.services.module_generation_v2.flow_service import (
create_hot_opening_project_v2,
delete_project_v2,
generate_video_from_prompt_v2,
get_v2_project_for_user,
is_project_idempotency_conflict,
mark_video_prompt_dispatch_failed_v2,
rebuild_video_prompt_step_v2,
update_video_prompt_schema_v2,
)
from app.services.upload_resource import cleanup_upload_resource_files_after_commit
router = APIRouter(prefix="/hot-opening-replications", tags=["hot-opening-replications-v2"])
async def _detail(db: AsyncSession, current_user: User, project_id: str) -> HotOpeningTaskDetailOut:
project = await get_v2_project_for_user(
db,
config=HOT_OPENING_V2,
project_id=project_id,
current_user=current_user,
)
return await project_to_detail_out(db, project)
async def _dispatch_or_mark_failed(
db: AsyncSession,
*,
project_id: str,
step_id: str,
) -> None:
dispatch = await dispatch_video_prompt_v2(
config=HOT_OPENING_V2,
project_id=project_id,
step_id=step_id,
)
if dispatch.recoverable:
return
error_message = "视频提词任务的 Redis 注册和 Celery 投递均失败,请重新执行步骤2"
await mark_video_prompt_dispatch_failed_v2(
db,
config=HOT_OPENING_V2,
project_id=project_id,
step_id=step_id,
error_message=error_message,
)
raise HTTPException(status_code=503, detail=error_message)
@router.post("/tasks", response_model=HotOpeningTaskDetailOut)
async def create_task_v2(
req: HotOpeningTaskCreateV2 = Body(...),
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
ensure_v2_celery_enabled()
try:
result = await create_hot_opening_project_v2(db, current_user=current_user, req=req)
project_id = str(result.project.id)
step_id = str(result.prompt_step.id)
created_new = bool(result.created_new)
await db.commit()
except IntegrityError as exc:
await db.rollback()
if not req.idempotency_key or not is_project_idempotency_conflict(exc):
raise HTTPException(status_code=500, detail="项目创建失败") from exc
# 同幂等键并发请求由唯一索引收敛;回查已提交项目并按幂等成功返回。
result = await create_hot_opening_project_v2(db, current_user=current_user, req=req)
project_id = str(result.project.id)
step_id = str(result.prompt_step.id)
created_new = bool(result.created_new)
await db.commit()
except HTTPException:
await db.rollback()
raise
except Exception as exc:
await db.rollback()
raise HTTPException(status_code=500, detail="创建爆款复刻 V2 项目失败") from exc
if created_new:
await _dispatch_or_mark_failed(db, project_id=project_id, step_id=step_id)
return await _detail(db, current_user, project_id)
@router.get("/tasks/{project_id}", response_model=HotOpeningTaskDetailOut)
async def get_task_v2(
project_id: str = Path(...),
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
return await _detail(db, current_user, project_id)
@router.post(
"/tasks/{project_id}/steps/{step_id}/retry-video-prompt",
response_model=HotOpeningActionOut,
)
async def retry_video_prompt_v2(
project_id: str,
step_id: str,
req: ModuleVideoPromptRetryV2,
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
ensure_v2_celery_enabled()
try:
project, new_step = await rebuild_video_prompt_step_v2(
db,
config=HOT_OPENING_V2,
current_user=current_user,
project_id=project_id,
source_prompt_step_id=step_id,
video_config=req.video_config,
)
project_id_value = str(project.id)
step_id_value = str(new_step.id)
await db.commit()
except HTTPException:
await db.rollback()
raise
await _dispatch_or_mark_failed(db, project_id=project_id_value, step_id=step_id_value)
return HotOpeningActionOut(
message="视频提词已重新提交",
project_id=project_id_value,
step_id=step_id_value,
detail=await _detail(db, current_user, project_id_value),
)
@router.put(
"/tasks/{project_id}/steps/{step_id}/video-prompt-schema",
response_model=HotOpeningActionOut,
)
async def update_video_prompt_schema_route_v2(
project_id: str,
step_id: str,
req: ModuleVideoPromptSchemaUpdateV2,
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
try:
project, step = await update_video_prompt_schema_v2(
db,
config=HOT_OPENING_V2,
current_user=current_user,
project_id=project_id,
step_id=step_id,
req=req,
)
project_id_value = str(project.id)
step_id_value = str(step.id)
await db.commit()
except HTTPException:
await db.rollback()
raise
return HotOpeningActionOut(
message="视频提词已保存",
project_id=project_id_value,
step_id=step_id_value,
detail=await _detail(db, current_user, project_id_value),
)
@router.post(
"/tasks/{project_id}/steps/{step_id}/generate-video",
response_model=HotOpeningActionOut,
)
async def generate_video_v2(
project_id: str,
step_id: str,
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
ensure_v2_celery_enabled()
try:
project, step, task = await generate_video_from_prompt_v2(
db,
config=HOT_OPENING_V2,
current_user=current_user,
project_id=project_id,
prompt_step_id=step_id,
)
project_id_value = str(project.id)
step_id_value = str(step.id)
task_id = str(task.id)
await db.commit()
except HTTPException:
await db.rollback()
raise
# commit 后重新读取,避免 ORM expire/lazy-load 风险。
queued_task = await db.get(ChatGenerationTask, task_id)
if queued_task is None:
raise HTTPException(status_code=500, detail="视频生成任务提交后无法重新读取")
try:
await enqueue_generation_create(queued_task, reason="hot_opening_v2_generate_video")
except Exception as exc:
# queued 状态已持久化,周期生成恢复任务会使用确定性 task_id 补投。
raise HTTPException(status_code=503, detail="视频生成任务暂未投递,将由恢复任务自动补投") from exc
return HotOpeningActionOut(
message="视频生成任务已提交",
project_id=project_id_value,
step_id=step_id_value,
detail=await _detail(db, current_user, project_id_value),
)
@router.delete("/tasks/{project_id}", response_model=HotOpeningDeleteOut)
async def delete_project_route_v2(
project_id: str,
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
try:
payload = await delete_project_v2(
db,
config=HOT_OPENING_V2,
current_user=current_user,
project_id=project_id,
)
pending_ids = list(payload.get("pending_delete_resource_ids") or [])
await db.commit()
except HTTPException:
await db.rollback()
raise
if pending_ids:
try:
await cleanup_upload_resource_files_after_commit(db, resource_ids=pending_ids)
await db.commit()
except Exception:
await db.rollback()
return HotOpeningDeleteOut(**payload)
+272
View File
@@ -0,0 +1,272 @@
from __future__ import annotations
from fastapi import APIRouter, Body, Depends, HTTPException, Path
from sqlalchemy.exc import IntegrityError
from sqlalchemy.ext.asyncio import AsyncSession
from app.dependencies import get_current_user, get_db
from app.models.chat_generation_task import ChatGenerationTask
from app.models.user import User
from app.schemas.module_generation_v2 import (
ModuleVideoPromptRetryV2,
ModuleVideoPromptSchemaUpdateV2,
ShotReplicateProjectCreateV2,
)
from app.schemas.shot_replicate import ShotReplicateActionOut, ShotReplicateDeleteOut, ShotReplicateTaskDetailOut
from app.services.generation.pipeline.enqueue_service import enqueue_generation_create
from app.services.module_generation_v2.config import SHOT_REPLICATE_V2
from app.services.module_generation_v2.dispatch_service import (
dispatch_video_prompt_v2,
ensure_v2_celery_enabled,
)
from app.services.module_generation_v2.flow_service import (
create_shot_replicate_project_v2,
delete_project_v2,
generate_video_from_prompt_v2,
get_v2_project_for_user,
is_project_idempotency_conflict,
mark_video_prompt_dispatch_failed_v2,
rebuild_video_prompt_step_v2,
update_video_prompt_schema_v2,
)
from app.services.shot_replicate_flow_service import project_to_detail_out
from app.services.shot_replicate_taskset_service import get_segment_for_user
from app.services.upload_resource import cleanup_upload_resource_files_after_commit
router = APIRouter(prefix="/shot-replications", tags=["shot-replications-v2"])
async def _detail(db: AsyncSession, current_user: User, project_id: str) -> ShotReplicateTaskDetailOut:
project = await get_v2_project_for_user(
db,
config=SHOT_REPLICATE_V2,
project_id=project_id,
current_user=current_user,
)
return await project_to_detail_out(db, project)
async def _dispatch_or_mark_failed(
db: AsyncSession,
*,
project_id: str,
step_id: str,
) -> None:
dispatch = await dispatch_video_prompt_v2(
config=SHOT_REPLICATE_V2,
project_id=project_id,
step_id=step_id,
)
if dispatch.recoverable:
return
error_message = "视频提词任务的 Redis 注册和 Celery 投递均失败,请重新执行步骤2"
await mark_video_prompt_dispatch_failed_v2(
db,
config=SHOT_REPLICATE_V2,
project_id=project_id,
step_id=step_id,
error_message=error_message,
)
raise HTTPException(status_code=503, detail=error_message)
@router.post(
"/segments/{segment_id}/replication-projects",
response_model=ShotReplicateActionOut,
)
async def create_project_v2(
segment_id: str = Path(...),
req: ShotReplicateProjectCreateV2 = Body(...),
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
ensure_v2_celery_enabled()
try:
segment = await get_segment_for_user(
db, segment_id=segment_id, user=current_user, for_update=True
)
result = await create_shot_replicate_project_v2(
db, current_user=current_user, segment=segment, req=req
)
project_id = str(result.project.id)
step_id = str(result.prompt_step.id)
created_new = bool(result.created_new)
await db.commit()
except IntegrityError as exc:
await db.rollback()
if not req.idempotency_key or not is_project_idempotency_conflict(exc):
raise HTTPException(status_code=500, detail="项目创建失败") from exc
segment = await get_segment_for_user(
db, segment_id=segment_id, user=current_user, for_update=True
)
result = await create_shot_replicate_project_v2(
db, current_user=current_user, segment=segment, req=req
)
project_id = str(result.project.id)
step_id = str(result.prompt_step.id)
created_new = bool(result.created_new)
await db.commit()
except HTTPException:
await db.rollback()
raise
except Exception as exc:
await db.rollback()
raise HTTPException(status_code=500, detail="创建拆镜复刻 V2 项目失败") from exc
if created_new:
await _dispatch_or_mark_failed(db, project_id=project_id, step_id=step_id)
return ShotReplicateActionOut(
message="V2 项目已创建,视频提词已自动提交" if created_new else "已返回现有幂等项目",
project_id=project_id,
step_id=step_id,
detail=await _detail(db, current_user, project_id),
)
@router.get("/projects/{project_id}", response_model=ShotReplicateTaskDetailOut)
async def get_project_v2(
project_id: str,
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
return await _detail(db, current_user, project_id)
@router.post(
"/projects/{project_id}/steps/{step_id}/retry-video-prompt",
response_model=ShotReplicateActionOut,
)
async def retry_video_prompt_v2(
project_id: str,
step_id: str,
req: ModuleVideoPromptRetryV2,
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
ensure_v2_celery_enabled()
try:
project, new_step = await rebuild_video_prompt_step_v2(
db,
config=SHOT_REPLICATE_V2,
current_user=current_user,
project_id=project_id,
source_prompt_step_id=step_id,
video_config=req.video_config,
)
project_id_value = str(project.id)
step_id_value = str(new_step.id)
await db.commit()
except HTTPException:
await db.rollback()
raise
await _dispatch_or_mark_failed(db, project_id=project_id_value, step_id=step_id_value)
return ShotReplicateActionOut(
message="视频提词已重新提交",
project_id=project_id_value,
step_id=step_id_value,
detail=await _detail(db, current_user, project_id_value),
)
@router.put(
"/projects/{project_id}/steps/{step_id}/video-prompt-schema",
response_model=ShotReplicateActionOut,
)
async def update_video_prompt_schema_route_v2(
project_id: str,
step_id: str,
req: ModuleVideoPromptSchemaUpdateV2,
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
try:
project, step = await update_video_prompt_schema_v2(
db,
config=SHOT_REPLICATE_V2,
current_user=current_user,
project_id=project_id,
step_id=step_id,
req=req,
)
project_id_value = str(project.id)
step_id_value = str(step.id)
await db.commit()
except HTTPException:
await db.rollback()
raise
return ShotReplicateActionOut(
message="视频提词已保存",
project_id=project_id_value,
step_id=step_id_value,
detail=await _detail(db, current_user, project_id_value),
)
@router.post(
"/projects/{project_id}/steps/{step_id}/generate-video",
response_model=ShotReplicateActionOut,
)
async def generate_video_v2(
project_id: str,
step_id: str,
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
ensure_v2_celery_enabled()
try:
project, step, task = await generate_video_from_prompt_v2(
db,
config=SHOT_REPLICATE_V2,
current_user=current_user,
project_id=project_id,
prompt_step_id=step_id,
)
project_id_value = str(project.id)
step_id_value = str(step.id)
task_id = str(task.id)
await db.commit()
except HTTPException:
await db.rollback()
raise
queued_task = await db.get(ChatGenerationTask, task_id)
if queued_task is None:
raise HTTPException(status_code=500, detail="视频生成任务提交后无法重新读取")
try:
await enqueue_generation_create(queued_task, reason="shot_replicate_v2_generate_video")
except Exception as exc:
raise HTTPException(status_code=503, detail="视频生成任务暂未投递,将由恢复任务自动补投") from exc
return ShotReplicateActionOut(
message="视频生成任务已提交",
project_id=project_id_value,
step_id=step_id_value,
detail=await _detail(db, current_user, project_id_value),
)
@router.delete("/projects/{project_id}", response_model=ShotReplicateDeleteOut)
async def delete_project_route_v2(
project_id: str,
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
try:
payload = await delete_project_v2(
db,
config=SHOT_REPLICATE_V2,
current_user=current_user,
project_id=project_id,
)
pending_ids = list(payload.get("pending_delete_resource_ids") or [])
await db.commit()
except HTTPException:
await db.rollback()
raise
if pending_ids:
try:
await cleanup_upload_resource_files_after_commit(db, resource_ids=pending_ids)
await db.commit()
except Exception:
await db.rollback()
return ShotReplicateDeleteOut(**payload)
+3
View File
@@ -206,6 +206,9 @@ class Settings(BaseSettings):
# - poll active 使用独立 Redis key,避免影响稳定的下载 active 注册表。
GENERATION_RECOVERY_BATCH_SIZE: int = 20
GENERATION_RECOVERY_MAX_ROUNDS: int = 1
GENERATION_CREATE_RECOVERY_INTERVAL_SECONDS: int = 60
GENERATION_CREATE_QUEUE_TIMEOUT_SECONDS: int = 5 * 60
MODULE_ASYNC_RECOVERY_INTERVAL_SECONDS: int = 60
POLL_RECOVERY_BATCH_SIZE: int = 20
POLL_TASK_LEASE_SECONDS: int = 5 * 60
POLL_TASK_QUEUE_TIMEOUT_SECONDS: int = 2 * 60
+1
View File
@@ -18,6 +18,7 @@ class CeleryTaskName(str, Enum):
DOWNLOAD_GENERATION_RESULT = "generation.download_generation_result_task"
RECOVER_DOWNLOAD = "generation.recover_download_tasks_once"
RECOVER_GENERATION = "generation.recover_generation_tasks_once"
RECOVER_CREATE = "generation.recover_create_tasks_once"
VIDEO_UPSCALE_EXECUTE_LOCAL = "video_upscale.execute_local"
VIDEO_UPSCALE_SUBMIT_REMOTE = "video_upscale.submit_remote"
VIDEO_UPSCALE_POLL_REMOTE = "video_upscale.poll_remote"
+20
View File
@@ -27,6 +27,13 @@ class LogSourceEnum(StrEnum):
RECOVERY = "recovery"
REMOTE_API = "remote_api"
class ModuleGenerationFlowVersionEnum(StrEnum):
"""模块生成项目流程版本。"""
V1 = "v1"
V2 = "v2"
class ModuleProjectStatusEnum(StrEnum):
"""通用模块项目状态。"""
@@ -72,6 +79,19 @@ class ModuleEventTypeEnum(StrEnum):
MEDIA_REFUND = "MEDIA_REFUND"
PROMPT_BILLING_SUCCESS = "PROMPT_BILLING_SUCCESS"
PROMPT_BILLING_FAILED = "PROMPT_BILLING_FAILED"
V1_CREATE_BLOCKED = "V1_CREATE_BLOCKED"
FLOW_VERSION_MISMATCH = "FLOW_VERSION_MISMATCH"
V2_PROJECT_CREATED = "V2_PROJECT_CREATED"
V2_VIDEO_PROMPT_AUTO_CREATED = "V2_VIDEO_PROMPT_AUTO_CREATED"
V2_VIDEO_PROMPT_REGENERATED = "V2_VIDEO_PROMPT_REGENERATED"
V2_VIDEO_PROMPT_DISPATCHED = "V2_VIDEO_PROMPT_DISPATCHED"
V2_VIDEO_PROMPT_REGISTRY_FAILED = "V2_VIDEO_PROMPT_REGISTRY_FAILED"
V2_VIDEO_PROMPT_DISPATCH_FAILED = "V2_VIDEO_PROMPT_DISPATCH_FAILED"
STEP_SUPERSEDED = "STEP_SUPERSEDED"
STALE_STEP_RESULT_DISCARDED = "STALE_STEP_RESULT_DISCARDED"
GENERATION_REFERENCE_OPTION_SAVED = "GENERATION_REFERENCE_OPTION_SAVED"
GENERATION_REFERENCE_INCLUDED = "GENERATION_REFERENCE_INCLUDED"
GENERATION_REFERENCE_EXCLUDED = "GENERATION_REFERENCE_EXCLUDED"
class ModulePromptTypeEnum(StrEnum):
@@ -30,6 +30,7 @@ class HotOpeningStepIOSchemaVersionEnum(StrEnum):
"""爆款开头复刻子任务 input_json/output_json 结构版本。"""
V1 = "hot_opening_step_io_v1"
V2 = "hot_opening_step_io_v2"
class HotOpeningLogEventEnum(StrEnum):
@@ -24,3 +24,4 @@ class ModuleGenerationFlowConfig:
cancel_chat_task_error_message: str
material_video_url_editable: bool = True
step_io_schema_version: str = "module_generation_step_io_v1"
expected_flow_version: str | None = None
@@ -29,6 +29,7 @@ class ShotReplicateStepIOSchemaVersionEnum(StrEnum):
"""拆镜复刻子任务 input_json/output_json 结构版本。"""
V1 = "shot_replicate_step_io_v1"
V2 = "shot_replicate_step_io_v2"
class ShotTaskSetStatusEnum(StrEnum):
+2
View File
@@ -12,6 +12,7 @@ from app.config import settings
from app.models import init_database, close_database
from app.utils.redis import init_redis, close_redis
from app.api.v1 import api_router
from app.api.v2 import api_router_v2
from app.middleware.logging import RequestLoggingMiddleware
from app.middleware.anti_crawler import AntiCrawlerMiddleware
from app.middleware.rate_limit import RateLimitMiddleware
@@ -545,6 +546,7 @@ def create_app() -> FastAPI:
# Routes
application.include_router(api_router, prefix="/api")
application.include_router(api_router_v2, prefix="/api/v2")
# Static files for uploads
upload_dir = os.path.abspath(settings.UPLOAD_LOCAL_PATH)
@@ -37,6 +37,9 @@ class GenerationRecord(Base, TimestampMixin, SoftDeleteMixin):
video_cover_url: Mapped[str | None] = mapped_column(String(512), nullable=True)
image_url: Mapped[str | None] = mapped_column(String(512), nullable=True)
media_references: Mapped[str | None] = mapped_column(Text, nullable=True)
include_media_references: Mapped[bool] = mapped_column(
Boolean, nullable=False, default=False, server_default="false"
)
video_url_expires_at: Mapped[datetime | None] = mapped_column(
DateTime(timezone=True), nullable=True
)
@@ -5,6 +5,7 @@ from datetime import datetime
from sqlalchemy import DateTime, ForeignKey, Index, String, Text, text
from sqlalchemy.orm import Mapped, mapped_column
from app.enums.common import ModuleGenerationFlowVersionEnum
from app.models.base import Base, SoftDeleteMixin, TimestampMixin
@@ -36,6 +37,12 @@ class ModuleGenerationProject(Base, TimestampMixin, SoftDeleteMixin):
String(32), ForeignKey("users.id", ondelete="CASCADE"), index=True, nullable=False
)
module: Mapped[str] = mapped_column(String(64), index=True, nullable=False)
flow_version: Mapped[str] = mapped_column(
String(16),
nullable=False,
default=ModuleGenerationFlowVersionEnum.V1.value,
server_default=ModuleGenerationFlowVersionEnum.V1.value,
)
title: Mapped[str | None] = mapped_column(String(160), nullable=True)
status: Mapped[str] = mapped_column(String(32), default="pending", index=True)
current_step_code: Mapped[str | None] = mapped_column(String(64), nullable=True, index=True)
@@ -15,12 +15,12 @@ _STEP_JSON_TYPE = JSON().with_variant(JSONB, "postgresql")
class ModuleGenerationStep(Base, TimestampMixin, SoftDeleteMixin):
"""通用模块生成步骤表。
爆款开头复刻固定步骤:
1 material_input
2 image_prompt_optimize
3 image_generate
4 video_prompt_optimize
5 video_generate
V1 固定五步:material_input / image_prompt_optimize / image_generate /
video_prompt_optimize / video_generate。
V2 固定三步:material_input / video_prompt_optimize / video_generate
version 只表示同一步骤的重建版本,不表示项目流程版本;项目版本由
ModuleGenerationProject.flow_version 保存。
input_json / output_json 使用 JSON/JSONB 存储。
建议结构:
+6
View File
@@ -25,6 +25,8 @@ class OptimizeParams(BaseModel):
class GenerateParams(BaseModel):
engine_id: str | None = Field(None, description="生成引擎ID;为空时优先沿用记录引擎,再回退默认引擎")
include_media_references: bool = Field(False, description="最终生成时是否携带提词阶段保存的附件")
aspect_ratio: str | None = None
resolution: str | None = None
image_size: str | None = None
@@ -56,6 +58,10 @@ class GenerationRecordOut(BaseModel):
video_cover_url: str | None = None
image_url: str | None = None
references: list[dict] | None = None
engine_id: str | None = None
engine_name: str | None = None
engine_snapshot: dict | None = None
include_media_references: bool = False
text_credits_cost: float = 0.0
# text_tokens_used: int = 0
credits_cost: float = 0.0
@@ -404,6 +404,8 @@ class HotOpeningMaterialOut(BaseModel):
source_project_name: str | None = Field(None, description="视频素材内容项目名称")
target_project_name: str | None = Field(None, description="生成项目名称")
core_content_point: str | None = Field(None, description="生成项目核心内容点")
project_description: str | None = Field(None, description="V2 可选项目描述")
video_config: dict[str, Any] | None = Field(None, description="V2 创建时保存的视频引擎、时长、比例和分辨率")
class HotOpeningImageGenerationOut(BaseModel):
@@ -445,6 +447,9 @@ class HotOpeningTaskDetailOut(BaseModel):
user_id: str | None = Field(None, description="所属用户ID;管理员后台排查使用")
user_name: str | None = Field(None, description="所属用户名;管理员后台排查使用")
module: str = Field(..., description="模块标识,爆款开头复刻固定为 hot_opening_replicate")
flow_version: str = Field("v1", description="流程版本:v1=历史五步骤,v2=简化三步骤")
step_count: int = Field(5, description="当前流程步骤数")
step_io_schema_version: str | None = Field(None, description="当前流程步骤 IO Schema 版本")
title: str | None = Field(None, description="项目标题,默认取生成项目名称")
status: str = Field(..., description="总任务状态:pending=已创建,waiting_user=等待用户操作,processing=处理中,completed=完成,failed=失败,cancelled=取消")
current_step_code: str | None = Field(None, description="当前所处步骤编码")
@@ -467,6 +472,8 @@ class HotOpeningTaskListItemOut(BaseModel):
user_id: str | None = Field(None, description="所属用户ID;管理员后台排查使用")
user_name: str | None = Field(None, description="所属用户名;管理员后台排查使用")
module: str = Field(..., description="模块标识")
flow_version: str = Field("v1", description="流程版本")
step_count: int = Field(5, description="流程步骤数")
title: str | None = Field(None, description="项目标题")
status: str = Field(..., description="总任务状态:pending=已创建,waiting_user=等待用户操作,processing=处理中,completed=完成,failed=失败,cancelled=取消")
current_step_code: str | None = Field(None, description="当前步骤")
@@ -0,0 +1,113 @@
from __future__ import annotations
from typing import Any
from pydantic import BaseModel, ConfigDict, Field, field_validator
class ModuleGenerationVideoConfigV2(BaseModel):
model_config = ConfigDict(extra="ignore")
engine_id: str = Field(..., min_length=1, max_length=32, description="视频引擎ID")
duration: int = Field(..., ge=1, le=120, description="视频时长,单位秒")
aspect_ratio: str = Field(..., min_length=1, max_length=16, description="视频比例")
resolution: str = Field(..., min_length=1, max_length=16, description="目标分辨率")
@field_validator("engine_id", "aspect_ratio", "resolution", mode="before")
@classmethod
def _strip_required(cls, value: str) -> str:
value = str(value or "").strip()
if not value:
raise ValueError("字段不能为空")
return value
class HotOpeningTaskCreateV2(BaseModel):
model_config = ConfigDict(extra="ignore")
material_video_url: str = Field(..., min_length=1, description="爆款参考视频,步骤2分析使用,步骤3绝不携带")
material_video_resource_id: str | None = Field(None, max_length=32)
material_video_duration_seconds: float | None = Field(None, gt=0)
material_image_url: str | None = Field(None, description="可选参考图片;仅支持本地上传/历史素材")
material_image_resource_id: str | None = Field(None, max_length=32)
source_project_name: str | None = Field(None, max_length=40)
target_project_name: str | None = Field(None, max_length=40)
core_content_point: str | None = Field(None, max_length=200)
project_description: str | None = Field(None, max_length=1000)
target_platform: str | None = Field(None, max_length=64)
video_config: ModuleGenerationVideoConfigV2
idempotency_key: str | None = Field(None, max_length=64)
@field_validator(
"material_video_url",
"material_image_url",
"source_project_name",
"target_project_name",
"core_content_point",
"project_description",
"target_platform",
"idempotency_key",
mode="before",
)
@classmethod
def _strip_text(cls, value: str | None) -> str | None:
if value is None:
return None
value = str(value).strip()
return value or None
class ShotReplicateProjectCreateV2(BaseModel):
model_config = ConfigDict(extra="ignore")
material_image_url: str | None = Field(None, description="可选参考图片;仅支持本地上传/历史素材")
material_image_resource_id: str | None = Field(None, max_length=32)
target_project_name: str | None = Field(None, max_length=40)
core_content_point: str | None = Field(None, max_length=200)
project_description: str | None = Field(None, max_length=1000)
target_platform: str | None = Field(None, max_length=64)
video_config: ModuleGenerationVideoConfigV2
idempotency_key: str | None = Field(None, max_length=64)
@field_validator(
"material_image_url",
"target_project_name",
"core_content_point",
"project_description",
"target_platform",
"idempotency_key",
mode="before",
)
@classmethod
def _strip_text(cls, value: str | None) -> str | None:
if value is None:
return None
value = str(value).strip()
return value or None
class ModuleVideoPromptRetryV2(BaseModel):
model_config = ConfigDict(extra="forbid")
video_config: ModuleGenerationVideoConfigV2
class ModuleVideoPromptSchemaUpdateV2(BaseModel):
model_config = ConfigDict(extra="ignore")
prompt_schema: dict[str, Any] = Field(..., description="允许编辑字段的 JSON patch")
@field_validator("prompt_schema")
@classmethod
def _non_empty_schema(cls, value: dict[str, Any]) -> dict[str, Any]:
if not isinstance(value, dict) or not value:
raise ValueError("prompt_schema 必须为非空对象")
return value
class ModuleGenerationV2ActionOut(BaseModel):
message: str
project_id: str
step_id: str | None = None
flow_version: str = "v2"
detail: Any | None = None
@@ -56,6 +56,10 @@ class RecentGenerationItemOut(BaseModel):
None,
description="关联拆镜片段ID。仅 shot_replicate 模块可能有值,来源 shot_replicate_segments.id;其他模块返回 null。",
)
module_project_flow_version: str | None = Field(
None,
description="关联模块项目流程版本。hot_opening_replicate/shot_replicate 返回 v1/v2;其他模块返回 null。",
)
module_project_id: str | None = Field(
None,
description="通用模块项目ID。hot_opening_replicate/shot_replicate 模块可能有值,来源 module_generation_steps.project_idproject/chat_ai 返回 null。",
@@ -395,6 +395,8 @@ class ShotReplicateMaterialOut(BaseModel):
source_project_name: str | None = Field(None, description="视频素材内容项目名称")
target_project_name: str | None = Field(None, description="生成项目名称")
core_content_point: str | None = Field(None, description="生成项目核心内容点")
project_description: str | None = Field(None, description="V2 可选项目描述")
video_config: dict[str, Any] | None = Field(None, description="V2 创建时保存的视频引擎、时长、比例和分辨率")
class ShotReplicateImageGenerationOut(BaseModel):
@@ -436,6 +438,9 @@ class ShotReplicateTaskDetailOut(BaseModel):
user_id: str | None = Field(None, description="所属用户ID;管理员后台排查使用")
user_name: str | None = Field(None, description="所属用户名;管理员后台排查使用")
module: str = Field(..., description="模块标识,拆镜复刻固定为 shot_replicate")
flow_version: str = Field("v1", description="流程版本:v1=历史五步骤,v2=简化三步骤")
step_count: int = Field(5, description="当前流程步骤数")
step_io_schema_version: str | None = Field(None, description="当前流程步骤 IO Schema 版本")
title: str | None = Field(None, description="项目标题,默认取生成项目名称")
status: str = Field(..., description="总任务状态:pending/waiting_user/processing/completed/failed/cancelled")
current_step_code: str | None = Field(None, description="当前所处步骤编码")
@@ -456,6 +461,8 @@ class ShotReplicateTaskListItemOut(BaseModel):
id: str = Field(..., description="总任务项目ID。这个ID就是前端项目ID")
project_id: str = Field(..., description="兼容前端命名,等同于 id")
module: str = Field(..., description="模块标识")
flow_version: str = Field("v1", description="流程版本")
step_count: int = Field(5, description="流程步骤数")
title: str | None = Field(None, description="项目标题")
status: str = Field(..., description="总任务状态")
current_step_code: str | None = Field(None, description="当前步骤")
@@ -702,6 +709,7 @@ class ShotSegmentOut(BaseModel):
module_project_title: str | None = Field(None, description="关联拆镜复刻项目标题;后台片段列表展示使用")
module_project_status: str | None = Field(None, description="关联拆镜复刻项目状态;后台片段列表展示使用")
module_project_current_step_code: str | None = Field(None, description="关联拆镜复刻项目当前步骤;后台片段列表展示使用")
module_project_flow_version: str | None = Field(None, description="关联复刻项目流程版本:v1/v2")
created_at: NaiveDatetimeOptional = Field(None, description="创建时间")
updated_at: NaiveDatetimeOptional = Field(None, description="更新时间")
@@ -113,7 +113,11 @@ def build_video_snapshot(engine: VideoEngine, ratio: str, resolution: str, durat
"supported_resolutions": parse_json_list(engine.supported_resolutions, []),
"supported_durations": parse_json_list(engine.supported_durations, []),
"max_duration": engine.max_duration,
"max_audio_count": engine.max_audio_count,
"max_image_count": int(getattr(engine, "max_image_count", 0) or 0),
"max_video_count": int(getattr(engine, "max_video_count", 0) or 0),
"max_audio_count": int(getattr(engine, "max_audio_count", 0) or 0),
"supports_universal_reference": bool(getattr(engine, "supports_universal_reference", False)),
"supports_first_last_frame": bool(getattr(engine, "supports_first_last_frame", False)),
"multi_generation_enabled": bool(getattr(engine, "multi_generation_enabled", False)),
"max_generation_count": normalize_generation_count(getattr(engine, "max_generation_count", 1)),
"selected_ratio": ratio,
@@ -10,6 +10,7 @@ from sqlalchemy.ext.asyncio import AsyncSession
from app.enums.credit_record import CreditRecordBillingScene, CreditRecordChargeKind, CreditRecordOwnerType, CreditRecordSourceModule
from app.models.credit_record import CreditRecord
from app.models.generation_record import GenerationRecord
from app.services.generation.media_reference_service import calculate_media_reference_usage
from app.models.module_generation_step import ModuleGenerationStep
from app.models.token_usage import TokenUsage
from app.models.system_config import SystemConfig
@@ -461,6 +462,7 @@ async def charge_generation_media_by_params(
resolution: str | None = None,
engine_id: str | None = None,
input_video_duration: float | None = None,
input_image_count: int | None = None,
project_name: str | None = None,
description_prefix: str = "AI创作-",
owner_type: str = OWNER_CHAT_GENERATION_TASK,
@@ -510,7 +512,9 @@ async def charge_generation_media_by_params(
if gen_type == "image":
size = image_size or "2K"
unit_amount = await calc_image_credits(db, size, engine_id=engine_id)
unit_amount = await calc_image_credits(
db, size, engine_id=engine_id, input_image_count=input_image_count
)
amount = round(unit_amount * quantity, 2)
items.append(
await deduct_credits_locked_once(
@@ -530,6 +534,7 @@ async def charge_generation_media_by_params(
db, duration or 5, resolution or "720p",
engine_id=engine_id,
input_video_duration=input_video_duration,
input_image_count=input_image_count,
)
amount = round(unit_amount * quantity, 2)
items.append(
@@ -560,6 +565,10 @@ async def charge_generation_media_for_record(
attempt_no: int | None = None,
engine_id: str | None = None,
) -> BillingSummary:
reference_usage = calculate_media_reference_usage(
record.media_references,
include=bool(record.include_media_references),
)
return await charge_generation_media_by_params(
db,
user_id=record.user_id,
@@ -569,6 +578,8 @@ async def charge_generation_media_for_record(
duration=record.duration,
resolution=record.resolution,
engine_id=engine_id or getattr(record, "engine_id", None),
input_video_duration=reference_usage.input_video_duration or None,
input_image_count=reference_usage.image_count or None,
project_name=project_name,
description_prefix=description_prefix,
owner_type=OWNER_GENERATION_RECORD,
@@ -0,0 +1,148 @@
from __future__ import annotations
import json
from dataclasses import dataclass
from typing import Any, Iterable
from fastapi import HTTPException
@dataclass(frozen=True, slots=True)
class MediaReferenceUsage:
image_count: int = 0
video_count: int = 0
audio_count: int = 0
input_video_duration: float = 0.0
input_audio_duration: float = 0.0
def parse_media_references(value: str | list[dict] | None) -> list[dict]:
if not value:
return []
data: Any = value
if isinstance(value, str):
try:
data = json.loads(value)
except (TypeError, json.JSONDecodeError):
return []
if not isinstance(data, list):
return []
return [item for item in data if isinstance(item, dict)]
def _reference_type(item: dict) -> str:
value = str(item.get("type") or item.get("media_type") or "").strip().lower()
if value in {"image", "video", "audio"}:
return value
mime = str(item.get("mime_type") or item.get("content_type") or "").lower()
if mime.startswith("image/"):
return "image"
if mime.startswith("video/"):
return "video"
if mime.startswith("audio/"):
return "audio"
url = str(item.get("url") or item.get("file_url") or "").lower().split("?", 1)[0]
if url.endswith((".png", ".jpg", ".jpeg", ".webp", ".gif", ".bmp")):
return "image"
if url.endswith((".mp4", ".mov", ".webm", ".mkv", ".avi")):
return "video"
if url.endswith((".mp3", ".wav", ".m4a", ".aac", ".ogg", ".flac")):
return "audio"
return ""
def _duration(item: dict) -> float:
for key in ("duration", "duration_seconds", "video_duration", "audio_duration"):
try:
value = float(item.get(key) or 0)
except (TypeError, ValueError):
continue
if value > 0:
return value
return 0.0
def calculate_media_reference_usage(
references: str | list[dict] | None,
*,
include: bool,
) -> MediaReferenceUsage:
if not include:
return MediaReferenceUsage()
image_count = video_count = audio_count = 0
video_duration = audio_duration = 0.0
for item in parse_media_references(references):
media_type = _reference_type(item)
if media_type == "image":
image_count += 1
elif media_type == "video":
video_count += 1
video_duration += _duration(item)
elif media_type == "audio":
audio_count += 1
audio_duration += _duration(item)
return MediaReferenceUsage(
image_count=image_count,
video_count=video_count,
audio_count=audio_count,
input_video_duration=round(video_duration, 3),
input_audio_duration=round(audio_duration, 3),
)
def filter_references_by_type(
references: str | list[dict] | None,
*,
allowed_types: Iterable[str],
max_count: int | None = None,
) -> list[dict]:
allowed = {str(item).lower() for item in allowed_types}
result = [item for item in parse_media_references(references) if _reference_type(item) in allowed]
if max_count is not None:
return result[: max(0, int(max_count))]
return result
def validate_media_reference_usage_for_engine(
usage: MediaReferenceUsage,
*,
gen_type: str,
engine: Any,
) -> None:
"""按所选引擎能力校验最终实际会发送的附件。"""
normalized = str(gen_type or "").strip().lower()
if normalized == "video":
if not bool(getattr(engine, "supports_universal_reference", True)) and (
usage.image_count or usage.video_count or usage.audio_count
):
raise HTTPException(status_code=400, detail="当前视频引擎不支持参考附件")
limits = {
"图片": int(getattr(engine, "max_image_count", 0) or 0),
"视频": int(getattr(engine, "max_video_count", 0) or 0),
"音频": int(getattr(engine, "max_audio_count", 0) or 0),
}
counts = {"图片": usage.image_count, "视频": usage.video_count, "音频": usage.audio_count}
for label, count in counts.items():
limit = limits[label]
if count > limit:
raise HTTPException(
status_code=400,
detail=f"当前视频引擎最多支持 {limit}{label}附件,当前为 {count}",
)
return
if normalized == "image":
if usage.video_count or usage.audio_count:
raise HTTPException(status_code=400, detail="图片生成只能携带图片附件")
limit = int(
getattr(engine, "max_reference_image_count", None)
if getattr(engine, "max_reference_image_count", None) is not None
else getattr(engine, "max_image_count", 0)
or 0
)
if usage.image_count > limit:
raise HTTPException(
status_code=400,
detail=f"当前图片引擎最多支持 {limit} 张参考图,当前为 {usage.image_count}",
)
return
raise HTTPException(status_code=400, detail="不支持的生成类型")
@@ -58,6 +58,18 @@ def owner_type_of(owner: GenerationOwner) -> str:
raise TypeError(f"不支持的生成任务对象: {type(owner)!r}")
def owner_include_media_references(owner: GenerationOwner) -> bool:
"""返回本次供应商创建是否应携带附件。
ChatGenerationTask 延续原有行为;GenerationRecord 使用用户提交并持久化的开关。
"""
if isinstance(owner, ChatGenerationTask):
return True
if isinstance(owner, GenerationRecord):
return bool(owner.include_media_references)
raise TypeError(f"不支持的生成任务对象: {type(owner)!r}")
def owner_mode(owner: GenerationOwner) -> str:
if isinstance(owner, ChatGenerationTask):
return str(owner.generation_mode or GenerationMode.CHATAPI_ASYNC.value)
@@ -10,8 +10,11 @@ from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.config import settings
from app.models.chat_generation_task import ChatGenerationTask
from app.services.generation.pipeline.owner_service import GenerationOwner, owner_provider_task_id
from app.services.generation.pipeline.owner_service import (
GenerationOwner,
owner_include_media_references,
owner_provider_task_id,
)
from app.models.image_engine import ImageEngine
from app.models.video_engine import VideoEngine
from app.services.generation.log_service import log_provider_call
@@ -105,7 +108,7 @@ async def _create_video_task(db: AsyncSession, task: GenerationOwner) -> dict:
started = time.perf_counter()
async with provider_limit("ark_video_create", settings.ARK_VIDEO_CREATE_MAX_CONCURRENCY):
try:
provider_task_id = await submit_video_task(None, engine, task, include_media_references=isinstance(task, ChatGenerationTask))
provider_task_id = await submit_video_task(None, engine, task, include_media_references=owner_include_media_references(task))
response = {"task_id": provider_task_id}
await log_provider_call(
task,
@@ -169,7 +172,7 @@ async def create_image_sync_batch_result_with_engine(
None,
engine,
task,
include_media_references=isinstance(task, ChatGenerationTask),
include_media_references=owner_include_media_references(task),
generation_count=count,
)
response_data = result.get("response_data") or result
@@ -606,6 +606,8 @@ async def recover_one_generation_task(
ChatGenerationPipelineStage.CREATING_PROVIDER_TASK.value,
):
task.pipeline_stage = ChatGenerationPipelineStage.QUEUED.value
# 刷新更新时间形成创建队列保护窗口,避免 Beat 在任务尚未消费时每轮重复补投。
task.updated_at = current_time
# Release the recovery row lock before writing an event through the
# independent logging session or talking to the broker.
await db.commit()
@@ -622,12 +624,17 @@ async def recover_one_generation_task(
kwargs={"owner_type": GenerationOwnerType.CHAT_GENERATION_TASK.value, "generation_attempt_no": int(task.generation_attempt_no or 1)},
queue=CeleryQueue.GEN_CHATAPI_CREATE.value,
countdown=0,
task_id=(
f"generation-create:{GenerationOwnerType.CHAT_GENERATION_TASK.value}:"
f"{task.id}:attempt:{int(task.generation_attempt_no or 1)}"
),
)
return "recover_create_no_remote_no_provider_before_deadline"
# result_ready 但没有 URL 是脏状态;未过 deadline 时回创建队列重新处理,过期上面已标记超时。
if task.pipeline_stage == ChatGenerationPipelineStage.RESULT_READY.value:
task.pipeline_stage = ChatGenerationPipelineStage.QUEUED.value
task.updated_at = current_time
await db.commit()
await _remove_poll_active(_chat_registry_id(task))
await log_task_event(
@@ -641,6 +648,10 @@ async def recover_one_generation_task(
kwargs={"owner_type": GenerationOwnerType.CHAT_GENERATION_TASK.value, "generation_attempt_no": int(task.generation_attempt_no or 1)},
queue=CeleryQueue.GEN_CHATAPI_CREATE.value,
countdown=0,
task_id=(
f"generation-create:{GenerationOwnerType.CHAT_GENERATION_TASK.value}:"
f"{task.id}:attempt:{int(task.generation_attempt_no or 1)}"
),
)
return "recover_create_result_ready_no_url_before_deadline"
@@ -887,6 +898,61 @@ async def recover_generation_tasks_once(db: AsyncSession) -> dict[str, Any]:
"results": results,
}
async def recover_stale_create_tasks_once(db: AsyncSession) -> dict[str, Any]:
"""轻量恢复长时间未消费的创建阶段 ChatGenerationTask。
只扫描 queued/preparing/creating_provider_task,避免周期任务重复执行完整
provider poll、下载和主任务汇总逻辑。
"""
current_time = _now()
cutoff = current_time - timedelta(
seconds=max(1, int(settings.GENERATION_CREATE_QUEUE_TIMEOUT_SECONDS or 300))
)
lease_expired_at = current_time
batch_size = max(1, int(settings.GENERATION_RECOVERY_BATCH_SIZE or 20))
result = await db.execute(
select(ChatGenerationTask.id)
.where(
ChatGenerationTask.deleted_at.is_(None),
ChatGenerationTask.generation_mode.in_(list(ALLOWED_GENERATION_MODES)),
ChatGenerationTask.status == ChatGenerationTaskStatus.GENERATING.value,
ChatGenerationTask.pipeline_stage.in_(
[
ChatGenerationPipelineStage.QUEUED.value,
ChatGenerationPipelineStage.PREPARING.value,
ChatGenerationPipelineStage.CREATING_PROVIDER_TASK.value,
]
),
ChatGenerationTask.remote_result_url.is_(None),
ChatGenerationTask.provider_task_id.is_(None),
ChatGenerationTask.seedance_task_id.is_(None),
ChatGenerationTask.updated_at <= cutoff,
(
ChatGenerationTask.provider_create_lease_until.is_(None)
| (ChatGenerationTask.provider_create_lease_until <= lease_expired_at)
),
)
.order_by(ChatGenerationTask.updated_at.asc(), ChatGenerationTask.id.asc())
.limit(batch_size)
)
task_ids = [str(value) for value in result.scalars().all()]
counts: dict[str, int] = {}
for task_id in task_ids:
task = await _load_chat_task_for_update(db, task_id)
if task is None:
await db.rollback()
action = "skip_missing_task"
else:
action = await recover_one_generation_task(
db,
task,
payload=None,
source="periodic_create_recovery",
)
counts[action] = counts.get(action, 0) + 1
return {"checked": len(task_ids), "results": counts}
async def dispatch_due_poll_tasks_once(db: AsyncSession) -> dict[str, Any]:
"""周期性轻量到期轮询调度。
@@ -26,6 +26,10 @@ from app.services.generation.ai.engine_service import (
parse_json_list as _parse_list,
)
from app.services.generation.billing_service import OWNER_CHAT_GENERATION_TASK, charge_generation_media_by_params
from app.services.generation.media_reference_service import (
calculate_media_reference_usage,
validate_media_reference_usage_for_engine,
)
from app.services.resource_capacity_service import assert_user_resource_capacity_available
from app.services.video_upscale.snapshot_service import build_video_upscale_snapshot
from app.services.private_portrait.reference_resolver import resolve_private_portrait_references
@@ -111,6 +115,8 @@ async def create_chat_generation_task_for_module(
raise HTTPException(status_code=400, detail=f"图片比例不支持: {proportion}")
px = px or normalize_px((sizes.get(size) or {}).get(proportion))
px = px or IMAGE_DEFAULT_PX
reference_usage = calculate_media_reference_usage(refs, include=True)
validate_media_reference_usage_for_engine(reference_usage, gen_type="image", engine=engine)
media_billing = await charge_generation_media_by_params(
db,
user_id=current_user.id,
@@ -118,6 +124,7 @@ async def create_chat_generation_task_for_module(
gen_type="image",
image_size=size,
engine_id=engine.id,
input_image_count=reference_usage.image_count or None,
project_name=billing_project_name,
description_prefix=billing_description_prefix,
owner_type=OWNER_CHAT_GENERATION_TASK,
@@ -174,6 +181,8 @@ async def create_chat_generation_task_for_module(
aspect_ratio=ratio,
supported_provider_resolutions=resolutions,
)
reference_usage = calculate_media_reference_usage(refs, include=True)
validate_media_reference_usage_for_engine(reference_usage, gen_type="video", engine=engine)
media_billing = await charge_generation_media_by_params(
db,
user_id=current_user.id,
@@ -182,6 +191,8 @@ async def create_chat_generation_task_for_module(
duration=selected_duration,
resolution=selected_resolution,
engine_id=engine.id,
input_video_duration=reference_usage.input_video_duration or None,
input_image_count=reference_usage.image_count or None,
project_name=billing_project_name,
description_prefix=billing_description_prefix,
owner_type=OWNER_CHAT_GENERATION_TASK,
@@ -117,6 +117,7 @@ FLOW_CONFIG = ModuleGenerationFlowConfig(
cancel_chat_task_error_message="爆款开头复刻步骤被重新生成或删除,旧生成任务已取消",
material_video_url_editable=True,
step_io_schema_version=STEP_IO_SCHEMA_VERSION,
expected_flow_version="v1",
)
@@ -422,12 +423,19 @@ async def project_to_detail_out(db: AsyncSession, project: ModuleGenerationProje
user_result = await db.execute(select(User.username).where(User.id == project.user_id).limit(1))
user_name = user_result.scalar_one_or_none()
flow_version = str(getattr(project, "flow_version", None) or "v1")
video_prompt_config = dict(video_prompt_input.get("video_config") or {})
video_prompt_engine_snapshot = dict(video_prompt_config.get("engine_snapshot") or {})
return HotOpeningTaskDetailOut(
id=project.id,
project_id=project.id,
user_id=project.user_id,
user_name=user_name,
module=project.module,
flow_version=flow_version,
step_count=3 if flow_version == "v2" else 5,
step_io_schema_version=("hot_opening_step_io_v2" if project.module == "hot_opening_replicate" else "shot_replicate_step_io_v2") if flow_version == "v2" else STEP_IO_SCHEMA_VERSION,
title=project.title,
status=project.status,
current_step_code=project.current_step_code,
@@ -442,6 +450,8 @@ async def project_to_detail_out(db: AsyncSession, project: ModuleGenerationProje
source_project_name=material_input.get("source_project_name"),
target_project_name=material_input.get("target_project_name"),
core_content_point=material_input.get("core_content_point"),
project_description=material_input.get("project_description"),
video_config=None if flow_version == "v2" else material_input.get("video_config"),
),
image_generation=HotOpeningImageGenerationOut(
prompt_step_id=image_prompt_step.id if image_prompt_step else None,
@@ -465,9 +475,9 @@ async def project_to_detail_out(db: AsyncSession, project: ModuleGenerationProje
schema_config_source=schema_config_source,
schema_config_version=schema_config_version,
schema_config_is_fallback=schema_config_is_fallback,
engine_id=video_snapshot.get("id") or video_generate_input.get("engine_id"),
engine_name=video_snapshot.get("name") or video_generate_input.get("engine_name"),
params=video_generate_input.get("params") or video_generate_input,
engine_id=video_snapshot.get("id") or video_generate_input.get("engine_id") or video_prompt_config.get("engine_id"),
engine_name=video_snapshot.get("name") or video_generate_input.get("engine_name") or video_prompt_engine_snapshot.get("name"),
params=video_generate_input.get("params") or video_prompt_config or video_generate_input,
chat_task_id=video_generate_step.chat_task_id if video_generate_step else None,
status=video_chat.status if video_chat else (video_generate_step.status if video_generate_step else None),
result_video_url=build_resource_signed_url(video_url) if video_url else None,
@@ -660,6 +670,8 @@ async def list_hot_opening_projects(
user_id=project.user_id,
user_name=user_name_map.get(project.user_id) if current_user.is_admin and project.user_id else None,
module=project.module,
flow_version=str(getattr(project, "flow_version", None) or "v1"),
step_count=3 if str(getattr(project, "flow_version", None) or "v1") == "v2" else 5,
title=project.title,
status=project.status,
current_step_code=project.current_step_code,
@@ -1532,6 +1544,9 @@ async def generate_video_from_prompt(
async def handle_chat_generation_task_completed(db: AsyncSession, task: ChatGenerationTask) -> None:
if not task or task.generation_mode != GENERATION_MODE:
return
from app.services.module_generation_v2.flow_service import handle_chat_generation_task_finished_v2
if await handle_chat_generation_task_finished_v2(db, task=task):
return
meta_result = await db.execute(
select(ModuleGenerationStep.id, ModuleGenerationStep.project_id).where(
ModuleGenerationStep.chat_task_id == task.id,
@@ -1614,6 +1629,9 @@ async def handle_chat_generation_task_completed(db: AsyncSession, task: ChatGene
async def handle_chat_generation_task_failed(db: AsyncSession, task: ChatGenerationTask) -> None:
if not task or task.generation_mode != GENERATION_MODE:
return
from app.services.module_generation_v2.flow_service import handle_chat_generation_task_finished_v2
if await handle_chat_generation_task_finished_v2(db, task=task):
return
meta_result = await db.execute(
select(ModuleGenerationStep.id, ModuleGenerationStep.project_id).where(
ModuleGenerationStep.chat_task_id == task.id,
@@ -1527,7 +1527,7 @@ async def optimize_hot_opening_video_prompt(
target_project_name: str,
core_content_point: str,
material_video_url: str,
generated_image_url: str,
generated_image_url: str | None,
video_config: dict[str, Any],
target_platform: str = "抖音",
schema_config_snapshot: Any | None = None,
@@ -1538,16 +1538,19 @@ async def optimize_hot_opening_video_prompt(
) -> tuple[dict[str, Any], str, dict[str, Any]]:
duration = int(video_config["duration"])
from app.utils.media import media_to_base64, get_llm_media_as_base64
if await get_llm_media_as_base64(db):
use_base64 = await get_llm_media_as_base64(db)
if use_base64:
video_url_final = await media_to_base64(material_video_url, "video/mp4")
image_url_final = await media_to_base64(generated_image_url, "image/png")
else:
video_url_final = _build_file_url_or_data_uri(material_video_url)
image_url_final = _build_file_url_or_data_uri(generated_image_url)
references = [
{"type": "video", "url": video_url_final},
{"type": "image", "url": image_url_final},
]
references = [{"type": "video", "url": video_url_final}]
if generated_image_url:
image_url_final = (
await media_to_base64(generated_image_url, "image/png")
if use_base64
else _build_file_url_or_data_uri(generated_image_url)
)
references.append({"type": "image", "url": image_url_final})
client_schema = build_dynamic_schema(video_config, schema_config_snapshot)
reference_video_fps = int(video_config.get("reference_video_fps") or DEFAULT_REFERENCE_VIDEO_FPS)
@@ -1690,7 +1693,7 @@ async def optimize_hot_opening_video_prompt(
"input_tokens": int(usage.get("prompt_tokens") or 0),
"output_tokens": int(usage.get("completion_tokens") or 0),
"total_tokens": int(usage.get("total_tokens") or 0),
"log_user_message": log_user_message,
# "log_user_message": log_user_message,
}
token_usage_id = generate_id()
db.add(
@@ -17,6 +17,7 @@ from app.enums.shot_replicate import (
ShotSegmentAnalysisStatusEnum,
ShotSplitStatusEnum,
)
from app.models.module_generation_project import ModuleGenerationProject
from app.models.module_generation_step import ModuleGenerationStep
from app.models.shot_replicate_segment import ShotReplicateSegment
from app.models.shot_replicate_task_set import ShotReplicateTaskSet
@@ -47,6 +48,7 @@ TASK_HOT_IMAGE_PROMPT = "hot_opening.start_image_prompt_optimize"
TASK_HOT_VIDEO_PROMPT = "hot_opening.start_video_prompt_optimize"
TASK_SHOT_IMAGE_PROMPT = "shot_replicate.start_image_prompt_optimize"
TASK_SHOT_VIDEO_PROMPT = "shot_replicate.start_video_prompt_optimize"
TASK_MODULE_V2_VIDEO_PROMPT = "module_generation_v2.start_video_prompt_optimize"
TASK_SHOT_ANALYZE_ORIGINAL = "shot_replicate.analyze_original_video"
TASK_SHOT_ANALYZE_CUSTOM_SEGMENT = "shot_replicate.analyze_custom_segment_video"
TASK_SHOT_SPLIT_ONE = "shot_replicate.split_one_segment"
@@ -511,9 +513,24 @@ async def _recover_stale_module_steps(db: AsyncSession, *, limit: int) -> dict[s
.with_for_update(skip_locked=True)
)
steps = list(result.scalars().all())
project_flow_map: dict[str, str] = {}
project_ids = list({step.project_id for step in steps})
if project_ids:
project_result = await db.execute(
select(ModuleGenerationProject.id, ModuleGenerationProject.flow_version).where(
ModuleGenerationProject.id.in_(project_ids),
ModuleGenerationProject.deleted_at.is_(None),
)
)
project_flow_map = {
str(project_id): str(flow_version or "v1")
for project_id, flow_version in project_result.all()
}
results: dict[str, int] = {}
for step in steps:
if step.module == HOT_MODULE and step.step_code == HotOpeningStepCodeEnum.IMAGE_PROMPT_OPTIMIZE.value:
if project_flow_map.get(step.project_id, "v1") == "v2" and step.step_code == HotOpeningStepCodeEnum.VIDEO_PROMPT_OPTIMIZE.value:
task_name = TASK_MODULE_V2_VIDEO_PROMPT
elif step.module == HOT_MODULE and step.step_code == HotOpeningStepCodeEnum.IMAGE_PROMPT_OPTIMIZE.value:
task_name = TASK_HOT_IMAGE_PROMPT
elif step.module == HOT_MODULE and step.step_code == HotOpeningStepCodeEnum.VIDEO_PROMPT_OPTIMIZE.value:
task_name = TASK_HOT_VIDEO_PROMPT
@@ -59,6 +59,18 @@ async def get_project_for_user(
project = result.scalar_one_or_none()
if not project:
raise HTTPException(status_code=404, detail=config.project_not_found_message)
if config.expected_flow_version:
actual_flow_version = str(getattr(project, "flow_version", None) or "v1")
if actual_flow_version != config.expected_flow_version:
raise HTTPException(
status_code=409,
detail={
"message": "项目流程版本与当前接口版本不匹配",
"project_id": project.id,
"flow_version": actual_flow_version,
"expected_flow_version": config.expected_flow_version,
},
)
return project
@@ -81,7 +81,19 @@ def build_step_output(
def is_wrapped_step_io(value: Any, *, schema_version: str = STEP_IO_SCHEMA_VERSION) -> bool:
return isinstance(value, dict) and value.get("schema_version") == schema_version
if not isinstance(value, dict):
return False
current = str(value.get("schema_version") or "")
if current == schema_version:
return True
# V1/V2 查询接口共用同一组解析器。只要结构符合统一步骤 IO 包装,
# 即按包装结构解析,避免 V2 被当成普通 payload。
return bool(
current
and value.get("step_code")
and "payload" in value
and ("status" in value or "source" in value)
)
def step_payload(value: Any, *, schema_version: str = STEP_IO_SCHEMA_VERSION) -> dict[str, Any]:
@@ -0,0 +1 @@
"""爆款复刻/拆镜复刻 V2 三步骤公共流程。"""
@@ -0,0 +1,75 @@
from __future__ import annotations
from dataclasses import dataclass
from app.enums.hot_opening_replicate import (
HotOpeningGenerationModeEnum,
HotOpeningStepIOSchemaVersionEnum,
ModuleCodeEnum as HotModuleCodeEnum,
)
from app.enums.shot_replicate import (
ModuleCodeEnum as ShotModuleCodeEnum,
ShotReplicateGenerationModeEnum,
ShotReplicateStepIOSchemaVersionEnum,
)
from app.enums.module_generation_flow import ModuleGenerationFlowConfig
MATERIAL_INPUT = "material_input"
VIDEO_PROMPT_OPTIMIZE = "video_prompt_optimize"
VIDEO_GENERATE = "video_generate"
@dataclass(frozen=True, slots=True)
class ModuleGenerationV2Config:
module: str
generation_mode: str
io_schema_version: str
display_name: str
project_not_found_message: str
material_video_locked: bool
@property
def flow_config(self) -> ModuleGenerationFlowConfig:
return ModuleGenerationFlowConfig(
module=self.module,
step_index_map={
MATERIAL_INPUT: 1,
VIDEO_PROMPT_OPTIMIZE: 2,
VIDEO_GENERATE: 3,
},
material_step_code=MATERIAL_INPUT,
image_prompt_step_code="__v2_no_image_prompt__",
image_generate_step_code="__v2_no_image_generate__",
video_prompt_step_code=VIDEO_PROMPT_OPTIMIZE,
video_generate_step_code=VIDEO_GENERATE,
project_not_found_message=self.project_not_found_message,
step_not_found_message="V2 子任务不存在",
cancel_chat_task_error_message=f"{self.display_name}步骤被重建或删除,旧生成任务已取消",
material_video_url_editable=not self.material_video_locked,
step_io_schema_version=self.io_schema_version,
expected_flow_version="v2",
)
HOT_OPENING_V2 = ModuleGenerationV2Config(
module=HotModuleCodeEnum.HOT_OPENING_REPLICATE.value,
generation_mode=HotOpeningGenerationModeEnum.HOT_OPENING_REPLICATE.value,
io_schema_version=HotOpeningStepIOSchemaVersionEnum.V2.value,
display_name="爆款开头复刻",
project_not_found_message="爆款开头复刻 V2 项目不存在",
material_video_locked=True,
)
SHOT_REPLICATE_V2 = ModuleGenerationV2Config(
module=ShotModuleCodeEnum.SHOT_REPLICATE.value,
generation_mode=ShotReplicateGenerationModeEnum.SHOT_REPLICATE.value,
io_schema_version=ShotReplicateStepIOSchemaVersionEnum.V2.value,
display_name="拆镜复刻",
project_not_found_message="拆镜复刻 V2 项目不存在",
material_video_locked=True,
)
CONFIG_BY_MODULE = {
HOT_OPENING_V2.module: HOT_OPENING_V2,
SHOT_REPLICATE_V2.module: SHOT_REPLICATE_V2,
}
@@ -0,0 +1,105 @@
from __future__ import annotations
from dataclasses import dataclass
from fastapi import HTTPException
from app.enums.celery_queue import CeleryQueue
from app.enums.common import ModuleEventTypeEnum
from app.services.module_async_recovery_service import (
TASK_MODULE_V2_VIDEO_PROMPT,
register_module_step_task,
)
from app.services.module_generation_log_service import log_module_error, log_module_event_file
from app.services.module_generation_v2.config import VIDEO_PROMPT_OPTIMIZE, ModuleGenerationV2Config
from app.tasks.celery_app import celery_app
from app.tasks.module_generation_v2_tasks import start_video_prompt_optimize_v2
@dataclass(frozen=True)
class VideoPromptDispatchResult:
registry_success: bool
celery_success: bool
registry_error: str | None = None
celery_error: str | None = None
@property
def recoverable(self) -> bool:
return self.registry_success or self.celery_success
def ensure_v2_celery_enabled() -> None:
if celery_app is None:
raise HTTPException(status_code=503, detail="Celery未启用")
async def dispatch_video_prompt_v2(
*,
config: ModuleGenerationV2Config,
project_id: str,
step_id: str,
) -> VideoPromptDispatchResult:
"""注册并投递 V2 视频提词任务。
Redis 注册成功但 Celery 直投失败时由周期恢复任务补投Celery 成功但
Redis 注册失败时任务仍可正常执行只有两个通道都失败时由 API 补偿落库为失败
"""
registry_error: Exception | None = None
try:
await register_module_step_task(
module=config.module,
project_id=project_id,
step_id=step_id,
step_code=VIDEO_PROMPT_OPTIMIZE,
task_name=TASK_MODULE_V2_VIDEO_PROMPT,
)
except Exception as exc:
registry_error = exc
log_module_error(
module=config.module,
event_type=ModuleEventTypeEnum.V2_VIDEO_PROMPT_REGISTRY_FAILED.value,
project_id=project_id,
step_id=step_id,
message="V2 视频提词 Redis 活跃注册失败,将继续尝试 Celery 直投",
exc=exc,
)
celery_error: Exception | None = None
try:
start_video_prompt_optimize_v2.apply_async(
args=[project_id, step_id],
queue=CeleryQueue.GEN_CHATAPI_CREATE.value,
countdown=0,
task_id=f"module-v2-video-prompt:{step_id}",
)
except Exception as exc:
celery_error = exc
log_module_error(
module=config.module,
event_type=ModuleEventTypeEnum.V2_VIDEO_PROMPT_DISPATCH_FAILED.value,
project_id=project_id,
step_id=step_id,
message="V2 视频提词 Celery 投递失败",
detail={"redis_registry_available": registry_error is None},
exc=exc,
)
result = VideoPromptDispatchResult(
registry_success=registry_error is None,
celery_success=celery_error is None,
registry_error=str(registry_error) if registry_error else None,
celery_error=str(celery_error) if celery_error else None,
)
if result.celery_success:
log_module_event_file(
module=config.module,
event_type=ModuleEventTypeEnum.V2_VIDEO_PROMPT_DISPATCHED.value,
project_id=project_id,
step_id=step_id,
message="V2 视频提词任务已投递",
detail={
"queue": CeleryQueue.GEN_CHATAPI_CREATE.value,
"redis_registry_available": result.registry_success,
},
)
return result
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,35 @@
from __future__ import annotations
from typing import Any
from fastapi import HTTPException
from app.services.module_generation_step_common_service import build_file_url_or_data_uri
def build_v2_video_generation_references(material_payload: dict[str, Any]) -> list[dict[str, Any]]:
"""V2 最终视频只能携带可选素材图片,绝不携带素材视频/音频。"""
image_url = str(material_payload.get("material_image_url") or "").strip()
if not image_url:
return []
return [
{
"type": "image",
"url": build_file_url_or_data_uri(image_url),
"name": "素材参考图片",
"upload_resource_id": material_payload.get("material_image_resource_id"),
"role": "reference_image",
}
]
def assert_v2_video_generation_references(references: list[dict[str, Any]]) -> None:
image_count = 0
for item in references:
media_type = str(item.get("type") or item.get("media_type") or "").lower()
if media_type == "image":
image_count += 1
continue
raise HTTPException(status_code=400, detail="V2 视频生成只允许携带一张素材图片,禁止视频或音频附件")
if image_count > 1:
raise HTTPException(status_code=400, detail="V2 视频生成最多携带一张素材图片")
@@ -19,6 +19,7 @@ from app.enums.recent_generation import (
)
from app.models.chat_generation_task import ChatGenerationTask
from app.models.generation_record import GenerationRecord
from app.models.module_generation_project import ModuleGenerationProject
from app.models.module_generation_step import ModuleGenerationStep
from app.models.shot_replicate_segment import ShotReplicateSegment
from app.schemas.recent_generation import RecentGenerationGroupOut, RecentGenerationItemOut
@@ -30,6 +31,7 @@ MAX_RECENT_GENERATION_LIMIT = 100
class _StepLinkInfo(TypedDict):
module_project_id: str | None
module_project_flow_version: str | None
module_step_id: str | None
module: str | None
@@ -115,6 +117,7 @@ def _build_item(
module=module,
shot_task_set_id=shot_info["shot_task_set_id"] if shot_info else None,
shot_segment_id=shot_info["shot_segment_id"] if shot_info else None,
module_project_flow_version=step_info["module_project_flow_version"] if step_info else None,
module_project_id=step_info["module_project_id"] if step_info else None,
module_step_id=step_info["module_step_id"] if step_info else None,
generation_id=generation_id,
@@ -248,11 +251,14 @@ async def _load_step_link_map(
ModuleGenerationStep.id.label("module_step_id"),
ModuleGenerationStep.project_id.label("module_project_id"),
ModuleGenerationStep.module.label("module"),
ModuleGenerationProject.flow_version.label("module_project_flow_version"),
ModuleGenerationStep.is_current.label("is_current"),
ModuleGenerationStep.updated_at.label("updated_at"),
)
.join(ModuleGenerationProject, ModuleGenerationProject.id == ModuleGenerationStep.project_id)
.where(
ModuleGenerationStep.deleted_at.is_(None),
ModuleGenerationProject.deleted_at.is_(None),
ModuleGenerationStep.chat_task_id.in_(chat_task_ids),
ModuleGenerationStep.module.in_(
[
@@ -276,6 +282,7 @@ async def _load_step_link_map(
continue
link_map[chat_task_id] = {
"module_project_id": row["module_project_id"],
"module_project_flow_version": str(row["module_project_flow_version"] or "v1"),
"module_step_id": row["module_step_id"],
"module": row["module"],
}
@@ -125,6 +125,7 @@ FLOW_CONFIG = ModuleGenerationFlowConfig(
cancel_chat_task_error_message="拆镜复刻步骤被重新生成或删除,旧生成任务已取消",
material_video_url_editable=False,
step_io_schema_version=STEP_IO_SCHEMA_VERSION,
expected_flow_version="v1",
)
@@ -434,12 +435,19 @@ async def project_to_detail_out(db: AsyncSession, project: ModuleGenerationProje
user_result = await db.execute(select(User.username).where(User.id == project.user_id).limit(1))
user_name = user_result.scalar_one_or_none()
flow_version = str(getattr(project, "flow_version", None) or "v1")
video_prompt_config = dict(video_prompt_input.get("video_config") or {})
video_prompt_engine_snapshot = dict(video_prompt_config.get("engine_snapshot") or {})
return ShotReplicateTaskDetailOut(
id=project.id,
project_id=project.id,
user_id=project.user_id,
user_name=user_name,
module=project.module,
flow_version=flow_version,
step_count=3 if flow_version == "v2" else 5,
step_io_schema_version=("hot_opening_step_io_v2" if project.module == "hot_opening_replicate" else "shot_replicate_step_io_v2") if flow_version == "v2" else STEP_IO_SCHEMA_VERSION,
title=project.title,
status=project.status,
current_step_code=project.current_step_code,
@@ -454,6 +462,8 @@ async def project_to_detail_out(db: AsyncSession, project: ModuleGenerationProje
source_project_name=material_input.get("source_project_name"),
target_project_name=material_input.get("target_project_name"),
core_content_point=material_input.get("core_content_point"),
project_description=material_input.get("project_description"),
video_config=None if flow_version == "v2" else material_input.get("video_config"),
),
image_generation=ShotReplicateImageGenerationOut(
prompt_step_id=image_prompt_step.id if image_prompt_step else None,
@@ -477,9 +487,9 @@ async def project_to_detail_out(db: AsyncSession, project: ModuleGenerationProje
schema_config_source=schema_config_source,
schema_config_version=schema_config_version,
schema_config_is_fallback=schema_config_is_fallback,
engine_id=video_snapshot.get("id") or video_generate_input.get("engine_id"),
engine_name=video_snapshot.get("name") or video_generate_input.get("engine_name"),
params=video_generate_input.get("params") or video_generate_input,
engine_id=video_snapshot.get("id") or video_generate_input.get("engine_id") or video_prompt_config.get("engine_id"),
engine_name=video_snapshot.get("name") or video_generate_input.get("engine_name") or video_prompt_engine_snapshot.get("name"),
params=video_generate_input.get("params") or video_prompt_config or video_generate_input,
chat_task_id=video_generate_step.chat_task_id if video_generate_step else None,
status=video_chat.status if video_chat else (video_generate_step.status if video_generate_step else None),
result_video_url=build_resource_signed_url(video_url) if video_url else None,
@@ -606,6 +616,8 @@ async def list_shot_replicate_projects(
id=project.id,
project_id=project.id,
module=project.module,
flow_version=str(getattr(project, "flow_version", None) or "v1"),
step_count=3 if str(getattr(project, "flow_version", None) or "v1") == "v2" else 5,
title=project.title,
status=project.status,
current_step_code=project.current_step_code,
@@ -1497,6 +1509,9 @@ async def generate_video_from_prompt(
async def handle_chat_generation_task_completed(db: AsyncSession, task: ChatGenerationTask) -> None:
if not task or task.generation_mode != GENERATION_MODE:
return
from app.services.module_generation_v2.flow_service import handle_chat_generation_task_finished_v2
if await handle_chat_generation_task_finished_v2(db, task=task):
return
meta_result = await db.execute(
select(ModuleGenerationStep.id, ModuleGenerationStep.project_id).where(
ModuleGenerationStep.chat_task_id == task.id,
@@ -1579,6 +1594,9 @@ async def handle_chat_generation_task_completed(db: AsyncSession, task: ChatGene
async def handle_chat_generation_task_failed(db: AsyncSession, task: ChatGenerationTask) -> None:
if not task or task.generation_mode != GENERATION_MODE:
return
from app.services.module_generation_v2.flow_service import handle_chat_generation_task_finished_v2
if await handle_chat_generation_task_finished_v2(db, task=task):
return
meta_result = await db.execute(
select(ModuleGenerationStep.id, ModuleGenerationStep.project_id).where(
ModuleGenerationStep.chat_task_id == task.id,
@@ -32,6 +32,7 @@ from app.schemas.shot_replicate import (
ShotSegmentListOut,
ShotSegmentSplitRetryOut,
ShotReanalyzeOut,
ShotReplicateDeleteOut,
ShotSegmentOut,
ShotSplitByAIOut,
ShotSplitByAIRequest,
@@ -116,6 +117,7 @@ def _segment_to_out(segment: ShotReplicateSegment, project: ModuleGenerationProj
data.module_project_title = project.title
data.module_project_status = project.status
data.module_project_current_step_code = project.current_step_code
data.module_project_flow_version = str(getattr(project, "flow_version", None) or "v1")
return data
@@ -126,6 +128,7 @@ def _segment_to_detail_out(segment: ShotReplicateSegment, project: ModuleGenerat
data.module_project_title = project.title
data.module_project_status = project.status
data.module_project_current_step_code = project.current_step_code
data.module_project_flow_version = str(getattr(project, "flow_version", None) or "v1")
return data
@@ -716,6 +719,44 @@ async def segment_detail(db: AsyncSession, *, current_user: User, segment_id: st
project = project_result.scalar_one_or_none()
return _segment_to_detail_out(segment, project)
async def _delete_linked_replication_project(
db: AsyncSession,
*,
current_user: User,
project_id: str,
) -> ShotReplicateDeleteOut:
"""按项目保存的流程版本分发 V1/V2 删除,供片段和任务集联动删除复用。"""
result = await db.execute(
select(ModuleGenerationProject)
.where(
ModuleGenerationProject.id == project_id,
ModuleGenerationProject.deleted_at.is_(None),
)
.limit(1)
)
project = result.scalar_one_or_none()
if project and str(getattr(project, "flow_version", None) or "v1") == "v2":
from app.services.module_generation_v2.config import SHOT_REPLICATE_V2
from app.services.module_generation_v2.flow_service import delete_project_v2
payload = await delete_project_v2(
db,
config=SHOT_REPLICATE_V2,
current_user=current_user,
project_id=project_id,
)
return ShotReplicateDeleteOut(**payload)
from app.services.shot_replicate_flow_service import delete_shot_replicate_project
return await delete_shot_replicate_project(
db,
current_user=current_user,
project_id=project_id,
refund_unfinished=False,
)
async def delete_segment(
db: AsyncSession,
*,
@@ -775,13 +816,10 @@ async def delete_segment(
deleted_module_project_id: str | None = None
if module_project_id:
from app.services.shot_replicate_flow_service import delete_shot_replicate_project
project_delete_out = await delete_shot_replicate_project(
project_delete_out = await _delete_linked_replication_project(
db,
current_user=current_user,
project_id=module_project_id,
refund_unfinished=False,
)
deleted_module_project_id = project_delete_out.project_id
released_size_bytes += int(project_delete_out.released_size_bytes or 0)
@@ -894,15 +932,12 @@ async def delete_task_set(
upload_resource_released += int(segment_upload_release.get("released") or 0)
pending_delete_resource_ids.extend(segment_upload_release.get("released_resource_ids") or [])
from app.services.shot_replicate_flow_service import delete_shot_replicate_project
deleted_module_project_count = 0
for module_project_id in dict.fromkeys(module_project_ids):
project_delete_out = await delete_shot_replicate_project(
project_delete_out = await _delete_linked_replication_project(
db,
current_user=current_user,
project_id=module_project_id,
refund_unfinished=False,
)
deleted_module_project_count += 1
released_size_bytes += int(project_delete_out.released_size_bytes or 0)
+22
View File
@@ -23,6 +23,7 @@ CELERY_TASK_IMPORTS = (
"app.tasks.shot_replicate_tasks",
"app.tasks.shot_replicate_flow_tasks",
"app.tasks.module_async_recovery_tasks",
"app.tasks.module_generation_v2_tasks",
"app.tasks.user_oauth_tasks",
"app.tasks.cleanup",
"app.tasks.private_portrait_asset_tasks",
@@ -53,6 +54,22 @@ def _beat_schedule() -> dict:
"priority": settings.DOWNLOAD_TASK_PRIORITY_RECOVER,
},
}
schedule["generation-create-recovery"] = {
"task": CeleryTaskName.RECOVER_CREATE.value,
"schedule": max(1, int(settings.GENERATION_CREATE_RECOVERY_INTERVAL_SECONDS or 60)),
"options": {
"queue": RECOVERY_QUEUE,
"priority": settings.DOWNLOAD_TASK_PRIORITY_RECOVER,
},
}
schedule["module-async-recovery"] = {
"task": CeleryTaskName.MODULE_ASYNC_RECOVERY.value,
"schedule": max(1, int(settings.MODULE_ASYNC_RECOVERY_INTERVAL_SECONDS or 60)),
"options": {
"queue": RECOVERY_QUEUE,
"priority": settings.DOWNLOAD_TASK_PRIORITY_RECOVER,
},
}
schedule["video-upscale-recovery-every-minute"] = {
"task": CeleryTaskName.VIDEO_UPSCALE_RECOVER.value,
"schedule": 60,
@@ -109,6 +126,7 @@ if broker_url:
"shot_replicate.split_one_segment": {"ignore_result": True},
"shot_replicate.start_image_prompt_optimize": {"ignore_result": True},
"shot_replicate.start_video_prompt_optimize": {"ignore_result": True},
"module_generation_v2.start_video_prompt_optimize": {"ignore_result": True},
CeleryTaskName.VIDEO_UPSCALE_EXECUTE_LOCAL.value: {
"ignore_result": True,
"soft_time_limit": max(60, int(settings.VIDEO_UPSCALE_LOCAL_TIMEOUT_SECONDS or 3600)) + 60,
@@ -119,6 +137,8 @@ if broker_url:
CeleryTaskName.VIDEO_UPSCALE_DOWNLOAD_REMOTE_RESULT.value: {"ignore_result": True},
CeleryTaskName.VIDEO_UPSCALE_FINALIZE.value: {"ignore_result": True},
CeleryTaskName.VIDEO_UPSCALE_RECOVER.value: {"ignore_result": True},
CeleryTaskName.RECOVER_CREATE.value: {"ignore_result": True},
CeleryTaskName.MODULE_ASYNC_RECOVERY.value: {"ignore_result": True},
},
worker_prefetch_multiplier=1,
broker_transport_options={
@@ -159,11 +179,13 @@ if broker_url:
"shot_replicate.split_one_segment": {"queue": CeleryQueue.GEN_RESULT_DOWNLOAD.value},
"shot_replicate.start_image_prompt_optimize": {"queue": CeleryQueue.GEN_CHATAPI_CREATE.value},
"shot_replicate.start_video_prompt_optimize": {"queue": CeleryQueue.GEN_CHATAPI_CREATE.value},
"module_generation_v2.start_video_prompt_optimize": {"queue": CeleryQueue.GEN_CHATAPI_CREATE.value},
# 恢复扫描统一走独立队列,避免占用下载/轮询/创建业务 worker。
CeleryTaskName.STARTUP_RECOVERY.value: {"queue": RECOVERY_QUEUE},
CeleryTaskName.SHOT_SPLIT_RECOVERY.value: {"queue": RECOVERY_QUEUE},
CeleryTaskName.RECOVER_DOWNLOAD.value: {"queue": RECOVERY_QUEUE},
CeleryTaskName.RECOVER_GENERATION.value: {"queue": RECOVERY_QUEUE},
CeleryTaskName.RECOVER_CREATE.value: {"queue": RECOVERY_QUEUE},
CeleryTaskName.MODULE_ASYNC_RECOVERY.value: {"queue": RECOVERY_QUEUE},
"user_oauth.update_oauth_accounts": {"queue": CeleryQueue.DEFAULT.value},
"app.tasks.cleanup.*": {"queue": CeleryQueue.DEFAULT.value},
@@ -40,6 +40,7 @@ async def _recover_generation_records_once(*, include_create: bool, include_poll
args=[ref.owner_id],
kwargs={"owner_type": ref.owner_type, "generation_attempt_no": ref.generation_attempt_no},
queue=CeleryQueue.GEN_CHATAPI_CREATE.value,
task_id=f"generation-create:{ref.owner_type}:{ref.owner_id}:attempt:{ref.generation_attempt_no}",
)
counts["create"] += 1
except Exception as exc:
@@ -116,6 +117,17 @@ async def _run_generation_once() -> Dict[str, Any]:
return {"chat_generation_task": chat_result, "generation_record": record_result}
async def _run_create_once() -> Dict[str, Any]:
from app.services.generation.recovery_service import recover_stale_create_tasks_once
async with async_session() as db:
chat_result = await recover_stale_create_tasks_once(db)
record_result = await _recover_generation_records_once(
include_create=True, include_poll=False, include_download=False
)
return {"chat_generation_task": chat_result, "generation_record": record_result}
async def _run_due_poll_dispatch_once() -> Dict[str, Any]:
from datetime import datetime, timedelta, timezone
@@ -396,6 +408,23 @@ if celery_app:
)
@celery_app.task(
name="generation.recover_create_tasks_once",
bind=True,
soft_time_limit=settings.CELERY_RECOVERY_SOFT_TIME_LIMIT_SECONDS,
time_limit=settings.CELERY_RECOVERY_TIME_LIMIT_SECONDS,
)
def recover_create_tasks_once(self) -> Dict[str, Any]:
return run_async(
_run_with_execution_lock(
lock_key=f"{settings.GENERATION_RECOVERY_LOCK_KEY}:create",
log_context="generation_create_recovery",
runner=_run_create_once,
ttl_seconds=max(55, int(settings.GENERATION_CREATE_RECOVERY_INTERVAL_SECONDS or 60) - 5),
)
)
@celery_app.task(
name="generation.dispatch_due_poll_tasks",
bind=True,
@@ -424,4 +453,5 @@ else:
startup_recovery_once = _DisabledTask()
recover_download_tasks_once = _DisabledTask()
recover_generation_tasks_once = _DisabledTask()
recover_create_tasks_once = _DisabledTask()
dispatch_due_poll_tasks = _DisabledTask()
@@ -0,0 +1,81 @@
from __future__ import annotations
from typing import Any
from sqlalchemy import select
from app.models.base import async_session
from app.models.module_generation_project import ModuleGenerationProject
from app.services.module_async_recovery_service import (
OBJECT_MODULE_STEP,
TASK_MODULE_V2_VIDEO_PROMPT,
acquire_object_lock,
cleanup_active_if_terminal,
mark_active_started,
register_module_step_task,
release_object_lock,
)
from app.services.module_generation_v2.config import VIDEO_PROMPT_OPTIMIZE
from app.services.module_generation_v2.flow_service import run_video_prompt_optimize_v2
from app.tasks.async_runner import run_async
from app.tasks.celery_app import celery_app
async def _run_video_prompt(project_id: str, step_id: str) -> None:
lock_token = await acquire_object_lock(object_type=OBJECT_MODULE_STEP, object_id=step_id)
if not lock_token:
return None
try:
async with async_session() as db:
result = await db.execute(
select(ModuleGenerationProject.module)
.where(
ModuleGenerationProject.id == project_id,
ModuleGenerationProject.deleted_at.is_(None),
)
.limit(1)
)
module = result.scalar_one_or_none()
if not module:
return None
await register_module_step_task(
module=str(module),
project_id=project_id,
step_id=step_id,
step_code=VIDEO_PROMPT_OPTIMIZE,
task_name=TASK_MODULE_V2_VIDEO_PROMPT,
)
await mark_active_started(object_type=OBJECT_MODULE_STEP, object_id=step_id)
await run_video_prompt_optimize_v2(db, project_id=project_id, step_id=step_id)
await cleanup_active_if_terminal(db, object_type=OBJECT_MODULE_STEP, object_id=step_id)
finally:
await release_object_lock(
object_type=OBJECT_MODULE_STEP,
object_id=step_id,
token=lock_token,
)
if celery_app:
@celery_app.task(
name="module_generation_v2.start_video_prompt_optimize",
bind=True,
max_retries=3,
default_retry_delay=30,
ignore_result=True,
)
def start_video_prompt_optimize_v2(self, project_id: str, step_id: str):
try:
run_async(_run_video_prompt(project_id, step_id))
return None
except Exception as exc:
raise self.retry(exc=exc) from exc
else:
class _DisabledTask:
def delay(self, *args: Any, **kwargs: Any):
raise RuntimeError("Celery is disabled")
def apply_async(self, *args: Any, **kwargs: Any):
raise RuntimeError("Celery is disabled")
start_video_prompt_optimize_v2 = _DisabledTask()