Files
video-gen/video-gen-api/app/tasks/generation_poll_tasks.py
T

506 lines
19 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
from app.tasks.async_runner import run_async
import json
from datetime import datetime, timedelta, timezone
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
from app.services.redis_registry_service import (
datetime_to_epoch,
ensure_aware_utc,
redis_remove_registry_item,
redis_upsert_registry_item,
utc_now,
)
from app.tasks.celery_app import celery_app
POLL_QUEUE = CeleryQueue.GEN_PROVIDER_POLL.value
def _now() -> datetime:
return datetime.now(timezone.utc)
def _is_success(status: str | None) -> bool:
return str(status or "").lower() in PROVIDER_SUCCESS_STATUSES
def _is_failed(status: str | None) -> bool:
return str(status or "").lower() in PROVIDER_FAILED_STATUSES
def _engine_snapshot(task: ChatGenerationTask) -> dict:
try:
return json.loads(task.engine_snapshot_json or "{}")
except Exception:
return {}
def _deadline_expired(task: ChatGenerationTask, now: datetime | None = None) -> bool:
return is_final_poll_due(task, now=now)
def _poll_check_at(*, delay_seconds: int | float | None = None, now: datetime | None = None) -> datetime:
current_time = now or _now()
delay = int(delay_seconds or settings.CHATAPI_ASYNC_POLL_INTERVAL_SECONDS or 30)
grace = int(settings.POLL_TASK_QUEUE_TIMEOUT_SECONDS or 120)
return current_time + timedelta(seconds=max(1, delay) + max(0, grace))
def _poll_lease_until(now: datetime | None = None) -> datetime:
current_time = now or _now()
return current_time + timedelta(seconds=int(settings.POLL_TASK_LEASE_SECONDS or 300))
def _build_poll_active_payload(
task: ChatGenerationTask,
*,
stage: str,
reason: str,
next_poll_at: datetime | None = None,
check_at: datetime | None = None,
) -> dict[str, Any]:
current_time = utc_now()
checked_next_poll_at = ensure_aware_utc(next_poll_at)
checked_check_at = ensure_aware_utc(check_at)
return {
"task_id": task.id,
"provider_task_id": task.provider_task_id,
"seedance_task_id": task.seedance_task_id,
"generation_mode": task.generation_mode,
"gen_type": task.gen_type,
"stage": stage,
"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,
"check_at": datetime_to_epoch(checked_check_at) if checked_check_at else None,
"updated_at": datetime_to_epoch(current_time),
"reason": reason,
}
async def register_poll_active(
task: ChatGenerationTask,
*,
check_at: datetime,
reason: str,
next_poll_at: datetime | None = None,
) -> None:
payload = _build_poll_active_payload(
task,
stage=task.pipeline_stage or "",
reason=reason,
next_poll_at=next_poll_at,
check_at=check_at,
)
await redis_upsert_registry_item(
hash_key=settings.POLL_ACTIVE_REDIS_HASH_KEY,
zset_key=settings.POLL_ACTIVE_REDIS_ZSET_KEY,
item_id=task.id,
payload=payload,
check_at=check_at,
log_context="poll_active",
)
async def remove_poll_active(task_id: str) -> None:
await redis_remove_registry_item(
hash_key=settings.POLL_ACTIVE_REDIS_HASH_KEY,
zset_key=settings.POLL_ACTIVE_REDIS_ZSET_KEY,
item_id=task_id,
log_context="poll_active",
)
async def _notify_finished(db, task: ChatGenerationTask) -> None:
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)
async def _reload_task(db, task_id: str) -> ChatGenerationTask | None:
"""
rollback 后重新查询任务对象。
说明:
- SQLAlchemy rollback 后,当前 ORM 对象可能进入过期状态。
- 后续继续访问旧 task.retry_count / task.id 等字段,有概率触发异步懒加载异常。
- 所以 poll/download 的异常分支统一 rollback 后重新 select。
"""
result = await db.execute(
select(ChatGenerationTask).where(
ChatGenerationTask.id == task_id,
ChatGenerationTask.deleted_at.is_(None),
).with_for_update().limit(1)
)
return result.scalar_one_or_none()
async def _mark_timeout(db, task: ChatGenerationTask, *, message: str = "任务轮询超时") -> None:
await mark_chat_generation_task_failed_and_refund_once(
db,
task=task,
error_message=message,
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=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:
await mark_chat_generation_task_failed_and_refund_once(
db,
task=task,
error_message=message,
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=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, *, force_due: bool = False):
async with async_session() as db:
result = await db.execute(select(ChatGenerationTask).where(
ChatGenerationTask.id == task_id,
ChatGenerationTask.deleted_at.is_(None),
).with_for_update().limit(1))
task = result.scalar_one_or_none()
if not task:
await remove_poll_active(task_id)
return
if task.generation_mode not in ALLOWED_GENERATION_MODES:
await remove_poll_active(task.id)
return
# 只处理正在生成,且处于远程等待/轮询中的任务。
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
current_time = _now()
if is_video_generation_task(task):
ensure_video_poll_fields(task, now=current_time)
# dispatcher / recovery 已经在投递前确认到期时,会传 force_due=True。
# 这样可以避免投递侧为了防重复消费临时写入的 next_poll_at
# 又被当前 worker 当成“业务下一次轮询时间”而误判未到期。
if not force_due and 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
if not (task.seedance_task_id or task.provider_task_id):
await _mark_failed(db, task, message="缺少外部任务ID")
return
if final_poll_before_timeout:
await log_task_event(
task,
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 = 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:
poll_result = await poll_provider_task(db, task)
status = poll_result.get("status")
response_data = poll_result.get("response_data")
try:
provider_response = json.loads(response_data or "{}")
except Exception:
provider_response = {"raw": response_data}
snapshot = _engine_snapshot(task)
await log_provider_call(
task,
provider=snapshot.get("provider") or "ark",
api_type=f"{task.gen_type}_poll",
model=snapshot.get("model_name"),
engine_id=task.engine_id,
status="success",
provider_task_id=task.seedance_task_id or task.provider_task_id,
response_data=provider_response,
)
if _is_success(status):
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:
task.remote_result_url = poll_result.get("video_url")
task.video_tokens_used = poll_result.get("video_tokens", 0) or 0
task.provider_response_json = response_data
await sync_chat_generation_task_media_token_snapshot(db, task, provider_response=response_data)
if not task.remote_result_url:
await _mark_failed(db, task, message="供应商任务成功但未返回结果URL", detail=poll_result)
return
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=ChatGenerationTaskEventType.POLL_SUCCESS.value, to_stage=ChatGenerationPipelineStage.RESULT_READY.value)
from app.tasks.generation_download_tasks import enqueue_download_task
await enqueue_download_task(db, task, reason="poll_success_result_ready")
return
if _is_failed(status):
task.provider_response_json = response_data
await _mark_failed(
db,
task,
message=poll_result.get("error") or f"供应商任务失败: {status}",
detail=poll_result,
)
return
if final_poll_before_timeout:
await log_task_event(
task,
event_type=ChatGenerationTaskEventType.FINAL_POLL_BEFORE_TIMEOUT_PENDING.value,
message=f"最终查询后供应商仍未完成,按超时处理。status={status}",
detail=poll_result,
)
await _mark_timeout(db, task, message="任务轮询超时")
return
# 供应商仍在 pending / running 时,把阶段从 polling 改回 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=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:
# 异常后先 rollback,再重新查询 task,不继续使用 rollback 前的旧 ORM 对象。
try:
await db.rollback()
except Exception:
pass
task = await _reload_task(db, task_id)
if not task:
await remove_poll_active(task_id)
return
if final_poll_before_timeout:
await log_task_event(
task,
event_type=ChatGenerationTaskEventType.FINAL_POLL_BEFORE_TIMEOUT_ERROR.value,
message=str(exc),
)
await _mark_timeout(db, task, message="任务轮询超时")
return
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 = 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()
await register_poll_active(
task,
check_at=_poll_check_at(delay_seconds=delay_seconds),
next_poll_at=default_schedule.next_poll_at,
reason="poll_exception_retry",
)
poll_generation_task.apply_async(
args=[task.id],
queue=POLL_QUEUE,
countdown=delay_seconds,
)
if celery_app:
@celery_app.task(name="generation.poll_generation_task", bind=True, max_retries=3, default_retry_delay=30)
def poll_generation_task(self, task_id: str, force_due: bool = False):
try:
return run_async(_run(task_id, force_due=bool(force_due)))
except Exception as exc:
# 只重试基础设施异常;供应商失败/业务失败已在 _run 内处理。
retries = int(getattr(self.request, "retries", 0) or 0) + 1
countdown = int(settings.CHATAPI_ASYNC_RETRY_BACKOFF_SECONDS or 30) * max(1, retries)
raise self.retry(exc=exc, countdown=countdown)
else:
class _DisabledTask:
def delay(self, *args, **kwargs):
raise RuntimeError("Celery is disabled")
def apply_async(self, *args, **kwargs):
raise RuntimeError("Celery is disabled")
poll_generation_task = _DisabledTask()