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
+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",
)