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

352 lines
13 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.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_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
ALLOWED_GENERATION_MODES = {"chatapi_async", "hot_opening_replicate", "shot_replicate"}
POLL_QUEUE = "gen_provider_poll"
def _now() -> datetime:
return datetime.now(timezone.utc)
def _is_success(status: str) -> bool:
return status in ("succeeded", "success", "completed", "done")
def _is_failed(status: str) -> bool:
return status in ("failed", "error", "canceled", "cancelled")
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:
deadline_at = ensure_aware_utc(task.deadline_at)
return bool(deadline_at and deadline_at <= (now or _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),
"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
await notify_chat_generation_task_finished(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="timeout",
)
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")
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="failed",
)
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)
async def _run(task_id: str):
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 != "generating" or task.pipeline_stage not in ("waiting_remote", "polling"):
await remove_poll_active(task.id)
return
if _deadline_expired(task):
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
# 标记本次正在轮询,并登记 poll lease。
# 如果 worker 在供应商接口调用过程中退出,启动恢复会在 lease 过期后重新投递。
task.pipeline_stage = "polling"
task.poll_count = (task.poll_count or 0) + 1
task.last_poll_at = _now()
await db.commit()
await register_poll_active(
task,
check_at=_poll_lease_until(task.last_poll_at),
reason="polling_lease",
)
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 == "image":
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 = "result_ready"
task.retry_count = 0
await db.commit()
await remove_poll_active(task.id)
await log_task_event(task, event_type="POLL_SUCCESS", to_stage="result_ready")
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
# 供应商仍在 pending / running 时,把阶段从 polling 改回 waiting_remote。
# 同时登记下一次 poll activeCelery countdown 丢失时可由恢复任务拉起。
task.pipeline_stage = "waiting_remote"
task.retry_count = 0
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(
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,
)
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
task.retry_count = (task.retry_count or 0) + 1
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"
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,
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):
return run_async(_run(task_id))
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()