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

236 lines
9.2 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, timezone
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.tasks.celery_app import celery_app
ALLOWED_GENERATION_MODES = {"chatapi_async", "hot_opening_replicate"}
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 {}
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 _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 or task.generation_mode not in ALLOWED_GENERATION_MODES:
return
# 只处理正在生成,且处于远程等待/轮询中的任务。
if task.status != "generating" or task.pipeline_stage not in ("waiting_remote", "polling"):
return
if task.deadline_at and datetime.now(timezone.utc) > task.deadline_at:
await mark_chat_generation_task_failed_and_refund_once(
db,
task=task,
error_message="任务轮询超时",
pipeline_stage="timeout",
)
await _notify_finished(db, task)
await db.commit()
await log_task_event(task, event_type="TASK_TIMEOUT", to_status="failed", to_stage="timeout")
return
if not (task.seedance_task_id or task.provider_task_id):
await mark_chat_generation_task_failed_and_refund_once(
db,
task=task,
error_message="缺少外部任务ID",
pipeline_stage="failed",
)
await _notify_finished(db, task)
await db.commit()
await log_task_event(task, event_type="POLL_FAILED", message=task.error_message)
return
# 标记本次正在轮询。
# 注意:pending 后会再改回 waiting_remote,避免任务长期卡在 polling。
task.pipeline_stage = "polling"
task.poll_count = (task.poll_count or 0) + 1
task.last_poll_at = datetime.now(timezone.utc)
await db.commit()
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
if not task.remote_result_url:
await mark_chat_generation_task_failed_and_refund_once(
db,
task=task,
error_message="供应商任务成功但未返回结果URL",
pipeline_stage="failed",
)
await _notify_finished(db, task)
await db.commit()
await log_task_event(task, event_type="POLL_FAILED", message=task.error_message)
return
task.pipeline_stage = "result_ready"
task.retry_count = 0
await db.commit()
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_chat_generation_task_failed_and_refund_once(
db,
task=task,
error_message=poll_result.get("error") or f"供应商任务失败: {status}",
pipeline_stage="failed",
)
await _notify_finished(db, task)
await db.commit()
await log_task_event(task, event_type="POLL_FAILED", message=task.error_message, detail=poll_result)
return
# 关键修改 1
# 供应商仍在 pending / running 时,把阶段从 polling 改回 waiting_remote。
# 这样数据库状态表示“等待下一次轮询”,不会长期停在 polling。
# 同时可以降低重复 Celery 消息形成多条轮询链的概率。
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}")
poll_generation_task.apply_async(
args=[task.id],
countdown=settings.CHATAPI_ASYNC_POLL_INTERVAL_SECONDS,
)
except Exception as exc:
# 关键修改 2
# 异常后先 rollback,再重新查询 task,不继续使用 rollback 前的旧 ORM 对象。
try:
await db.rollback()
except Exception:
pass
task = await _reload_task(db, task_id)
if not task:
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_chat_generation_task_failed_and_refund_once(
db,
task=task,
error_message=error_message,
pipeline_stage="failed",
)
await _notify_finished(db, task)
await db.commit()
await log_task_event(task, event_type="POLL_FAILED", message=task.error_message)
else:
# 临时轮询异常时,不让任务停在 polling。
# 回到 waiting_remote,等待下一次重试轮询。
task.pipeline_stage = "waiting_remote"
await db.commit()
poll_generation_task.apply_async(
args=[task.id],
countdown=settings.CHATAPI_ASYNC_RETRY_BACKOFF_SECONDS * task.retry_count,
)
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()