Files
video-gen/video-gen-api/app/tasks/generation_poll_tasks.py
T
2026-07-20 13:48:17 +08:00

724 lines
25 KiB
Python

from __future__ import annotations
import json
from datetime import datetime, timedelta, timezone
from typing import Any
from app.config import settings
from app.enums.celery_queue import CeleryQueue
from app.enums.generation_status import GenerationRecordPipelineStage
from app.enums.generation_task import (
ALLOWED_GENERATION_MODES,
ChatGenerationPipelineStage,
ChatGenerationTaskEventType,
GenerationOwnerType,
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_provider_call, log_task_event
from app.services.generation.pipeline.db_lock_service import DatabaseRowLockBusy
from app.services.generation.pipeline.lifecycle_service import (
mark_owner_failed_and_refund_once,
notify_owner_finished,
)
from app.services.generation.pipeline.owner_service import (
GenerationOwner,
is_attempt_current,
load_generation_owner,
load_generation_owner_for_update_retry,
normalize_owner_type,
owner_is_generating,
owner_mode,
owner_provider_task_id,
owner_type_of,
redis_owner_item_id,
)
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.provider_service import poll_provider_task
from app.services.media_token_usage_snapshot_service import (
sync_chat_generation_task_media_token_snapshot,
sync_generation_record_media_token_snapshot,
)
from app.services.redis_registry_service import (
RedisExecutionLockError,
RedisExecutionLockLease,
datetime_to_epoch,
ensure_aware_utc,
redis_remove_registry_item,
redis_upsert_registry_item,
utc_now,
)
from app.tasks.async_runner import run_async
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 _stage(owner: GenerationOwner, chat_stage: ChatGenerationPipelineStage) -> str:
if isinstance(owner, ChatGenerationTask):
return chat_stage.value
try:
return GenerationRecordPipelineStage(chat_stage.value).value
except ValueError:
return chat_stage.value
def _engine_snapshot(owner: GenerationOwner) -> dict:
try:
data = json.loads(owner.engine_snapshot_json or "{}")
return data if isinstance(data, dict) else {}
except Exception:
return {}
async def _log_poll_provider_call_after_commit(
owner: GenerationOwner,
*,
provider_response: Any,
) -> None:
"""Provider logs use an independent session, so the owner row must be committed first."""
snapshot = _engine_snapshot(owner)
await log_provider_call(
owner,
provider=snapshot.get("provider") or "ark",
api_type=f"{owner.gen_type}_poll",
model=snapshot.get("model_name"),
engine_id=owner.engine_id,
status="success",
provider_task_id=owner_provider_task_id(owner),
response_data=provider_response,
)
def _registry_id(owner: GenerationOwner) -> str:
return redis_owner_item_id(
owner_type_of(owner), owner.id, int(owner.generation_attempt_no or 1)
)
def _poll_lease_until(now: datetime | None = None) -> datetime:
return (now or _now()) + timedelta(
seconds=int(settings.POLL_TASK_LEASE_SECONDS or 300)
)
def _lock_key(owner_type: str, owner_id: str, attempt_no: int) -> str:
return (
f"{settings.GENERATION_POLL_LOCK_KEY_PREFIX}:"
f"{owner_type}:{owner_id}:attempt:{attempt_no}"
)
def _build_poll_active_payload(
owner: GenerationOwner,
*,
reason: str,
next_poll_at: datetime | None,
check_at: datetime | None,
) -> dict[str, Any]:
return {
"owner_type": owner_type_of(owner),
"owner_id": owner.id,
"task_id": owner.id,
"generation_attempt_no": int(owner.generation_attempt_no or 1),
"provider_task_id": owner_provider_task_id(owner),
"generation_mode": owner_mode(owner),
"gen_type": owner.gen_type,
"stage": owner.pipeline_stage or "",
"queue": POLL_QUEUE,
"poll_count": int(owner.poll_count or 0),
"poll_error_count": int(getattr(owner, "poll_error_count", 0) or 0),
"manual_retry_count": int(getattr(owner, "manual_retry_count", 0) or 0),
"poll_started_at": datetime_to_epoch(owner.poll_started_at)
if owner.poll_started_at
else None,
"poll_interval_seconds": int(owner.poll_interval_seconds or 0),
"last_poll_at": datetime_to_epoch(owner.last_poll_at)
if owner.last_poll_at
else None,
"next_poll_at": datetime_to_epoch(ensure_aware_utc(next_poll_at))
if next_poll_at
else None,
"poll_lease_until": datetime_to_epoch(owner.poll_lease_until)
if owner.poll_lease_until
else None,
"deadline_at": datetime_to_epoch(owner.deadline_at)
if owner.deadline_at
else None,
"check_at": datetime_to_epoch(ensure_aware_utc(check_at))
if check_at
else None,
"updated_at": datetime_to_epoch(utc_now()),
"reason": reason,
}
async def register_poll_active(
owner: GenerationOwner,
*,
check_at: datetime,
reason: str,
next_poll_at: datetime | None = None,
) -> None:
await redis_upsert_registry_item(
hash_key=settings.POLL_ACTIVE_REDIS_HASH_KEY,
zset_key=settings.POLL_ACTIVE_REDIS_ZSET_KEY,
item_id=_registry_id(owner),
payload=_build_poll_active_payload(
owner,
reason=reason,
next_poll_at=next_poll_at,
check_at=check_at,
),
check_at=check_at,
log_context="poll_active",
)
async def remove_poll_active(
owner: GenerationOwner | None = None,
*,
owner_type: str | None = None,
owner_id: str | None = None,
attempt_no: int | None = None,
) -> None:
if owner is not None:
item_id = _registry_id(owner)
else:
item_id = redis_owner_item_id(owner_type, str(owner_id or ""), attempt_no)
await redis_remove_registry_item(
hash_key=settings.POLL_ACTIVE_REDIS_HASH_KEY,
zset_key=settings.POLL_ACTIVE_REDIS_ZSET_KEY,
item_id=item_id,
log_context="poll_active",
)
async def _sync_snapshot(
db, owner: GenerationOwner, provider_response: Any = None
) -> None:
if isinstance(owner, ChatGenerationTask):
await sync_chat_generation_task_media_token_snapshot(
db, owner, provider_response=provider_response
)
else:
await sync_generation_record_media_token_snapshot(
db, owner, provider_response=provider_response
)
async def _mark_failed(
db,
owner: GenerationOwner,
*,
message: str,
stage: ChatGenerationPipelineStage,
event_type: str,
detail: Any = None,
) -> None:
await mark_owner_failed_and_refund_once(
db,
owner,
error_message=message,
pipeline_stage=_stage(owner, stage),
)
owner.next_poll_at = None
owner.poll_claim_token = None
owner.poll_lease_until = None
await db.commit()
await notify_owner_finished(db, owner)
await db.commit()
await remove_poll_active(owner)
await log_task_event(
owner,
event_type=event_type,
message=message,
detail=detail,
to_stage=owner.pipeline_stage,
)
async def _schedule_next_poll(
db,
owner: GenerationOwner,
*,
reason: str,
default_delay_seconds: int | None = None,
) -> None:
current = _now()
schedule = (
build_video_pending_poll_schedule(owner, now=current)
if is_video_generation_task(owner)
else build_default_poll_schedule(
owner,
now=current,
delay_seconds=default_delay_seconds,
reason=reason,
)
)
owner.pipeline_stage = _stage(owner, ChatGenerationPipelineStage.WAITING_REMOTE)
owner.next_poll_at = schedule.next_poll_at
owner.poll_interval_seconds = schedule.poll_interval_seconds
owner.poll_claim_token = None
owner.poll_lease_until = None
await db.commit()
await register_poll_active(
owner,
check_at=schedule.next_poll_at,
next_poll_at=schedule.next_poll_at,
reason=schedule.reason,
)
await log_task_event(
owner,
event_type=ChatGenerationTaskEventType.POLL_SCHEDULED.value,
message=f"已登记下一次轮询。reason={schedule.reason}",
detail={
"next_poll_at": schedule.next_poll_at,
"delay_seconds": schedule.delay_seconds,
},
)
if schedule.direct_countdown:
poll_generation_task.apply_async(
args=[owner.id],
kwargs={
"owner_type": owner_type_of(owner),
"generation_attempt_no": int(owner.generation_attempt_no or 1),
"force_due": False,
},
queue=POLL_QUEUE,
countdown=max(0, int(schedule.delay_seconds)),
)
async def _resolve_attempt(
task_id: str, *, owner_type: str, message_attempt: int | None
) -> int | None:
async with async_session() as db:
owner = await load_generation_owner(
db, owner_type=owner_type, owner_id=task_id, for_update=False
)
if not owner:
return None
if not is_attempt_current(owner, message_attempt):
await remove_poll_active(
owner_type=owner_type,
owner_id=task_id,
attempt_no=message_attempt,
)
await log_task_event(
owner,
event_type=ChatGenerationTaskEventType.STALE_ATTEMPT_MESSAGE_SKIPPED.value,
message="轮询消息属于旧生成轮次,已跳过",
detail={
"message_attempt": message_attempt,
"current_attempt": owner.generation_attempt_no,
},
)
return None
return int(owner.generation_attempt_no or 1)
async def _restore_after_lock_error(
*,
owner_type: str,
owner_id: str,
attempt_no: int,
claim_token: str,
) -> None:
async with async_session() as db:
owner = await load_generation_owner(
db,
owner_type=owner_type,
owner_id=owner_id,
for_update=True,
)
if not owner or not is_attempt_current(owner, attempt_no):
return
if owner.poll_claim_token != claim_token:
return
owner.pipeline_stage = _stage(owner, ChatGenerationPipelineStage.WAITING_REMOTE)
owner.poll_claim_token = None
owner.poll_lease_until = None
owner.next_poll_at = _now()
await db.commit()
await register_poll_active(
owner,
check_at=owner.next_poll_at,
next_poll_at=owner.next_poll_at,
reason="poll_execution_lock_error",
)
async def _run(
task_id: str,
*,
owner_type: str = GenerationOwnerType.CHAT_GENERATION_TASK.value,
generation_attempt_no: int | None = None,
force_due: bool = False,
):
normalized_owner_type = normalize_owner_type(owner_type)
effective_attempt = await _resolve_attempt(
task_id,
owner_type=normalized_owner_type,
message_attempt=generation_attempt_no,
)
if effective_attempt is None:
return
lease = await RedisExecutionLockLease.acquire(
lock_key=_lock_key(normalized_owner_type, task_id, effective_attempt),
ttl_seconds=int(settings.GENERATION_POLL_LOCK_TTL_SECONDS or 300),
log_context="generation_poll",
renew_interval_seconds=int(
settings.REDIS_EXECUTION_LOCK_RENEW_INTERVAL_SECONDS or 30
),
)
if lease is None:
return
claim_started = False
try:
async with async_session() as db:
owner = await load_generation_owner(
db,
owner_type=normalized_owner_type,
owner_id=task_id,
for_update=True,
)
if not owner:
await remove_poll_active(
owner_type=normalized_owner_type,
owner_id=task_id,
attempt_no=effective_attempt,
)
return
if not is_attempt_current(owner, effective_attempt):
await remove_poll_active(owner)
return
if (
isinstance(owner, ChatGenerationTask)
and owner.generation_mode not in ALLOWED_GENERATION_MODES
):
await remove_poll_active(owner)
return
if not owner_is_generating(owner) or owner.pipeline_stage not in {
_stage(owner, ChatGenerationPipelineStage.WAITING_REMOTE),
_stage(owner, ChatGenerationPipelineStage.POLLING),
}:
await remove_poll_active(owner)
return
current = _now()
if is_video_generation_task(owner):
ensure_video_poll_fields(owner, now=current)
# Redis execution lock is the authoritative single-worker guard.
# A stale database lease left by a crashed worker must not block the
# worker that successfully acquired the current Redis lock.
if not force_due and is_poll_not_due(owner, now=current):
await db.commit()
await register_poll_active(
owner,
check_at=owner.next_poll_at,
next_poll_at=owner.next_poll_at,
reason="poll_task_not_due",
)
return
final_poll = is_final_poll_due(owner, now=current)
if not owner_provider_task_id(owner):
await _mark_failed(
db,
owner,
message="任务轮询超时" if final_poll else "缺少外部任务ID",
stage=ChatGenerationPipelineStage.TIMEOUT
if final_poll
else ChatGenerationPipelineStage.FAILED,
event_type=ChatGenerationTaskEventType.TASK_TIMEOUT.value
if final_poll
else ChatGenerationTaskEventType.POLL_FAILED.value,
)
return
owner.pipeline_stage = _stage(owner, ChatGenerationPipelineStage.POLLING)
owner.poll_claim_token = lease.token
owner.poll_lease_until = _poll_lease_until(current)
owner.poll_count = int(owner.poll_count or 0) + 1
owner.last_poll_at = current
await db.commit()
claim_started = True
await register_poll_active(
owner,
check_at=owner.poll_lease_until,
next_poll_at=owner.next_poll_at,
reason="polling_lease",
)
try:
poll_result = await poll_provider_task(db, owner)
await lease.ensure_owned()
owner = await load_generation_owner_for_update_retry(
db,
owner_type=normalized_owner_type,
owner_id=task_id,
)
if not owner or not is_attempt_current(owner, effective_attempt):
await remove_poll_active(
owner_type=normalized_owner_type,
owner_id=task_id,
attempt_no=effective_attempt,
)
return
if owner.poll_claim_token != lease.token:
return
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}
if _is_success(status):
if owner.gen_type == GenerationType.IMAGE.value:
owner.remote_result_url = poll_result.get("image_url")
owner.image_tokens_used = int(
poll_result.get("image_tokens", 0) or 0
)
else:
owner.remote_result_url = poll_result.get("video_url")
owner.video_tokens_used = int(
poll_result.get("video_tokens", 0) or 0
)
owner.provider_response_json = response_data
await _sync_snapshot(db, owner, response_data)
if not owner.remote_result_url:
await _mark_failed(
db,
owner,
message="供应商任务成功但未返回结果URL",
stage=ChatGenerationPipelineStage.FAILED,
event_type=ChatGenerationTaskEventType.POLL_FAILED.value,
detail=poll_result,
)
await _log_poll_provider_call_after_commit(
owner, provider_response=provider_response
)
return
owner.pipeline_stage = _stage(
owner, ChatGenerationPipelineStage.RESULT_READY
)
owner.poll_error_count = 0
owner.poll_claim_token = None
owner.poll_lease_until = None
owner.next_poll_at = None
await db.commit()
await _log_poll_provider_call_after_commit(
owner, provider_response=provider_response
)
await remove_poll_active(owner)
await log_task_event(
owner,
event_type=ChatGenerationTaskEventType.POLL_SUCCESS.value,
to_stage=owner.pipeline_stage,
)
from app.tasks.generation_download_tasks import (
enqueue_download_task,
)
await enqueue_download_task(
db, owner, reason="poll_success_result_ready"
)
return
if _is_failed(status):
owner.provider_response_json = response_data
await _mark_failed(
db,
owner,
message=poll_result.get("error")
or f"供应商任务失败: {status}",
stage=ChatGenerationPipelineStage.FAILED,
event_type=ChatGenerationTaskEventType.POLL_FAILED.value,
detail=poll_result,
)
await _log_poll_provider_call_after_commit(
owner, provider_response=provider_response
)
return
if final_poll:
await _mark_failed(
db,
owner,
message="任务轮询超时",
stage=ChatGenerationPipelineStage.TIMEOUT,
event_type=ChatGenerationTaskEventType.TASK_TIMEOUT.value,
detail=poll_result,
)
await _log_poll_provider_call_after_commit(
owner, provider_response=provider_response
)
return
owner.poll_error_count = 0
await _schedule_next_poll(
db, owner, reason="poll_pending_next"
)
await _log_poll_provider_call_after_commit(
owner, provider_response=provider_response
)
await log_task_event(
owner,
event_type=ChatGenerationTaskEventType.POLL_PENDING.value,
message=f"status={status}",
)
except (RedisExecutionLockError, DatabaseRowLockBusy):
try:
await db.rollback()
except Exception:
pass
raise
except Exception as exc:
try:
await db.rollback()
except Exception:
pass
await lease.ensure_owned()
owner = await load_generation_owner_for_update_retry(
db,
owner_type=normalized_owner_type,
owner_id=task_id,
)
if not owner or not is_attempt_current(owner, effective_attempt):
await remove_poll_active(
owner_type=normalized_owner_type,
owner_id=task_id,
attempt_no=effective_attempt,
)
return
if owner.poll_claim_token != lease.token:
return
if is_final_poll_due(owner, now=_now()):
await _mark_failed(
db,
owner,
message="任务轮询超时",
stage=ChatGenerationPipelineStage.TIMEOUT,
event_type=ChatGenerationTaskEventType.TASK_TIMEOUT.value,
)
return
owner.poll_error_count = int(owner.poll_error_count or 0) + 1
if is_video_generation_task(owner):
await _schedule_next_poll(
db, owner, reason="poll_exception_retry"
)
return
if owner.poll_error_count > int(
settings.CHATAPI_ASYNC_MAX_RETRIES or 3
):
error_message = (
extract_error_message(exc, "轮询")
if callable(extract_error_message)
else str(exc)
)
await _mark_failed(
db,
owner,
message=error_message,
stage=ChatGenerationPipelineStage.FAILED,
event_type=ChatGenerationTaskEventType.POLL_FAILED.value,
)
else:
delay = int(
settings.CHATAPI_ASYNC_RETRY_BACKOFF_SECONDS or 30
) * owner.poll_error_count
await _schedule_next_poll(
db,
owner,
reason="poll_exception_retry",
default_delay_seconds=delay,
)
except (RedisExecutionLockError, DatabaseRowLockBusy):
if claim_started:
try:
await _restore_after_lock_error(
owner_type=normalized_owner_type,
owner_id=task_id,
attempt_no=effective_attempt,
claim_token=lease.token,
)
except Exception:
logger.exception(
"轮询执行锁异常后的状态恢复失败 owner_type=%s owner_id=%s attempt=%s",
normalized_owner_type,
task_id,
effective_attempt,
)
raise
finally:
await lease.close()
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,
force_due: bool = False,
owner_type: str = GenerationOwnerType.CHAT_GENERATION_TASK.value,
generation_attempt_no: int | None = None,
):
try:
return run_async(
_run(
task_id,
owner_type=owner_type,
generation_attempt_no=generation_attempt_no,
force_due=bool(force_due),
)
)
except Exception as exc:
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()