celery worker poll 频次机制调整| celery beat 设置poll检测任务| 拆镜复刻切片删除API | 交易流水时区BUG
This commit is contained in:
@@ -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 active,Celery 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",
|
||||
)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user