350 lines
13 KiB
Python
350 lines
13 KiB
Python
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.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
|
||
|
||
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 active,Celery 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()
|