236 lines
9.2 KiB
Python
236 lines
9.2 KiB
Python
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() |