1
This commit is contained in:
@@ -18,10 +18,10 @@ from app.enums.generation_task import (
|
||||
from app.models.base import async_session
|
||||
from app.models.chat_generation_task import ChatGenerationTask
|
||||
from app.services.error_codes import extract_error_message
|
||||
from app.services.generation_log_service import log_task_event
|
||||
from app.services.generation_poll_schedule_service import ensure_video_poll_fields
|
||||
from app.services.generation_refund_service import mark_chat_generation_task_failed_and_refund_once
|
||||
from app.services.generation_provider_service import create_provider_task
|
||||
from app.services.generation.log_service import log_task_event
|
||||
from app.services.generation.poll_schedule_service import ensure_video_poll_fields
|
||||
from app.services.generation.refund_service import mark_chat_generation_task_failed_and_refund_once
|
||||
from app.services.generation.provider_service import create_provider_task
|
||||
from app.services.media_token_usage_snapshot_service import sync_chat_generation_task_media_token_snapshot
|
||||
from app.services.redis_registry_service import ensure_aware_utc
|
||||
from app.tasks.celery_app import celery_app
|
||||
@@ -138,14 +138,22 @@ async def _run(task_id: str):
|
||||
).with_for_update().limit(1))
|
||||
task = result.scalar_one_or_none()
|
||||
|
||||
if not task or task.generation_mode not in ALLOWED_GENERATION_MODES:
|
||||
is_image_main = bool(
|
||||
task
|
||||
and task.generation_mode == GenerationMode.CHATAPI_MAIN.value
|
||||
and task.gen_type == GenerationType.IMAGE.value
|
||||
and int(task.generation_count or 1) > 1
|
||||
)
|
||||
if not task or (task.generation_mode not in ALLOWED_GENERATION_MODES and not is_image_main):
|
||||
return
|
||||
|
||||
if task.status != ChatGenerationTaskStatus.GENERATING.value:
|
||||
return
|
||||
|
||||
deadline_at = ensure_aware_utc(task.deadline_at)
|
||||
if deadline_at and datetime.now(timezone.utc) > deadline_at:
|
||||
# 图片 main 的 deadline 与 provider claim 由 image_batch_service 原子处理,
|
||||
# 避免重复 Celery 消息在有效租约期间把正在执行的批次错误退款。
|
||||
if not is_image_main and deadline_at and datetime.now(timezone.utc) > deadline_at:
|
||||
await mark_chat_generation_task_failed_and_refund_once(
|
||||
db,
|
||||
task=task,
|
||||
@@ -159,8 +167,10 @@ async def _run(task_id: str):
|
||||
to_status=ChatGenerationTaskStatus.FAILED.value,
|
||||
to_stage=ChatGenerationPipelineStage.TIMEOUT.value,
|
||||
)
|
||||
from app.services.generation_module_hook_service import notify_chat_generation_task_finished
|
||||
from app.services.generation.module_hook_service import notify_chat_generation_task_finished
|
||||
from app.services.generation.ai.task_group_service import aggregate_parent_for_child
|
||||
await notify_chat_generation_task_finished(db, task)
|
||||
await aggregate_parent_for_child(db, task)
|
||||
await db.commit()
|
||||
return
|
||||
|
||||
@@ -202,6 +212,12 @@ async def _run(task_id: str):
|
||||
},
|
||||
)
|
||||
|
||||
if is_image_main:
|
||||
from app.services.generation.ai.image_batch_service import run_image_main_batch
|
||||
|
||||
await run_image_main_batch(db, task)
|
||||
return
|
||||
|
||||
if task.seedance_task_id or task.provider_task_id:
|
||||
task.pipeline_stage = ChatGenerationPipelineStage.WAITING_REMOTE.value
|
||||
if task.gen_type == GenerationType.VIDEO.value:
|
||||
@@ -302,16 +318,41 @@ async def _run(task_id: str):
|
||||
|
||||
if task:
|
||||
error_message = extract_error_message(exc, "生成任务") if callable(extract_error_message) else str(exc)
|
||||
await mark_chat_generation_task_failed_and_refund_once(
|
||||
db,
|
||||
task=task,
|
||||
error_message=error_message,
|
||||
pipeline_stage=ChatGenerationPipelineStage.FAILED.value,
|
||||
)
|
||||
if is_image_main:
|
||||
# image_batch_service 负责供应商/拆分失败退款。若 child 已落库,
|
||||
# 顶层兜底绝不能再把 main 退款。
|
||||
child_result = await db.execute(
|
||||
select(ChatGenerationTask.id).where(
|
||||
ChatGenerationTask.parent_task_id == task.id,
|
||||
ChatGenerationTask.generation_mode == GenerationMode.CHATAPI_CHILD.value,
|
||||
).limit(1)
|
||||
)
|
||||
has_children = child_result.scalar_one_or_none() is not None
|
||||
if not has_children:
|
||||
task.provider_create_claim_token = None
|
||||
task.provider_create_lease_until = None
|
||||
await mark_chat_generation_task_failed_and_refund_once(
|
||||
db,
|
||||
task=task,
|
||||
error_message=error_message,
|
||||
pipeline_stage=ChatGenerationPipelineStage.FAILED.value,
|
||||
)
|
||||
else:
|
||||
from app.services.generation.ai.task_group_service import aggregate_main_task_status
|
||||
await aggregate_main_task_status(db, parent_task_id=str(task.id))
|
||||
else:
|
||||
await mark_chat_generation_task_failed_and_refund_once(
|
||||
db,
|
||||
task=task,
|
||||
error_message=error_message,
|
||||
pipeline_stage=ChatGenerationPipelineStage.FAILED.value,
|
||||
)
|
||||
await db.commit()
|
||||
await log_task_event(task, event_type=ChatGenerationTaskEventType.TASK_FAILED.value, message=task.error_message)
|
||||
from app.services.generation_module_hook_service import notify_chat_generation_task_finished
|
||||
await log_task_event(task, event_type=ChatGenerationTaskEventType.TASK_FAILED.value, message=error_message)
|
||||
from app.services.generation.module_hook_service import notify_chat_generation_task_finished
|
||||
from app.services.generation.ai.task_group_service import aggregate_parent_for_child
|
||||
await notify_chat_generation_task_finished(db, task)
|
||||
await aggregate_parent_for_child(db, task)
|
||||
await db.commit()
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user