生成记录接口兼容图片类型|CHAT对话生成视频/图片任务API+celery异步模块完成

This commit is contained in:
2026-05-27 13:11:55 +08:00
parent d61dcdc8db
commit b8eda53c0b
30 changed files with 3392 additions and 45 deletions
@@ -0,0 +1,199 @@
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_provider_service import poll_provider_task
from app.tasks.celery_app import celery_app
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 _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)
)
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))
task = result.scalar_one_or_none()
if not task or task.generation_mode != "chatapi_async":
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:
task.status = "failed"
task.pipeline_stage = "timeout"
task.error_message = "任务轮询超时"
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):
task.status = "failed"
task.pipeline_stage = "failed"
task.error_message = "缺少外部任务ID"
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:
task.status = "failed"
task.pipeline_stage = "failed"
task.error_message = "供应商任务成功但未返回结果URL"
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 download_generation_result_task
download_generation_result_task.delay(task.id)
return
if _is_failed(status):
task.status = "failed"
task.pipeline_stage = "failed"
task.error_message = poll_result.get("error") or f"供应商任务失败: {status}"
task.provider_response_json = response_data
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:
task.status = "failed"
task.pipeline_stage = "failed"
task.error_message = extract_error_message(exc, "轮询") if callable(extract_error_message) else str(exc)
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()