celery worker poll 频次机制调整| celery beat 设置poll检测任务| 拆镜复刻切片删除API | 交易流水时区BUG

This commit is contained in:
2026-07-02 09:36:38 +08:00
parent c6df04a895
commit f6c5032c7b
22 changed files with 1039 additions and 124 deletions
@@ -6,21 +6,31 @@ from typing import Any, Optional
from sqlalchemy import select
from app.config import settings
from app.enums.celery_queue import CeleryQueue
from app.enums.generation_task import (
ALLOWED_GENERATION_MODES,
ChatGenerationPipelineStage,
ChatGenerationTaskEventType,
ChatGenerationTaskStatus,
GenerationMode,
GenerationType,
)
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.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
ALLOWED_GENERATION_MODES = {"chatapi_async", "hot_opening_replicate", "shot_replicate"}
def _now() -> datetime:
return datetime.now(timezone.utc)
def _get_first_value(obj: Any, *field_names: str) -> Optional[Any]:
"""
兼容不同版本字段名,避免字段调整后 Celery 任务直接报错。
@@ -74,7 +84,7 @@ def _build_optimized_prompt_by_params(task: ChatGenerationTask) -> str:
# 爆款开头复刻第5步的视频生成,original_prompt 已经是视频提词 JSON schema。
# 不能再追加“时长/比例/分辨率”中文参数,否则会污染 schema。
if generation_mode in {"hot_opening_replicate", "shot_replicate"} and gen_type == "video":
if generation_mode in {GenerationMode.HOT_OPENING_REPLICATE.value, GenerationMode.SHOT_REPLICATE.value} and gen_type == GenerationType.VIDEO.value:
stripped = base_prompt.strip()
if stripped.startswith("{") or stripped.startswith("["):
return base_prompt
@@ -88,25 +98,25 @@ def _build_optimized_prompt_by_params(task: ChatGenerationTask) -> str:
parts = []
if gen_type == "video":
if gen_type == GenerationType.VIDEO.value:
# 时长:4秒,画面比例:16:9,分辨率:480p
if duration:
parts.append(f"时长:{duration}")
parts.append(f"画面比例:{aspect_ratio}")
parts.append(f"分辨率:{resolution}")
else:
parts.append(f"时长:4秒")
parts.append(f"画面比例:16:9")
parts.append(f"分辨率:480p")
elif gen_type == "image":
if image_size :
parts.append("时长:4秒")
parts.append("画面比例:16:9")
parts.append("分辨率:480p")
elif gen_type == GenerationType.IMAGE.value:
if image_size:
parts.append(f"分辨率:{image_size}")
parts.append(f"画布比例:{image_proportion}")
parts.append(f"像素尺寸:{image_px}")
else:
parts.append(f"分辨率:2K")
parts.append(f"画布比例:1:1")
parts.append(f"像素尺寸:2048x2048")
parts.append("分辨率:2K")
parts.append("画布比例:1:1")
parts.append("像素尺寸:2048x2048")
else:
# 未知类型时返回原始字符
return base_prompt
@@ -131,7 +141,7 @@ async def _run(task_id: str):
if not task or task.generation_mode not in ALLOWED_GENERATION_MODES:
return
if task.status != "generating":
if task.status != ChatGenerationTaskStatus.GENERATING.value:
return
deadline_at = ensure_aware_utc(task.deadline_at)
@@ -140,28 +150,37 @@ async def _run(task_id: str):
db,
task=task,
error_message="任务超时",
pipeline_stage="timeout",
pipeline_stage=ChatGenerationPipelineStage.TIMEOUT.value,
)
await db.commit()
await log_task_event(task, event_type="TASK_TIMEOUT", to_status="failed", to_stage="timeout")
await log_task_event(
task,
event_type=ChatGenerationTaskEventType.TASK_TIMEOUT.value,
to_status=ChatGenerationTaskStatus.FAILED.value,
to_stage=ChatGenerationPipelineStage.TIMEOUT.value,
)
from app.services.generation_module_hook_service import notify_chat_generation_task_finished
await notify_chat_generation_task_finished(db, task)
await db.commit()
return
if task.pipeline_stage not in ("queued", "preparing", "creating_provider_task"):
if task.pipeline_stage not in (
ChatGenerationPipelineStage.QUEUED.value,
ChatGenerationPipelineStage.PREPARING.value,
ChatGenerationPipelineStage.CREATING_PROVIDER_TASK.value,
):
return
try:
if not task.optimized_prompt:
old_stage = task.pipeline_stage
task.pipeline_stage = "preparing"
task.pipeline_stage = ChatGenerationPipelineStage.PREPARING.value
await db.commit()
await log_task_event(
task,
event_type="PROMPT_CONCAT_START",
event_type=ChatGenerationTaskEventType.PROMPT_CONCAT_START.value,
from_stage=old_stage,
to_stage="preparing",
to_stage=ChatGenerationPipelineStage.PREPARING.value,
message="开始本地拼接提示词,不调用提词优化API",
)
@@ -174,8 +193,8 @@ async def _run(task_id: str):
await db.commit()
await log_task_event(
task,
event_type="PROMPT_CONCAT_SUCCESS",
to_stage="preparing",
event_type=ChatGenerationTaskEventType.PROMPT_CONCAT_SUCCESS.value,
to_stage=ChatGenerationPipelineStage.PREPARING.value,
detail={
"optimized_prompt": optimized_prompt,
"gen_type": task.gen_type,
@@ -184,11 +203,14 @@ async def _run(task_id: str):
)
if task.seedance_task_id or task.provider_task_id:
task.pipeline_stage = "waiting_remote"
task.pipeline_stage = ChatGenerationPipelineStage.WAITING_REMOTE.value
if task.gen_type == GenerationType.VIDEO.value:
ensure_video_poll_fields(task, now=_now())
task.next_poll_at = _now()
await db.commit()
elif task.remote_result_url:
task.pipeline_stage = "result_ready"
task.pipeline_stage = ChatGenerationPipelineStage.RESULT_READY.value
await db.commit()
from app.tasks.generation_download_tasks import enqueue_download_task
@@ -198,13 +220,13 @@ async def _run(task_id: str):
else:
old_stage = task.pipeline_stage
task.pipeline_stage = "creating_provider_task"
task.pipeline_stage = ChatGenerationPipelineStage.CREATING_PROVIDER_TASK.value
await db.commit()
await log_task_event(
task,
event_type="PROVIDER_CREATE_START",
event_type=ChatGenerationTaskEventType.PROVIDER_CREATE_START.value,
from_stage=old_stage,
to_stage="creating_provider_task",
to_stage=ChatGenerationPipelineStage.CREATING_PROVIDER_TASK.value,
)
created = await create_provider_task(db, task)
@@ -216,7 +238,7 @@ async def _run(task_id: str):
task.remote_result_url = created.get("remote_result_url") or task.remote_result_url
if task.gen_type == "image":
if task.gen_type == GenerationType.IMAGE.value:
task.image_tokens_used = created.get("image_tokens", task.image_tokens_used or 0) or 0
task.provider_response_json = json.dumps(
@@ -228,21 +250,26 @@ async def _run(task_id: str):
if task.remote_result_url and not task.seedance_task_id:
# 同步图片路径:原 SDK 已经返回最终 URL。
task.pipeline_stage = "result_ready"
task.pipeline_stage = ChatGenerationPipelineStage.RESULT_READY.value
else:
# 视频路径:provider 返回 task id,后续轮询。
task.pipeline_stage = "waiting_remote"
task.pipeline_stage = ChatGenerationPipelineStage.WAITING_REMOTE.value
if task.gen_type == GenerationType.VIDEO.value:
current_time = _now()
ensure_video_poll_fields(task, now=current_time)
task.next_poll_at = current_time
task.poll_interval_seconds = int(task.poll_interval_seconds or 0)
task.status = "generating"
task.status = ChatGenerationTaskStatus.GENERATING.value
await db.commit()
await log_task_event(
task,
event_type="PROVIDER_CREATE_SUCCESS",
event_type=ChatGenerationTaskEventType.PROVIDER_CREATE_SUCCESS.value,
to_stage=task.pipeline_stage,
detail=created,
)
if task.pipeline_stage == "result_ready":
if task.pipeline_stage == ChatGenerationPipelineStage.RESULT_READY.value:
from app.tasks.generation_download_tasks import enqueue_download_task
await enqueue_download_task(db, task, reason="create_result_ready")
@@ -253,10 +280,11 @@ async def _run(task_id: str):
task,
reason="create_provider_success",
check_at=_now() + timedelta(seconds=int(settings.POLL_TASK_LEASE_SECONDS or 300)),
next_poll_at=getattr(task, "next_poll_at", None),
)
poll_generation_task.apply_async(
args=[task.id],
queue="gen_provider_poll",
queue=CeleryQueue.GEN_PROVIDER_POLL.value,
countdown=0,
)
@@ -278,10 +306,10 @@ async def _run(task_id: str):
db,
task=task,
error_message=error_message,
pipeline_stage="failed",
pipeline_stage=ChatGenerationPipelineStage.FAILED.value,
)
await db.commit()
await log_task_event(task, event_type="TASK_FAILED", message=task.error_message)
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 notify_chat_generation_task_finished(db, task)
await db.commit()
@@ -305,4 +333,4 @@ else:
def apply_async(self, *args, **kwargs):
raise RuntimeError("Celery is disabled")
chatapi_create_generation_task = _DisabledTask()
chatapi_create_generation_task = _DisabledTask()