爆款/拆镜生成简化3个步骤 | 项目生成可携带附件控制
This commit is contained in:
+87
@@ -0,0 +1,87 @@
|
||||
"""add module flow version and generation reference option
|
||||
|
||||
Revision ID: d8ebe79ab575
|
||||
Revises: e6eac828ff61
|
||||
Create Date: 2026-07-21 09:28:04.318458
|
||||
"""
|
||||
from typing import Sequence, Union
|
||||
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
|
||||
|
||||
revision: str = "d8ebe79ab575"
|
||||
down_revision: Union[str, None] = "e6eac828ff61"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def _column_names(table_name: str) -> set[str]:
|
||||
inspector = sa.inspect(op.get_bind())
|
||||
return {str(column["name"]) for column in inspector.get_columns(table_name)}
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
generation_record_columns = _column_names("generation_records")
|
||||
if "include_media_references" not in generation_record_columns:
|
||||
op.add_column(
|
||||
"generation_records",
|
||||
sa.Column(
|
||||
"include_media_references",
|
||||
sa.Boolean(),
|
||||
server_default=sa.text("false"),
|
||||
nullable=False,
|
||||
),
|
||||
)
|
||||
else:
|
||||
op.execute(
|
||||
sa.text(
|
||||
"UPDATE generation_records "
|
||||
"SET include_media_references = false "
|
||||
"WHERE include_media_references IS NULL"
|
||||
)
|
||||
)
|
||||
op.alter_column(
|
||||
"generation_records",
|
||||
"include_media_references",
|
||||
existing_type=sa.Boolean(),
|
||||
nullable=False,
|
||||
server_default=sa.text("false"),
|
||||
)
|
||||
|
||||
project_columns = _column_names("module_generation_projects")
|
||||
if "flow_version" not in project_columns:
|
||||
op.add_column(
|
||||
"module_generation_projects",
|
||||
sa.Column(
|
||||
"flow_version",
|
||||
sa.String(length=16),
|
||||
server_default=sa.text("'v1'"),
|
||||
nullable=False,
|
||||
),
|
||||
)
|
||||
else:
|
||||
op.execute(
|
||||
sa.text(
|
||||
"UPDATE module_generation_projects "
|
||||
"SET flow_version = 'v1' "
|
||||
"WHERE flow_version IS NULL OR btrim(flow_version) = ''"
|
||||
)
|
||||
)
|
||||
op.alter_column(
|
||||
"module_generation_projects",
|
||||
"flow_version",
|
||||
existing_type=sa.String(length=16),
|
||||
nullable=False,
|
||||
server_default=sa.text("'v1'"),
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
project_columns = _column_names("module_generation_projects")
|
||||
if "flow_version" in project_columns:
|
||||
op.drop_column("module_generation_projects", "flow_version")
|
||||
|
||||
generation_record_columns = _column_names("generation_records")
|
||||
if "include_media_references" in generation_record_columns:
|
||||
op.drop_column("generation_records", "include_media_references")
|
||||
@@ -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="请选择文件")
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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",
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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 存储。
|
||||
建议结构:
|
||||
|
||||
@@ -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_id;project/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)
|
||||
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user