501 lines
19 KiB
Python
501 lines
19 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.enums.celery_queue import CeleryQueue
|
|
from app.enums.generation_task import (
|
|
ALLOWED_GENERATION_MODES,
|
|
ChatGenerationPipelineStage,
|
|
ChatGenerationTaskEventType,
|
|
ChatGenerationTaskStatus,
|
|
GenerationType,
|
|
PROVIDER_FAILED_STATUSES,
|
|
PROVIDER_SUCCESS_STATUSES,
|
|
)
|
|
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_poll_schedule_service import (
|
|
build_default_poll_schedule,
|
|
build_video_pending_poll_schedule,
|
|
ensure_video_poll_fields,
|
|
is_final_poll_due,
|
|
is_poll_not_due,
|
|
is_video_generation_task,
|
|
)
|
|
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.media_token_usage_snapshot_service import sync_chat_generation_task_media_token_snapshot
|
|
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
|
|
|
|
POLL_QUEUE = CeleryQueue.GEN_PROVIDER_POLL.value
|
|
|
|
|
|
def _now() -> datetime:
|
|
return datetime.now(timezone.utc)
|
|
|
|
|
|
def _is_success(status: str | None) -> bool:
|
|
return str(status or "").lower() in PROVIDER_SUCCESS_STATUSES
|
|
|
|
|
|
def _is_failed(status: str | None) -> bool:
|
|
return str(status or "").lower() in PROVIDER_FAILED_STATUSES
|
|
|
|
|
|
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:
|
|
return is_final_poll_due(task, now=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),
|
|
"poll_started_at": datetime_to_epoch(task.poll_started_at) if getattr(task, "poll_started_at", None) else None,
|
|
"poll_interval_seconds": int(getattr(task, "poll_interval_seconds", 0) 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=ChatGenerationPipelineStage.TIMEOUT.value,
|
|
)
|
|
task.next_poll_at = None
|
|
await _notify_finished(db, task)
|
|
await db.commit()
|
|
await remove_poll_active(task.id)
|
|
await log_task_event(
|
|
task,
|
|
event_type=ChatGenerationTaskEventType.TASK_TIMEOUT.value,
|
|
to_status=ChatGenerationTaskStatus.FAILED.value,
|
|
to_stage=ChatGenerationPipelineStage.TIMEOUT.value,
|
|
)
|
|
|
|
|
|
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=ChatGenerationPipelineStage.FAILED.value,
|
|
)
|
|
task.next_poll_at = None
|
|
await _notify_finished(db, task)
|
|
await db.commit()
|
|
await remove_poll_active(task.id)
|
|
await log_task_event(
|
|
task,
|
|
event_type=ChatGenerationTaskEventType.POLL_FAILED.value,
|
|
message=task.error_message,
|
|
detail=detail,
|
|
)
|
|
|
|
|
|
async def _skip_not_due(task: ChatGenerationTask) -> None:
|
|
next_poll_at = ensure_aware_utc(task.next_poll_at)
|
|
if next_poll_at is None:
|
|
return
|
|
await register_poll_active(
|
|
task,
|
|
check_at=next_poll_at,
|
|
next_poll_at=next_poll_at,
|
|
reason="poll_task_not_due",
|
|
)
|
|
await log_task_event(
|
|
task,
|
|
event_type=ChatGenerationTaskEventType.POLL_SKIP_NOT_DUE.value,
|
|
message="视频任务尚未到下一次轮询时间,本次 poll 跳过",
|
|
detail={"next_poll_at": next_poll_at, "pipeline_stage": task.pipeline_stage},
|
|
)
|
|
|
|
|
|
async def _schedule_next_poll(
|
|
task: ChatGenerationTask,
|
|
*,
|
|
reason: str,
|
|
default_delay_seconds: int | None = None,
|
|
) -> None:
|
|
current_time = _now()
|
|
if is_video_generation_task(task):
|
|
schedule = build_video_pending_poll_schedule(task, now=current_time)
|
|
else:
|
|
schedule = build_default_poll_schedule(
|
|
task,
|
|
now=current_time,
|
|
delay_seconds=default_delay_seconds,
|
|
reason=reason,
|
|
)
|
|
|
|
task.next_poll_at = schedule.next_poll_at
|
|
task.poll_interval_seconds = schedule.poll_interval_seconds
|
|
|
|
await log_task_event(
|
|
task,
|
|
event_type=ChatGenerationTaskEventType.POLL_SCHEDULED.value,
|
|
message=f"已登记下一次轮询。reason={schedule.reason}",
|
|
detail={
|
|
"delay_seconds": schedule.delay_seconds,
|
|
"next_poll_at": schedule.next_poll_at,
|
|
"direct_countdown": schedule.direct_countdown,
|
|
"poll_interval_seconds": schedule.poll_interval_seconds,
|
|
"source_reason": reason,
|
|
},
|
|
)
|
|
|
|
await register_poll_active(
|
|
task,
|
|
check_at=schedule.next_poll_at,
|
|
next_poll_at=schedule.next_poll_at,
|
|
reason=schedule.reason,
|
|
)
|
|
|
|
if schedule.direct_countdown:
|
|
poll_generation_task.apply_async(
|
|
args=[task.id],
|
|
queue=POLL_QUEUE,
|
|
countdown=max(0, int(schedule.delay_seconds)),
|
|
)
|
|
|
|
|
|
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 != ChatGenerationTaskStatus.GENERATING.value or task.pipeline_stage not in (
|
|
ChatGenerationPipelineStage.WAITING_REMOTE.value,
|
|
ChatGenerationPipelineStage.POLLING.value,
|
|
):
|
|
await remove_poll_active(task.id)
|
|
return
|
|
|
|
current_time = _now()
|
|
if is_video_generation_task(task):
|
|
ensure_video_poll_fields(task, now=current_time)
|
|
if is_poll_not_due(task, now=current_time):
|
|
await db.commit()
|
|
await _skip_not_due(task)
|
|
return
|
|
await db.commit()
|
|
|
|
final_poll_before_timeout = _deadline_expired(task, current_time)
|
|
if final_poll_before_timeout and not (task.seedance_task_id or task.provider_task_id):
|
|
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
|
|
|
|
if final_poll_before_timeout:
|
|
await log_task_event(
|
|
task,
|
|
event_type=ChatGenerationTaskEventType.FINAL_POLL_BEFORE_TIMEOUT.value,
|
|
message="任务已到 deadline,执行最后一次供应商查询后再判定超时",
|
|
detail={"deadline_at": task.deadline_at, "stage": task.pipeline_stage},
|
|
)
|
|
|
|
# 标记本次正在轮询,并登记 poll lease。
|
|
# 如果 worker 在供应商接口调用过程中退出,启动恢复会在 lease 过期后重新投递。
|
|
task.pipeline_stage = ChatGenerationPipelineStage.POLLING.value
|
|
task.poll_count = (task.poll_count or 0) + 1
|
|
task.last_poll_at = _now()
|
|
task.next_poll_at = _poll_lease_until(task.last_poll_at)
|
|
await db.commit()
|
|
await register_poll_active(
|
|
task,
|
|
check_at=_poll_lease_until(task.last_poll_at),
|
|
reason="polling_lease",
|
|
next_poll_at=task.next_poll_at,
|
|
)
|
|
|
|
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 == GenerationType.IMAGE.value:
|
|
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
|
|
await sync_chat_generation_task_media_token_snapshot(db, task, provider_response=response_data)
|
|
|
|
if not task.remote_result_url:
|
|
await _mark_failed(db, task, message="供应商任务成功但未返回结果URL", detail=poll_result)
|
|
return
|
|
|
|
task.pipeline_stage = ChatGenerationPipelineStage.RESULT_READY.value
|
|
task.retry_count = 0
|
|
task.next_poll_at = None
|
|
await db.commit()
|
|
await remove_poll_active(task.id)
|
|
|
|
await log_task_event(task, event_type=ChatGenerationTaskEventType.POLL_SUCCESS.value, to_stage=ChatGenerationPipelineStage.RESULT_READY.value)
|
|
|
|
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
|
|
|
|
if final_poll_before_timeout:
|
|
await log_task_event(
|
|
task,
|
|
event_type=ChatGenerationTaskEventType.FINAL_POLL_BEFORE_TIMEOUT_PENDING.value,
|
|
message=f"最终查询后供应商仍未完成,按超时处理。status={status}",
|
|
detail=poll_result,
|
|
)
|
|
await _mark_timeout(db, task, message="任务轮询超时")
|
|
return
|
|
|
|
# 供应商仍在 pending / running 时,把阶段从 polling 改回 waiting_remote。
|
|
# 视频任务写入 next_poll_at,由 Beat dispatcher 到期投递;短间隔可保留 countdown 兼容。
|
|
task.pipeline_stage = ChatGenerationPipelineStage.WAITING_REMOTE.value
|
|
task.retry_count = 0
|
|
await _schedule_next_poll(task, reason="poll_pending_next")
|
|
await db.commit()
|
|
|
|
await log_task_event(
|
|
task,
|
|
event_type=ChatGenerationTaskEventType.POLL_PENDING.value,
|
|
message=f"status={status}",
|
|
detail={"next_poll_at": task.next_poll_at, "poll_interval_seconds": task.poll_interval_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
|
|
|
|
if final_poll_before_timeout:
|
|
await log_task_event(
|
|
task,
|
|
event_type=ChatGenerationTaskEventType.FINAL_POLL_BEFORE_TIMEOUT_ERROR.value,
|
|
message=str(exc),
|
|
)
|
|
await _mark_timeout(db, task, message="任务轮询超时")
|
|
return
|
|
|
|
task.retry_count = (task.retry_count or 0) + 1
|
|
|
|
# 视频轮询的临时异常不再 3 次内直接退款;继续降频到 24 小时最终 deadline。
|
|
if is_video_generation_task(task):
|
|
task.pipeline_stage = ChatGenerationPipelineStage.WAITING_REMOTE.value
|
|
await _schedule_next_poll(task, reason="poll_exception_retry")
|
|
await db.commit()
|
|
return
|
|
|
|
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 = ChatGenerationPipelineStage.WAITING_REMOTE.value
|
|
delay_seconds = int(settings.CHATAPI_ASYNC_RETRY_BACKOFF_SECONDS or 30) * int(task.retry_count or 1)
|
|
default_schedule = build_default_poll_schedule(
|
|
task,
|
|
now=_now(),
|
|
delay_seconds=delay_seconds,
|
|
reason="poll_exception_retry",
|
|
)
|
|
task.next_poll_at = default_schedule.next_poll_at
|
|
await db.commit()
|
|
|
|
await register_poll_active(
|
|
task,
|
|
check_at=_poll_check_at(delay_seconds=delay_seconds),
|
|
next_poll_at=default_schedule.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):
|
|
try:
|
|
return run_async(_run(task_id))
|
|
except Exception as exc:
|
|
# 只重试基础设施异常;供应商失败/业务失败已在 _run 内处理。
|
|
retries = int(getattr(self.request, "retries", 0) or 0) + 1
|
|
countdown = int(settings.CHATAPI_ASYNC_RETRY_BACKOFF_SECONDS or 30) * max(1, retries)
|
|
raise self.retry(exc=exc, countdown=countdown)
|
|
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()
|