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
+39 -21
View File
@@ -4,6 +4,7 @@ from celery import Celery
from celery.signals import worker_process_init, worker_process_shutdown, worker_ready
from app.config import settings
from app.enums.celery_queue import CeleryQueue, CeleryTaskName
from app.models.base import engine
from app.tasks.async_runner import close_loop, run_async
@@ -26,7 +27,7 @@ CELERY_TASK_IMPORTS = (
)
RECOVERY_QUEUE = settings.CELERY_RECOVERY_QUEUE or "gen_recovery"
RECOVERY_QUEUE = settings.CELERY_RECOVERY_QUEUE or CeleryQueue.GEN_RECOVERY.value
def _derive_redis_db(url: str, db_no: int) -> str:
@@ -39,6 +40,21 @@ def _derive_redis_db(url: str, db_no: int) -> str:
return url.rstrip("/") + f"/{db_no}"
def _beat_schedule() -> dict:
if not bool(getattr(settings, "POLL_DUE_DISPATCH_ENABLED", True)):
return {}
return {
"dispatch-due-poll-tasks-every-minute": {
"task": CeleryTaskName.DISPATCH_DUE_POLL.value,
"schedule": max(1, int(settings.POLL_DUE_DISPATCH_INTERVAL_SECONDS or 60)),
"options": {
"queue": RECOVERY_QUEUE,
"priority": settings.DOWNLOAD_TASK_PRIORITY_RECOVER,
},
}
}
broker_url = settings.CELERY_BROKER_URL or (_derive_redis_db(settings.REDIS_URL, 1) if settings.REDIS_URL else "")
backend_url = settings.CELERY_RESULT_BACKEND or (_derive_redis_db(settings.REDIS_URL, 2) if settings.REDIS_URL else "")
@@ -58,6 +74,7 @@ if broker_url:
task_acks_late=True,
task_reject_on_worker_lost=True,
task_track_started=True,
beat_schedule=_beat_schedule(),
task_annotations={
# 生成链路任务以数据库状态为准,不依赖 Celery result backend。
# 这里忽略结果可避免任务误返回 ORM / 非 JSON 对象时触发结果序列化失败。
@@ -80,24 +97,25 @@ if broker_url:
"sep": ":",
},
task_routes={
"generation.chatapi_create_generation_task": {"queue": "gen_chatapi_create"},
"generation.poll_generation_task": {"queue": "gen_provider_poll"},
"generation.download_generation_result_task": {"queue": "gen_result_download"},
"hot_opening.start_image_prompt_optimize": {"queue": "gen_chatapi_create"},
"hot_opening.start_video_prompt_optimize": {"queue": "gen_chatapi_create"},
"shot_replicate.analyze_original_video": {"queue": "gen_chatapi_create"},
"shot_replicate.analyze_custom_segment_video": {"queue": "gen_chatapi_create"},
"shot_replicate.split_one_segment": {"queue": "gen_result_download"},
"shot_replicate.start_image_prompt_optimize": {"queue": "gen_chatapi_create"},
"shot_replicate.start_video_prompt_optimize": {"queue": "gen_chatapi_create"},
CeleryTaskName.CHATAPI_CREATE.value: {"queue": CeleryQueue.GEN_CHATAPI_CREATE.value},
CeleryTaskName.POLL_GENERATION.value: {"queue": CeleryQueue.GEN_PROVIDER_POLL.value},
CeleryTaskName.DOWNLOAD_GENERATION_RESULT.value: {"queue": CeleryQueue.GEN_RESULT_DOWNLOAD.value},
CeleryTaskName.DISPATCH_DUE_POLL.value: {"queue": RECOVERY_QUEUE},
"hot_opening.start_image_prompt_optimize": {"queue": CeleryQueue.GEN_CHATAPI_CREATE.value},
"hot_opening.start_video_prompt_optimize": {"queue": CeleryQueue.GEN_CHATAPI_CREATE.value},
"shot_replicate.analyze_original_video": {"queue": CeleryQueue.GEN_CHATAPI_CREATE.value},
"shot_replicate.analyze_custom_segment_video": {"queue": CeleryQueue.GEN_CHATAPI_CREATE.value},
"shot_replicate.split_one_segment": {"queue": CeleryQueue.GEN_RESULT_DOWNLOAD.value},
"shot_replicate.start_image_prompt_optimize": {"queue": CeleryQueue.GEN_CHATAPI_CREATE.value},
"shot_replicate.start_video_prompt_optimize": {"queue": CeleryQueue.GEN_CHATAPI_CREATE.value},
# 恢复扫描统一走独立队列,避免占用下载/轮询/创建业务 worker。
"recovery.startup_recovery_once": {"queue": RECOVERY_QUEUE},
"shot_replicate.recover_split_tasks_once": {"queue": RECOVERY_QUEUE},
"generation.recover_download_tasks_once": {"queue": RECOVERY_QUEUE},
"generation.recover_generation_tasks_once": {"queue": RECOVERY_QUEUE},
"module_async.recover_module_async_tasks_once": {"queue": RECOVERY_QUEUE},
"user_oauth.update_oauth_accounts": {"queue": "default"},
"app.tasks.cleanup.*": {"queue": "default"},
CeleryTaskName.STARTUP_RECOVERY.value: {"queue": RECOVERY_QUEUE},
CeleryTaskName.SHOT_SPLIT_RECOVERY.value: {"queue": RECOVERY_QUEUE},
CeleryTaskName.RECOVER_DOWNLOAD.value: {"queue": RECOVERY_QUEUE},
CeleryTaskName.RECOVER_GENERATION.value: {"queue": RECOVERY_QUEUE},
CeleryTaskName.MODULE_ASYNC_RECOVERY.value: {"queue": RECOVERY_QUEUE},
"user_oauth.update_oauth_accounts": {"queue": CeleryQueue.DEFAULT.value},
"app.tasks.cleanup.*": {"queue": CeleryQueue.DEFAULT.value},
},
)
else:
@@ -121,8 +139,8 @@ def on_worker_ready(sender=None, **kwargs):
"""Celery worker 启动时做一次容灾恢复。
注意:
- 不启用 Celery beat
- 启动容灾保留,但只投递一个 recovery.startup_recovery_once 协调任务
- 启动容灾只投递一个 recovery.startup_recovery_once 协调任务
- Celery Beat 只用于每分钟触发轻量 generation.dispatch_due_poll_tasks,不跑完整启动容灾
- 协调任务走独立 gen_recovery 队列,串行扫描并把真实业务任务投回原队列。
- 所有 worker 都尝试抢 Redis 投递锁,只有抢到锁的 worker 投递恢复任务。
"""
@@ -182,4 +200,4 @@ def on_worker_process_shutdown(**kwargs):
except Exception:
pass
finally:
close_loop()
close_loop()
@@ -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()
+156 -41
View File
@@ -6,10 +6,28 @@ from typing import Any
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,
GenerationType,
PROVIDER_FAILED_STATUSES,
PROVIDER_SUCCESS_STATUSES,
)
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, log_provider_call
from app.services.generation_poll_schedule_service import (
build_default_poll_schedule,
build_video_pending_poll_schedule,
ensure_video_poll_fields,
is_final_poll_due,
is_poll_not_due,
is_video_generation_task,
)
from app.services.generation_refund_service import mark_chat_generation_task_failed_and_refund_once
from app.services.generation_provider_service import poll_provider_task
from app.services.media_token_usage_snapshot_service import sync_chat_generation_task_media_token_snapshot
@@ -22,20 +40,19 @@ from app.services.redis_registry_service import (
)
from app.tasks.celery_app import celery_app
ALLOWED_GENERATION_MODES = {"chatapi_async", "hot_opening_replicate", "shot_replicate"}
POLL_QUEUE = "gen_provider_poll"
POLL_QUEUE = CeleryQueue.GEN_PROVIDER_POLL.value
def _now() -> datetime:
return datetime.now(timezone.utc)
def _is_success(status: str) -> bool:
return status in ("succeeded", "success", "completed", "done")
def _is_success(status: str | None) -> bool:
return str(status or "").lower() in PROVIDER_SUCCESS_STATUSES
def _is_failed(status: str) -> bool:
return status in ("failed", "error", "canceled", "cancelled")
def _is_failed(status: str | None) -> bool:
return str(status or "").lower() in PROVIDER_FAILED_STATUSES
def _engine_snapshot(task: ChatGenerationTask) -> dict:
@@ -46,8 +63,7 @@ def _engine_snapshot(task: ChatGenerationTask) -> dict:
def _deadline_expired(task: ChatGenerationTask, now: datetime | None = None) -> bool:
deadline_at = ensure_aware_utc(task.deadline_at)
return bool(deadline_at and deadline_at <= (now or _now()))
return is_final_poll_due(task, now=now)
def _poll_check_at(*, delay_seconds: int | float | None = None, now: datetime | None = None) -> datetime:
@@ -83,6 +99,8 @@ def _build_poll_active_payload(
"queue": POLL_QUEUE,
"poll_count": int(task.poll_count or 0),
"retry_count": int(task.retry_count or 0),
"poll_started_at": datetime_to_epoch(task.poll_started_at) if getattr(task, "poll_started_at", None) else None,
"poll_interval_seconds": int(getattr(task, "poll_interval_seconds", 0) or 0),
"last_poll_at": datetime_to_epoch(task.last_poll_at) if task.last_poll_at else None,
"next_poll_at": datetime_to_epoch(checked_next_poll_at) if checked_next_poll_at else None,
"deadline_at": datetime_to_epoch(task.deadline_at) if task.deadline_at else None,
@@ -154,12 +172,18 @@ async def _mark_timeout(db, task: ChatGenerationTask, *, message: str = "任务
db,
task=task,
error_message=message,
pipeline_stage="timeout",
pipeline_stage=ChatGenerationPipelineStage.TIMEOUT.value,
)
task.next_poll_at = None
await _notify_finished(db, task)
await db.commit()
await remove_poll_active(task.id)
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,
)
async def _mark_failed(db, task: ChatGenerationTask, *, message: str, detail: Any = None) -> None:
@@ -167,12 +191,84 @@ async def _mark_failed(db, task: ChatGenerationTask, *, message: str, detail: An
db,
task=task,
error_message=message,
pipeline_stage="failed",
pipeline_stage=ChatGenerationPipelineStage.FAILED.value,
)
task.next_poll_at = None
await _notify_finished(db, task)
await db.commit()
await remove_poll_active(task.id)
await log_task_event(task, event_type="POLL_FAILED", message=task.error_message, detail=detail)
await log_task_event(
task,
event_type=ChatGenerationTaskEventType.POLL_FAILED.value,
message=task.error_message,
detail=detail,
)
async def _skip_not_due(task: ChatGenerationTask) -> None:
next_poll_at = ensure_aware_utc(task.next_poll_at)
if next_poll_at is None:
return
await register_poll_active(
task,
check_at=next_poll_at,
next_poll_at=next_poll_at,
reason="poll_task_not_due",
)
await log_task_event(
task,
event_type=ChatGenerationTaskEventType.POLL_SKIP_NOT_DUE.value,
message="视频任务尚未到下一次轮询时间,本次 poll 跳过",
detail={"next_poll_at": next_poll_at, "pipeline_stage": task.pipeline_stage},
)
async def _schedule_next_poll(
task: ChatGenerationTask,
*,
reason: str,
default_delay_seconds: int | None = None,
) -> None:
current_time = _now()
if is_video_generation_task(task):
schedule = build_video_pending_poll_schedule(task, now=current_time)
else:
schedule = build_default_poll_schedule(
task,
now=current_time,
delay_seconds=default_delay_seconds,
reason=reason,
)
task.next_poll_at = schedule.next_poll_at
task.poll_interval_seconds = schedule.poll_interval_seconds
await log_task_event(
task,
event_type=ChatGenerationTaskEventType.POLL_SCHEDULED.value,
message=f"已登记下一次轮询。reason={schedule.reason}",
detail={
"delay_seconds": schedule.delay_seconds,
"next_poll_at": schedule.next_poll_at,
"direct_countdown": schedule.direct_countdown,
"poll_interval_seconds": schedule.poll_interval_seconds,
"source_reason": reason,
},
)
await register_poll_active(
task,
check_at=schedule.next_poll_at,
next_poll_at=schedule.next_poll_at,
reason=schedule.reason,
)
if schedule.direct_countdown:
poll_generation_task.apply_async(
args=[task.id],
queue=POLL_QUEUE,
countdown=max(0, int(schedule.delay_seconds)),
)
async def _run(task_id: str):
@@ -190,11 +286,23 @@ async def _run(task_id: str):
return
# 只处理正在生成,且处于远程等待/轮询中的任务。
if task.status != "generating" or task.pipeline_stage not in ("waiting_remote", "polling"):
if task.status != ChatGenerationTaskStatus.GENERATING.value or task.pipeline_stage not in (
ChatGenerationPipelineStage.WAITING_REMOTE.value,
ChatGenerationPipelineStage.POLLING.value,
):
await remove_poll_active(task.id)
return
final_poll_before_timeout = _deadline_expired(task)
current_time = _now()
if is_video_generation_task(task):
ensure_video_poll_fields(task, now=current_time)
if is_poll_not_due(task, now=current_time):
await db.commit()
await _skip_not_due(task)
return
await db.commit()
final_poll_before_timeout = _deadline_expired(task, current_time)
if final_poll_before_timeout and not (task.seedance_task_id or task.provider_task_id):
await _mark_timeout(db, task, message="任务轮询超时")
return
@@ -206,21 +314,23 @@ async def _run(task_id: str):
if final_poll_before_timeout:
await log_task_event(
task,
event_type="FINAL_POLL_BEFORE_TIMEOUT",
event_type=ChatGenerationTaskEventType.FINAL_POLL_BEFORE_TIMEOUT.value,
message="任务已到 deadline,执行最后一次供应商查询后再判定超时",
detail={"deadline_at": task.deadline_at, "stage": task.pipeline_stage},
)
# 标记本次正在轮询,并登记 poll lease。
# 如果 worker 在供应商接口调用过程中退出,启动恢复会在 lease 过期后重新投递。
task.pipeline_stage = "polling"
task.pipeline_stage = ChatGenerationPipelineStage.POLLING.value
task.poll_count = (task.poll_count or 0) + 1
task.last_poll_at = _now()
task.next_poll_at = _poll_lease_until(task.last_poll_at)
await db.commit()
await register_poll_active(
task,
check_at=_poll_lease_until(task.last_poll_at),
reason="polling_lease",
next_poll_at=task.next_poll_at,
)
try:
@@ -247,7 +357,7 @@ async def _run(task_id: str):
)
if _is_success(status):
if task.gen_type == "image":
if task.gen_type == GenerationType.IMAGE.value:
task.remote_result_url = poll_result.get("image_url")
task.image_tokens_used = poll_result.get("image_tokens", 0) or 0
else:
@@ -261,12 +371,13 @@ async def _run(task_id: str):
await _mark_failed(db, task, message="供应商任务成功但未返回结果URL", detail=poll_result)
return
task.pipeline_stage = "result_ready"
task.pipeline_stage = ChatGenerationPipelineStage.RESULT_READY.value
task.retry_count = 0
task.next_poll_at = None
await db.commit()
await remove_poll_active(task.id)
await log_task_event(task, event_type="POLL_SUCCESS", to_stage="result_ready")
await log_task_event(task, event_type=ChatGenerationTaskEventType.POLL_SUCCESS.value, to_stage=ChatGenerationPipelineStage.RESULT_READY.value)
from app.tasks.generation_download_tasks import enqueue_download_task
@@ -286,7 +397,7 @@ async def _run(task_id: str):
if final_poll_before_timeout:
await log_task_event(
task,
event_type="FINAL_POLL_BEFORE_TIMEOUT_PENDING",
event_type=ChatGenerationTaskEventType.FINAL_POLL_BEFORE_TIMEOUT_PENDING.value,
message=f"最终查询后供应商仍未完成,按超时处理。status={status}",
detail=poll_result,
)
@@ -294,26 +405,17 @@ async def _run(task_id: str):
return
# 供应商仍在 pending / running 时,把阶段从 polling 改回 waiting_remote。
# 同时登记下一次 poll activeCelery countdown 丢失时可由恢复任务拉起
task.pipeline_stage = "waiting_remote"
# 视频任务写入 next_poll_at,由 Beat dispatcher 到期投递;短间隔可保留 countdown 兼容
task.pipeline_stage = ChatGenerationPipelineStage.WAITING_REMOTE.value
task.retry_count = 0
await _schedule_next_poll(task, reason="poll_pending_next")
await db.commit()
await log_task_event(task, event_type="POLL_PENDING", message=f"status={status}")
delay_seconds = int(settings.CHATAPI_ASYNC_POLL_INTERVAL_SECONDS or 30)
next_poll_at = _now() + timedelta(seconds=max(1, delay_seconds))
await register_poll_active(
await log_task_event(
task,
check_at=_poll_check_at(delay_seconds=delay_seconds),
next_poll_at=next_poll_at,
reason="poll_pending_next",
)
poll_generation_task.apply_async(
args=[task.id],
queue=POLL_QUEUE,
countdown=delay_seconds,
event_type=ChatGenerationTaskEventType.POLL_PENDING.value,
message=f"status={status}",
detail={"next_poll_at": task.next_poll_at, "poll_interval_seconds": task.poll_interval_seconds},
)
except Exception as exc:
@@ -331,7 +433,7 @@ async def _run(task_id: str):
if final_poll_before_timeout:
await log_task_event(
task,
event_type="FINAL_POLL_BEFORE_TIMEOUT_ERROR",
event_type=ChatGenerationTaskEventType.FINAL_POLL_BEFORE_TIMEOUT_ERROR.value,
message=str(exc),
)
await _mark_timeout(db, task, message="任务轮询超时")
@@ -339,21 +441,34 @@ async def _run(task_id: str):
task.retry_count = (task.retry_count or 0) + 1
# 视频轮询的临时异常不再 3 次内直接退款;继续降频到 24 小时最终 deadline。
if is_video_generation_task(task):
task.pipeline_stage = ChatGenerationPipelineStage.WAITING_REMOTE.value
await _schedule_next_poll(task, reason="poll_exception_retry")
await db.commit()
return
if task.retry_count > settings.CHATAPI_ASYNC_MAX_RETRIES:
error_message = extract_error_message(exc, "轮询") if callable(extract_error_message) else str(exc)
await _mark_failed(db, task, message=error_message)
else:
# 临时轮询异常时,不让任务停在 polling。
# 回到 waiting_remote,等待下一次重试轮询。
task.pipeline_stage = "waiting_remote"
task.pipeline_stage = ChatGenerationPipelineStage.WAITING_REMOTE.value
delay_seconds = int(settings.CHATAPI_ASYNC_RETRY_BACKOFF_SECONDS or 30) * int(task.retry_count or 1)
default_schedule = build_default_poll_schedule(
task,
now=_now(),
delay_seconds=delay_seconds,
reason="poll_exception_retry",
)
task.next_poll_at = default_schedule.next_poll_at
await db.commit()
delay_seconds = int(settings.CHATAPI_ASYNC_RETRY_BACKOFF_SECONDS or 30) * int(task.retry_count or 1)
next_poll_at = _now() + timedelta(seconds=max(1, delay_seconds))
await register_poll_active(
task,
check_at=_poll_check_at(delay_seconds=delay_seconds),
next_poll_at=next_poll_at,
next_poll_at=default_schedule.next_poll_at,
reason="poll_exception_retry",
)
@@ -5,6 +5,7 @@ import logging
from typing import Any, Awaitable, Callable, Dict
from app.config import settings
from app.enums.celery_queue import CeleryQueue
from app.models.base import async_session
from app.services.redis_registry_service import get_registry_redis, redis_acquire_lock, redis_release_lock
from app.tasks.async_runner import run_async
@@ -12,7 +13,7 @@ from app.tasks.celery_app import celery_app
logger = logging.getLogger("video_gen")
RECOVERY_QUEUE = settings.CELERY_RECOVERY_QUEUE or "gen_recovery"
RECOVERY_QUEUE = settings.CELERY_RECOVERY_QUEUE or CeleryQueue.GEN_RECOVERY.value
RecoveryRunner = Callable[[], Awaitable[Dict[str, Any]]]
@@ -30,6 +31,13 @@ async def _run_generation_once() -> Dict[str, Any]:
return await recover_generation_tasks_once(db)
async def _run_due_poll_dispatch_once() -> Dict[str, Any]:
from app.services.generation_recovery_service import dispatch_due_poll_tasks_once
async with async_session() as db:
return await dispatch_due_poll_tasks_once(db)
async def _run_module_async_once() -> Dict[str, Any]:
from app.services.module_async_recovery_service import recover_module_async_tasks_once
@@ -49,6 +57,7 @@ async def _run_with_execution_lock(
lock_key: str,
log_context: str,
runner: RecoveryRunner,
ttl_seconds: int | None = None,
) -> Dict[str, Any]:
"""恢复任务执行锁。
@@ -61,7 +70,7 @@ async def _run_with_execution_lock(
if redis is not None:
token = await redis_acquire_lock(
lock_key=lock_key,
ttl_seconds=int(settings.CELERY_RECOVERY_TASK_LOCK_TTL_SECONDS or 600),
ttl_seconds=int(ttl_seconds or settings.CELERY_RECOVERY_TASK_LOCK_TTL_SECONDS or 600),
log_context=log_context,
)
if not token:
@@ -79,6 +88,36 @@ async def _run_with_execution_lock(
await redis_release_lock(lock_key=lock_key, token=token, log_context=log_context)
async def _is_lock_held(lock_key: str) -> bool:
redis = await get_registry_redis()
if redis is None:
return False
try:
return bool(await redis.exists(lock_key))
except Exception:
return False
async def _startup_or_generation_recovery_running() -> str | None:
# Beat 触发 dispatcher 时,如果启动容灾或完整生成容灾还在跑,直接跳过本轮。
# gen_recovery concurrency=1 已经能串行;这里是多机部署、残留消息、手动触发时的双保险。
lock_checks = [
("startup_recovery", settings.CELERY_RECOVERY_STARTUP_TASK_LOCK_KEY),
("generation_recovery", settings.GENERATION_RECOVERY_LOCK_KEY),
]
for name, lock_key in lock_checks:
if await _is_lock_held(lock_key):
return name
return None
async def _run_due_poll_dispatch_with_guard() -> Dict[str, Any]:
running = await _startup_or_generation_recovery_running()
if running:
return {"skipped": "recovery_lock_held", "lock": running}
return await _run_due_poll_dispatch_once()
async def _acquire_download_recovery_loop_lock() -> tuple[bool, str]:
"""下载恢复循环锁。
@@ -220,6 +259,23 @@ if celery_app:
)
)
@celery_app.task(
name="generation.dispatch_due_poll_tasks",
bind=True,
soft_time_limit=settings.CELERY_RECOVERY_SOFT_TIME_LIMIT_SECONDS,
time_limit=settings.CELERY_RECOVERY_TIME_LIMIT_SECONDS,
)
def dispatch_due_poll_tasks(self) -> Dict[str, Any]:
return run_async(
_run_with_execution_lock(
lock_key=settings.POLL_DUE_DISPATCH_LOCK_KEY,
log_context="due_poll_dispatch",
runner=_run_due_poll_dispatch_with_guard,
ttl_seconds=int(settings.POLL_DUE_DISPATCH_LOCK_TTL_SECONDS or 55),
)
)
else:
class _DisabledTask:
@@ -232,3 +288,4 @@ else:
startup_recovery_once = _DisabledTask()
recover_download_tasks_once = _DisabledTask()
recover_generation_tasks_once = _DisabledTask()
dispatch_due_poll_tasks = _DisabledTask()