项目/AI生成链路合并
This commit is contained in:
@@ -1,5 +1,6 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
from datetime import datetime
|
||||
from typing import Any
|
||||
@@ -42,6 +43,11 @@ from app.services.generation.ai.engine_service import (
|
||||
)
|
||||
from app.services.generation.billing_service import charge_module_prompt_usage
|
||||
from app.services.generation.refund_service import mark_chat_generation_task_failed_and_refund_once
|
||||
from app.services.generation.pipeline.db_lock_service import (
|
||||
DatabaseRowLockBusy,
|
||||
apply_short_lock_timeout,
|
||||
execute_with_lock_timeout,
|
||||
)
|
||||
from app.services.generation.task_factory_service import create_chat_generation_task_for_module
|
||||
from app.services.hot_opening_video_prompt_service import (
|
||||
build_final_video_prompt,
|
||||
@@ -338,6 +344,18 @@ async def _soft_delete_steps_from_index(
|
||||
refund_unfinished: bool = False,
|
||||
release_stats: dict[str, int] | None = None,
|
||||
) -> None:
|
||||
processing_result = await db.execute(
|
||||
select(ModuleGenerationStep.id).where(
|
||||
ModuleGenerationStep.project_id == project.id,
|
||||
ModuleGenerationStep.module == MODULE,
|
||||
ModuleGenerationStep.deleted_at.is_(None),
|
||||
ModuleGenerationStep.is_current == True,
|
||||
ModuleGenerationStep.status == ModuleStepStatusEnum.PROCESSING.value,
|
||||
ModuleGenerationStep.step_index >= start_index,
|
||||
).limit(1)
|
||||
)
|
||||
if processing_result.scalar_one_or_none() is not None:
|
||||
raise HTTPException(status_code=409, detail="当前步骤正在处理中,请等待完成后再操作")
|
||||
await _base_soft_delete_steps_from_index(
|
||||
db,
|
||||
project=project,
|
||||
@@ -693,6 +711,69 @@ async def update_shot_replicate_video_prompt_schema(
|
||||
)
|
||||
|
||||
|
||||
async def _reload_prompt_context_for_update(
|
||||
db: AsyncSession,
|
||||
*,
|
||||
project_id: str,
|
||||
step_id: str,
|
||||
step_code: str,
|
||||
) -> tuple[ModuleGenerationProject | None, ModuleGenerationStep | None]:
|
||||
last_error: DatabaseRowLockBusy | None = None
|
||||
for retry_index in range(3):
|
||||
try:
|
||||
project_result = await execute_with_lock_timeout(
|
||||
db,
|
||||
select(ModuleGenerationProject)
|
||||
.where(
|
||||
ModuleGenerationProject.id == project_id,
|
||||
ModuleGenerationProject.module == MODULE,
|
||||
ModuleGenerationProject.deleted_at.is_(None),
|
||||
)
|
||||
.with_for_update()
|
||||
.execution_options(populate_existing=True)
|
||||
.limit(1)
|
||||
)
|
||||
project = project_result.scalar_one_or_none()
|
||||
if project is None:
|
||||
return None, None
|
||||
step_result = await execute_with_lock_timeout(
|
||||
db,
|
||||
select(ModuleGenerationStep)
|
||||
.where(
|
||||
ModuleGenerationStep.id == step_id,
|
||||
ModuleGenerationStep.project_id == project_id,
|
||||
ModuleGenerationStep.module == MODULE,
|
||||
ModuleGenerationStep.step_code == step_code,
|
||||
ModuleGenerationStep.deleted_at.is_(None),
|
||||
ModuleGenerationStep.is_current == True,
|
||||
)
|
||||
.with_for_update()
|
||||
.execution_options(populate_existing=True)
|
||||
.limit(1)
|
||||
)
|
||||
return project, step_result.scalar_one_or_none()
|
||||
except DatabaseRowLockBusy as exc:
|
||||
last_error = exc
|
||||
await db.rollback()
|
||||
if retry_index < 2:
|
||||
await asyncio.sleep(1 + retry_index)
|
||||
raise last_error or DatabaseRowLockBusy()
|
||||
|
||||
|
||||
def _prompt_context_matches(
|
||||
step: ModuleGenerationStep | None,
|
||||
*,
|
||||
expected_version: int,
|
||||
expected_input_json: str,
|
||||
) -> bool:
|
||||
if step is None or step.status != ModuleStepStatusEnum.PROCESSING.value:
|
||||
return False
|
||||
if int(step.version or 1) != int(expected_version):
|
||||
return False
|
||||
current_input = json.dumps(step.input_json, ensure_ascii=False, sort_keys=True, default=str)
|
||||
return current_input == expected_input_json
|
||||
|
||||
|
||||
async def submit_image_prompt_optimize(
|
||||
db: AsyncSession,
|
||||
*,
|
||||
@@ -732,6 +813,7 @@ async def submit_image_prompt_optimize(
|
||||
|
||||
|
||||
async def run_image_prompt_optimize(db: AsyncSession, *, project_id: str, step_id: str | None = None) -> ModuleGenerationStep | None:
|
||||
await apply_short_lock_timeout(db)
|
||||
project_result = await db.execute(
|
||||
select(ModuleGenerationProject)
|
||||
.where(ModuleGenerationProject.id == project_id, ModuleGenerationProject.module == MODULE, ModuleGenerationProject.deleted_at.is_(None))
|
||||
@@ -749,6 +831,7 @@ async def run_image_prompt_optimize(db: AsyncSession, *, project_id: str, step_i
|
||||
return None
|
||||
|
||||
if step_id:
|
||||
await apply_short_lock_timeout(db)
|
||||
result = await db.execute(
|
||||
select(ModuleGenerationStep)
|
||||
.where(
|
||||
@@ -795,24 +878,45 @@ async def run_image_prompt_optimize(db: AsyncSession, *, project_id: str, step_i
|
||||
{"type": "video", "url": material.get("material_video_url"), "name": "参考素材视频"},
|
||||
{"type": "image", "url": material.get("material_image_url"), "name": "新产品图片"},
|
||||
]
|
||||
project_id_value = str(project.id)
|
||||
step_id_value = str(step.id)
|
||||
user_id_value = str(project.user_id)
|
||||
module_value = str(project.module)
|
||||
expected_step_version = int(step.version or 1)
|
||||
expected_input_json = json.dumps(step.input_json, ensure_ascii=False, sort_keys=True, default=str)
|
||||
await db.commit()
|
||||
|
||||
try:
|
||||
request_log = {"original_prompt": prompt_text, "references": references, "gen_type": "image"}
|
||||
log_module_prompt_event(
|
||||
event_type="module_prompt_request",
|
||||
project_id=project.id,
|
||||
step_id=step.id,
|
||||
user_id=project.user_id,
|
||||
module=project.module,
|
||||
project_id=project_id_value,
|
||||
step_id=step_id_value,
|
||||
user_id=user_id_value,
|
||||
module=module_value,
|
||||
prompt_type=ModulePromptTypeEnum.IMAGE_PROMPT.value,
|
||||
request=request_log,
|
||||
)
|
||||
optimized, token_usage = await optimize_prompt(
|
||||
db,
|
||||
original_prompt=prompt_text,
|
||||
user_id=project.user_id,
|
||||
user_id=user_id_value,
|
||||
references=references,
|
||||
gen_type="image",
|
||||
)
|
||||
project, step = await _reload_prompt_context_for_update(
|
||||
db,
|
||||
project_id=project_id_value,
|
||||
step_id=step_id_value,
|
||||
step_code=ShotReplicateStepCodeEnum.IMAGE_PROMPT_OPTIMIZE.value,
|
||||
)
|
||||
if not _prompt_context_matches(
|
||||
step,
|
||||
expected_version=expected_step_version,
|
||||
expected_input_json=expected_input_json,
|
||||
):
|
||||
await db.rollback()
|
||||
return None
|
||||
billing = await charge_module_prompt_usage(
|
||||
db,
|
||||
user_id=project.user_id,
|
||||
@@ -857,7 +961,25 @@ async def run_image_prompt_optimize(db: AsyncSession, *, project_id: str, step_i
|
||||
token_usage=usage,
|
||||
)
|
||||
await log_module_event(db, project=project, step=step, event_type=ModuleEventTypeEnum.IMAGE_PROMPT_SUCCESS.value, message="图片 AI 提词生成成功")
|
||||
await db.commit()
|
||||
except DatabaseRowLockBusy:
|
||||
await db.rollback()
|
||||
raise
|
||||
except Exception as exc:
|
||||
await db.rollback()
|
||||
project, step = await _reload_prompt_context_for_update(
|
||||
db,
|
||||
project_id=project_id_value,
|
||||
step_id=step_id_value,
|
||||
step_code=ShotReplicateStepCodeEnum.IMAGE_PROMPT_OPTIMIZE.value,
|
||||
)
|
||||
if not _prompt_context_matches(
|
||||
step,
|
||||
expected_version=expected_step_version,
|
||||
expected_input_json=expected_input_json,
|
||||
):
|
||||
await db.rollback()
|
||||
return None
|
||||
step.status = ModuleStepStatusEnum.FAILED.value
|
||||
step.error_message = str(exc)
|
||||
step.completed_at = _now()
|
||||
@@ -875,6 +997,7 @@ async def run_image_prompt_optimize(db: AsyncSession, *, project_id: str, step_i
|
||||
)
|
||||
_log_project_error(project=project, step=step, event_type="IMAGE_PROMPT_FAILED", message=project.error_message, exc=exc)
|
||||
await log_module_event(db, project=project, step=step, event_type=ModuleEventTypeEnum.IMAGE_PROMPT_FAILED.value, message=project.error_message)
|
||||
await db.commit()
|
||||
return step
|
||||
|
||||
|
||||
@@ -1053,6 +1176,7 @@ async def submit_video_prompt_optimize(
|
||||
|
||||
|
||||
async def run_video_prompt_optimize(db: AsyncSession, *, project_id: str, step_id: str | None = None) -> ModuleGenerationStep | None:
|
||||
await apply_short_lock_timeout(db)
|
||||
project_result = await db.execute(
|
||||
select(ModuleGenerationProject)
|
||||
.where(ModuleGenerationProject.id == project_id, ModuleGenerationProject.module == MODULE, ModuleGenerationProject.deleted_at.is_(None))
|
||||
@@ -1072,6 +1196,7 @@ async def run_video_prompt_optimize(db: AsyncSession, *, project_id: str, step_i
|
||||
return None
|
||||
|
||||
if step_id:
|
||||
await apply_short_lock_timeout(db)
|
||||
result = await db.execute(
|
||||
select(ModuleGenerationStep)
|
||||
.where(
|
||||
@@ -1111,8 +1236,17 @@ async def run_video_prompt_optimize(db: AsyncSession, *, project_id: str, step_i
|
||||
step.error_message = "缺少新项目图片结果,不能生成视频提词"
|
||||
project.status = ModuleProjectStatusEnum.FAILED.value
|
||||
project.error_message = step.error_message
|
||||
await db.commit()
|
||||
return step
|
||||
|
||||
project_id_value = str(project.id)
|
||||
step_id_value = str(step.id)
|
||||
user_id_value = str(project.user_id)
|
||||
module_value = str(project.module)
|
||||
expected_step_version = int(step.version or 1)
|
||||
expected_input_json = json.dumps(step.input_json, ensure_ascii=False, sort_keys=True, default=str)
|
||||
await db.commit()
|
||||
|
||||
try:
|
||||
request_log = {
|
||||
"source_project_name": material.get("source_project_name") or "无",
|
||||
@@ -1125,10 +1259,10 @@ async def run_video_prompt_optimize(db: AsyncSession, *, project_id: str, step_i
|
||||
}
|
||||
log_module_prompt_event(
|
||||
event_type="module_prompt_request",
|
||||
project_id=project.id,
|
||||
step_id=step.id,
|
||||
user_id=project.user_id,
|
||||
module=project.module,
|
||||
project_id=project_id_value,
|
||||
step_id=step_id_value,
|
||||
user_id=user_id_value,
|
||||
module=module_value,
|
||||
prompt_type=ModulePromptTypeEnum.VIDEO_PROMPT.value,
|
||||
request=request_log,
|
||||
)
|
||||
@@ -1136,7 +1270,7 @@ async def run_video_prompt_optimize(db: AsyncSession, *, project_id: str, step_i
|
||||
request_log["schema_config_source"] = schema_config_snapshot.get("source")
|
||||
prompt_schema, final_prompt, token_usage = await optimize_shot_replicate_video_prompt(
|
||||
db,
|
||||
user_id=project.user_id,
|
||||
user_id=user_id_value,
|
||||
source_project_name=request_log["source_project_name"],
|
||||
target_project_name=request_log["target_project_name"],
|
||||
core_content_point=request_log["core_content_point"],
|
||||
@@ -1145,11 +1279,24 @@ async def run_video_prompt_optimize(db: AsyncSession, *, project_id: str, step_i
|
||||
video_config=video_config,
|
||||
target_platform=target_platform,
|
||||
schema_config_snapshot=schema_config_snapshot,
|
||||
module=project.module,
|
||||
project_id=project.id,
|
||||
step_id=step.id,
|
||||
trace_id=f"shot-video-prompt:{step.id}",
|
||||
module=module_value,
|
||||
project_id=project_id_value,
|
||||
step_id=step_id_value,
|
||||
trace_id=f"shot-video-prompt:{step_id_value}",
|
||||
)
|
||||
project, step = await _reload_prompt_context_for_update(
|
||||
db,
|
||||
project_id=project_id_value,
|
||||
step_id=step_id_value,
|
||||
step_code=ShotReplicateStepCodeEnum.VIDEO_PROMPT_OPTIMIZE.value,
|
||||
)
|
||||
if not _prompt_context_matches(
|
||||
step,
|
||||
expected_version=expected_step_version,
|
||||
expected_input_json=expected_input_json,
|
||||
):
|
||||
await db.rollback()
|
||||
return None
|
||||
billing = await charge_module_prompt_usage(
|
||||
db,
|
||||
user_id=project.user_id,
|
||||
@@ -1197,7 +1344,25 @@ async def run_video_prompt_optimize(db: AsyncSession, *, project_id: str, step_i
|
||||
token_usage=usage,
|
||||
)
|
||||
await log_module_event(db, project=project, step=step, event_type=ModuleEventTypeEnum.VIDEO_PROMPT_SUCCESS.value, message="视频 AI 提词生成成功")
|
||||
await db.commit()
|
||||
except DatabaseRowLockBusy:
|
||||
await db.rollback()
|
||||
raise
|
||||
except Exception as exc:
|
||||
await db.rollback()
|
||||
project, step = await _reload_prompt_context_for_update(
|
||||
db,
|
||||
project_id=project_id_value,
|
||||
step_id=step_id_value,
|
||||
step_code=ShotReplicateStepCodeEnum.VIDEO_PROMPT_OPTIMIZE.value,
|
||||
)
|
||||
if not _prompt_context_matches(
|
||||
step,
|
||||
expected_version=expected_step_version,
|
||||
expected_input_json=expected_input_json,
|
||||
):
|
||||
await db.rollback()
|
||||
return None
|
||||
step.status = ModuleStepStatusEnum.FAILED.value
|
||||
step.error_message = str(exc)
|
||||
step.completed_at = _now()
|
||||
@@ -1215,6 +1380,7 @@ async def run_video_prompt_optimize(db: AsyncSession, *, project_id: str, step_i
|
||||
)
|
||||
_log_project_error(project=project, step=step, event_type="VIDEO_PROMPT_FAILED", message=project.error_message, exc=exc)
|
||||
await log_module_event(db, project=project, step=step, event_type=ModuleEventTypeEnum.VIDEO_PROMPT_FAILED.value, message=project.error_message)
|
||||
await db.commit()
|
||||
return step
|
||||
|
||||
|
||||
@@ -1331,9 +1497,37 @@ 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
|
||||
result = await db.execute(
|
||||
meta_result = await db.execute(
|
||||
select(ModuleGenerationStep.id, ModuleGenerationStep.project_id).where(
|
||||
ModuleGenerationStep.chat_task_id == task.id,
|
||||
ModuleGenerationStep.module == MODULE,
|
||||
ModuleGenerationStep.is_current == True,
|
||||
ModuleGenerationStep.deleted_at.is_(None),
|
||||
).limit(1)
|
||||
)
|
||||
meta = meta_result.first()
|
||||
if not meta:
|
||||
return
|
||||
step_id_value, project_id_value = str(meta.id), str(meta.project_id)
|
||||
project_result = await execute_with_lock_timeout(
|
||||
db,
|
||||
select(ModuleGenerationProject)
|
||||
.where(
|
||||
ModuleGenerationProject.id == project_id_value,
|
||||
ModuleGenerationProject.deleted_at.is_(None),
|
||||
)
|
||||
.with_for_update()
|
||||
.limit(1)
|
||||
)
|
||||
project = project_result.scalar_one_or_none()
|
||||
if not project:
|
||||
return
|
||||
step_result = await execute_with_lock_timeout(
|
||||
db,
|
||||
select(ModuleGenerationStep)
|
||||
.where(
|
||||
ModuleGenerationStep.id == step_id_value,
|
||||
ModuleGenerationStep.project_id == project_id_value,
|
||||
ModuleGenerationStep.chat_task_id == task.id,
|
||||
ModuleGenerationStep.module == MODULE,
|
||||
ModuleGenerationStep.is_current == True,
|
||||
@@ -1342,18 +1536,9 @@ async def handle_chat_generation_task_completed(db: AsyncSession, task: ChatGene
|
||||
.with_for_update()
|
||||
.limit(1)
|
||||
)
|
||||
step = result.scalar_one_or_none()
|
||||
step = step_result.scalar_one_or_none()
|
||||
if not step:
|
||||
return
|
||||
project_result = await db.execute(
|
||||
select(ModuleGenerationProject)
|
||||
.where(ModuleGenerationProject.id == step.project_id, ModuleGenerationProject.deleted_at.is_(None))
|
||||
.with_for_update()
|
||||
.limit(1)
|
||||
)
|
||||
project = project_result.scalar_one_or_none()
|
||||
if not project:
|
||||
return
|
||||
|
||||
if step.step_code == ShotReplicateStepCodeEnum.IMAGE_GENERATE.value:
|
||||
step.status = ModuleStepStatusEnum.COMPLETED.value
|
||||
@@ -1394,19 +1579,48 @@ 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
|
||||
result = await db.execute(
|
||||
select(ModuleGenerationStep)
|
||||
.where(ModuleGenerationStep.chat_task_id == task.id, ModuleGenerationStep.module == MODULE, ModuleGenerationStep.is_current == True, ModuleGenerationStep.deleted_at.is_(None))
|
||||
meta_result = await db.execute(
|
||||
select(ModuleGenerationStep.id, ModuleGenerationStep.project_id).where(
|
||||
ModuleGenerationStep.chat_task_id == task.id,
|
||||
ModuleGenerationStep.module == MODULE,
|
||||
ModuleGenerationStep.is_current == True,
|
||||
ModuleGenerationStep.deleted_at.is_(None),
|
||||
).limit(1)
|
||||
)
|
||||
meta = meta_result.first()
|
||||
if not meta:
|
||||
return
|
||||
step_id_value, project_id_value = str(meta.id), str(meta.project_id)
|
||||
project_result = await execute_with_lock_timeout(
|
||||
db,
|
||||
select(ModuleGenerationProject)
|
||||
.where(
|
||||
ModuleGenerationProject.id == project_id_value,
|
||||
ModuleGenerationProject.deleted_at.is_(None),
|
||||
)
|
||||
.with_for_update()
|
||||
.limit(1)
|
||||
)
|
||||
step = result.scalar_one_or_none()
|
||||
if not step:
|
||||
return
|
||||
project_result = await db.execute(select(ModuleGenerationProject).where(ModuleGenerationProject.id == step.project_id).with_for_update().limit(1))
|
||||
project = project_result.scalar_one_or_none()
|
||||
if not project:
|
||||
return
|
||||
step_result = await execute_with_lock_timeout(
|
||||
db,
|
||||
select(ModuleGenerationStep)
|
||||
.where(
|
||||
ModuleGenerationStep.id == step_id_value,
|
||||
ModuleGenerationStep.project_id == project_id_value,
|
||||
ModuleGenerationStep.chat_task_id == task.id,
|
||||
ModuleGenerationStep.module == MODULE,
|
||||
ModuleGenerationStep.is_current == True,
|
||||
ModuleGenerationStep.deleted_at.is_(None),
|
||||
)
|
||||
.with_for_update()
|
||||
.limit(1)
|
||||
)
|
||||
step = step_result.scalar_one_or_none()
|
||||
if not step:
|
||||
return
|
||||
step.status = ModuleStepStatusEnum.FAILED.value
|
||||
step.error_message = task.error_message
|
||||
step.completed_at = _now()
|
||||
@@ -1439,7 +1653,8 @@ async def mark_shot_replicate_step_dispatch_failed(
|
||||
error_message: str,
|
||||
) -> None:
|
||||
project = await _get_project_for_user(db, project_id=project_id, user=current_user, for_update=True)
|
||||
result = await db.execute(
|
||||
result = await execute_with_lock_timeout(
|
||||
db,
|
||||
select(ModuleGenerationStep)
|
||||
.where(
|
||||
ModuleGenerationStep.id == step_id,
|
||||
|
||||
Reference in New Issue
Block a user