生成记录接口兼容图片类型|CHAT对话生成视频/图片任务API+celery异步模块完成
This commit is contained in:
@@ -0,0 +1,113 @@
|
||||
from app.tasks.async_runner import run_async
|
||||
import json
|
||||
from datetime import datetime, timezone
|
||||
|
||||
from sqlalchemy import select
|
||||
|
||||
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_billing_service import charge_chatapi_prompt_usage
|
||||
from app.services.generation_log_service import log_task_event
|
||||
from app.services.generation_prompt_service import build_prompt_with_chatapi
|
||||
from app.services.generation_provider_service import create_provider_task
|
||||
from app.tasks.celery_app import celery_app
|
||||
|
||||
|
||||
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":
|
||||
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 task.pipeline_stage not in ("queued", "preparing", "creating_provider_task"):
|
||||
return
|
||||
|
||||
try:
|
||||
if not task.optimized_prompt:
|
||||
old_stage = task.pipeline_stage
|
||||
task.pipeline_stage = "preparing"
|
||||
await db.commit()
|
||||
await log_task_event(task, event_type="CHATAPI_START", from_stage=old_stage, to_stage="preparing")
|
||||
|
||||
optimized, usage = await build_prompt_with_chatapi(db, task)
|
||||
await charge_chatapi_prompt_usage(db, record=task, usage=usage, project_name="AI生成任务")
|
||||
task.optimized_prompt = optimized
|
||||
task.text_tokens_used = usage["total_tokens"]
|
||||
await db.commit()
|
||||
await log_task_event(task, event_type="CHATAPI_SUCCESS", to_stage="preparing", detail=usage)
|
||||
|
||||
if task.seedance_task_id or task.provider_task_id:
|
||||
task.pipeline_stage = "waiting_remote"
|
||||
await db.commit()
|
||||
elif task.remote_result_url:
|
||||
task.pipeline_stage = "result_ready"
|
||||
await db.commit()
|
||||
from app.tasks.generation_download_tasks import download_generation_result_task
|
||||
download_generation_result_task.delay(task.id)
|
||||
return
|
||||
else:
|
||||
old_stage = task.pipeline_stage
|
||||
task.pipeline_stage = "creating_provider_task"
|
||||
await db.commit()
|
||||
await log_task_event(task, event_type="PROVIDER_CREATE_START", from_stage=old_stage, to_stage="creating_provider_task")
|
||||
|
||||
created = await create_provider_task(db, task)
|
||||
provider_task_id = created.get("task_id")
|
||||
if provider_task_id:
|
||||
task.provider_task_id = provider_task_id
|
||||
task.seedance_task_id = provider_task_id
|
||||
task.remote_result_url = created.get("remote_result_url") or task.remote_result_url
|
||||
if task.gen_type == "image":
|
||||
task.image_tokens_used = created.get("image_tokens", task.image_tokens_used or 0) or 0
|
||||
task.provider_response_json = json.dumps(created.get("response_data") or {}, ensure_ascii=False, default=str)
|
||||
if task.remote_result_url and not task.seedance_task_id:
|
||||
# Sync image path: original SDK already returned final URL.
|
||||
task.pipeline_stage = "result_ready"
|
||||
else:
|
||||
# Video path: provider returns task id, poll later.
|
||||
task.pipeline_stage = "waiting_remote"
|
||||
task.status = "generating"
|
||||
await db.commit()
|
||||
await log_task_event(task, event_type="PROVIDER_CREATE_SUCCESS", to_stage=task.pipeline_stage, detail=created)
|
||||
|
||||
if task.pipeline_stage == "result_ready":
|
||||
from app.tasks.generation_download_tasks import download_generation_result_task
|
||||
download_generation_result_task.delay(task.id)
|
||||
else:
|
||||
from app.tasks.generation_poll_tasks import poll_generation_task
|
||||
poll_generation_task.delay(task.id)
|
||||
except Exception as exc:
|
||||
try:
|
||||
await db.rollback()
|
||||
except Exception:
|
||||
pass
|
||||
result = await db.execute(select(ChatGenerationTask).where(ChatGenerationTask.id == task_id))
|
||||
task = result.scalar_one_or_none()
|
||||
if task:
|
||||
task.status = "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="TASK_FAILED", message=task.error_message)
|
||||
|
||||
|
||||
if celery_app:
|
||||
@celery_app.task(name="generation.chatapi_create_generation_task", bind=True, max_retries=3, default_retry_delay=30)
|
||||
def chatapi_create_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")
|
||||
chatapi_create_generation_task = _DisabledTask()
|
||||
Reference in New Issue
Block a user