From 0be8091181f75e133bd71cff9716c41e928cb552 Mon Sep 17 00:00:00 2001 From: GinHa <15201596918@163.com> Date: Thu, 18 Jun 2026 13:03:11 +0800 Subject: [PATCH] =?UTF-8?q?=E6=8B=86=E9=95=9C=E5=A4=8D=E5=88=BB=E3=80=81?= =?UTF-8?q?=E7=88=86=E6=AC=BE=E5=BC=80=E5=A4=B4=E5=A4=8D=E5=88=BB=E6=8A=BD?= =?UTF-8?q?=E8=B1=A1=E5=85=AC=E5=85=B1=E9=80=BB=E8=BE=91,=E6=8B=86?= =?UTF-8?q?=E9=95=9C=E5=A4=8D=E5=88=BB=E5=85=B3=E8=81=94=E9=A1=B9=E7=9B=AE?= =?UTF-8?q?=E6=AD=A5=E9=AA=A4=E7=8A=B6=E6=80=81BUG=E4=BF=AE=E5=A4=8D?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../app/enums/module_generation_flow.py | 26 + .../services/hot_opening_replicate_service.py | 615 ++++------------- .../module_generation_flow_base_service.py | 258 ++++++++ .../module_generation_step_common_service.py | 168 +++++ .../module_generation_step_update_service.py | 363 ++++++++++ .../services/shot_replicate_flow_service.py | 625 ++++-------------- 6 files changed, 1093 insertions(+), 962 deletions(-) create mode 100644 video-gen-api/app/enums/module_generation_flow.py create mode 100644 video-gen-api/app/services/module_generation_flow_base_service.py create mode 100644 video-gen-api/app/services/module_generation_step_common_service.py create mode 100644 video-gen-api/app/services/module_generation_step_update_service.py diff --git a/video-gen-api/app/enums/module_generation_flow.py b/video-gen-api/app/enums/module_generation_flow.py new file mode 100644 index 00000000..cc66cac2 --- /dev/null +++ b/video-gen-api/app/enums/module_generation_flow.py @@ -0,0 +1,26 @@ +from __future__ import annotations + +from dataclasses import dataclass +from typing import Mapping + + +@dataclass(frozen=True, slots=True) +class ModuleGenerationFlowConfig: + """模块生成流程公共配置。 + + 这个配置只描述两个复刻模块的差异点,不承载具体业务逻辑。 + 具体日志、退款、生成、回调等仍保留在各自模块 service 中。 + """ + + module: str + step_index_map: Mapping[str, int] + material_step_code: str + image_prompt_step_code: str + image_generate_step_code: str + video_prompt_step_code: str + video_generate_step_code: str + project_not_found_message: str + step_not_found_message: str + cancel_chat_task_error_message: str + material_video_url_editable: bool = True + step_io_schema_version: str = "module_generation_step_io_v1" diff --git a/video-gen-api/app/services/hot_opening_replicate_service.py b/video-gen-api/app/services/hot_opening_replicate_service.py index 34051e1a..248dbee8 100644 --- a/video-gen-api/app/services/hot_opening_replicate_service.py +++ b/video-gen-api/app/services/hot_opening_replicate_service.py @@ -49,6 +49,39 @@ from app.services.hot_opening_video_prompt_service import build_final_video_prom from app.services.module_generation_log_service import log_module_error, log_module_event_file, log_module_prompt_event from app.services.llm import optimize_prompt from app.services.resource_accounting_service import soft_delete_chat_task_resources +from app.services.module_generation_flow_base_service import ( + chat_tasks_by_id as _base_chat_tasks_by_id, + create_module_step as _base_create_step, + get_current_step_by_code as _base_get_current_step_by_code, + get_current_steps as _base_get_current_steps, + get_project_for_user as _base_get_project_for_user, + get_step_for_user as _base_get_step_for_user, + next_version as _base_next_version, + soft_delete_steps_from_index as _base_soft_delete_steps_from_index, +) +from app.enums.module_generation_flow import ModuleGenerationFlowConfig +from app.services.module_generation_step_common_service import ( + build_file_url_or_data_uri as _common_build_file_url_or_data_uri, + build_step_input as _common_build_step_input, + build_step_output as _common_build_step_output, + force_set_json as _common_force_set_json, + is_wrapped_step_io as _common_is_wrapped_step_io, + json_dumps as _common_json, + merge_dict as _common_merge_dict, + parse_json as _common_parse_json, + snapshot_from_chat as _common_snapshot_from_chat, + step_payload as _common_step_payload, + step_result as _common_step_result, + step_usage as _common_step_usage, + unwrap_step_output as _common_unwrap_step_output, + utc_now as _common_now, +) +from app.services.module_generation_step_update_service import ( + update_module_image_prompt, + update_module_material_input, + update_module_step, + update_module_video_prompt_schema, +) from app.services.resource_signed_url_service import build_resource_signed_url from app.utils.id_gen import generate_id @@ -65,28 +98,27 @@ STEP_INDEX_MAP = { STEP_IO_SCHEMA_VERSION = "hot_opening_step_io_v1" - -def _now() -> datetime: - return datetime.now(timezone.utc) +FLOW_CONFIG = ModuleGenerationFlowConfig( + module=MODULE, + step_index_map=STEP_INDEX_MAP, + material_step_code=HotOpeningStepCodeEnum.MATERIAL_INPUT.value, + image_prompt_step_code=HotOpeningStepCodeEnum.IMAGE_PROMPT_OPTIMIZE.value, + image_generate_step_code=HotOpeningStepCodeEnum.IMAGE_GENERATE.value, + video_prompt_step_code=HotOpeningStepCodeEnum.VIDEO_PROMPT_OPTIMIZE.value, + video_generate_step_code=HotOpeningStepCodeEnum.VIDEO_GENERATE.value, + project_not_found_message="爆款开头复刻项目不存在", + step_not_found_message="子任务不存在", + cancel_chat_task_error_message="爆款开头复刻步骤被重新生成或删除,旧生成任务已取消", + material_video_url_editable=True, + step_io_schema_version=STEP_IO_SCHEMA_VERSION, +) -def _json(data: Any) -> str | None: - if data is None: - return None - return json.dumps(data, ensure_ascii=False, default=str) - - -def _parse_json(value: Any, fallback: Any = None) -> Any: - if value is None or value == "": - return fallback - if isinstance(value, (dict, list)): - return value - if isinstance(value, str): - try: - return json.loads(value) - except Exception: - return fallback - return fallback +_now = _common_now +_json = _common_json +_parse_json = _common_parse_json +_merge_dict = _common_merge_dict +_force_set_json = _common_force_set_json def _step_input( @@ -97,16 +129,14 @@ def _step_input( parent_step_id: str | None = None, context: dict[str, Any] | None = None, ) -> dict[str, Any]: - return { - "schema_version": STEP_IO_SCHEMA_VERSION, - "step_code": step_code, - "source": { - "source_step_id": source_step_id, - "parent_step_id": parent_step_id, - }, - "payload": payload or {}, - "context": context or {}, - } + return _common_build_step_input( + step_code=step_code, + payload=payload, + source_step_id=source_step_id, + parent_step_id=parent_step_id, + context=context, + schema_version=STEP_IO_SCHEMA_VERSION, + ) def _step_output( @@ -118,85 +148,35 @@ def _step_output( usage: dict[str, Any] | None = None, error: dict[str, Any] | None = None, ) -> dict[str, Any]: - return { - "schema_version": STEP_IO_SCHEMA_VERSION, - "step_code": step_code, - "status": status, - "payload": payload or {}, - "result": result or {}, - "usage": usage or {}, - "error": error or {}, - } + return _common_build_step_output( + step_code=step_code, + status=status, + payload=payload, + result=result, + usage=usage, + error=error, + schema_version=STEP_IO_SCHEMA_VERSION, + ) def _is_wrapped_step_io(value: Any) -> bool: - return isinstance(value, dict) and value.get("schema_version") == STEP_IO_SCHEMA_VERSION + return _common_is_wrapped_step_io(value, schema_version=STEP_IO_SCHEMA_VERSION) def _step_payload(value: Any) -> dict[str, Any]: - data = _parse_json(value, {}) or {} - if _is_wrapped_step_io(data): - payload = data.get("payload") - return payload if isinstance(payload, dict) else {} - return data if isinstance(data, dict) else {} + return _common_step_payload(value, schema_version=STEP_IO_SCHEMA_VERSION) def _step_result(value: Any) -> dict[str, Any]: - data = _parse_json(value, {}) or {} - if _is_wrapped_step_io(data): - result = data.get("result") - if isinstance(result, dict) and result: - return result - payload = data.get("payload") - return payload if isinstance(payload, dict) else {} - return data if isinstance(data, dict) else {} + return _common_step_result(value, schema_version=STEP_IO_SCHEMA_VERSION) def _step_usage(value: Any) -> dict[str, Any]: - data = _parse_json(value, {}) or {} - if _is_wrapped_step_io(data): - usage = data.get("usage") - return usage if isinstance(usage, dict) else {} - usage = data.get("token_usage") if isinstance(data, dict) else {} - return usage if isinstance(usage, dict) else {} + return _common_step_usage(value, schema_version=STEP_IO_SCHEMA_VERSION) def _unwrap_step_output(value: Any) -> dict[str, Any]: - data = _parse_json(value, {}) or {} - if not _is_wrapped_step_io(data): - return data if isinstance(data, dict) else {} - merged: dict[str, Any] = {} - payload = data.get("payload") - result = data.get("result") - usage = data.get("usage") - if isinstance(payload, dict): - merged.update(payload) - if isinstance(result, dict): - merged.update(result) - if isinstance(usage, dict) and usage: - merged["token_usage"] = usage - return merged - - -def _merge_dict(old: dict[str, Any] | None, new: dict[str, Any] | None) -> dict[str, Any]: - merged = dict(old or {}) - for key, value in (new or {}).items(): - if value is not None: - merged[key] = value - return merged - - -def _force_set_json(model_obj: Any, field_name: str, value: Any) -> None: - """强制持久化 JSON / JSONB 字段。 - - SQLAlchemy 对 dict/list 的嵌套原地修改不会稳定触发 dirty 判定。 - 所有编辑类接口在写入 input_json / output_json 时统一走这里: - 1. deepcopy 断开旧引用; - 2. 整体重新赋值; - 3. flag_modified 显式标记字段已变更。 - """ - setattr(model_obj, field_name, deepcopy(value)) - flag_modified(model_obj, field_name) + return _common_unwrap_step_output(value, schema_version=STEP_IO_SCHEMA_VERSION) async def log_module_event( @@ -251,22 +231,14 @@ async def _get_project_for_user( for_update: bool = False, populate_existing: bool = False, ) -> ModuleGenerationProject: - query = select(ModuleGenerationProject).where( - ModuleGenerationProject.id == project_id, - ModuleGenerationProject.module == MODULE, - ModuleGenerationProject.deleted_at.is_(None), + return await _base_get_project_for_user( + db, + project_id=project_id, + user=user, + config=FLOW_CONFIG, + for_update=for_update, + populate_existing=populate_existing, ) - if not user.is_admin: - query = query.where(ModuleGenerationProject.user_id == user.id) - if populate_existing: - query = query.execution_options(populate_existing=True) - if for_update: - query = query.with_for_update() - result = await db.execute(query.limit(1)) - project = result.scalar_one_or_none() - if not project: - raise HTTPException(status_code=404, detail="爆款开头复刻项目不存在") - return project async def _get_step_for_user( @@ -277,64 +249,26 @@ async def _get_step_for_user( user: User, for_update: bool = False, ) -> ModuleGenerationStep: - await _get_project_for_user(db, project_id=project_id, user=user, for_update=for_update) - query = select(ModuleGenerationStep).where( - ModuleGenerationStep.id == step_id, - ModuleGenerationStep.project_id == project_id, - ModuleGenerationStep.module == MODULE, - ModuleGenerationStep.deleted_at.is_(None), - ModuleGenerationStep.is_current == True, + return await _base_get_step_for_user( + db, + project_id=project_id, + step_id=step_id, + user=user, + config=FLOW_CONFIG, + for_update=for_update, ) - if not user.is_admin: - query = query.where(ModuleGenerationStep.user_id == user.id) - if for_update: - query = query.with_for_update() - result = await db.execute(query.limit(1)) - step = result.scalar_one_or_none() - if not step: - raise HTTPException(status_code=404, detail="子任务不存在") - return step async def _get_current_steps(db: AsyncSession, project_id: str) -> list[ModuleGenerationStep]: - result = await db.execute( - select(ModuleGenerationStep) - .where( - ModuleGenerationStep.project_id == project_id, - ModuleGenerationStep.module == MODULE, - ModuleGenerationStep.deleted_at.is_(None), - ModuleGenerationStep.is_current == True, - ) - .order_by(ModuleGenerationStep.step_index.asc(), ModuleGenerationStep.created_at.asc()) - ) - return list(result.scalars().all()) + return await _base_get_current_steps(db, project_id=project_id, config=FLOW_CONFIG) async def _get_current_step_by_code(db: AsyncSession, project_id: str, step_code: str) -> ModuleGenerationStep | None: - result = await db.execute( - select(ModuleGenerationStep) - .where( - ModuleGenerationStep.project_id == project_id, - ModuleGenerationStep.module == MODULE, - ModuleGenerationStep.step_code == step_code, - ModuleGenerationStep.is_current == True, - ModuleGenerationStep.deleted_at.is_(None), - ) - .order_by(ModuleGenerationStep.version.desc(), ModuleGenerationStep.created_at.desc()) - .limit(1) - ) - return result.scalar_one_or_none() + return await _base_get_current_step_by_code(db, project_id=project_id, step_code=step_code, config=FLOW_CONFIG) async def _next_version(db: AsyncSession, project_id: str, step_code: str) -> int: - result = await db.execute( - select(func.max(ModuleGenerationStep.version)).where( - ModuleGenerationStep.project_id == project_id, - ModuleGenerationStep.module == MODULE, - ModuleGenerationStep.step_code == step_code, - ) - ) - return int(result.scalar_one_or_none() or 0) + 1 + return await _base_next_version(db, project_id=project_id, step_code=step_code, config=FLOW_CONFIG) async def _create_step( @@ -349,39 +283,19 @@ async def _create_step( input_data: dict[str, Any] | None = None, output_data: dict[str, Any] | None = None, ) -> ModuleGenerationStep: - version = await _next_version(db, project.id, step_code) - step = ModuleGenerationStep( - id=generate_id(), - project_id=project.id, - user_id=project.user_id, - module=project.module, - step_index=STEP_INDEX_MAP[step_code], + return await _base_create_step( + db, + project=project, step_code=step_code, + config=FLOW_CONFIG, + log_module_event=log_module_event, status=status, - version=version, - is_current=True, parent_step_id=parent_step_id, source_step_id=source_step_id, chat_task_id=chat_task_id, - input_json=_step_input( - step_code=step_code, - payload=input_data, - source_step_id=source_step_id, - parent_step_id=parent_step_id, - ) if input_data is not None else None, - output_json=_step_output( - step_code=step_code, - status=status, - result=output_data, - ) if output_data is not None else None, - started_at=_now() if status == ModuleStepStatusEnum.PROCESSING.value else None, - completed_at=_now() if status == ModuleStepStatusEnum.COMPLETED.value else None, + input_data=input_data, + output_data=output_data, ) - db.add(step) - project.current_step_code = step_code - await db.flush() - await log_module_event(db, project=project, step=step, event_type=ModuleEventTypeEnum.STEP_CREATED.value, detail={"step_code": step_code, "version": version}) - return step async def _soft_delete_steps_from_index( @@ -391,49 +305,14 @@ async def _soft_delete_steps_from_index( start_index: int, deleted_at: datetime | None = None, ) -> None: - deleted_at = deleted_at or _now() - result = await db.execute( - select(ModuleGenerationStep) - .where( - ModuleGenerationStep.project_id == project.id, - ModuleGenerationStep.module == MODULE, - ModuleGenerationStep.is_current == True, - ModuleGenerationStep.deleted_at.is_(None), - ModuleGenerationStep.step_index >= start_index, - ) - .with_for_update() + await _base_soft_delete_steps_from_index( + db, + project=project, + start_index=start_index, + config=FLOW_CONFIG, + log_module_event=log_module_event, + deleted_at=deleted_at, ) - steps = list(result.scalars().all()) - for step in steps: - step.is_current = False - step.deleted_at = deleted_at - if step.chat_task_id: - chat_result = await db.execute( - select(ChatGenerationTask) - .where(ChatGenerationTask.id == step.chat_task_id, ChatGenerationTask.deleted_at.is_(None)) - .with_for_update() - .limit(1) - ) - chat_task = chat_result.scalar_one_or_none() - if chat_task: - if chat_task.status == "completed": - await soft_delete_chat_task_resources(db, chat_task.id, deleted_at=deleted_at) - elif chat_task.status != "failed": - await mark_chat_generation_task_failed_and_refund_once( - db, - task=chat_task, - error_message="爆款开头复刻步骤被重新生成或删除,旧生成任务已取消", - pipeline_stage="failed", - ) - chat_task.deleted_at = deleted_at - if steps: - await log_module_event( - db, - project=project, - event_type=ModuleEventTypeEnum.SOFT_DELETE_STEPS.value, - message=f"软删除第 {start_index} 步及之后的旧子任务", - detail={"step_ids": [step.id for step in steps]}, - ) def _step_to_out(step: ModuleGenerationStep) -> HotOpeningStepOut: @@ -459,17 +338,11 @@ def _step_to_out(step: ModuleGenerationStep) -> HotOpeningStepOut: def _snapshot_from_chat(chat_task: ChatGenerationTask | None) -> dict[str, Any]: - if not chat_task: - return {} - return _parse_json(chat_task.engine_snapshot_json, {}) or {} + return _common_snapshot_from_chat(chat_task) async def _chat_tasks_by_id(db: AsyncSession, steps: list[ModuleGenerationStep]) -> dict[str, ChatGenerationTask]: - ids = [step.chat_task_id for step in steps if step.chat_task_id] - if not ids: - return {} - result = await db.execute(select(ChatGenerationTask).where(ChatGenerationTask.id.in_(ids))) - return {task.id: task for task in result.scalars().all()} + return await _base_chat_tasks_by_id(db, steps) async def project_to_detail_out(db: AsyncSession, project: ModuleGenerationProject) -> HotOpeningTaskDetailOut: @@ -763,73 +636,15 @@ async def update_hot_opening_step( step_id: str, req: HotOpeningStepUpdate, ) -> tuple[ModuleGenerationProject, ModuleGenerationStep]: - project = await _get_project_for_user(db, project_id=project_id, user=current_user, for_update=True) - step = await _get_step_for_user(db, project_id=project_id, step_id=step_id, user=current_user, for_update=True) - if step.status == ModuleStepStatusEnum.PROCESSING.value: - raise HTTPException(status_code=400, detail="当前子任务正在处理中,暂不能修改") - - input_data = _step_payload(step.input_json) - output_data = _unwrap_step_output(step.output_json) - - if step.step_code == HotOpeningStepCodeEnum.MATERIAL_INPUT.value: - input_data = _merge_dict( - input_data, - { - "material_video_url": req.material_video_url, - "material_image_url": req.material_image_url, - "source_project_name": req.source_project_name, - "target_project_name": req.target_project_name, - "core_content_point": req.core_content_point, - }, - ) - if req.target_project_name: - project.title = req.target_project_name - elif step.step_code == HotOpeningStepCodeEnum.IMAGE_PROMPT_OPTIMIZE.value: - if req.prompt is not None: - output_data["optimized_prompt"] = req.prompt - output_data["prompt"] = req.prompt - elif step.step_code == HotOpeningStepCodeEnum.VIDEO_PROMPT_OPTIMIZE.value: - if req.prompt_schema is not None: - output_data["prompt_schema"] = req.prompt_schema - if req.prompt is not None: - output_data["final_prompt"] = req.prompt - else: - if req.input_json: - input_data = _merge_dict(input_data, req.input_json) - if req.output_json: - output_data = _merge_dict(output_data, req.output_json) - - if req.input_json: - input_data = _merge_dict(input_data, req.input_json) - if req.output_json: - output_data = _merge_dict(output_data, req.output_json) - - _force_set_json( - step, - "input_json", - _step_input( - step_code=step.step_code, - payload=input_data, - source_step_id=step.source_step_id, - parent_step_id=step.parent_step_id, - ), + return await update_module_step( + db, + current_user=current_user, + project_id=project_id, + step_id=step_id, + req=req, + config=FLOW_CONFIG, + log_module_event=log_module_event, ) - _force_set_json( - step, - "output_json", - _step_output(step_code=step.step_code, status=ModuleStepStatusEnum.COMPLETED.value, payload=output_data), - ) - step.status = ModuleStepStatusEnum.COMPLETED.value - step.error_message = None - step.completed_at = _now() - project.status = ModuleProjectStatusEnum.WAITING_USER.value - project.current_step_code = step.step_code - project.error_message = None - - await _soft_delete_steps_from_index(db, project=project, start_index=step.step_index + 1) - await log_module_event(db, project=project, step=step, event_type=ModuleEventTypeEnum.STEP_UPDATED.value, message="用户修改子任务内容") - await db.flush() - return project, step async def update_hot_opening_material_input( @@ -844,54 +659,14 @@ async def update_hot_opening_material_input( 采用方案 B:软删除旧第1步及之后的当前有效步骤,然后新建第1步 version+1。 未传字段沿用旧第1步素材输入,避免前端只改一个字段时丢失其它素材信息。 """ - project = await _get_project_for_user(db, project_id=project_id, user=current_user, for_update=True) - old_material_step = await _get_current_step_by_code(db, project.id, HotOpeningStepCodeEnum.MATERIAL_INPUT.value) - old_material = _step_payload(old_material_step.input_json if old_material_step else None) - - material = { - "material_video_url": req.material_video_url if req.material_video_url is not None else old_material.get("material_video_url"), - "material_image_url": req.material_image_url if req.material_image_url is not None else old_material.get("material_image_url"), - "source_project_name": req.source_project_name if req.source_project_name is not None else old_material.get("source_project_name"), - "target_project_name": req.target_project_name if req.target_project_name is not None else old_material.get("target_project_name"), - "core_content_point": req.core_content_point if req.core_content_point is not None else old_material.get("core_content_point"), - } - - missing_fields = [key for key, value in material.items() if value is None or str(value).strip() == ""] - if missing_fields: - raise HTTPException(status_code=400, detail=f"素材输入缺少必要字段: {', '.join(missing_fields)}") - - await _soft_delete_steps_from_index(db, project=project, start_index=STEP_INDEX_MAP[HotOpeningStepCodeEnum.MATERIAL_INPUT.value]) - - project.title = str(material["target_project_name"]) - project.status = ModuleProjectStatusEnum.WAITING_USER.value - project.current_step_code = HotOpeningStepCodeEnum.MATERIAL_INPUT.value - project.final_image_url = None - project.final_video_url = None - project.final_video_cover_url = None - project.error_message = None - project.completed_at = None - - new_step = await _create_step( + return await update_module_material_input( db, - project=project, - step_code=HotOpeningStepCodeEnum.MATERIAL_INPUT.value, - status=ModuleStepStatusEnum.COMPLETED.value, - input_data=material, - output_data={"message": "素材输入已修改,旧步骤已软删除。下一步请重新生成图片AI提词。"}, + current_user=current_user, + project_id=project_id, + req=req, + config=FLOW_CONFIG, + log_module_event=log_module_event, ) - await log_module_event( - db, - project=project, - step=new_step, - event_type=ModuleEventTypeEnum.STEP_UPDATED.value, - message="用户修改素材输入并重建第1步新版本", - detail={ - "old_material_step_id": old_material_step.id if old_material_step else None, - "new_material_step_id": new_step.id, - "version": new_step.version, - }, - ) - return project.id, new_step.id async def update_hot_opening_image_prompt( @@ -906,55 +681,15 @@ async def update_hot_opening_image_prompt( 修改后软删除第3、4、5步当前有效任务,让用户从图片生成开始重新执行。 """ - project = await _get_project_for_user(db, project_id=project_id, user=current_user, for_update=True) - step = await _get_step_for_user(db, project_id=project_id, step_id=step_id, user=current_user, for_update=True) - if step.step_code != HotOpeningStepCodeEnum.IMAGE_PROMPT_OPTIMIZE.value: - raise HTTPException(status_code=400, detail="只能修改第2步图片 AI 提词子任务") - if step.status != ModuleStepStatusEnum.COMPLETED.value: - raise HTTPException(status_code=400, detail="图片 AI 提词未完成,不能直接修改") - - output_data = _step_payload(step.output_json) - usage = _step_usage(step.output_json) - new_prompt = req.prompt.strip() - output_data["optimized_prompt"] = new_prompt - output_data["prompt"] = new_prompt - output_data["manual_edited"] = True - output_data["manual_edited_at"] = _now().isoformat() - - _force_set_json( - step, - "output_json", - _step_output( - step_code=HotOpeningStepCodeEnum.IMAGE_PROMPT_OPTIMIZE.value, - status=ModuleStepStatusEnum.COMPLETED.value, - payload=output_data, - usage=usage, - ), - ) - step.status = ModuleStepStatusEnum.COMPLETED.value - step.error_message = None - step.completed_at = _now() - - await _soft_delete_steps_from_index(db, project=project, start_index=STEP_INDEX_MAP[HotOpeningStepCodeEnum.IMAGE_GENERATE.value]) - - project.status = ModuleProjectStatusEnum.WAITING_USER.value - project.current_step_code = HotOpeningStepCodeEnum.IMAGE_PROMPT_OPTIMIZE.value - project.final_image_url = None - project.final_video_url = None - project.final_video_cover_url = None - project.completed_at = None - project.error_message = None - - await log_module_event( + return await update_module_image_prompt( db, - project=project, - step=step, - event_type=ModuleEventTypeEnum.STEP_UPDATED.value, - message="用户直接修改图片 AI 优化提词,已软删除后续步骤", - detail={"start_deleted_step_index": STEP_INDEX_MAP[HotOpeningStepCodeEnum.IMAGE_GENERATE.value]}, + current_user=current_user, + project_id=project_id, + step_id=step_id, + req=req, + config=FLOW_CONFIG, + log_module_event=log_module_event, ) - await db.flush() - return project, step async def update_hot_opening_video_prompt_schema( @@ -970,84 +705,17 @@ async def update_hot_opening_video_prompt_schema( 服务端已有 schema 为基准:视频规格、数组长度、时间段、合规控制、质量控制、协议字段均锁定。 最终提示词允许修改,但保存前会清洗视频时长、比例、分辨率、帧率等参数。 """ - project = await _get_project_for_user(db, project_id=project_id, user=current_user, for_update=True) - step = await _get_step_for_user(db, project_id=project_id, step_id=step_id, user=current_user, for_update=True) - if step.step_code != HotOpeningStepCodeEnum.VIDEO_PROMPT_OPTIMIZE.value: - raise HTTPException(status_code=400, detail="只能修改第4步视频 AI 提词 JSON schema 子任务") - if step.status != ModuleStepStatusEnum.COMPLETED.value: - raise HTTPException(status_code=400, detail="视频 AI 提词未完成,不能直接修改") - - output_data = _step_payload(step.output_json) - usage = _step_usage(step.output_json) - input_data = _step_payload(step.input_json) - server_schema = output_data.get("prompt_schema") if isinstance(output_data.get("prompt_schema"), dict) else {} - video_config = output_data.get("params_used_for_prompt") or input_data.get("video_config") or {} - if not isinstance(video_config, dict) or not video_config.get("duration") or not video_config.get("aspect_ratio") or not video_config.get("resolution"): - raise HTTPException(status_code=400, detail="缺少第4步视频参数快照,不能安全修改视频 schema") - - patched_schema = patch_video_prompt_schema_from_client( - server_schema=server_schema, - client_schema=req.prompt_schema, - video_config=video_config, - ) - final_prompt = build_final_video_prompt(patched_schema) - - output_data["prompt_schema"] = patched_schema - output_data["final_prompt"] = final_prompt - output_data["params_used_for_prompt"] = video_config - output_data["manual_edited"] = True - output_data["manual_edited_at"] = _now().isoformat() - - _force_set_json( - step, - "output_json", - _step_output( - step_code=HotOpeningStepCodeEnum.VIDEO_PROMPT_OPTIMIZE.value, - status=ModuleStepStatusEnum.COMPLETED.value, - payload=output_data, - usage=usage, - ), - ) - step.status = ModuleStepStatusEnum.COMPLETED.value - step.error_message = None - step.completed_at = _now() - - await _soft_delete_steps_from_index(db, project=project, start_index=STEP_INDEX_MAP[HotOpeningStepCodeEnum.VIDEO_GENERATE.value]) - - project.status = ModuleProjectStatusEnum.WAITING_USER.value - project.current_step_code = HotOpeningStepCodeEnum.VIDEO_PROMPT_OPTIMIZE.value - project.final_video_url = None - project.final_video_cover_url = None - project.completed_at = None - project.error_message = None - - await log_module_event( + return await update_module_video_prompt_schema( db, - project=project, - step=step, - event_type=ModuleEventTypeEnum.STEP_UPDATED.value, - message="用户修改视频 AI 提词 schema,已软删除视频生成步骤", - detail={ - "start_deleted_step_index": STEP_INDEX_MAP[HotOpeningStepCodeEnum.VIDEO_GENERATE.value], - "locked_fields": [ - "schema_version", - "schema_usage", - "画面属性.视频时长", - "画面属性.视频比例", - "画面属性.清晰度", - "画面属性.帧率", - "画面属性.推荐分辨率", - "动作流程[*].时间段", - "镜头流程[*].时间段", - "动态时间规划", - "输出规格限制", - "质量控制", - "合规控制", - ], - }, + current_user=current_user, + project_id=project_id, + step_id=step_id, + req=req, + config=FLOW_CONFIG, + log_module_event=log_module_event, + patch_video_prompt_schema_from_client=patch_video_prompt_schema_from_client, + build_final_video_prompt=build_final_video_prompt, ) - await db.flush() - return project, step async def submit_image_prompt_optimize( @@ -1754,11 +1422,4 @@ async def delete_hot_opening_project(db: AsyncSession, *, current_user: User, pr return HotOpeningDeleteOut(message="项目已删除", project_id=project.id, deleted=True) def _build_file_url_or_data_uri(file_url: str) -> str: - """ - Convert local upload path to base64 data URI. - Keep remote http/https/data URLs as-is. - """ - if file_url.startswith(("http://", "https://", "data:")): - return file_url - file_url_sign = build_resource_signed_url(resource_url=file_url, expire_seconds=86400) - return f"{settings.BASE_URL}{file_url_sign}" \ No newline at end of file + return _common_build_file_url_or_data_uri(file_url) diff --git a/video-gen-api/app/services/module_generation_flow_base_service.py b/video-gen-api/app/services/module_generation_flow_base_service.py new file mode 100644 index 00000000..2d876d72 --- /dev/null +++ b/video-gen-api/app/services/module_generation_flow_base_service.py @@ -0,0 +1,258 @@ +from __future__ import annotations + +from collections.abc import Awaitable, Callable +from datetime import datetime +from typing import Any + +from fastapi import HTTPException +from sqlalchemy import func, select +from sqlalchemy.ext.asyncio import AsyncSession + +from app.enums.common import ModuleEventTypeEnum, ModuleStepStatusEnum +from app.models.chat_generation_task import ChatGenerationTask +from app.models.module_generation_project import ModuleGenerationProject +from app.models.module_generation_step import ModuleGenerationStep +from app.models.user import User +from app.services.generation_refund_service import mark_chat_generation_task_failed_and_refund_once +from app.enums.module_generation_flow import ModuleGenerationFlowConfig +from app.services.module_generation_step_common_service import build_step_input, build_step_output, utc_now +from app.services.resource_accounting_service import soft_delete_chat_task_resources +from app.utils.id_gen import generate_id + +LogModuleEventCallable = Callable[..., Awaitable[None]] + + +async def get_project_for_user( + db: AsyncSession, + *, + project_id: str, + user: User, + config: ModuleGenerationFlowConfig, + for_update: bool = False, + populate_existing: bool = False, +) -> ModuleGenerationProject: + query = select(ModuleGenerationProject).where( + ModuleGenerationProject.id == project_id, + ModuleGenerationProject.module == config.module, + ModuleGenerationProject.deleted_at.is_(None), + ) + if not user.is_admin: + query = query.where(ModuleGenerationProject.user_id == user.id) + if populate_existing: + query = query.execution_options(populate_existing=True) + if for_update: + query = query.with_for_update() + result = await db.execute(query.limit(1)) + project = result.scalar_one_or_none() + if not project: + raise HTTPException(status_code=404, detail=config.project_not_found_message) + return project + + +async def get_step_for_user( + db: AsyncSession, + *, + project_id: str, + step_id: str, + user: User, + config: ModuleGenerationFlowConfig, + for_update: bool = False, +) -> ModuleGenerationStep: + await get_project_for_user(db, project_id=project_id, user=user, config=config, for_update=for_update) + query = select(ModuleGenerationStep).where( + ModuleGenerationStep.id == step_id, + ModuleGenerationStep.project_id == project_id, + ModuleGenerationStep.module == config.module, + ModuleGenerationStep.deleted_at.is_(None), + ModuleGenerationStep.is_current == True, + ) + if not user.is_admin: + query = query.where(ModuleGenerationStep.user_id == user.id) + if for_update: + query = query.with_for_update() + result = await db.execute(query.limit(1)) + step = result.scalar_one_or_none() + if not step: + raise HTTPException(status_code=404, detail=config.step_not_found_message) + return step + + +async def get_current_steps(db: AsyncSession, *, project_id: str, config: ModuleGenerationFlowConfig) -> list[ModuleGenerationStep]: + result = await db.execute( + select(ModuleGenerationStep) + .where( + ModuleGenerationStep.project_id == project_id, + ModuleGenerationStep.module == config.module, + ModuleGenerationStep.deleted_at.is_(None), + ModuleGenerationStep.is_current == True, + ) + .order_by(ModuleGenerationStep.step_index.asc(), ModuleGenerationStep.created_at.asc()) + ) + return list(result.scalars().all()) + + +async def get_current_step_by_code( + db: AsyncSession, + *, + project_id: str, + step_code: str, + config: ModuleGenerationFlowConfig, +) -> ModuleGenerationStep | None: + result = await db.execute( + select(ModuleGenerationStep) + .where( + ModuleGenerationStep.project_id == project_id, + ModuleGenerationStep.module == config.module, + ModuleGenerationStep.step_code == step_code, + ModuleGenerationStep.is_current == True, + ModuleGenerationStep.deleted_at.is_(None), + ) + .order_by(ModuleGenerationStep.version.desc(), ModuleGenerationStep.created_at.desc()) + .limit(1) + ) + return result.scalar_one_or_none() + + +async def next_version( + db: AsyncSession, + *, + project_id: str, + step_code: str, + config: ModuleGenerationFlowConfig, +) -> int: + result = await db.execute( + select(func.max(ModuleGenerationStep.version)).where( + ModuleGenerationStep.project_id == project_id, + ModuleGenerationStep.module == config.module, + ModuleGenerationStep.step_code == step_code, + ) + ) + return int(result.scalar_one_or_none() or 0) + 1 + + +async def create_module_step( + db: AsyncSession, + *, + project: ModuleGenerationProject, + step_code: str, + config: ModuleGenerationFlowConfig, + log_module_event: LogModuleEventCallable, + status: str = ModuleStepStatusEnum.PENDING.value, + parent_step_id: str | None = None, + source_step_id: str | None = None, + chat_task_id: str | None = None, + input_data: dict[str, Any] | None = None, + output_data: dict[str, Any] | None = None, +) -> ModuleGenerationStep: + version = await next_version(db, project_id=project.id, step_code=step_code, config=config) + step = ModuleGenerationStep( + id=generate_id(), + project_id=project.id, + user_id=project.user_id, + module=project.module, + step_index=config.step_index_map[step_code], + step_code=step_code, + status=status, + version=version, + is_current=True, + parent_step_id=parent_step_id, + source_step_id=source_step_id, + chat_task_id=chat_task_id, + input_json=( + build_step_input( + step_code=step_code, + payload=input_data, + source_step_id=source_step_id, + parent_step_id=parent_step_id, + schema_version=config.step_io_schema_version, + ) + if input_data is not None + else None + ), + output_json=( + build_step_output( + step_code=step_code, + status=status, + result=output_data, + schema_version=config.step_io_schema_version, + ) + if output_data is not None + else None + ), + started_at=utc_now() if status == ModuleStepStatusEnum.PROCESSING.value else None, + completed_at=utc_now() if status == ModuleStepStatusEnum.COMPLETED.value else None, + ) + db.add(step) + project.current_step_code = step_code + await db.flush() + await log_module_event( + db, + project=project, + step=step, + event_type=ModuleEventTypeEnum.STEP_CREATED.value, + detail={"step_code": step_code, "version": version}, + ) + return step + + +async def soft_delete_steps_from_index( + db: AsyncSession, + *, + project: ModuleGenerationProject, + start_index: int, + config: ModuleGenerationFlowConfig, + log_module_event: LogModuleEventCallable, + deleted_at: datetime | None = None, +) -> list[ModuleGenerationStep]: + deleted_at = deleted_at or utc_now() + result = await db.execute( + select(ModuleGenerationStep) + .where( + ModuleGenerationStep.project_id == project.id, + ModuleGenerationStep.module == config.module, + ModuleGenerationStep.is_current == True, + ModuleGenerationStep.deleted_at.is_(None), + ModuleGenerationStep.step_index >= start_index, + ) + .with_for_update() + ) + steps = list(result.scalars().all()) + for step in steps: + step.is_current = False + step.deleted_at = deleted_at + if step.chat_task_id: + chat_result = await db.execute( + select(ChatGenerationTask) + .where(ChatGenerationTask.id == step.chat_task_id, ChatGenerationTask.deleted_at.is_(None)) + .with_for_update() + .limit(1) + ) + chat_task = chat_result.scalar_one_or_none() + if chat_task: + if chat_task.status == "completed": + await soft_delete_chat_task_resources(db, chat_task.id, deleted_at=deleted_at) + elif chat_task.status != "failed": + await mark_chat_generation_task_failed_and_refund_once( + db, + task=chat_task, + error_message=config.cancel_chat_task_error_message, + pipeline_stage="failed", + ) + chat_task.deleted_at = deleted_at + if steps: + await log_module_event( + db, + project=project, + event_type=ModuleEventTypeEnum.SOFT_DELETE_STEPS.value, + message=f"软删除第 {start_index} 步及之后的旧子任务", + detail={"step_ids": [step.id for step in steps]}, + ) + return steps + + +async def chat_tasks_by_id(db: AsyncSession, steps: list[ModuleGenerationStep]) -> dict[str, ChatGenerationTask]: + ids = [step.chat_task_id for step in steps if step.chat_task_id] + if not ids: + return {} + result = await db.execute(select(ChatGenerationTask).where(ChatGenerationTask.id.in_(ids))) + return {task.id: task for task in result.scalars().all()} diff --git a/video-gen-api/app/services/module_generation_step_common_service.py b/video-gen-api/app/services/module_generation_step_common_service.py new file mode 100644 index 00000000..ca731dc6 --- /dev/null +++ b/video-gen-api/app/services/module_generation_step_common_service.py @@ -0,0 +1,168 @@ +from __future__ import annotations + +import base64 +import json +from copy import deepcopy +from datetime import datetime, timezone +from typing import Any + +from sqlalchemy.orm.attributes import flag_modified + +from app.config import settings +from app.models.chat_generation_task import ChatGenerationTask +from app.services.resource_signed_url_service import build_resource_signed_url + +STEP_IO_SCHEMA_VERSION = "module_generation_step_io_v1" + + +def utc_now() -> datetime: + return datetime.now(timezone.utc) + + +def json_dumps(data: Any) -> str | None: + if data is None: + return None + return json.dumps(data, ensure_ascii=False, default=str) + + +def parse_json(value: Any, fallback: Any = None) -> Any: + if value is None or value == "": + return fallback + if isinstance(value, (dict, list)): + return value + if isinstance(value, str): + try: + return json.loads(value) + except Exception: + return fallback + return fallback + + +def build_step_input( + *, + step_code: str, + payload: dict[str, Any] | None = None, + source_step_id: str | None = None, + parent_step_id: str | None = None, + context: dict[str, Any] | None = None, + schema_version: str = STEP_IO_SCHEMA_VERSION, +) -> dict[str, Any]: + return { + "schema_version": schema_version, + "step_code": step_code, + "source": { + "source_step_id": source_step_id, + "parent_step_id": parent_step_id, + }, + "payload": payload or {}, + "context": context or {}, + } + + +def build_step_output( + *, + step_code: str, + status: str, + payload: dict[str, Any] | None = None, + result: dict[str, Any] | None = None, + usage: dict[str, Any] | None = None, + error: dict[str, Any] | None = None, + schema_version: str = STEP_IO_SCHEMA_VERSION, +) -> dict[str, Any]: + return { + "schema_version": schema_version, + "step_code": step_code, + "status": status, + "payload": payload or {}, + "result": result or {}, + "usage": usage or {}, + "error": error or {}, + } + + +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 + + +def step_payload(value: Any, *, schema_version: str = STEP_IO_SCHEMA_VERSION) -> dict[str, Any]: + data = parse_json(value, {}) or {} + if is_wrapped_step_io(data, schema_version=schema_version): + payload = data.get("payload") + return payload if isinstance(payload, dict) else {} + return data if isinstance(data, dict) else {} + + +def step_result(value: Any, *, schema_version: str = STEP_IO_SCHEMA_VERSION) -> dict[str, Any]: + data = parse_json(value, {}) or {} + if is_wrapped_step_io(data, schema_version=schema_version): + result = data.get("result") + if isinstance(result, dict) and result: + return result + payload = data.get("payload") + return payload if isinstance(payload, dict) else {} + return data if isinstance(data, dict) else {} + + +def step_usage(value: Any, *, schema_version: str = STEP_IO_SCHEMA_VERSION) -> dict[str, Any]: + data = parse_json(value, {}) or {} + if is_wrapped_step_io(data, schema_version=schema_version): + usage = data.get("usage") + return usage if isinstance(usage, dict) else {} + usage = data.get("token_usage") if isinstance(data, dict) else {} + return usage if isinstance(usage, dict) else {} + + +def unwrap_step_output(value: Any, *, schema_version: str = STEP_IO_SCHEMA_VERSION) -> dict[str, Any]: + data = parse_json(value, {}) or {} + if not is_wrapped_step_io(data, schema_version=schema_version): + return data if isinstance(data, dict) else {} + merged: dict[str, Any] = {} + payload = data.get("payload") + result = data.get("result") + usage = data.get("usage") + if isinstance(payload, dict): + merged.update(payload) + if isinstance(result, dict): + merged.update(result) + if isinstance(usage, dict) and usage: + merged["token_usage"] = usage + return merged + + +def merge_dict(old: dict[str, Any] | None, new: dict[str, Any] | None) -> dict[str, Any]: + merged = dict(old or {}) + for key, value in (new or {}).items(): + if value is not None: + merged[key] = value + return merged + + +def force_set_json(model_obj: Any, field_name: str, value: Any) -> None: + """强制持久化 JSON / JSONB 字段。 + + SQLAlchemy 对 dict/list 的嵌套原地修改不会稳定触发 dirty 判定。 + 所有编辑类接口在写入 input_json / output_json 时统一走这里: + 1. deepcopy 断开旧引用; + 2. 整体重新赋值; + 3. flag_modified 显式标记字段已变更。 + """ + setattr(model_obj, field_name, deepcopy(value)) + flag_modified(model_obj, field_name) + + +def snapshot_from_chat(chat_task: ChatGenerationTask | None) -> dict[str, Any]: + if not chat_task: + return {} + return {"chat_task_id": chat_task.id, "status": chat_task.status, "pipeline_stage": chat_task.pipeline_stage} + + +def build_file_url_or_data_uri(file_url: str) -> str: + if file_url.startswith("http://") or file_url.startswith("https://") or file_url.startswith("data:"): + return file_url + file_url_sign = build_resource_signed_url(resource_url=file_url, expire_seconds=86400) + return f"{settings.BASE_URL}{file_url_sign}" + # try: + # with open(file_url, "rb") as f: + # return "data:image/png;base64," + base64.b64encode(f.read()).decode("utf-8") + # except Exception: + # return f"{settings.STORAGE_BASE_URL.rstrip('/')}/{file_url.lstrip('/')}" diff --git a/video-gen-api/app/services/module_generation_step_update_service.py b/video-gen-api/app/services/module_generation_step_update_service.py new file mode 100644 index 00000000..ce0ae576 --- /dev/null +++ b/video-gen-api/app/services/module_generation_step_update_service.py @@ -0,0 +1,363 @@ +from __future__ import annotations + +from collections.abc import Awaitable, Callable +from typing import Any + +from fastapi import HTTPException +from sqlalchemy.ext.asyncio import AsyncSession + +from app.enums.common import ModuleEventTypeEnum, ModuleProjectStatusEnum, ModuleStepStatusEnum +from app.models.module_generation_project import ModuleGenerationProject +from app.models.module_generation_step import ModuleGenerationStep +from app.models.user import User +from app.services.module_generation_flow_base_service import ( + create_module_step, + get_current_step_by_code, + get_project_for_user, + get_step_for_user, + soft_delete_steps_from_index, +) +from app.enums.module_generation_flow import ModuleGenerationFlowConfig +from app.services.module_generation_step_common_service import ( + build_step_input, + build_step_output, + force_set_json, + merge_dict, + step_payload, + step_usage, + unwrap_step_output, + utc_now, +) + +LogModuleEventCallable = Callable[..., Awaitable[None]] +PatchVideoPromptSchemaCallable = Callable[..., dict[str, Any]] +BuildFinalVideoPromptCallable = Callable[[dict[str, Any]], str] + + +async def update_module_step( + db: AsyncSession, + *, + current_user: User, + project_id: str, + step_id: str, + req: Any, + config: ModuleGenerationFlowConfig, + log_module_event: LogModuleEventCallable, +) -> tuple[ModuleGenerationProject, ModuleGenerationStep]: + project = await get_project_for_user(db, project_id=project_id, user=current_user, config=config, for_update=True) + step = await get_step_for_user(db, project_id=project_id, step_id=step_id, user=current_user, config=config, for_update=True) + if step.status == ModuleStepStatusEnum.PROCESSING.value: + raise HTTPException(status_code=400, detail="当前子任务正在处理中,暂不能修改") + + input_data = step_payload(step.input_json, schema_version=config.step_io_schema_version) + output_data = unwrap_step_output(step.output_json, schema_version=config.step_io_schema_version) + + if step.step_code == config.material_step_code: + input_data = merge_dict( + input_data, + { + "material_video_url": getattr(req, "material_video_url", None), + "material_image_url": getattr(req, "material_image_url", None), + "source_project_name": getattr(req, "source_project_name", None), + "target_project_name": getattr(req, "target_project_name", None), + "core_content_point": getattr(req, "core_content_point", None), + }, + ) + target_project_name = getattr(req, "target_project_name", None) + if target_project_name: + project.title = target_project_name + elif step.step_code == config.image_prompt_step_code: + prompt = getattr(req, "prompt", None) + if prompt is not None: + output_data["optimized_prompt"] = prompt + output_data["prompt"] = prompt + elif step.step_code == config.video_prompt_step_code: + prompt_schema = getattr(req, "prompt_schema", None) + prompt = getattr(req, "prompt", None) + if prompt_schema is not None: + output_data["prompt_schema"] = prompt_schema + if prompt is not None: + output_data["final_prompt"] = prompt + else: + if getattr(req, "input_json", None): + input_data = merge_dict(input_data, req.input_json) + if getattr(req, "output_json", None): + output_data = merge_dict(output_data, req.output_json) + + if getattr(req, "input_json", None): + input_data = merge_dict(input_data, req.input_json) + if getattr(req, "output_json", None): + output_data = merge_dict(output_data, req.output_json) + + force_set_json( + step, + "input_json", + build_step_input( + step_code=step.step_code, + payload=input_data, + source_step_id=step.source_step_id, + parent_step_id=step.parent_step_id, + schema_version=config.step_io_schema_version, + ), + ) + force_set_json( + step, + "output_json", + build_step_output(step_code=step.step_code, status=ModuleStepStatusEnum.COMPLETED.value, payload=output_data, schema_version=config.step_io_schema_version), + ) + step.status = ModuleStepStatusEnum.COMPLETED.value + step.error_message = None + step.completed_at = utc_now() + project.status = ModuleProjectStatusEnum.WAITING_USER.value + project.current_step_code = step.step_code + project.error_message = None + + await soft_delete_steps_from_index( + db, + project=project, + start_index=step.step_index + 1, + config=config, + log_module_event=log_module_event, + ) + await log_module_event(db, project=project, step=step, event_type=ModuleEventTypeEnum.STEP_UPDATED.value, message="用户修改子任务内容") + await db.flush() + return project, step + + +async def update_module_material_input( + db: AsyncSession, + *, + current_user: User, + project_id: str, + req: Any, + config: ModuleGenerationFlowConfig, + log_module_event: LogModuleEventCallable, +) -> tuple[str, str]: + project = await get_project_for_user(db, project_id=project_id, user=current_user, config=config, for_update=True) + old_material_step = await get_current_step_by_code(db, project_id=project.id, step_code=config.material_step_code, config=config) + old_material = step_payload(old_material_step.input_json if old_material_step else None, schema_version=config.step_io_schema_version) + + material_video_url = ( + getattr(req, "material_video_url", None) + if config.material_video_url_editable and getattr(req, "material_video_url", None) is not None + else old_material.get("material_video_url") + ) + material = { + "material_video_url": material_video_url, + "material_image_url": getattr(req, "material_image_url", None) if getattr(req, "material_image_url", None) is not None else old_material.get("material_image_url"), + "source_project_name": getattr(req, "source_project_name", None) if getattr(req, "source_project_name", None) is not None else old_material.get("source_project_name"), + "target_project_name": getattr(req, "target_project_name", None) if getattr(req, "target_project_name", None) is not None else old_material.get("target_project_name"), + "core_content_point": getattr(req, "core_content_point", None) if getattr(req, "core_content_point", None) is not None else old_material.get("core_content_point"), + } + + missing_fields = [key for key, value in material.items() if value is None or str(value).strip() == ""] + if missing_fields: + raise HTTPException(status_code=400, detail=f"素材输入缺少必要字段: {', '.join(missing_fields)}") + + await soft_delete_steps_from_index( + db, + project=project, + start_index=config.step_index_map[config.material_step_code], + config=config, + log_module_event=log_module_event, + ) + + project.title = str(material["target_project_name"]) + project.status = ModuleProjectStatusEnum.WAITING_USER.value + project.current_step_code = config.material_step_code + project.final_image_url = None + project.final_video_url = None + project.final_video_cover_url = None + project.error_message = None + project.completed_at = None + + new_step = await create_module_step( + db, + project=project, + step_code=config.material_step_code, + config=config, + log_module_event=log_module_event, + status=ModuleStepStatusEnum.COMPLETED.value, + input_data=material, + output_data={"message": "素材输入已修改,旧步骤已软删除。下一步请重新生成图片AI提词。"}, + ) + await log_module_event( + db, + project=project, + step=new_step, + event_type=ModuleEventTypeEnum.STEP_UPDATED.value, + message="用户修改素材输入并重建第1步新版本", + detail={ + "old_material_step_id": old_material_step.id if old_material_step else None, + "new_material_step_id": new_step.id, + "version": new_step.version, + }, + ) + return project.id, new_step.id + + +async def update_module_image_prompt( + db: AsyncSession, + *, + current_user: User, + project_id: str, + step_id: str, + req: Any, + config: ModuleGenerationFlowConfig, + log_module_event: LogModuleEventCallable, +) -> tuple[ModuleGenerationProject, ModuleGenerationStep]: + project = await get_project_for_user(db, project_id=project_id, user=current_user, config=config, for_update=True) + step = await get_step_for_user(db, project_id=project_id, step_id=step_id, user=current_user, config=config, for_update=True) + if step.step_code != config.image_prompt_step_code: + raise HTTPException(status_code=400, detail="只能修改第2步图片 AI 提词子任务") + if step.status != ModuleStepStatusEnum.COMPLETED.value: + raise HTTPException(status_code=400, detail="图片 AI 提词未完成,不能直接修改") + + output_data = step_payload(step.output_json, schema_version=config.step_io_schema_version) + usage = step_usage(step.output_json, schema_version=config.step_io_schema_version) + new_prompt = req.prompt.strip() + output_data["optimized_prompt"] = new_prompt + output_data["prompt"] = new_prompt + output_data["manual_edited"] = True + output_data["manual_edited_at"] = utc_now().isoformat() + + force_set_json( + step, + "output_json", + build_step_output( + step_code=config.image_prompt_step_code, + status=ModuleStepStatusEnum.COMPLETED.value, + payload=output_data, + usage=usage, + schema_version=config.step_io_schema_version, + ), + ) + step.status = ModuleStepStatusEnum.COMPLETED.value + step.error_message = None + step.completed_at = utc_now() + + await soft_delete_steps_from_index( + db, + project=project, + start_index=config.step_index_map[config.image_generate_step_code], + config=config, + log_module_event=log_module_event, + ) + + project.status = ModuleProjectStatusEnum.WAITING_USER.value + project.current_step_code = config.image_prompt_step_code + project.final_image_url = None + project.final_video_url = None + project.final_video_cover_url = None + project.completed_at = None + project.error_message = None + + await log_module_event( + db, + project=project, + step=step, + event_type=ModuleEventTypeEnum.STEP_UPDATED.value, + message="用户直接修改图片 AI 优化提词,已软删除后续步骤", + detail={"start_deleted_step_index": config.step_index_map[config.image_generate_step_code]}, + ) + await db.flush() + return project, step + + +async def update_module_video_prompt_schema( + db: AsyncSession, + *, + current_user: User, + project_id: str, + step_id: str, + req: Any, + config: ModuleGenerationFlowConfig, + log_module_event: LogModuleEventCallable, + patch_video_prompt_schema_from_client: PatchVideoPromptSchemaCallable, + build_final_video_prompt: BuildFinalVideoPromptCallable, +) -> tuple[ModuleGenerationProject, ModuleGenerationStep]: + project = await get_project_for_user(db, project_id=project_id, user=current_user, config=config, for_update=True) + step = await get_step_for_user(db, project_id=project_id, step_id=step_id, user=current_user, config=config, for_update=True) + if step.step_code != config.video_prompt_step_code: + raise HTTPException(status_code=400, detail="只能修改第4步视频 AI 提词 JSON schema 子任务") + if step.status != ModuleStepStatusEnum.COMPLETED.value: + raise HTTPException(status_code=400, detail="视频 AI 提词未完成,不能直接修改") + + output_data = step_payload(step.output_json, schema_version=config.step_io_schema_version) + usage = step_usage(step.output_json, schema_version=config.step_io_schema_version) + input_data = step_payload(step.input_json, schema_version=config.step_io_schema_version) + server_schema = output_data.get("prompt_schema") if isinstance(output_data.get("prompt_schema"), dict) else {} + video_config = output_data.get("params_used_for_prompt") or input_data.get("video_config") or {} + if not isinstance(video_config, dict) or not video_config.get("duration") or not video_config.get("aspect_ratio") or not video_config.get("resolution"): + raise HTTPException(status_code=400, detail="缺少第4步视频参数快照,不能安全修改视频 schema") + + patched_schema = patch_video_prompt_schema_from_client( + server_schema=server_schema, + client_schema=req.prompt_schema, + video_config=video_config, + ) + final_prompt = build_final_video_prompt(patched_schema) + + output_data["prompt_schema"] = patched_schema + output_data["final_prompt"] = final_prompt + output_data["params_used_for_prompt"] = video_config + output_data["manual_edited"] = True + output_data["manual_edited_at"] = utc_now().isoformat() + + force_set_json( + step, + "output_json", + build_step_output( + step_code=config.video_prompt_step_code, + status=ModuleStepStatusEnum.COMPLETED.value, + payload=output_data, + usage=usage, + schema_version=config.step_io_schema_version, + ), + ) + step.status = ModuleStepStatusEnum.COMPLETED.value + step.error_message = None + step.completed_at = utc_now() + + await soft_delete_steps_from_index( + db, + project=project, + start_index=config.step_index_map[config.video_generate_step_code], + config=config, + log_module_event=log_module_event, + ) + + project.status = ModuleProjectStatusEnum.WAITING_USER.value + project.current_step_code = config.video_prompt_step_code + project.final_video_url = None + project.final_video_cover_url = None + project.completed_at = None + project.error_message = None + + await log_module_event( + db, + project=project, + step=step, + event_type=ModuleEventTypeEnum.STEP_UPDATED.value, + message="用户修改视频 AI 提词 schema,已软删除视频生成步骤", + detail={ + "start_deleted_step_index": config.step_index_map[config.video_generate_step_code], + "locked_fields": [ + "schema_version", + "schema_usage", + "画面属性.视频时长", + "画面属性.视频比例", + "画面属性.清晰度", + "画面属性.帧率", + "画面属性.推荐分辨率", + "动作流程[*].时间段", + "镜头流程[*].时间段", + "动态时间规划", + "输出规格限制", + "质量控制", + "合规控制", + ], + }, + ) + await db.flush() + return project, step diff --git a/video-gen-api/app/services/shot_replicate_flow_service.py b/video-gen-api/app/services/shot_replicate_flow_service.py index 7383ee9c..19300367 100644 --- a/video-gen-api/app/services/shot_replicate_flow_service.py +++ b/video-gen-api/app/services/shot_replicate_flow_service.py @@ -53,6 +53,39 @@ from app.services.hot_opening_video_prompt_service import ( from app.services.module_generation_log_service import log_module_error, log_module_event_file, log_module_prompt_event from app.services.llm import optimize_prompt from app.services.resource_accounting_service import soft_delete_chat_task_resources +from app.services.module_generation_flow_base_service import ( + chat_tasks_by_id as _base_chat_tasks_by_id, + create_module_step as _base_create_step, + get_current_step_by_code as _base_get_current_step_by_code, + get_current_steps as _base_get_current_steps, + get_project_for_user as _base_get_project_for_user, + get_step_for_user as _base_get_step_for_user, + next_version as _base_next_version, + soft_delete_steps_from_index as _base_soft_delete_steps_from_index, +) +from app.enums.module_generation_flow import ModuleGenerationFlowConfig +from app.services.module_generation_step_common_service import ( + build_file_url_or_data_uri as _common_build_file_url_or_data_uri, + build_step_input as _common_build_step_input, + build_step_output as _common_build_step_output, + force_set_json as _common_force_set_json, + is_wrapped_step_io as _common_is_wrapped_step_io, + json_dumps as _common_json, + merge_dict as _common_merge_dict, + parse_json as _common_parse_json, + snapshot_from_chat as _common_snapshot_from_chat, + step_payload as _common_step_payload, + step_result as _common_step_result, + step_usage as _common_step_usage, + unwrap_step_output as _common_unwrap_step_output, + utc_now as _common_now, +) +from app.services.module_generation_step_update_service import ( + update_module_image_prompt, + update_module_material_input, + update_module_step, + update_module_video_prompt_schema, +) from app.services.resource_signed_url_service import build_resource_signed_url from app.utils.id_gen import generate_id from app.models.shot_replicate_segment import ShotReplicateSegment @@ -72,28 +105,27 @@ STEP_INDEX_MAP = { STEP_IO_SCHEMA_VERSION = "shot_replicate_step_io_v1" - -def _now() -> datetime: - return datetime.now(timezone.utc) +FLOW_CONFIG = ModuleGenerationFlowConfig( + module=MODULE, + step_index_map=STEP_INDEX_MAP, + material_step_code=ShotReplicateStepCodeEnum.MATERIAL_INPUT.value, + image_prompt_step_code=ShotReplicateStepCodeEnum.IMAGE_PROMPT_OPTIMIZE.value, + image_generate_step_code=ShotReplicateStepCodeEnum.IMAGE_GENERATE.value, + video_prompt_step_code=ShotReplicateStepCodeEnum.VIDEO_PROMPT_OPTIMIZE.value, + video_generate_step_code=ShotReplicateStepCodeEnum.VIDEO_GENERATE.value, + project_not_found_message="拆镜复刻项目不存在", + step_not_found_message="子任务不存在", + cancel_chat_task_error_message="拆镜复刻步骤被重新生成或删除,旧生成任务已取消", + material_video_url_editable=False, + step_io_schema_version=STEP_IO_SCHEMA_VERSION, +) -def _json(data: Any) -> str | None: - if data is None: - return None - return json.dumps(data, ensure_ascii=False, default=str) - - -def _parse_json(value: Any, fallback: Any = None) -> Any: - if value is None or value == "": - return fallback - if isinstance(value, (dict, list)): - return value - if isinstance(value, str): - try: - return json.loads(value) - except Exception: - return fallback - return fallback +_now = _common_now +_json = _common_json +_parse_json = _common_parse_json +_merge_dict = _common_merge_dict +_force_set_json = _common_force_set_json def _step_input( @@ -104,16 +136,14 @@ def _step_input( parent_step_id: str | None = None, context: dict[str, Any] | None = None, ) -> dict[str, Any]: - return { - "schema_version": STEP_IO_SCHEMA_VERSION, - "step_code": step_code, - "source": { - "source_step_id": source_step_id, - "parent_step_id": parent_step_id, - }, - "payload": payload or {}, - "context": context or {}, - } + return _common_build_step_input( + step_code=step_code, + payload=payload, + source_step_id=source_step_id, + parent_step_id=parent_step_id, + context=context, + schema_version=STEP_IO_SCHEMA_VERSION, + ) def _step_output( @@ -125,85 +155,35 @@ def _step_output( usage: dict[str, Any] | None = None, error: dict[str, Any] | None = None, ) -> dict[str, Any]: - return { - "schema_version": STEP_IO_SCHEMA_VERSION, - "step_code": step_code, - "status": status, - "payload": payload or {}, - "result": result or {}, - "usage": usage or {}, - "error": error or {}, - } + return _common_build_step_output( + step_code=step_code, + status=status, + payload=payload, + result=result, + usage=usage, + error=error, + schema_version=STEP_IO_SCHEMA_VERSION, + ) def _is_wrapped_step_io(value: Any) -> bool: - return isinstance(value, dict) and value.get("schema_version") == STEP_IO_SCHEMA_VERSION + return _common_is_wrapped_step_io(value, schema_version=STEP_IO_SCHEMA_VERSION) def _step_payload(value: Any) -> dict[str, Any]: - data = _parse_json(value, {}) or {} - if _is_wrapped_step_io(data): - payload = data.get("payload") - return payload if isinstance(payload, dict) else {} - return data if isinstance(data, dict) else {} + return _common_step_payload(value, schema_version=STEP_IO_SCHEMA_VERSION) def _step_result(value: Any) -> dict[str, Any]: - data = _parse_json(value, {}) or {} - if _is_wrapped_step_io(data): - result = data.get("result") - if isinstance(result, dict) and result: - return result - payload = data.get("payload") - return payload if isinstance(payload, dict) else {} - return data if isinstance(data, dict) else {} + return _common_step_result(value, schema_version=STEP_IO_SCHEMA_VERSION) def _step_usage(value: Any) -> dict[str, Any]: - data = _parse_json(value, {}) or {} - if _is_wrapped_step_io(data): - usage = data.get("usage") - return usage if isinstance(usage, dict) else {} - usage = data.get("token_usage") if isinstance(data, dict) else {} - return usage if isinstance(usage, dict) else {} + return _common_step_usage(value, schema_version=STEP_IO_SCHEMA_VERSION) def _unwrap_step_output(value: Any) -> dict[str, Any]: - data = _parse_json(value, {}) or {} - if not _is_wrapped_step_io(data): - return data if isinstance(data, dict) else {} - merged: dict[str, Any] = {} - payload = data.get("payload") - result = data.get("result") - usage = data.get("usage") - if isinstance(payload, dict): - merged.update(payload) - if isinstance(result, dict): - merged.update(result) - if isinstance(usage, dict) and usage: - merged["token_usage"] = usage - return merged - - -def _merge_dict(old: dict[str, Any] | None, new: dict[str, Any] | None) -> dict[str, Any]: - merged = dict(old or {}) - for key, value in (new or {}).items(): - if value is not None: - merged[key] = value - return merged - - -def _force_set_json(model_obj: Any, field_name: str, value: Any) -> None: - """强制持久化 JSON / JSONB 字段。 - - SQLAlchemy 对 dict/list 的嵌套原地修改不会稳定触发 dirty 判定。 - 所有编辑类接口在写入 input_json / output_json 时统一走这里: - 1. deepcopy 断开旧引用; - 2. 整体重新赋值; - 3. flag_modified 显式标记字段已变更。 - """ - setattr(model_obj, field_name, deepcopy(value)) - flag_modified(model_obj, field_name) + return _common_unwrap_step_output(value, schema_version=STEP_IO_SCHEMA_VERSION) async def log_module_event( @@ -258,22 +238,14 @@ async def _get_project_for_user( for_update: bool = False, populate_existing: bool = False, ) -> ModuleGenerationProject: - query = select(ModuleGenerationProject).where( - ModuleGenerationProject.id == project_id, - ModuleGenerationProject.module == MODULE, - ModuleGenerationProject.deleted_at.is_(None), + return await _base_get_project_for_user( + db, + project_id=project_id, + user=user, + config=FLOW_CONFIG, + for_update=for_update, + populate_existing=populate_existing, ) - if not user.is_admin: - query = query.where(ModuleGenerationProject.user_id == user.id) - if populate_existing: - query = query.execution_options(populate_existing=True) - if for_update: - query = query.with_for_update() - result = await db.execute(query.limit(1)) - project = result.scalar_one_or_none() - if not project: - raise HTTPException(status_code=404, detail="拆镜复刻项目不存在") - return project async def _get_step_for_user( @@ -284,64 +256,26 @@ async def _get_step_for_user( user: User, for_update: bool = False, ) -> ModuleGenerationStep: - await _get_project_for_user(db, project_id=project_id, user=user, for_update=for_update) - query = select(ModuleGenerationStep).where( - ModuleGenerationStep.id == step_id, - ModuleGenerationStep.project_id == project_id, - ModuleGenerationStep.module == MODULE, - ModuleGenerationStep.deleted_at.is_(None), - ModuleGenerationStep.is_current == True, + return await _base_get_step_for_user( + db, + project_id=project_id, + step_id=step_id, + user=user, + config=FLOW_CONFIG, + for_update=for_update, ) - if not user.is_admin: - query = query.where(ModuleGenerationStep.user_id == user.id) - if for_update: - query = query.with_for_update() - result = await db.execute(query.limit(1)) - step = result.scalar_one_or_none() - if not step: - raise HTTPException(status_code=404, detail="子任务不存在") - return step async def _get_current_steps(db: AsyncSession, project_id: str) -> list[ModuleGenerationStep]: - result = await db.execute( - select(ModuleGenerationStep) - .where( - ModuleGenerationStep.project_id == project_id, - ModuleGenerationStep.module == MODULE, - ModuleGenerationStep.deleted_at.is_(None), - ModuleGenerationStep.is_current == True, - ) - .order_by(ModuleGenerationStep.step_index.asc(), ModuleGenerationStep.created_at.asc()) - ) - return list(result.scalars().all()) + return await _base_get_current_steps(db, project_id=project_id, config=FLOW_CONFIG) async def _get_current_step_by_code(db: AsyncSession, project_id: str, step_code: str) -> ModuleGenerationStep | None: - result = await db.execute( - select(ModuleGenerationStep) - .where( - ModuleGenerationStep.project_id == project_id, - ModuleGenerationStep.module == MODULE, - ModuleGenerationStep.step_code == step_code, - ModuleGenerationStep.is_current == True, - ModuleGenerationStep.deleted_at.is_(None), - ) - .order_by(ModuleGenerationStep.version.desc(), ModuleGenerationStep.created_at.desc()) - .limit(1) - ) - return result.scalar_one_or_none() + return await _base_get_current_step_by_code(db, project_id=project_id, step_code=step_code, config=FLOW_CONFIG) async def _next_version(db: AsyncSession, project_id: str, step_code: str) -> int: - result = await db.execute( - select(func.max(ModuleGenerationStep.version)).where( - ModuleGenerationStep.project_id == project_id, - ModuleGenerationStep.module == MODULE, - ModuleGenerationStep.step_code == step_code, - ) - ) - return int(result.scalar_one_or_none() or 0) + 1 + return await _base_next_version(db, project_id=project_id, step_code=step_code, config=FLOW_CONFIG) async def _create_step( @@ -356,39 +290,19 @@ async def _create_step( input_data: dict[str, Any] | None = None, output_data: dict[str, Any] | None = None, ) -> ModuleGenerationStep: - version = await _next_version(db, project.id, step_code) - step = ModuleGenerationStep( - id=generate_id(), - project_id=project.id, - user_id=project.user_id, - module=project.module, - step_index=STEP_INDEX_MAP[step_code], + return await _base_create_step( + db, + project=project, step_code=step_code, + config=FLOW_CONFIG, + log_module_event=log_module_event, status=status, - version=version, - is_current=True, parent_step_id=parent_step_id, source_step_id=source_step_id, chat_task_id=chat_task_id, - input_json=_step_input( - step_code=step_code, - payload=input_data, - source_step_id=source_step_id, - parent_step_id=parent_step_id, - ) if input_data is not None else None, - output_json=_step_output( - step_code=step_code, - status=status, - result=output_data, - ) if output_data is not None else None, - started_at=_now() if status == ModuleStepStatusEnum.PROCESSING.value else None, - completed_at=_now() if status == ModuleStepStatusEnum.COMPLETED.value else None, + input_data=input_data, + output_data=output_data, ) - db.add(step) - project.current_step_code = step_code - await db.flush() - await log_module_event(db, project=project, step=step, event_type=ModuleEventTypeEnum.STEP_CREATED.value, detail={"step_code": step_code, "version": version}) - return step async def _soft_delete_steps_from_index( @@ -398,49 +312,14 @@ async def _soft_delete_steps_from_index( start_index: int, deleted_at: datetime | None = None, ) -> None: - deleted_at = deleted_at or _now() - result = await db.execute( - select(ModuleGenerationStep) - .where( - ModuleGenerationStep.project_id == project.id, - ModuleGenerationStep.module == MODULE, - ModuleGenerationStep.is_current == True, - ModuleGenerationStep.deleted_at.is_(None), - ModuleGenerationStep.step_index >= start_index, - ) - .with_for_update() + await _base_soft_delete_steps_from_index( + db, + project=project, + start_index=start_index, + config=FLOW_CONFIG, + log_module_event=log_module_event, + deleted_at=deleted_at, ) - steps = list(result.scalars().all()) - for step in steps: - step.is_current = False - step.deleted_at = deleted_at - if step.chat_task_id: - chat_result = await db.execute( - select(ChatGenerationTask) - .where(ChatGenerationTask.id == step.chat_task_id, ChatGenerationTask.deleted_at.is_(None)) - .with_for_update() - .limit(1) - ) - chat_task = chat_result.scalar_one_or_none() - if chat_task: - if chat_task.status == "completed": - await soft_delete_chat_task_resources(db, chat_task.id, deleted_at=deleted_at) - elif chat_task.status != "failed": - await mark_chat_generation_task_failed_and_refund_once( - db, - task=chat_task, - error_message="拆镜复刻步骤被重新生成或删除,旧生成任务已取消", - pipeline_stage="failed", - ) - chat_task.deleted_at = deleted_at - if steps: - await log_module_event( - db, - project=project, - event_type=ModuleEventTypeEnum.SOFT_DELETE_STEPS.value, - message=f"软删除第 {start_index} 步及之后的旧子任务", - detail={"step_ids": [step.id for step in steps]}, - ) def _step_to_out(step: ModuleGenerationStep) -> ShotReplicateStepOut: @@ -466,17 +345,11 @@ def _step_to_out(step: ModuleGenerationStep) -> ShotReplicateStepOut: def _snapshot_from_chat(chat_task: ChatGenerationTask | None) -> dict[str, Any]: - if not chat_task: - return {} - return _parse_json(chat_task.engine_snapshot_json, {}) or {} + return _common_snapshot_from_chat(chat_task) async def _chat_tasks_by_id(db: AsyncSession, steps: list[ModuleGenerationStep]) -> dict[str, ChatGenerationTask]: - ids = [step.chat_task_id for step in steps if step.chat_task_id] - if not ids: - return {} - result = await db.execute(select(ChatGenerationTask).where(ChatGenerationTask.id.in_(ids))) - return {task.id: task for task in result.scalars().all()} + return await _base_chat_tasks_by_id(db, steps) async def project_to_detail_out(db: AsyncSession, project: ModuleGenerationProject) -> ShotReplicateTaskDetailOut: @@ -703,73 +576,15 @@ async def update_shot_replicate_step( step_id: str, req: ShotReplicateStepUpdate, ) -> tuple[ModuleGenerationProject, ModuleGenerationStep]: - project = await _get_project_for_user(db, project_id=project_id, user=current_user, for_update=True) - step = await _get_step_for_user(db, project_id=project_id, step_id=step_id, user=current_user, for_update=True) - if step.status == ModuleStepStatusEnum.PROCESSING.value: - raise HTTPException(status_code=400, detail="当前子任务正在处理中,暂不能修改") - - input_data = _step_payload(step.input_json) - output_data = _unwrap_step_output(step.output_json) - - if step.step_code == ShotReplicateStepCodeEnum.MATERIAL_INPUT.value: - input_data = _merge_dict( - input_data, - { - "material_video_url": req.material_video_url, - "material_image_url": req.material_image_url, - "source_project_name": req.source_project_name, - "target_project_name": req.target_project_name, - "core_content_point": req.core_content_point, - }, - ) - if req.target_project_name: - project.title = req.target_project_name - elif step.step_code == ShotReplicateStepCodeEnum.IMAGE_PROMPT_OPTIMIZE.value: - if req.prompt is not None: - output_data["optimized_prompt"] = req.prompt - output_data["prompt"] = req.prompt - elif step.step_code == ShotReplicateStepCodeEnum.VIDEO_PROMPT_OPTIMIZE.value: - if req.prompt_schema is not None: - output_data["prompt_schema"] = req.prompt_schema - if req.prompt is not None: - output_data["final_prompt"] = req.prompt - else: - if req.input_json: - input_data = _merge_dict(input_data, req.input_json) - if req.output_json: - output_data = _merge_dict(output_data, req.output_json) - - if req.input_json: - input_data = _merge_dict(input_data, req.input_json) - if req.output_json: - output_data = _merge_dict(output_data, req.output_json) - - _force_set_json( - step, - "input_json", - _step_input( - step_code=step.step_code, - payload=input_data, - source_step_id=step.source_step_id, - parent_step_id=step.parent_step_id, - ), + return await update_module_step( + db, + current_user=current_user, + project_id=project_id, + step_id=step_id, + req=req, + config=FLOW_CONFIG, + log_module_event=log_module_event, ) - _force_set_json( - step, - "output_json", - _step_output(step_code=step.step_code, status=ModuleStepStatusEnum.COMPLETED.value, payload=output_data), - ) - step.status = ModuleStepStatusEnum.COMPLETED.value - step.error_message = None - step.completed_at = _now() - project.status = ModuleProjectStatusEnum.WAITING_USER.value - project.current_step_code = step.step_code - project.error_message = None - - await _soft_delete_steps_from_index(db, project=project, start_index=step.step_index + 1) - await log_module_event(db, project=project, step=step, event_type=ModuleEventTypeEnum.STEP_UPDATED.value, message="用户修改子任务内容") - await db.flush() - return project, step async def update_shot_replicate_material_input( @@ -784,56 +599,14 @@ async def update_shot_replicate_material_input( 采用方案 B:软删除旧第1步及之后的当前有效步骤,然后新建第1步 version+1。 未传字段沿用旧第1步素材输入,避免前端只改一个字段时丢失其它素材信息。 """ - project = await _get_project_for_user(db, project_id=project_id, user=current_user, for_update=True) - old_material_step = await _get_current_step_by_code(db, project.id, ShotReplicateStepCodeEnum.MATERIAL_INPUT.value) - old_material = _step_payload(old_material_step.input_json if old_material_step else None) - - # 拆镜复刻的素材视频来自 shot_replicate_segments.segment_video_url, - # 与片段强绑定,不允许前端在项目素材修改接口中覆盖。 - material = { - "material_video_url": old_material.get("material_video_url"), - "material_image_url": req.material_image_url if req.material_image_url is not None else old_material.get("material_image_url"), - "source_project_name": req.source_project_name if req.source_project_name is not None else old_material.get("source_project_name"), - "target_project_name": req.target_project_name if req.target_project_name is not None else old_material.get("target_project_name"), - "core_content_point": req.core_content_point if req.core_content_point is not None else old_material.get("core_content_point"), - } - - missing_fields = [key for key, value in material.items() if value is None or str(value).strip() == ""] - if missing_fields: - raise HTTPException(status_code=400, detail=f"素材输入缺少必要字段: {', '.join(missing_fields)}") - - await _soft_delete_steps_from_index(db, project=project, start_index=STEP_INDEX_MAP[ShotReplicateStepCodeEnum.MATERIAL_INPUT.value]) - - project.title = str(material["target_project_name"]) - project.status = ModuleProjectStatusEnum.WAITING_USER.value - project.current_step_code = ShotReplicateStepCodeEnum.MATERIAL_INPUT.value - project.final_image_url = None - project.final_video_url = None - project.final_video_cover_url = None - project.error_message = None - project.completed_at = None - - new_step = await _create_step( + return await update_module_material_input( db, - project=project, - step_code=ShotReplicateStepCodeEnum.MATERIAL_INPUT.value, - status=ModuleStepStatusEnum.COMPLETED.value, - input_data=material, - output_data={"message": "素材输入已修改,旧步骤已软删除。下一步请重新生成图片AI提词。"}, + current_user=current_user, + project_id=project_id, + req=req, + config=FLOW_CONFIG, + log_module_event=log_module_event, ) - await log_module_event( - db, - project=project, - step=new_step, - event_type=ModuleEventTypeEnum.STEP_UPDATED.value, - message="用户修改素材输入并重建第1步新版本", - detail={ - "old_material_step_id": old_material_step.id if old_material_step else None, - "new_material_step_id": new_step.id, - "version": new_step.version, - }, - ) - return project.id, new_step.id async def update_shot_replicate_image_prompt( @@ -848,55 +621,15 @@ async def update_shot_replicate_image_prompt( 修改后软删除第3、4、5步当前有效任务,让用户从图片生成开始重新执行。 """ - project = await _get_project_for_user(db, project_id=project_id, user=current_user, for_update=True) - step = await _get_step_for_user(db, project_id=project_id, step_id=step_id, user=current_user, for_update=True) - if step.step_code != ShotReplicateStepCodeEnum.IMAGE_PROMPT_OPTIMIZE.value: - raise HTTPException(status_code=400, detail="只能修改第2步图片 AI 提词子任务") - if step.status != ModuleStepStatusEnum.COMPLETED.value: - raise HTTPException(status_code=400, detail="图片 AI 提词未完成,不能直接修改") - - output_data = _step_payload(step.output_json) - usage = _step_usage(step.output_json) - new_prompt = req.prompt.strip() - output_data["optimized_prompt"] = new_prompt - output_data["prompt"] = new_prompt - output_data["manual_edited"] = True - output_data["manual_edited_at"] = _now().isoformat() - - _force_set_json( - step, - "output_json", - _step_output( - step_code=ShotReplicateStepCodeEnum.IMAGE_PROMPT_OPTIMIZE.value, - status=ModuleStepStatusEnum.COMPLETED.value, - payload=output_data, - usage=usage, - ), - ) - step.status = ModuleStepStatusEnum.COMPLETED.value - step.error_message = None - step.completed_at = _now() - - await _soft_delete_steps_from_index(db, project=project, start_index=STEP_INDEX_MAP[ShotReplicateStepCodeEnum.IMAGE_GENERATE.value]) - - project.status = ModuleProjectStatusEnum.WAITING_USER.value - project.current_step_code = ShotReplicateStepCodeEnum.IMAGE_PROMPT_OPTIMIZE.value - project.final_image_url = None - project.final_video_url = None - project.final_video_cover_url = None - project.completed_at = None - project.error_message = None - - await log_module_event( + return await update_module_image_prompt( db, - project=project, - step=step, - event_type=ModuleEventTypeEnum.STEP_UPDATED.value, - message="用户直接修改图片 AI 优化提词,已软删除后续步骤", - detail={"start_deleted_step_index": STEP_INDEX_MAP[ShotReplicateStepCodeEnum.IMAGE_GENERATE.value]}, + current_user=current_user, + project_id=project_id, + step_id=step_id, + req=req, + config=FLOW_CONFIG, + log_module_event=log_module_event, ) - await db.flush() - return project, step async def update_shot_replicate_video_prompt_schema( @@ -912,84 +645,17 @@ async def update_shot_replicate_video_prompt_schema( 服务端已有 schema 为基准:视频规格、数组长度、时间段、合规控制、质量控制、协议字段均锁定。 最终提示词允许修改,但保存前会清洗视频时长、比例、分辨率、帧率等参数。 """ - project = await _get_project_for_user(db, project_id=project_id, user=current_user, for_update=True) - step = await _get_step_for_user(db, project_id=project_id, step_id=step_id, user=current_user, for_update=True) - if step.step_code != ShotReplicateStepCodeEnum.VIDEO_PROMPT_OPTIMIZE.value: - raise HTTPException(status_code=400, detail="只能修改第4步视频 AI 提词 JSON schema 子任务") - if step.status != ModuleStepStatusEnum.COMPLETED.value: - raise HTTPException(status_code=400, detail="视频 AI 提词未完成,不能直接修改") - - output_data = _step_payload(step.output_json) - usage = _step_usage(step.output_json) - input_data = _step_payload(step.input_json) - server_schema = output_data.get("prompt_schema") if isinstance(output_data.get("prompt_schema"), dict) else {} - video_config = output_data.get("params_used_for_prompt") or input_data.get("video_config") or {} - if not isinstance(video_config, dict) or not video_config.get("duration") or not video_config.get("aspect_ratio") or not video_config.get("resolution"): - raise HTTPException(status_code=400, detail="缺少第4步视频参数快照,不能安全修改视频 schema") - - patched_schema = patch_video_prompt_schema_from_client( - server_schema=server_schema, - client_schema=req.prompt_schema, - video_config=video_config, - ) - final_prompt = build_final_video_prompt(patched_schema) - - output_data["prompt_schema"] = patched_schema - output_data["final_prompt"] = final_prompt - output_data["params_used_for_prompt"] = video_config - output_data["manual_edited"] = True - output_data["manual_edited_at"] = _now().isoformat() - - _force_set_json( - step, - "output_json", - _step_output( - step_code=ShotReplicateStepCodeEnum.VIDEO_PROMPT_OPTIMIZE.value, - status=ModuleStepStatusEnum.COMPLETED.value, - payload=output_data, - usage=usage, - ), - ) - step.status = ModuleStepStatusEnum.COMPLETED.value - step.error_message = None - step.completed_at = _now() - - await _soft_delete_steps_from_index(db, project=project, start_index=STEP_INDEX_MAP[ShotReplicateStepCodeEnum.VIDEO_GENERATE.value]) - - project.status = ModuleProjectStatusEnum.WAITING_USER.value - project.current_step_code = ShotReplicateStepCodeEnum.VIDEO_PROMPT_OPTIMIZE.value - project.final_video_url = None - project.final_video_cover_url = None - project.completed_at = None - project.error_message = None - - await log_module_event( + return await update_module_video_prompt_schema( db, - project=project, - step=step, - event_type=ModuleEventTypeEnum.STEP_UPDATED.value, - message="用户修改视频 AI 提词 schema,已软删除视频生成步骤", - detail={ - "start_deleted_step_index": STEP_INDEX_MAP[ShotReplicateStepCodeEnum.VIDEO_GENERATE.value], - "locked_fields": [ - "schema_version", - "schema_usage", - "画面属性.视频时长", - "画面属性.视频比例", - "画面属性.清晰度", - "画面属性.帧率", - "画面属性.推荐分辨率", - "动作流程[*].时间段", - "镜头流程[*].时间段", - "动态时间规划", - "输出规格限制", - "质量控制", - "合规控制", - ], - }, + current_user=current_user, + project_id=project_id, + step_id=step_id, + req=req, + config=FLOW_CONFIG, + log_module_event=log_module_event, + patch_video_prompt_schema_from_client=patch_video_prompt_schema_from_client, + build_final_video_prompt=build_final_video_prompt, ) - await db.flush() - return project, step async def submit_image_prompt_optimize( @@ -1764,7 +1430,7 @@ async def create_shot_replicate_project_from_segment( module=MODULE, title=req.target_project_name, status=ModuleProjectStatusEnum.WAITING_USER.value, - current_step_code=ShotReplicateStepCodeEnum.IMAGE_PROMPT_OPTIMIZE.value, + current_step_code=ShotReplicateStepCodeEnum.MATERIAL_INPUT.value, idempotency_key=req.idempotency_key, ) db.add(project) @@ -1815,7 +1481,7 @@ async def create_shot_replicate_project_from_segment( **context["analysis"], }, ) - project.current_step_code = ShotReplicateStepCodeEnum.IMAGE_PROMPT_OPTIMIZE.value + project.current_step_code = ShotReplicateStepCodeEnum.MATERIAL_INPUT.value project.status = ModuleProjectStatusEnum.WAITING_USER.value segment.module_project_id = project.id segment.replicate_status = ShotSegmentReplicateStatusEnum.PROJECT_CREATED.value @@ -1832,15 +1498,4 @@ async def create_shot_replicate_project_from_segment( def _build_file_url_or_data_uri(file_url: str) -> str: - """构建模型可访问的素材 URL。 - - - /uploads 与 /uploads/shot_segments 属于上传/素材域,不走 generate token 签名。 - - /generate 下的 AI 生成资源继续走 token 签名。 - - 远程 URL / data URI 原样返回。 - """ - if file_url.startswith(("http://", "https://", "data:")): - return file_url - if file_url.startswith("/generate/"): - file_url = build_resource_signed_url(resource_url=file_url, expire_seconds=86400) - base_url = str(settings.BASE_URL or "").rstrip("/") - return f"{base_url}{file_url}" if base_url else file_url + return _common_build_file_url_or_data_uri(file_url) \ No newline at end of file