Merge branch 'main' of gitee.com:wg123/video-gen into main
This commit is contained in:
@@ -43,6 +43,7 @@ from app.services.generation_billing_service import (
|
||||
get_next_credit_attempt_no,
|
||||
)
|
||||
from app.services.generation_refund_service import mark_generation_record_failed_and_refund_once
|
||||
from app.services.media_token_usage_snapshot_service import sync_generation_record_media_token_snapshot
|
||||
from app.services.credit_record_meta_service import build_generation_record_prompt_meta
|
||||
from app.services.video_cover_service import async_create_video_cover_for_local_video
|
||||
from app.utils.id_gen import generate_id
|
||||
@@ -724,6 +725,7 @@ async def seedance_callback(request: Request, db: AsyncSession = Depends(get_db)
|
||||
usage = data.get("usage", {})
|
||||
if usage:
|
||||
record.video_tokens_used = usage.get("total_tokens", 0)
|
||||
await sync_generation_record_media_token_snapshot(db, record, provider_response=data)
|
||||
# Log callback response
|
||||
from app.services.video_gen import _log_video_response
|
||||
_log_video_response(record.id, data)
|
||||
|
||||
+29
-11
@@ -113,8 +113,8 @@ class Settings(BaseSettings):
|
||||
PROVIDER_LIMIT_TOKEN_TTL_SECONDS: int = 600
|
||||
|
||||
CELERY_DB_POOL_SIZE: int = 1
|
||||
CELERY_DB_MAX_OVERFLOW: int = 1
|
||||
CELERY_DB_POOL_TIMEOUT: int = 30
|
||||
CELERY_DB_MAX_OVERFLOW: int = 2
|
||||
CELERY_DB_POOL_TIMEOUT: int = 60
|
||||
CELERY_DB_POOL_RECYCLE: int = 1800
|
||||
|
||||
# Celery 图片/视频下载容灾配置。
|
||||
@@ -125,30 +125,48 @@ class Settings(BaseSettings):
|
||||
DOWNLOAD_TASK_RETRY_BACKOFF_SECONDS: int = 30
|
||||
DOWNLOAD_TASK_LEASE_SECONDS: int = 10 * 60
|
||||
DOWNLOAD_TASK_QUEUE_TIMEOUT_SECONDS: int = 5 * 60
|
||||
DOWNLOAD_RECOVERY_BATCH_SIZE: int = 100
|
||||
DOWNLOAD_RECOVERY_BATCH_SIZE: int = 20
|
||||
DOWNLOAD_RECOVERY_STARTUP_DELAY_SECONDS: int = 3
|
||||
# 下载恢复自循环:不依赖 Celery beat,不新增 worker;由 gen_result_download 队列周期扫描 DB/Redis。
|
||||
DOWNLOAD_RECOVERY_LOOP_ENABLED: bool = False
|
||||
DOWNLOAD_RECOVERY_INTERVAL_SECONDS: int = 60
|
||||
DOWNLOAD_RECOVERY_LOOP_LOCK_KEY: str = "vg:celery:download_recovery_loop_lock"
|
||||
DOWNLOAD_RECOVERY_LOOP_LOCK_TTL_SECONDS: int = 55
|
||||
DOWNLOAD_RETRY_COUNTDOWN_EXTRA_SECONDS: int = 1
|
||||
DOWNLOAD_NON_RETRYABLE_LOCAL_ERRORS: bool = True
|
||||
DOWNLOAD_EVENT_VERBOSE_ENABLED: bool = True
|
||||
MEDIA_TOKEN_SNAPSHOT_ENABLED: bool = True
|
||||
DOWNLOAD_ACTIVE_REDIS_HASH_KEY: str = "vg:celery:download:active"
|
||||
DOWNLOAD_ACTIVE_REDIS_ZSET_KEY: str = "vg:celery:download:active_index"
|
||||
|
||||
# Celery 生成链路 / provider poll 容灾配置。
|
||||
# 说明:
|
||||
# - 不新增 Celery worker;恢复任务仍投递到 gen_result_download。
|
||||
# - worker_ready 每个 worker 都会尝试抢启动恢复锁,只有抢到锁的 worker 投递恢复任务。
|
||||
# - 启动容灾保留,但恢复扫描独立投递到 CELERY_RECOVERY_QUEUE。
|
||||
# - worker_ready 每个 worker 都会尝试抢启动恢复锁,只有抢到锁的 worker 投递恢复协调任务。
|
||||
# - poll active 使用独立 Redis key,避免影响稳定的下载 active 注册表。
|
||||
GENERATION_RECOVERY_BATCH_SIZE: int = 100
|
||||
GENERATION_RECOVERY_MAX_ROUNDS: int = 5
|
||||
POLL_RECOVERY_BATCH_SIZE: int = 100
|
||||
GENERATION_RECOVERY_BATCH_SIZE: int = 20
|
||||
GENERATION_RECOVERY_MAX_ROUNDS: int = 1
|
||||
POLL_RECOVERY_BATCH_SIZE: int = 20
|
||||
POLL_TASK_LEASE_SECONDS: int = 5 * 60
|
||||
POLL_TASK_QUEUE_TIMEOUT_SECONDS: int = 2 * 60
|
||||
POLL_ACTIVE_REDIS_HASH_KEY: str = "vg:celery:poll:active"
|
||||
POLL_ACTIVE_REDIS_ZSET_KEY: str = "vg:celery:poll:active_index"
|
||||
CELERY_STARTUP_RECOVERY_LOCK_KEY: str = "vg:celery:startup_recovery_lock"
|
||||
CELERY_STARTUP_RECOVERY_LOCK_TTL_SECONDS: int = 120
|
||||
CELERY_RECOVERY_QUEUE: str = "gen_recovery"
|
||||
CELERY_RECOVERY_TASK_LOCK_TTL_SECONDS: int = 10 * 60
|
||||
CELERY_RECOVERY_SOFT_TIME_LIMIT_SECONDS: int = 300
|
||||
CELERY_RECOVERY_TIME_LIMIT_SECONDS: int = 420
|
||||
CELERY_RECOVERY_STARTUP_TASK_LOCK_KEY: str = "vg:celery:startup_recovery_task_lock"
|
||||
GENERATION_RECOVERY_LOCK_KEY: str = "vg:celery:generation_recovery_lock"
|
||||
DOWNLOAD_RECOVERY_LOCK_KEY: str = "vg:celery:download_recovery_lock"
|
||||
MODULE_ASYNC_RECOVERY_LOCK_KEY: str = "vg:celery:module_async_recovery_lock"
|
||||
SHOT_SPLIT_RECOVERY_LOCK_KEY: str = "vg:celery:shot_split_recovery_lock"
|
||||
|
||||
# 模块异步任务容灾配置。
|
||||
# 覆盖 ModuleGenerationStep 提词任务、shot 原视频/片段分析、shot ffmpeg 切割 active 注册。
|
||||
# 不新增 worker 队列:恢复扫描仍走 gen_result_download,真实业务任务回到原始队列。
|
||||
MODULE_ASYNC_RECOVERY_BATCH_SIZE: int = 100
|
||||
# 恢复扫描走 CELERY_RECOVERY_QUEUE,真实业务任务回到原始队列。
|
||||
MODULE_ASYNC_RECOVERY_BATCH_SIZE: int = 20
|
||||
MODULE_ASYNC_QUEUE_TIMEOUT_SECONDS: int = 5 * 60
|
||||
MODULE_ASYNC_LEASE_SECONDS: int = 10 * 60
|
||||
MODULE_ASYNC_LOCK_TTL_SECONDS: int = 10 * 60
|
||||
@@ -192,7 +210,7 @@ class Settings(BaseSettings):
|
||||
SHOT_SPLIT_RETRY_BACKOFF_SECONDS: int = 30
|
||||
SHOT_SPLIT_LEASE_SECONDS: int = 10 * 60
|
||||
SHOT_SPLIT_PENDING_TIMEOUT_SECONDS: int = 5 * 60
|
||||
SHOT_SPLIT_RECOVERY_BATCH_SIZE: int = 50
|
||||
SHOT_SPLIT_RECOVERY_BATCH_SIZE: int = 20
|
||||
SHOT_SPLIT_LOCK_KEY_PREFIX: str = "vg:shot_replicate:split:lock"
|
||||
SHOT_SPLIT_SEMAPHORE_KEY_PREFIX: str = "vg:shot_replicate:split:semaphore"
|
||||
|
||||
|
||||
@@ -6,3 +6,4 @@ from app.enums.video_prompt_schema import *
|
||||
from app.enums.user import *
|
||||
from app.enums.credit_record import *
|
||||
from app.enums.token_usage import *
|
||||
from app.enums.generation_task import *
|
||||
|
||||
@@ -0,0 +1,101 @@
|
||||
from enum import Enum
|
||||
|
||||
|
||||
class GenerationMode(str, Enum):
|
||||
CHATAPI_ASYNC = "chatapi_async"
|
||||
HOT_OPENING_REPLICATE = "hot_opening_replicate"
|
||||
SHOT_REPLICATE = "shot_replicate"
|
||||
|
||||
|
||||
class GenerationType(str, Enum):
|
||||
IMAGE = "image"
|
||||
VIDEO = "video"
|
||||
|
||||
|
||||
class ChatGenerationTaskStatus(str, Enum):
|
||||
PENDING = "pending"
|
||||
GENERATING = "generating"
|
||||
COMPLETED = "completed"
|
||||
FAILED = "failed"
|
||||
|
||||
|
||||
class ChatGenerationPipelineStage(str, Enum):
|
||||
QUEUED = "queued"
|
||||
PREPARING = "preparing"
|
||||
CREATING_PROVIDER_TASK = "creating_provider_task"
|
||||
WAITING_REMOTE = "waiting_remote"
|
||||
POLLING = "polling"
|
||||
RESULT_READY = "result_ready"
|
||||
DOWNLOAD_QUEUED = "download_queued"
|
||||
DOWNLOADING = "downloading"
|
||||
RETRY_WAITING = "retry_waiting"
|
||||
DONE = "done"
|
||||
FAILED = "failed"
|
||||
TIMEOUT = "timeout"
|
||||
DOWNLOAD_FAILED = "download_failed"
|
||||
|
||||
|
||||
class ChatGenerationTaskEventType(str, Enum):
|
||||
PROMPT_CONCAT_START = "PROMPT_CONCAT_START"
|
||||
PROMPT_CONCAT_SUCCESS = "PROMPT_CONCAT_SUCCESS"
|
||||
|
||||
PROVIDER_CREATE_START = "PROVIDER_CREATE_START"
|
||||
PROVIDER_CREATE_SUCCESS = "PROVIDER_CREATE_SUCCESS"
|
||||
PROVIDER_CREATE_FAILED = "PROVIDER_CREATE_FAILED"
|
||||
|
||||
POLL_START = "POLL_START"
|
||||
POLL_PENDING = "POLL_PENDING"
|
||||
POLL_SUCCESS = "POLL_SUCCESS"
|
||||
POLL_FAILED = "POLL_FAILED"
|
||||
POLL_SUCCESS_AFTER_TIMEOUT_RECOVERY = "POLL_SUCCESS_AFTER_TIMEOUT_RECOVERY"
|
||||
FINAL_POLL_BEFORE_TIMEOUT_ERROR = "FINAL_POLL_BEFORE_TIMEOUT_ERROR"
|
||||
FINAL_POLL_BEFORE_TIMEOUT_PENDING = "FINAL_POLL_BEFORE_TIMEOUT_PENDING"
|
||||
GENERATION_RECOVERY_ENQUEUE = "GENERATION_RECOVERY_ENQUEUE"
|
||||
|
||||
DOWNLOAD_ENQUEUE = "DOWNLOAD_ENQUEUE"
|
||||
DOWNLOAD_ENQUEUE_FAILED = "DOWNLOAD_ENQUEUE_FAILED"
|
||||
DOWNLOAD_START = "DOWNLOAD_START"
|
||||
DOWNLOAD_SUCCESS = "DOWNLOAD_SUCCESS"
|
||||
DOWNLOAD_RETRY_WAITING = "DOWNLOAD_RETRY_WAITING"
|
||||
DOWNLOAD_RETRY_ENQUEUE = "DOWNLOAD_RETRY_ENQUEUE"
|
||||
DOWNLOAD_RETRY_ENQUEUE_FAILED = "DOWNLOAD_RETRY_ENQUEUE_FAILED"
|
||||
DOWNLOAD_RETRY_NOT_DUE = "DOWNLOAD_RETRY_NOT_DUE"
|
||||
DOWNLOAD_STUCK_RECOVER = "DOWNLOAD_STUCK_RECOVER"
|
||||
DOWNLOAD_RECOVERY_ENQUEUE = "DOWNLOAD_RECOVERY_ENQUEUE"
|
||||
DOWNLOAD_RECOVERY_ENQUEUE_FAILED = "DOWNLOAD_RECOVERY_ENQUEUE_FAILED"
|
||||
DOWNLOAD_FAILED = "DOWNLOAD_FAILED"
|
||||
DOWNLOAD_FAILED_NON_RETRYABLE = "DOWNLOAD_FAILED_NON_RETRYABLE"
|
||||
|
||||
DOWNLOAD_SKIP_TASK_MISSING = "DOWNLOAD_SKIP_TASK_MISSING"
|
||||
DOWNLOAD_SKIP_INVALID_MODE = "DOWNLOAD_SKIP_INVALID_MODE"
|
||||
DOWNLOAD_SKIP_NOT_GENERATING = "DOWNLOAD_SKIP_NOT_GENERATING"
|
||||
DOWNLOAD_SKIP_ALREADY_COMPLETED = "DOWNLOAD_SKIP_ALREADY_COMPLETED"
|
||||
DOWNLOAD_SKIP_NO_REMOTE_RESULT_URL = "DOWNLOAD_SKIP_NO_REMOTE_RESULT_URL"
|
||||
DOWNLOAD_SKIP_STAGE_NOT_ALLOWED = "DOWNLOAD_SKIP_STAGE_NOT_ALLOWED"
|
||||
DOWNLOAD_SKIP_DOWNLOADING_LEASE_ALIVE = "DOWNLOAD_SKIP_DOWNLOADING_LEASE_ALIVE"
|
||||
DOWNLOAD_SKIP_RETRY_NOT_DUE = "DOWNLOAD_SKIP_RETRY_NOT_DUE"
|
||||
DOWNLOAD_SKIP_FINAL_STATE = "DOWNLOAD_SKIP_FINAL_STATE"
|
||||
DOWNLOAD_SKIP_DISABLED = "DOWNLOAD_SKIP_DISABLED"
|
||||
|
||||
TASK_TIMEOUT = "TASK_TIMEOUT"
|
||||
|
||||
|
||||
ALLOWED_GENERATION_MODES = {
|
||||
GenerationMode.CHATAPI_ASYNC.value,
|
||||
GenerationMode.HOT_OPENING_REPLICATE.value,
|
||||
GenerationMode.SHOT_REPLICATE.value,
|
||||
}
|
||||
|
||||
FINAL_CHAT_GENERATION_STAGES = {
|
||||
ChatGenerationPipelineStage.DONE.value,
|
||||
ChatGenerationPipelineStage.FAILED.value,
|
||||
ChatGenerationPipelineStage.TIMEOUT.value,
|
||||
ChatGenerationPipelineStage.DOWNLOAD_FAILED.value,
|
||||
}
|
||||
|
||||
DOWNLOAD_RECOVERABLE_STAGES = {
|
||||
ChatGenerationPipelineStage.RESULT_READY.value,
|
||||
ChatGenerationPipelineStage.DOWNLOAD_QUEUED.value,
|
||||
ChatGenerationPipelineStage.DOWNLOADING.value,
|
||||
ChatGenerationPipelineStage.RETRY_WAITING.value,
|
||||
}
|
||||
@@ -2,36 +2,40 @@ from __future__ import annotations
|
||||
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.enums.generation_task import ChatGenerationTaskStatus, GenerationMode
|
||||
from app.models.chat_generation_task import ChatGenerationTask
|
||||
|
||||
|
||||
async def notify_chat_generation_task_finished(db: AsyncSession, task: ChatGenerationTask) -> None:
|
||||
"""通知业务模块 ChatGenerationTask 已进入终态。
|
||||
|
||||
当前用于爆款开头复刻:
|
||||
- image_generate 完成后自动进入 video_prompt_optimize
|
||||
- video_generate 完成后总任务完成
|
||||
该方法必须幂等:下载恢复任务、重试任务、服务重启补偿都可能重复调用。
|
||||
具体模块服务需要自行判断 step/project 是否已经完成或失败,避免重复推进。
|
||||
"""
|
||||
if not task:
|
||||
return
|
||||
if task.generation_mode == "hot_opening_replicate":
|
||||
|
||||
status = getattr(task, "status", None)
|
||||
generation_mode = getattr(task, "generation_mode", None)
|
||||
|
||||
if generation_mode == GenerationMode.HOT_OPENING_REPLICATE.value:
|
||||
from app.services.hot_opening_replicate_service import (
|
||||
handle_chat_generation_task_completed,
|
||||
handle_chat_generation_task_failed,
|
||||
)
|
||||
if task.status == "completed":
|
||||
if status == ChatGenerationTaskStatus.COMPLETED.value:
|
||||
await handle_chat_generation_task_completed(db, task)
|
||||
elif task.status == "failed":
|
||||
elif status == ChatGenerationTaskStatus.FAILED.value:
|
||||
await handle_chat_generation_task_failed(db, task)
|
||||
return
|
||||
|
||||
if task.generation_mode == "shot_replicate":
|
||||
if generation_mode == GenerationMode.SHOT_REPLICATE.value:
|
||||
from app.services.shot_replicate_flow_service import (
|
||||
handle_chat_generation_task_completed,
|
||||
handle_chat_generation_task_failed,
|
||||
)
|
||||
if task.status == "completed":
|
||||
if status == ChatGenerationTaskStatus.COMPLETED.value:
|
||||
await handle_chat_generation_task_completed(db, task)
|
||||
elif task.status == "failed":
|
||||
elif status == ChatGenerationTaskStatus.FAILED.value:
|
||||
await handle_chat_generation_task_failed(db, task)
|
||||
return
|
||||
|
||||
@@ -9,6 +9,12 @@ from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.config import settings
|
||||
from app.enums.generation_task import (
|
||||
ALLOWED_GENERATION_MODES,
|
||||
ChatGenerationPipelineStage,
|
||||
ChatGenerationTaskEventType,
|
||||
ChatGenerationTaskStatus,
|
||||
)
|
||||
from app.models.chat_generation_task import ChatGenerationTask
|
||||
from app.services.celery_download_recovery_service import (
|
||||
ensure_aware_utc,
|
||||
@@ -17,9 +23,8 @@ from app.services.celery_download_recovery_service import (
|
||||
postpone_download_active_check,
|
||||
remove_download_active,
|
||||
)
|
||||
from app.services.generation_log_service import log_provider_call, log_task_event
|
||||
from app.services.generation_log_service import log_task_event
|
||||
from app.services.generation_module_hook_service import notify_chat_generation_task_finished
|
||||
from app.services.generation_provider_service import poll_provider_task
|
||||
from app.services.generation_refund_service import mark_chat_generation_task_failed_and_refund_once
|
||||
from app.services.redis_registry_service import (
|
||||
redis_get_due_registry_ids,
|
||||
@@ -30,7 +35,6 @@ from app.services.redis_registry_service import (
|
||||
|
||||
logger = logging.getLogger("video_gen")
|
||||
|
||||
ALLOWED_GENERATION_MODES = {"chatapi_async", "hot_opening_replicate", "shot_replicate"}
|
||||
POLL_QUEUE = "gen_provider_poll"
|
||||
|
||||
|
||||
@@ -59,11 +63,11 @@ def _is_queue_timeout(task: ChatGenerationTask, now: datetime | None = None) ->
|
||||
|
||||
|
||||
def _is_final_task_state(task: ChatGenerationTask) -> bool:
|
||||
return task.status in ("completed", "failed") or task.pipeline_stage in (
|
||||
"done",
|
||||
"failed",
|
||||
"timeout",
|
||||
"download_failed",
|
||||
return task.status in (ChatGenerationTaskStatus.COMPLETED.value, ChatGenerationTaskStatus.FAILED.value) or task.pipeline_stage in (
|
||||
ChatGenerationPipelineStage.DONE.value,
|
||||
ChatGenerationPipelineStage.FAILED.value,
|
||||
ChatGenerationPipelineStage.TIMEOUT.value,
|
||||
ChatGenerationPipelineStage.DOWNLOAD_FAILED.value,
|
||||
)
|
||||
|
||||
|
||||
@@ -137,19 +141,31 @@ async def recover_one_download_task(
|
||||
if _is_final_task_state(task):
|
||||
await remove_download_active(task.id)
|
||||
return "clean_final_state"
|
||||
if task.status != "generating":
|
||||
if task.status != ChatGenerationTaskStatus.GENERATING.value:
|
||||
await remove_download_active(task.id)
|
||||
await log_task_event(
|
||||
task,
|
||||
event_type=ChatGenerationTaskEventType.DOWNLOAD_SKIP_NOT_GENERATING.value,
|
||||
message=f"{source} 下载恢复跳过:任务不是 generating",
|
||||
detail={"status": task.status, "stage": task.pipeline_stage},
|
||||
)
|
||||
return "clean_not_generating"
|
||||
if not task.remote_result_url:
|
||||
await log_task_event(
|
||||
task,
|
||||
event_type=ChatGenerationTaskEventType.DOWNLOAD_SKIP_NO_REMOTE_RESULT_URL.value,
|
||||
message=f"{source} 下载恢复跳过:缺少 remote_result_url",
|
||||
detail={"status": task.status, "stage": task.pipeline_stage},
|
||||
)
|
||||
return "skip_no_remote_result_url"
|
||||
|
||||
stage = task.pipeline_stage
|
||||
redis_payload = payload or {}
|
||||
|
||||
if stage == "result_ready":
|
||||
if stage == ChatGenerationPipelineStage.RESULT_READY.value:
|
||||
await log_task_event(
|
||||
task,
|
||||
event_type="DOWNLOAD_RECOVERY_ENQUEUE",
|
||||
event_type=ChatGenerationTaskEventType.DOWNLOAD_RECOVERY_ENQUEUE.value,
|
||||
message=f"{source} 发现 result_ready 未完成下载,启动时恢复投递下载任务",
|
||||
detail={"payload": redis_payload},
|
||||
)
|
||||
@@ -165,7 +181,7 @@ async def recover_one_download_task(
|
||||
if _is_queue_timeout(task, current_time):
|
||||
await log_task_event(
|
||||
task,
|
||||
event_type="DOWNLOAD_RECOVERY_ENQUEUE",
|
||||
event_type=ChatGenerationTaskEventType.DOWNLOAD_RECOVERY_ENQUEUE.value,
|
||||
message=f"{source} 发现 download_queued 长时间未消费,启动时恢复投递下载任务",
|
||||
detail={"payload": redis_payload},
|
||||
)
|
||||
@@ -188,7 +204,7 @@ async def recover_one_download_task(
|
||||
if _is_expired(task.download_lease_until, current_time):
|
||||
await log_task_event(
|
||||
task,
|
||||
event_type="DOWNLOAD_RECOVERY_ENQUEUE",
|
||||
event_type=ChatGenerationTaskEventType.DOWNLOAD_RECOVERY_ENQUEUE.value,
|
||||
message=f"{source} 发现 downloading lease 过期,启动时恢复投递下载任务",
|
||||
detail={"payload": redis_payload},
|
||||
)
|
||||
@@ -211,7 +227,7 @@ async def recover_one_download_task(
|
||||
if _is_expired(task.download_next_retry_at, current_time):
|
||||
await log_task_event(
|
||||
task,
|
||||
event_type="DOWNLOAD_RECOVERY_ENQUEUE",
|
||||
event_type=ChatGenerationTaskEventType.DOWNLOAD_RECOVERY_ENQUEUE.value,
|
||||
message=f"{source} 发现 retry_waiting 到期,启动时恢复投递下载任务",
|
||||
detail={"payload": redis_payload},
|
||||
)
|
||||
@@ -276,11 +292,16 @@ async def recover_download_tasks_once(db: AsyncSession) -> dict[str, Any]:
|
||||
select(ChatGenerationTask)
|
||||
.where(
|
||||
ChatGenerationTask.deleted_at.is_(None),
|
||||
ChatGenerationTask.generation_mode.in_(["chatapi_async", "hot_opening_replicate", "shot_replicate"]),
|
||||
ChatGenerationTask.status == "generating",
|
||||
ChatGenerationTask.generation_mode.in_(list(ALLOWED_GENERATION_MODES)),
|
||||
ChatGenerationTask.status == ChatGenerationTaskStatus.GENERATING.value,
|
||||
ChatGenerationTask.remote_result_url.is_not(None),
|
||||
ChatGenerationTask.pipeline_stage.in_(
|
||||
["result_ready", "download_queued", "downloading", "retry_waiting"]
|
||||
[
|
||||
ChatGenerationPipelineStage.RESULT_READY.value,
|
||||
ChatGenerationPipelineStage.DOWNLOAD_QUEUED.value,
|
||||
ChatGenerationPipelineStage.DOWNLOADING.value,
|
||||
ChatGenerationPipelineStage.RETRY_WAITING.value,
|
||||
]
|
||||
),
|
||||
)
|
||||
.order_by(ChatGenerationTask.updated_at.asc())
|
||||
@@ -314,7 +335,7 @@ async def _mark_timeout(
|
||||
db,
|
||||
task=task,
|
||||
error_message=error_message,
|
||||
pipeline_stage="timeout",
|
||||
pipeline_stage=ChatGenerationPipelineStage.TIMEOUT.value,
|
||||
)
|
||||
await notify_chat_generation_task_finished(db, task)
|
||||
await db.commit()
|
||||
@@ -323,7 +344,7 @@ async def _mark_timeout(
|
||||
task,
|
||||
event_type="TASK_TIMEOUT",
|
||||
to_status="failed",
|
||||
to_stage="timeout",
|
||||
to_stage=ChatGenerationPipelineStage.TIMEOUT.value,
|
||||
)
|
||||
return "mark_timeout"
|
||||
|
||||
@@ -340,7 +361,7 @@ async def _mark_failed(
|
||||
db,
|
||||
task=task,
|
||||
error_message=error_message,
|
||||
pipeline_stage="failed",
|
||||
pipeline_stage=ChatGenerationPipelineStage.FAILED.value,
|
||||
)
|
||||
await notify_chat_generation_task_finished(db, task)
|
||||
await db.commit()
|
||||
@@ -349,93 +370,6 @@ async def _mark_failed(
|
||||
return "mark_failed"
|
||||
|
||||
|
||||
async def _try_final_poll_before_timeout(db: AsyncSession, task: ChatGenerationTask) -> str:
|
||||
"""超时前最后查一次供应商,避免 Celery 中断导致本地假超时。
|
||||
|
||||
如果供应商已经成功,继续进入下载;如果仍 running 或查询失败,再按超时处理。
|
||||
"""
|
||||
from app.tasks.generation_download_tasks import enqueue_download_task
|
||||
|
||||
if not (task.provider_task_id or task.seedance_task_id):
|
||||
return await _mark_timeout(db, task)
|
||||
|
||||
try:
|
||||
poll_result = await poll_provider_task(db, task)
|
||||
status = poll_result.get("status")
|
||||
response_data = poll_result.get("response_data")
|
||||
except Exception as exc:
|
||||
await log_task_event(
|
||||
task,
|
||||
event_type="FINAL_POLL_BEFORE_TIMEOUT_ERROR",
|
||||
message=str(exc),
|
||||
)
|
||||
return await _mark_timeout(db, task)
|
||||
|
||||
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}_final_poll_before_timeout",
|
||||
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:
|
||||
return await _mark_failed(
|
||||
db,
|
||||
task,
|
||||
error_message="供应商任务成功但未返回结果URL",
|
||||
detail=poll_result,
|
||||
)
|
||||
|
||||
task.pipeline_stage = "result_ready"
|
||||
task.retry_count = 0
|
||||
await db.commit()
|
||||
await _remove_poll_active(task.id)
|
||||
await log_task_event(
|
||||
task,
|
||||
event_type="POLL_SUCCESS_AFTER_TIMEOUT_RECOVERY",
|
||||
to_stage="result_ready",
|
||||
detail=poll_result,
|
||||
)
|
||||
await enqueue_download_task(db, task, recover=True, reason="final_poll_before_timeout_success")
|
||||
return "recover_timeout_success_to_download"
|
||||
|
||||
if _is_failed(status):
|
||||
task.provider_response_json = response_data
|
||||
return await _mark_failed(
|
||||
db,
|
||||
task,
|
||||
error_message=poll_result.get("error") or f"供应商任务失败: {status}",
|
||||
detail=poll_result,
|
||||
)
|
||||
|
||||
await log_task_event(
|
||||
task,
|
||||
event_type="FINAL_POLL_BEFORE_TIMEOUT_PENDING",
|
||||
message=f"status={status}",
|
||||
detail=poll_result,
|
||||
)
|
||||
return await _mark_timeout(db, task)
|
||||
|
||||
|
||||
async def recover_one_generation_task(
|
||||
db: AsyncSession,
|
||||
task: ChatGenerationTask,
|
||||
@@ -443,6 +377,14 @@ async def recover_one_generation_task(
|
||||
payload: dict[str, Any] | None = None,
|
||||
source: str = "startup_db",
|
||||
) -> str:
|
||||
"""恢复单个生成任务。
|
||||
|
||||
分流原则:
|
||||
1. 已有 remote_result_url:只恢复下载,不 poll,不重新 create。
|
||||
2. 已有 provider_task_id/seedance_task_id:恢复 poll。
|
||||
3. 无结果 URL、无供应商任务 ID:deadline 未过才恢复 create。
|
||||
4. 无结果 URL、无供应商任务 ID:deadline 已过直接超时失败,不再补救生成。
|
||||
"""
|
||||
from app.tasks.generation_create_tasks import chatapi_create_generation_task
|
||||
from app.tasks.generation_download_tasks import enqueue_download_task
|
||||
from app.tasks.generation_poll_tasks import poll_generation_task, register_poll_active
|
||||
@@ -458,37 +400,107 @@ async def recover_one_generation_task(
|
||||
if _is_final_task_state(task):
|
||||
await _remove_poll_active(task.id)
|
||||
return "clean_final_state"
|
||||
if task.status != "generating":
|
||||
if task.status != ChatGenerationTaskStatus.GENERATING.value:
|
||||
await _remove_poll_active(task.id)
|
||||
return "clean_not_generating"
|
||||
|
||||
if task.deadline_at and _is_expired(task.deadline_at, current_time):
|
||||
if task.pipeline_stage in ("waiting_remote", "polling"):
|
||||
return await _try_final_poll_before_timeout(db, task)
|
||||
return await _mark_timeout(db, task)
|
||||
has_remote_result = bool(str(task.remote_result_url or "").strip())
|
||||
has_provider_task_id = bool(str(task.provider_task_id or "").strip() or str(task.seedance_task_id or "").strip())
|
||||
is_deadline_expired = bool(task.deadline_at and _is_expired(task.deadline_at, current_time))
|
||||
|
||||
if task.pipeline_stage in ("queued", "preparing", "creating_provider_task"):
|
||||
if task.provider_task_id or task.seedance_task_id:
|
||||
task.pipeline_stage = "waiting_remote"
|
||||
# 最高优先级:只要远程结果 URL 已经落库,说明生成侧已经成功。
|
||||
# 不管当前 pipeline_stage 是 queued/creating/waiting/result_ready/download_*,恢复时都不能重复 create 或 poll。
|
||||
if has_remote_result:
|
||||
await _remove_poll_active(task.id)
|
||||
await log_task_event(
|
||||
task,
|
||||
event_type="GENERATION_RECOVERY_ENQUEUE",
|
||||
message=f"{source} 发现任务已存在 remote_result_url,恢复投递下载队列",
|
||||
detail={
|
||||
"pipeline_stage": task.pipeline_stage,
|
||||
"payload": redis_payload,
|
||||
"deadline_expired": is_deadline_expired,
|
||||
},
|
||||
)
|
||||
await enqueue_download_task(
|
||||
db,
|
||||
task,
|
||||
recover=True,
|
||||
reason=f"{source}_has_remote_result_url",
|
||||
)
|
||||
return "recover_download_has_remote_result"
|
||||
|
||||
# 已经过 deadline 且没有结果 URL:
|
||||
# - 有供应商任务 ID:交给 poll worker 做最后一次状态确认;
|
||||
# - 没有供应商任务 ID:说明没有可查询的远程任务,直接按超时失败处理,不再重新 create。
|
||||
if is_deadline_expired:
|
||||
if has_provider_task_id:
|
||||
task.pipeline_stage = ChatGenerationPipelineStage.WAITING_REMOTE.value
|
||||
await db.commit()
|
||||
await log_task_event(
|
||||
task,
|
||||
event_type="GENERATION_RECOVERY_ENQUEUE",
|
||||
message=f"{source} 发现创建阶段已存在供应商任务ID,恢复投递轮询队列",
|
||||
message=f"{source} 发现任务已到 deadline 且存在供应商任务ID,投递 poll 队列做最终查询",
|
||||
detail={"pipeline_stage": task.pipeline_stage, "payload": redis_payload},
|
||||
)
|
||||
poll_generation_task.apply_async(args=[task.id], queue=POLL_QUEUE, countdown=0)
|
||||
await register_poll_active(
|
||||
task,
|
||||
check_at=_poll_queue_timeout_at(),
|
||||
reason=f"{source}_create_stage_has_provider_id",
|
||||
reason=f"{source}_deadline_final_poll",
|
||||
)
|
||||
return "recover_poll_from_create_stage"
|
||||
return "recover_deadline_final_poll"
|
||||
|
||||
await log_task_event(
|
||||
task,
|
||||
event_type="GENERATION_RECOVERY_TIMEOUT",
|
||||
message=f"{source} 发现任务已到 deadline,且没有 remote_result_url/供应商任务ID,按超时失败处理",
|
||||
detail={"pipeline_stage": task.pipeline_stage, "payload": redis_payload},
|
||||
)
|
||||
return await _mark_timeout(db, task)
|
||||
|
||||
# 未过 deadline:有供应商任务 ID 才允许恢复到 poll 队列。
|
||||
if has_provider_task_id:
|
||||
task.pipeline_stage = ChatGenerationPipelineStage.WAITING_REMOTE.value
|
||||
await db.commit()
|
||||
await log_task_event(
|
||||
task,
|
||||
event_type="GENERATION_RECOVERY_ENQUEUE",
|
||||
message=f"{source} 发现创建阶段任务未完成,恢复投递创建队列",
|
||||
message=f"{source} 发现任务存在供应商任务ID,恢复投递轮询队列",
|
||||
detail={"pipeline_stage": task.pipeline_stage, "payload": redis_payload},
|
||||
)
|
||||
poll_generation_task.apply_async(args=[task.id], queue=POLL_QUEUE, countdown=0)
|
||||
await register_poll_active(
|
||||
task,
|
||||
check_at=_poll_queue_timeout_at(),
|
||||
reason=f"{source}_has_provider_task_id",
|
||||
)
|
||||
return "recover_poll_has_provider_id"
|
||||
|
||||
# 未过 deadline,且没有结果 URL / 供应商任务 ID:
|
||||
# 图片同步任务会重新进入 submit_image_task;视频/其它任务会重新创建供应商任务。
|
||||
# 这里不能投 poll,因为没有 provider_task_id/seedance_task_id 可查询。
|
||||
recoverable_create_stages = {
|
||||
ChatGenerationPipelineStage.QUEUED.value,
|
||||
ChatGenerationPipelineStage.PREPARING.value,
|
||||
ChatGenerationPipelineStage.CREATING_PROVIDER_TASK.value,
|
||||
ChatGenerationPipelineStage.WAITING_REMOTE.value,
|
||||
ChatGenerationPipelineStage.POLLING.value,
|
||||
}
|
||||
if task.pipeline_stage in recoverable_create_stages:
|
||||
if task.pipeline_stage not in (
|
||||
ChatGenerationPipelineStage.QUEUED.value,
|
||||
ChatGenerationPipelineStage.PREPARING.value,
|
||||
ChatGenerationPipelineStage.CREATING_PROVIDER_TASK.value,
|
||||
):
|
||||
task.pipeline_stage = ChatGenerationPipelineStage.QUEUED.value
|
||||
await db.commit()
|
||||
|
||||
await _remove_poll_active(task.id)
|
||||
await log_task_event(
|
||||
task,
|
||||
event_type="GENERATION_RECOVERY_ENQUEUE",
|
||||
message=f"{source} 发现任务未超时且缺少 remote_result_url/供应商任务ID,恢复投递创建队列",
|
||||
detail={"pipeline_stage": task.pipeline_stage, "payload": redis_payload},
|
||||
)
|
||||
chatapi_create_generation_task.apply_async(
|
||||
@@ -496,67 +508,25 @@ async def recover_one_generation_task(
|
||||
queue="gen_chatapi_create",
|
||||
countdown=0,
|
||||
)
|
||||
return "recover_create"
|
||||
return "recover_create_no_remote_no_provider_before_deadline"
|
||||
|
||||
if task.pipeline_stage in ("waiting_remote", "polling"):
|
||||
if task.remote_result_url:
|
||||
await _remove_poll_active(task.id)
|
||||
await enqueue_download_task(
|
||||
db,
|
||||
task,
|
||||
recover=True,
|
||||
reason=f"{source}_waiting_remote_has_result",
|
||||
)
|
||||
return "recover_waiting_has_result"
|
||||
|
||||
if task.provider_task_id or task.seedance_task_id:
|
||||
await log_task_event(
|
||||
task,
|
||||
event_type="GENERATION_RECOVERY_ENQUEUE",
|
||||
message=f"{source} 发现远程等待/轮询阶段任务未完成,恢复投递轮询队列",
|
||||
detail={"pipeline_stage": task.pipeline_stage, "payload": redis_payload},
|
||||
)
|
||||
task.pipeline_stage = "waiting_remote"
|
||||
await db.commit()
|
||||
poll_generation_task.apply_async(
|
||||
args=[task.id],
|
||||
queue=POLL_QUEUE,
|
||||
countdown=0,
|
||||
)
|
||||
await register_poll_active(
|
||||
task,
|
||||
check_at=_poll_queue_timeout_at(),
|
||||
reason=f"{source}_recover_poll",
|
||||
)
|
||||
return "recover_poll"
|
||||
|
||||
await log_task_event(
|
||||
task,
|
||||
event_type="GENERATION_RECOVERY_ENQUEUE",
|
||||
message=f"{source} 发现任务缺少供应商任务ID,恢复投递创建队列",
|
||||
detail={"pipeline_stage": task.pipeline_stage, "payload": redis_payload},
|
||||
)
|
||||
task.pipeline_stage = "queued"
|
||||
# result_ready 但没有 URL 是脏状态;未过 deadline 时回创建队列重新处理,过期上面已标记超时。
|
||||
if task.pipeline_stage == ChatGenerationPipelineStage.RESULT_READY.value:
|
||||
task.pipeline_stage = ChatGenerationPipelineStage.QUEUED.value
|
||||
await db.commit()
|
||||
await _remove_poll_active(task.id)
|
||||
await log_task_event(
|
||||
task,
|
||||
event_type="GENERATION_RECOVERY_ENQUEUE",
|
||||
message=f"{source} 发现 result_ready 但缺少 remote_result_url,未超时,恢复投递创建队列",
|
||||
detail={"pipeline_stage": task.pipeline_stage, "payload": redis_payload},
|
||||
)
|
||||
chatapi_create_generation_task.apply_async(
|
||||
args=[task.id],
|
||||
queue="gen_chatapi_create",
|
||||
countdown=0,
|
||||
)
|
||||
return "recover_create_missing_provider_id"
|
||||
|
||||
if task.pipeline_stage == "result_ready":
|
||||
await _remove_poll_active(task.id)
|
||||
if task.remote_result_url:
|
||||
await enqueue_download_task(
|
||||
db,
|
||||
task,
|
||||
recover=True,
|
||||
reason=f"{source}_generation_result_ready",
|
||||
)
|
||||
return "recover_result_ready"
|
||||
return "skip_result_ready_no_url"
|
||||
return "recover_create_result_ready_no_url_before_deadline"
|
||||
|
||||
return f"skip_stage_{task.pipeline_stage}"
|
||||
|
||||
@@ -617,8 +587,8 @@ async def recover_generation_tasks_once(db: AsyncSession) -> dict[str, Any]:
|
||||
select(ChatGenerationTask)
|
||||
.where(
|
||||
ChatGenerationTask.deleted_at.is_(None),
|
||||
ChatGenerationTask.generation_mode.in_(["chatapi_async", "hot_opening_replicate", "shot_replicate"]),
|
||||
ChatGenerationTask.status == "generating",
|
||||
ChatGenerationTask.generation_mode.in_(list(ALLOWED_GENERATION_MODES)),
|
||||
ChatGenerationTask.status == ChatGenerationTaskStatus.GENERATING.value,
|
||||
ChatGenerationTask.pipeline_stage.in_(
|
||||
[
|
||||
"queued",
|
||||
@@ -656,11 +626,8 @@ async def recover_generation_tasks_once(db: AsyncSession) -> dict[str, Any]:
|
||||
if len(tasks) < batch_size or progressed_this_round <= 0:
|
||||
break
|
||||
|
||||
# 下载阶段单独跑 DB fallback。
|
||||
download_result = await recover_download_tasks_once(db)
|
||||
return {
|
||||
"checked": len(checked_ids),
|
||||
"db_checked": total_db_checked,
|
||||
"results": results,
|
||||
"download_recovery": download_result,
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,253 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from typing import Any, Mapping
|
||||
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.config import settings
|
||||
from app.enums.credit_record import (
|
||||
CreditRecordAction,
|
||||
CreditRecordChargeKind,
|
||||
CreditRecordOwnerType,
|
||||
CreditRecordSourceModule,
|
||||
)
|
||||
from app.models.chat_generation_task import ChatGenerationTask
|
||||
from app.models.credit_record import CreditRecord
|
||||
from app.models.generation_record import GenerationRecord
|
||||
from app.models.token_usage import TokenUsage
|
||||
from app.utils.id_gen import generate_id
|
||||
|
||||
|
||||
def _safe_int(value: Any, default: int = 0) -> int:
|
||||
try:
|
||||
if value is None or value == "":
|
||||
return default
|
||||
return int(value)
|
||||
except Exception:
|
||||
return default
|
||||
|
||||
|
||||
def _safe_json_dict(value: Any) -> dict[str, Any]:
|
||||
if not value:
|
||||
return {}
|
||||
if isinstance(value, dict):
|
||||
return value
|
||||
try:
|
||||
parsed = json.loads(value)
|
||||
return parsed if isinstance(parsed, dict) else {}
|
||||
except Exception:
|
||||
return {}
|
||||
|
||||
|
||||
def _extract_usage(provider_response: Any) -> dict[str, Any]:
|
||||
data = _safe_json_dict(provider_response)
|
||||
usage = data.get("usage")
|
||||
return usage if isinstance(usage, dict) else {}
|
||||
|
||||
|
||||
def _normalize_media_tokens(
|
||||
*,
|
||||
gen_type: str | None,
|
||||
provider_response: Any = None,
|
||||
fallback_total: int | None = None,
|
||||
) -> tuple[int, int, int]:
|
||||
usage = _extract_usage(provider_response)
|
||||
input_tokens = _safe_int(usage.get("input_tokens"), 0)
|
||||
output_tokens = _safe_int(
|
||||
usage.get("output_tokens"),
|
||||
_safe_int(usage.get("generated_tokens"), 0),
|
||||
)
|
||||
total_tokens = _safe_int(usage.get("total_tokens"), 0)
|
||||
|
||||
if total_tokens <= 0:
|
||||
total_tokens = _safe_int(fallback_total, 0)
|
||||
if output_tokens <= 0:
|
||||
output_tokens = max(0, total_tokens - input_tokens)
|
||||
if total_tokens <= 0:
|
||||
total_tokens = input_tokens + output_tokens
|
||||
|
||||
# 图片生成多数供应商只返回 output/total,没有 input;保持 input=0。视频同理兼容缺字段。
|
||||
return input_tokens, output_tokens, total_tokens
|
||||
|
||||
|
||||
def _engine_model_from_provider_response(provider_response: Any) -> str | None:
|
||||
data = _safe_json_dict(provider_response)
|
||||
model = data.get("model")
|
||||
return str(model) if model else None
|
||||
|
||||
|
||||
async def _find_latest_media_charge(
|
||||
db: AsyncSession,
|
||||
*,
|
||||
user_id: str,
|
||||
owner_type: str,
|
||||
owner_id: str,
|
||||
media_type: str | None,
|
||||
) -> CreditRecord | None:
|
||||
query = (
|
||||
select(CreditRecord)
|
||||
.where(CreditRecord.user_id == user_id)
|
||||
.where(CreditRecord.type == "consume")
|
||||
.where(CreditRecord.owner_type == owner_type)
|
||||
.where(CreditRecord.owner_id == owner_id)
|
||||
.where(CreditRecord.charge_kind == CreditRecordChargeKind.MEDIA.value)
|
||||
.where(CreditRecord.charge_action == CreditRecordAction.CHARGE.value)
|
||||
)
|
||||
if media_type:
|
||||
query = query.where(CreditRecord.media_type == media_type)
|
||||
query = query.order_by(CreditRecord.attempt_no.desc().nullslast(), CreditRecord.created_at.desc()).limit(1)
|
||||
result = await db.execute(query)
|
||||
return result.scalar_one_or_none()
|
||||
|
||||
|
||||
async def _get_or_create_token_usage(
|
||||
db: AsyncSession,
|
||||
*,
|
||||
charge: CreditRecord,
|
||||
input_tokens: int,
|
||||
output_tokens: int,
|
||||
total_tokens: int,
|
||||
model_config_id: str | None = None,
|
||||
) -> TokenUsage:
|
||||
token_usage: TokenUsage | None = None
|
||||
if charge.token_usage_id:
|
||||
result = await db.execute(select(TokenUsage).where(TokenUsage.id == charge.token_usage_id).limit(1))
|
||||
token_usage = result.scalar_one_or_none()
|
||||
if token_usage is None and charge.biz_key:
|
||||
result = await db.execute(select(TokenUsage).where(TokenUsage.biz_key == charge.biz_key).limit(1))
|
||||
token_usage = result.scalar_one_or_none()
|
||||
if token_usage is None:
|
||||
token_usage = TokenUsage(
|
||||
id=generate_id(),
|
||||
user_id=charge.user_id,
|
||||
model_config_id=model_config_id,
|
||||
owner_type=charge.owner_type,
|
||||
owner_id=charge.owner_id,
|
||||
biz_key=charge.biz_key,
|
||||
source_module=charge.source_module,
|
||||
source_step_code=charge.source_step_code,
|
||||
input_tokens=input_tokens,
|
||||
output_tokens=output_tokens,
|
||||
total_tokens=total_tokens,
|
||||
)
|
||||
db.add(token_usage)
|
||||
await db.flush()
|
||||
else:
|
||||
token_usage.user_id = token_usage.user_id or charge.user_id
|
||||
token_usage.model_config_id = token_usage.model_config_id or model_config_id
|
||||
token_usage.owner_type = token_usage.owner_type or charge.owner_type
|
||||
token_usage.owner_id = token_usage.owner_id or charge.owner_id
|
||||
token_usage.biz_key = token_usage.biz_key or charge.biz_key
|
||||
token_usage.source_module = token_usage.source_module or charge.source_module
|
||||
token_usage.source_step_code = token_usage.source_step_code or charge.source_step_code
|
||||
token_usage.input_tokens = input_tokens
|
||||
token_usage.output_tokens = output_tokens
|
||||
token_usage.total_tokens = total_tokens
|
||||
return token_usage
|
||||
|
||||
|
||||
async def _sync_charge_snapshot(
|
||||
db: AsyncSession,
|
||||
*,
|
||||
charge: CreditRecord | None,
|
||||
gen_type: str | None,
|
||||
provider_response: Any = None,
|
||||
fallback_total: int | None = None,
|
||||
) -> CreditRecord | None:
|
||||
if not charge:
|
||||
return None
|
||||
|
||||
input_tokens, output_tokens, total_tokens = _normalize_media_tokens(
|
||||
gen_type=gen_type,
|
||||
provider_response=provider_response,
|
||||
fallback_total=fallback_total,
|
||||
)
|
||||
if total_tokens <= 0:
|
||||
return charge
|
||||
|
||||
token_usage = await _get_or_create_token_usage(
|
||||
db,
|
||||
charge=charge,
|
||||
input_tokens=input_tokens,
|
||||
output_tokens=output_tokens,
|
||||
total_tokens=total_tokens,
|
||||
model_config_id=None,
|
||||
)
|
||||
|
||||
charge.token_usage_id = token_usage.id
|
||||
charge.input_tokens = input_tokens
|
||||
charge.output_tokens = output_tokens
|
||||
charge.total_tokens = total_tokens
|
||||
|
||||
# 兼容旧流水扣费时未冷备 engine_model_name 的场景,能从 provider response 推出来就补充。
|
||||
provider_model = _engine_model_from_provider_response(provider_response)
|
||||
if provider_model and not charge.engine_model_name:
|
||||
charge.engine_model_name = provider_model
|
||||
return charge
|
||||
|
||||
|
||||
async def sync_chat_generation_task_media_token_snapshot(
|
||||
db: AsyncSession,
|
||||
task: ChatGenerationTask,
|
||||
*,
|
||||
provider_response: Any = None,
|
||||
) -> CreditRecord | None:
|
||||
"""把 ChatGenerationTask 图片/视频媒体生成 token 后置快照回填到积分流水。
|
||||
|
||||
媒体扣费发生在创建任务前,供应商 usage 只能在创建/轮询成功后拿到,
|
||||
所以这里按 owner_type + owner_id + media_type 找到对应 media charge 流水并回填。
|
||||
"""
|
||||
if not bool(getattr(settings, "MEDIA_TOKEN_SNAPSHOT_ENABLED", True)):
|
||||
return None
|
||||
if not task:
|
||||
return None
|
||||
|
||||
gen_type = (getattr(task, "gen_type", None) or "").lower().strip()
|
||||
fallback_total = task.image_tokens_used if gen_type == "image" else task.video_tokens_used
|
||||
response = provider_response if provider_response is not None else getattr(task, "provider_response_json", None)
|
||||
charge = await _find_latest_media_charge(
|
||||
db,
|
||||
user_id=task.user_id,
|
||||
owner_type=CreditRecordOwnerType.CHAT_GENERATION_TASK.value,
|
||||
owner_id=task.id,
|
||||
media_type=gen_type or None,
|
||||
)
|
||||
return await _sync_charge_snapshot(
|
||||
db,
|
||||
charge=charge,
|
||||
gen_type=gen_type,
|
||||
provider_response=response,
|
||||
fallback_total=fallback_total,
|
||||
)
|
||||
|
||||
|
||||
async def sync_generation_record_media_token_snapshot(
|
||||
db: AsyncSession,
|
||||
record: GenerationRecord,
|
||||
*,
|
||||
provider_response: Any = None,
|
||||
) -> CreditRecord | None:
|
||||
"""把旧 GenerationRecord 图片/视频媒体生成 token 后置快照回填到积分流水。"""
|
||||
if not bool(getattr(settings, "MEDIA_TOKEN_SNAPSHOT_ENABLED", True)):
|
||||
return None
|
||||
if not record:
|
||||
return None
|
||||
|
||||
gen_type = (getattr(record, "gen_type", None) or "").lower().strip()
|
||||
fallback_total = record.image_tokens_used if gen_type == "image" else record.video_tokens_used
|
||||
charge = await _find_latest_media_charge(
|
||||
db,
|
||||
user_id=record.user_id,
|
||||
owner_type=CreditRecordOwnerType.GENERATION_RECORD.value,
|
||||
owner_id=record.id,
|
||||
media_type=gen_type or None,
|
||||
)
|
||||
return await _sync_charge_snapshot(
|
||||
db,
|
||||
charge=charge,
|
||||
gen_type=gen_type,
|
||||
provider_response=provider_response,
|
||||
fallback_total=fallback_total,
|
||||
)
|
||||
@@ -10,6 +10,7 @@ from app.models.base import async_session
|
||||
from app.models.generation_record import GenerationRecord
|
||||
from app.services.video_gen import get_active_engine, poll_task_status, download_video, _log_video_response
|
||||
from app.services.image_gen import get_active_image_engine, download_image
|
||||
from app.services.media_token_usage_snapshot_service import sync_generation_record_media_token_snapshot
|
||||
from app.services.resource_accounting_service import (
|
||||
record_generation_record_generated_resource,
|
||||
safe_file_size,
|
||||
@@ -157,6 +158,7 @@ class TaskQueue:
|
||||
else:
|
||||
record.video_url = file_url
|
||||
record.video_tokens_used = poll_result.get("video_tokens", 0)
|
||||
await sync_generation_record_media_token_snapshot(db, record, provider_response=resp_data)
|
||||
record.status = "completed"
|
||||
record.generated_at = datetime.now()
|
||||
if record.video_url:
|
||||
@@ -235,6 +237,7 @@ class TaskQueue:
|
||||
else:
|
||||
record.image_url = remote_url
|
||||
record.image_tokens_used = poll_result.get("image_tokens", 0)
|
||||
await sync_generation_record_media_token_snapshot(db, record, provider_response=poll_result)
|
||||
record.status = "completed"
|
||||
record.generated_at = datetime.now()
|
||||
if record.image_url:
|
||||
|
||||
@@ -1,13 +1,16 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
import os
|
||||
import threading
|
||||
from concurrent.futures import Future
|
||||
from concurrent.futures import Future, TimeoutError as FutureTimeoutError
|
||||
from typing import Awaitable, TypeVar
|
||||
|
||||
from app.config import settings
|
||||
|
||||
logger = logging.getLogger("video_gen")
|
||||
|
||||
T = TypeVar("T")
|
||||
|
||||
_thread_local = threading.local()
|
||||
@@ -18,6 +21,28 @@ _single_loop_pid: int | None = None
|
||||
_single_loop_ready: threading.Event | None = None
|
||||
|
||||
|
||||
async def _dispose_async_resources() -> None:
|
||||
"""释放当前 async loop 内缓存的异步资源。
|
||||
|
||||
Celery soft time limit 会打断同步等待 future.result() 的线程;如果不主动
|
||||
cancel coroutine 并释放 engine/redis,后台 loop 里残留的协程可能继续占用
|
||||
SQLAlchemy QueuePool 连接,后续任务就会出现 QueuePool timeout。
|
||||
"""
|
||||
try:
|
||||
from app.services.redis_registry_service import close_registry_redis
|
||||
|
||||
await close_registry_redis()
|
||||
except Exception:
|
||||
logger.debug("关闭 Celery Redis registry 连接失败", exc_info=True)
|
||||
|
||||
try:
|
||||
from app.models.base import engine
|
||||
|
||||
await engine.dispose()
|
||||
except Exception:
|
||||
logger.debug("dispose Celery SQLAlchemy engine 失败", exc_info=True)
|
||||
|
||||
|
||||
def _runner_mode() -> str:
|
||||
mode = str(getattr(settings, "CELERY_ASYNC_RUNNER_MODE", "single_loop") or "single_loop").strip().lower()
|
||||
if mode not in {"single_loop", "direct"}:
|
||||
@@ -46,16 +71,18 @@ def _get_or_create_thread_local_loop() -> asyncio.AbstractEventLoop:
|
||||
def _single_loop_worker(loop: asyncio.AbstractEventLoop, ready: threading.Event) -> None:
|
||||
asyncio.set_event_loop(loop)
|
||||
ready.set()
|
||||
loop.run_forever()
|
||||
try:
|
||||
loop.run_forever()
|
||||
finally:
|
||||
pending = [task for task in asyncio.all_tasks(loop) if not task.done()]
|
||||
if pending:
|
||||
for task in pending:
|
||||
task.cancel()
|
||||
loop.run_until_complete(asyncio.gather(*pending, return_exceptions=True))
|
||||
|
||||
pending = [task for task in asyncio.all_tasks(loop) if not task.done()]
|
||||
if pending:
|
||||
for task in pending:
|
||||
task.cancel()
|
||||
loop.run_until_complete(asyncio.gather(*pending, return_exceptions=True))
|
||||
|
||||
loop.run_until_complete(loop.shutdown_asyncgens())
|
||||
loop.close()
|
||||
loop.run_until_complete(_dispose_async_resources())
|
||||
loop.run_until_complete(loop.shutdown_asyncgens())
|
||||
loop.close()
|
||||
|
||||
|
||||
def _get_or_create_single_loop() -> asyncio.AbstractEventLoop:
|
||||
@@ -92,13 +119,28 @@ def _get_or_create_single_loop() -> asyncio.AbstractEventLoop:
|
||||
return _single_loop
|
||||
|
||||
|
||||
def _cancel_future_and_reset_loop(future: Future[T] | None, *, reason: str) -> None:
|
||||
"""取消当前协程并重置当前进程内 event loop。"""
|
||||
if future is not None and not future.done():
|
||||
future.cancel()
|
||||
try:
|
||||
future.result(timeout=2)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
logger.warning("Celery async_runner 正在重置 event loop。reason=%s", reason)
|
||||
close_loop()
|
||||
|
||||
|
||||
def run_async(coro: Awaitable[T]) -> T:
|
||||
"""Celery 同步 task 调用异步协程的统一入口。
|
||||
|
||||
默认 single_loop 模式:
|
||||
- 一个 Celery 子进程只有一个专用 event loop;
|
||||
- 所有 asyncpg / redis.asyncio 操作都在这个 loop 内创建和使用;
|
||||
- 避免 got Future attached to a different loop。
|
||||
- 避免 got Future attached to a different loop;
|
||||
- 当 Celery soft time limit 打断 future.result() 时,主动 cancel 后台协程并
|
||||
释放连接池,避免 QueuePool 被残留任务长期占用。
|
||||
|
||||
降级 direct 模式:
|
||||
- 兼容旧的线程本地 loop 方案;
|
||||
@@ -106,7 +148,15 @@ def run_async(coro: Awaitable[T]) -> T:
|
||||
"""
|
||||
if _runner_mode() == "direct":
|
||||
loop = _get_or_create_thread_local_loop()
|
||||
return loop.run_until_complete(coro)
|
||||
try:
|
||||
return loop.run_until_complete(coro)
|
||||
except BaseException:
|
||||
try:
|
||||
if not loop.is_closed():
|
||||
loop.run_until_complete(_dispose_async_resources())
|
||||
finally:
|
||||
close_loop()
|
||||
raise
|
||||
|
||||
loop = _get_or_create_single_loop()
|
||||
try:
|
||||
@@ -118,7 +168,16 @@ def run_async(coro: Awaitable[T]) -> T:
|
||||
raise RuntimeError("run_async() 不能在 Celery async_runner 的事件循环内部被同步调用")
|
||||
|
||||
future: Future[T] = asyncio.run_coroutine_threadsafe(coro, loop)
|
||||
return future.result()
|
||||
try:
|
||||
return future.result()
|
||||
except FutureTimeoutError:
|
||||
_cancel_future_and_reset_loop(future, reason="future_result_timeout")
|
||||
raise
|
||||
except BaseException:
|
||||
# Celery SoftTimeLimitExceeded/worker shutdown 等异常会从这里抛出。
|
||||
# 必须重置 loop,否则后台协程继续运行会拖住 DB 连接池。
|
||||
_cancel_future_and_reset_loop(future, reason="base_exception")
|
||||
raise
|
||||
|
||||
|
||||
def close_loop() -> None:
|
||||
@@ -130,8 +189,16 @@ def close_loop() -> None:
|
||||
loop = _single_loop
|
||||
thread = _single_loop_thread
|
||||
if loop is not None and not loop.is_closed() and thread is not None and thread.is_alive():
|
||||
loop.call_soon_threadsafe(loop.stop)
|
||||
thread.join(timeout=5)
|
||||
try:
|
||||
cleanup_future = asyncio.run_coroutine_threadsafe(_dispose_async_resources(), loop)
|
||||
cleanup_future.result(timeout=5)
|
||||
except Exception:
|
||||
logger.debug("关闭 loop 前清理 async 资源失败", exc_info=True)
|
||||
try:
|
||||
loop.call_soon_threadsafe(loop.stop)
|
||||
thread.join(timeout=5)
|
||||
except Exception:
|
||||
logger.debug("关闭 Celery async_runner loop 失败", exc_info=True)
|
||||
|
||||
_single_loop = None
|
||||
_single_loop_thread = None
|
||||
@@ -141,6 +208,17 @@ def close_loop() -> None:
|
||||
# 关闭 direct 降级模式的线程本地 loop。
|
||||
loop = getattr(_thread_local, "loop", None)
|
||||
if loop is not None and not loop.is_closed():
|
||||
loop.close()
|
||||
try:
|
||||
pending = [task for task in asyncio.all_tasks(loop) if not task.done()]
|
||||
for task in pending:
|
||||
task.cancel()
|
||||
if pending:
|
||||
loop.run_until_complete(asyncio.gather(*pending, return_exceptions=True))
|
||||
loop.run_until_complete(_dispose_async_resources())
|
||||
loop.run_until_complete(loop.shutdown_asyncgens())
|
||||
except Exception:
|
||||
logger.debug("关闭 direct loop 前清理失败", exc_info=True)
|
||||
finally:
|
||||
loop.close()
|
||||
_thread_local.loop = None
|
||||
_thread_local.pid = None
|
||||
|
||||
@@ -26,6 +26,9 @@ CELERY_TASK_IMPORTS = (
|
||||
)
|
||||
|
||||
|
||||
RECOVERY_QUEUE = settings.CELERY_RECOVERY_QUEUE or "gen_recovery"
|
||||
|
||||
|
||||
def _derive_redis_db(url: str, db_no: int) -> str:
|
||||
if not url:
|
||||
return url
|
||||
@@ -55,6 +58,20 @@ if broker_url:
|
||||
task_acks_late=True,
|
||||
task_reject_on_worker_lost=True,
|
||||
task_track_started=True,
|
||||
task_annotations={
|
||||
# 生成链路任务以数据库状态为准,不依赖 Celery result backend。
|
||||
# 这里忽略结果可避免任务误返回 ORM / 非 JSON 对象时触发结果序列化失败。
|
||||
# "generation.chatapi_create_generation_task": {"ignore_result": True},
|
||||
# "generation.poll_generation_task": {"ignore_result": True},
|
||||
# "generation.download_generation_result_task": {"ignore_result": True},
|
||||
"hot_opening.start_image_prompt_optimize": {"ignore_result": True},
|
||||
"hot_opening.start_video_prompt_optimize": {"ignore_result": True},
|
||||
"shot_replicate.analyze_original_video": {"ignore_result": True},
|
||||
"shot_replicate.analyze_custom_segment_video": {"ignore_result": True},
|
||||
"shot_replicate.split_one_segment": {"ignore_result": True},
|
||||
"shot_replicate.start_image_prompt_optimize": {"ignore_result": True},
|
||||
"shot_replicate.start_video_prompt_optimize": {"ignore_result": True},
|
||||
},
|
||||
worker_prefetch_multiplier=1,
|
||||
broker_transport_options={
|
||||
"visibility_timeout": 3600,
|
||||
@@ -73,10 +90,12 @@ if broker_url:
|
||||
"shot_replicate.split_one_segment": {"queue": "gen_result_download"},
|
||||
"shot_replicate.start_image_prompt_optimize": {"queue": "gen_chatapi_create"},
|
||||
"shot_replicate.start_video_prompt_optimize": {"queue": "gen_chatapi_create"},
|
||||
"shot_replicate.recover_split_tasks_once": {"queue": "gen_result_download"},
|
||||
"generation.recover_download_tasks_once": {"queue": "gen_result_download"},
|
||||
"generation.recover_generation_tasks_once": {"queue": "gen_result_download"},
|
||||
"module_async.recover_module_async_tasks_once": {"queue": "gen_result_download"},
|
||||
# 恢复扫描统一走独立队列,避免占用下载/轮询/创建业务 worker。
|
||||
"recovery.startup_recovery_once": {"queue": RECOVERY_QUEUE},
|
||||
"shot_replicate.recover_split_tasks_once": {"queue": RECOVERY_QUEUE},
|
||||
"generation.recover_download_tasks_once": {"queue": RECOVERY_QUEUE},
|
||||
"generation.recover_generation_tasks_once": {"queue": RECOVERY_QUEUE},
|
||||
"module_async.recover_module_async_tasks_once": {"queue": RECOVERY_QUEUE},
|
||||
"user_oauth.update_oauth_accounts": {"queue": "default"},
|
||||
"app.tasks.cleanup.*": {"queue": "default"},
|
||||
},
|
||||
@@ -86,7 +105,7 @@ else:
|
||||
|
||||
|
||||
async def _try_acquire_startup_recovery_lock() -> bool:
|
||||
"""任意 worker 启动时都可尝试抢恢复锁,避免依赖 hostname 命名。"""
|
||||
"""任意 worker 启动时都可尝试抢恢复投递锁,避免依赖 hostname 命名。"""
|
||||
from app.services.redis_registry_service import redis_acquire_lock
|
||||
|
||||
token = await redis_acquire_lock(
|
||||
@@ -103,9 +122,9 @@ def on_worker_ready(sender=None, **kwargs):
|
||||
|
||||
注意:
|
||||
- 不启用 Celery beat。
|
||||
- 不要求新增第四条启动命令。
|
||||
- 不再依赖 worker hostname 是否包含 gen_result_download。
|
||||
- 所有 worker 都尝试抢 Redis 锁,只有抢到锁的 worker 投递恢复任务。
|
||||
- 启动容灾保留,但只投递一个 recovery.startup_recovery_once 协调任务。
|
||||
- 协调任务走独立 gen_recovery 队列,串行扫描并把真实业务任务投回原队列。
|
||||
- 所有 worker 都尝试抢 Redis 投递锁,只有抢到锁的 worker 投递恢复任务。
|
||||
"""
|
||||
if celery_app is None:
|
||||
return
|
||||
@@ -122,39 +141,21 @@ def on_worker_ready(sender=None, **kwargs):
|
||||
return
|
||||
|
||||
try:
|
||||
from app.tasks.generation_recovery_tasks import (
|
||||
recover_download_tasks_once,
|
||||
recover_generation_tasks_once,
|
||||
)
|
||||
from app.tasks.shot_replicate_tasks import recover_split_tasks_once
|
||||
from app.tasks.module_async_recovery_tasks import recover_module_async_tasks_once_task
|
||||
from app.tasks.generation_recovery_tasks import startup_recovery_once
|
||||
|
||||
countdown = max(0, int(settings.DOWNLOAD_RECOVERY_STARTUP_DELAY_SECONDS or 0))
|
||||
|
||||
recover_generation_tasks_once.apply_async(
|
||||
startup_recovery_once.apply_async(
|
||||
countdown=countdown,
|
||||
queue="gen_result_download",
|
||||
queue=RECOVERY_QUEUE,
|
||||
priority=settings.DOWNLOAD_TASK_PRIORITY_RECOVER,
|
||||
)
|
||||
recover_download_tasks_once.apply_async(
|
||||
countdown=countdown + 5,
|
||||
queue="gen_result_download",
|
||||
priority=settings.DOWNLOAD_TASK_PRIORITY_RECOVER,
|
||||
logger.info(
|
||||
"启动容灾恢复协调任务已投递。queue=%s countdown=%s",
|
||||
RECOVERY_QUEUE,
|
||||
countdown,
|
||||
)
|
||||
recover_split_tasks_once.apply_async(
|
||||
countdown=countdown + 10,
|
||||
queue="gen_result_download",
|
||||
priority=settings.DOWNLOAD_TASK_PRIORITY_RECOVER,
|
||||
)
|
||||
recover_module_async_tasks_once_task.apply_async(
|
||||
countdown=countdown + 15,
|
||||
queue="gen_result_download",
|
||||
priority=settings.DOWNLOAD_TASK_PRIORITY_RECOVER,
|
||||
)
|
||||
|
||||
logger.info("启动容灾恢复任务已投递。countdown=%s", countdown)
|
||||
except Exception:
|
||||
logger.exception("启动容灾恢复任务投递失败")
|
||||
logger.exception("启动容灾恢复协调任务投递失败")
|
||||
|
||||
|
||||
@worker_process_init.connect
|
||||
|
||||
@@ -12,6 +12,7 @@ from app.services.error_codes import extract_error_message
|
||||
from app.services.generation_log_service import log_task_event
|
||||
from app.services.generation_refund_service import mark_chat_generation_task_failed_and_refund_once
|
||||
from app.services.generation_provider_service import create_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 ensure_aware_utc
|
||||
from app.tasks.celery_app import celery_app
|
||||
|
||||
@@ -223,6 +224,7 @@ async def _run(task_id: str):
|
||||
ensure_ascii=False,
|
||||
default=str,
|
||||
)
|
||||
await sync_chat_generation_task_media_token_snapshot(db, task, provider_response=task.provider_response_json)
|
||||
|
||||
if task.remote_result_url and not task.seedance_task_id:
|
||||
# 同步图片路径:原 SDK 已经返回最终 URL。
|
||||
@@ -288,7 +290,13 @@ async def _run(task_id: str):
|
||||
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))
|
||||
try:
|
||||
return run_async(_run(task_id))
|
||||
except Exception as exc:
|
||||
# 只处理 run_async/连接池/worker 中断等基础设施异常;业务异常已在 _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):
|
||||
|
||||
@@ -1,11 +1,24 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from app.tasks.async_runner import run_async
|
||||
from datetime import datetime, timezone, timedelta
|
||||
|
||||
import errno
|
||||
import math
|
||||
import uuid
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.config import settings
|
||||
from app.enums.generation_task import (
|
||||
ALLOWED_GENERATION_MODES,
|
||||
ChatGenerationPipelineStage,
|
||||
ChatGenerationTaskEventType,
|
||||
ChatGenerationTaskStatus,
|
||||
GenerationType,
|
||||
)
|
||||
from app.models.base import async_session
|
||||
from app.models.chat_generation_task import ChatGenerationTask
|
||||
from app.services.celery_download_recovery_service import (
|
||||
@@ -18,17 +31,17 @@ from app.services.error_codes import extract_error_message
|
||||
from app.services.generation_download_service import download_generation_result
|
||||
from app.services.generation_log_service import log_task_event
|
||||
from app.services.generation_refund_service import mark_chat_generation_task_failed_and_refund_once
|
||||
from app.services.media_token_usage_snapshot_service import sync_chat_generation_task_media_token_snapshot
|
||||
from app.services.resource_accounting_service import record_chat_task_generated_resource
|
||||
from app.tasks.celery_app import celery_app
|
||||
|
||||
ALLOWED_GENERATION_MODES = {"chatapi_async", "hot_opening_replicate", "shot_replicate"}
|
||||
|
||||
DOWNLOAD_QUEUE = "gen_result_download"
|
||||
DOWNLOAD_STAGE_QUEUED = "download_queued"
|
||||
DOWNLOAD_STAGE_DOWNLOADING = "downloading"
|
||||
DOWNLOAD_STAGE_RETRY_WAITING = "retry_waiting"
|
||||
DOWNLOAD_STAGE_DONE = "done"
|
||||
DOWNLOAD_STAGE_FAILED = "download_failed"
|
||||
DOWNLOAD_STAGE_QUEUED = ChatGenerationPipelineStage.DOWNLOAD_QUEUED.value
|
||||
DOWNLOAD_STAGE_DOWNLOADING = ChatGenerationPipelineStage.DOWNLOADING.value
|
||||
DOWNLOAD_STAGE_RETRY_WAITING = ChatGenerationPipelineStage.RETRY_WAITING.value
|
||||
DOWNLOAD_STAGE_DONE = ChatGenerationPipelineStage.DONE.value
|
||||
DOWNLOAD_STAGE_FAILED = ChatGenerationPipelineStage.DOWNLOAD_FAILED.value
|
||||
RESULT_READY_STAGE = ChatGenerationPipelineStage.RESULT_READY.value
|
||||
|
||||
|
||||
def _now() -> datetime:
|
||||
@@ -56,6 +69,15 @@ def _retry_at(attempt: int, now: datetime | None = None) -> datetime:
|
||||
return now + timedelta(seconds=max(1, base * max(1, attempt)))
|
||||
|
||||
|
||||
def _countdown_until(value: datetime | None, *, minimum: int = 1) -> int:
|
||||
target = ensure_aware_utc(value)
|
||||
if target is None:
|
||||
return minimum
|
||||
extra = int(getattr(settings, "DOWNLOAD_RETRY_COUNTDOWN_EXTRA_SECONDS", 1) or 0)
|
||||
seconds = (target - _now()).total_seconds()
|
||||
return max(minimum, math.ceil(seconds) + extra)
|
||||
|
||||
|
||||
def _is_expired(value: datetime | None, now: datetime | None = None) -> bool:
|
||||
value = ensure_aware_utc(value)
|
||||
if value is None:
|
||||
@@ -64,10 +86,10 @@ def _is_expired(value: datetime | None, now: datetime | None = None) -> bool:
|
||||
|
||||
|
||||
def _is_already_completed(task: ChatGenerationTask) -> bool:
|
||||
if task.status == "completed" or task.pipeline_stage == DOWNLOAD_STAGE_DONE:
|
||||
if task.gen_type == "image" and task.image_url:
|
||||
if task.status == ChatGenerationTaskStatus.COMPLETED.value or task.pipeline_stage == DOWNLOAD_STAGE_DONE:
|
||||
if task.gen_type == GenerationType.IMAGE.value and task.image_url:
|
||||
return True
|
||||
if task.gen_type == "video" and task.video_url:
|
||||
if task.gen_type == GenerationType.VIDEO.value and task.video_url:
|
||||
return True
|
||||
return False
|
||||
|
||||
@@ -77,6 +99,61 @@ def _build_celery_task_id(task_id: str, attempt: int | None = None, reason: str
|
||||
return f"download:{task_id}:{int(attempt or 0)}:{safe_reason}:{uuid.uuid4().hex[:12]}"
|
||||
|
||||
|
||||
def _is_non_retryable_download_error(exc: Exception) -> bool:
|
||||
if not bool(getattr(settings, "DOWNLOAD_NON_RETRYABLE_LOCAL_ERRORS", True)):
|
||||
return False
|
||||
if isinstance(exc, PermissionError):
|
||||
return True
|
||||
if isinstance(exc, OSError) and getattr(exc, "errno", None) in {errno.EACCES, errno.EPERM, errno.ENOSPC, errno.EROFS, errno.ENAMETOOLONG}:
|
||||
return True
|
||||
message = str(exc).lower()
|
||||
non_retryable_fragments = (
|
||||
"permission denied",
|
||||
"no space left on device",
|
||||
"read-only file system",
|
||||
"file name too long",
|
||||
"invalid argument",
|
||||
)
|
||||
return any(fragment in message for fragment in non_retryable_fragments)
|
||||
|
||||
|
||||
async def _log_download_event(
|
||||
task: ChatGenerationTask | None = None,
|
||||
*,
|
||||
task_id: str | None = None,
|
||||
event_type: ChatGenerationTaskEventType | str,
|
||||
from_status: str | None = None,
|
||||
to_status: str | None = None,
|
||||
from_stage: str | None = None,
|
||||
to_stage: str | None = None,
|
||||
message: str | None = None,
|
||||
detail: Any = None,
|
||||
) -> None:
|
||||
if not bool(getattr(settings, "DOWNLOAD_EVENT_VERBOSE_ENABLED", True)):
|
||||
# 成功/失败关键事件仍保留;只关闭 verbose skip 事件。
|
||||
critical = {
|
||||
ChatGenerationTaskEventType.DOWNLOAD_START.value,
|
||||
ChatGenerationTaskEventType.DOWNLOAD_SUCCESS.value,
|
||||
ChatGenerationTaskEventType.DOWNLOAD_FAILED.value,
|
||||
ChatGenerationTaskEventType.DOWNLOAD_FAILED_NON_RETRYABLE.value,
|
||||
ChatGenerationTaskEventType.DOWNLOAD_RETRY_WAITING.value,
|
||||
}
|
||||
event_value = event_type.value if hasattr(event_type, "value") else str(event_type)
|
||||
if event_value not in critical:
|
||||
return
|
||||
await log_task_event(
|
||||
task,
|
||||
task_id=task_id,
|
||||
event_type=event_type.value if hasattr(event_type, "value") else str(event_type),
|
||||
from_status=from_status,
|
||||
to_status=to_status,
|
||||
from_stage=from_stage,
|
||||
to_stage=to_stage,
|
||||
message=message,
|
||||
detail=detail,
|
||||
)
|
||||
|
||||
|
||||
async def _register_active_from_task(
|
||||
task: ChatGenerationTask,
|
||||
*,
|
||||
@@ -101,6 +178,62 @@ async def _register_active_from_task(
|
||||
await upsert_download_active(record_id=task.id, payload=payload, check_at=check_at)
|
||||
|
||||
|
||||
async def _apply_download_async(
|
||||
task: ChatGenerationTask,
|
||||
*,
|
||||
priority: int,
|
||||
countdown: int | None,
|
||||
reason: str,
|
||||
event_type: ChatGenerationTaskEventType,
|
||||
failed_event_type: ChatGenerationTaskEventType,
|
||||
) -> bool:
|
||||
if not celery_app:
|
||||
await _log_download_event(
|
||||
task,
|
||||
event_type=failed_event_type,
|
||||
message="Celery 未启用,下载任务无法投递",
|
||||
detail={"queue": DOWNLOAD_QUEUE, "priority": priority, "countdown": countdown, "reason": reason},
|
||||
)
|
||||
return False
|
||||
try:
|
||||
download_generation_result_task.apply_async(
|
||||
args=[task.id],
|
||||
queue=DOWNLOAD_QUEUE,
|
||||
priority=priority,
|
||||
countdown=countdown,
|
||||
task_id=task.download_celery_task_id,
|
||||
)
|
||||
await _log_download_event(
|
||||
task,
|
||||
event_type=event_type,
|
||||
to_stage=task.pipeline_stage,
|
||||
detail={
|
||||
"queue": DOWNLOAD_QUEUE,
|
||||
"priority": priority,
|
||||
"countdown": countdown,
|
||||
"download_celery_task_id": task.download_celery_task_id,
|
||||
"reason": reason,
|
||||
"attempt": task.download_attempt_count,
|
||||
},
|
||||
)
|
||||
return True
|
||||
except Exception as exc:
|
||||
await _log_download_event(
|
||||
task,
|
||||
event_type=failed_event_type,
|
||||
message=str(exc),
|
||||
detail={
|
||||
"queue": DOWNLOAD_QUEUE,
|
||||
"priority": priority,
|
||||
"countdown": countdown,
|
||||
"download_celery_task_id": task.download_celery_task_id,
|
||||
"reason": reason,
|
||||
"attempt": task.download_attempt_count,
|
||||
},
|
||||
)
|
||||
return False
|
||||
|
||||
|
||||
async def enqueue_download_task(
|
||||
db: AsyncSession,
|
||||
task: ChatGenerationTask,
|
||||
@@ -110,11 +243,17 @@ async def enqueue_download_task(
|
||||
countdown: int | None = None,
|
||||
) -> str | None:
|
||||
"""统一投递图片/视频下载任务,并同步 DB + Redis active 注册表。"""
|
||||
if not task or task.generation_mode not in ALLOWED_GENERATION_MODES:
|
||||
reason = reason or ("recover" if recover else "normal")
|
||||
if not task:
|
||||
return None
|
||||
if task.status != "generating":
|
||||
if task.generation_mode not in ALLOWED_GENERATION_MODES:
|
||||
await _log_download_event(task, event_type=ChatGenerationTaskEventType.DOWNLOAD_SKIP_INVALID_MODE, message="不支持的 generation_mode", detail={"reason": reason})
|
||||
return None
|
||||
if task.status != ChatGenerationTaskStatus.GENERATING.value:
|
||||
await _log_download_event(task, event_type=ChatGenerationTaskEventType.DOWNLOAD_SKIP_NOT_GENERATING, message="任务不是 generating 状态", detail={"reason": reason, "status": task.status})
|
||||
return None
|
||||
if not task.remote_result_url:
|
||||
await _log_download_event(task, event_type=ChatGenerationTaskEventType.DOWNLOAD_SKIP_NO_REMOTE_RESULT_URL, message="缺少 remote_result_url", detail={"reason": reason})
|
||||
return None
|
||||
|
||||
now = _now()
|
||||
@@ -122,9 +261,10 @@ async def enqueue_download_task(
|
||||
celery_task_id = _build_celery_task_id(
|
||||
task.id,
|
||||
attempt=task.download_attempt_count or task.retry_count or 0,
|
||||
reason=reason or ("recover" if recover else "normal"),
|
||||
reason=reason,
|
||||
)
|
||||
|
||||
old_stage = task.pipeline_stage
|
||||
task.pipeline_stage = DOWNLOAD_STAGE_QUEUED
|
||||
task.download_celery_task_id = celery_task_id
|
||||
task.download_enqueued_at = now
|
||||
@@ -137,13 +277,22 @@ async def enqueue_download_task(
|
||||
check_at = _queue_timeout_at(now)
|
||||
await _register_active_from_task(task, check_at=check_at, priority=priority, reason=reason)
|
||||
|
||||
if celery_app:
|
||||
download_generation_result_task.apply_async(
|
||||
args=[task.id],
|
||||
queue=DOWNLOAD_QUEUE,
|
||||
priority=priority,
|
||||
countdown=countdown,
|
||||
task_id=celery_task_id,
|
||||
await _apply_download_async(
|
||||
task,
|
||||
priority=priority,
|
||||
countdown=countdown,
|
||||
reason=reason,
|
||||
event_type=ChatGenerationTaskEventType.DOWNLOAD_RECOVERY_ENQUEUE if recover else ChatGenerationTaskEventType.DOWNLOAD_ENQUEUE,
|
||||
failed_event_type=ChatGenerationTaskEventType.DOWNLOAD_RECOVERY_ENQUEUE_FAILED if recover else ChatGenerationTaskEventType.DOWNLOAD_ENQUEUE_FAILED,
|
||||
)
|
||||
if old_stage != DOWNLOAD_STAGE_QUEUED:
|
||||
# 独立记录阶段变化的上下文,便于和真正投递事件对照。
|
||||
await _log_download_event(
|
||||
task,
|
||||
event_type=ChatGenerationTaskEventType.DOWNLOAD_ENQUEUE,
|
||||
from_stage=old_stage,
|
||||
to_stage=DOWNLOAD_STAGE_QUEUED,
|
||||
detail={"reason": reason, "recover": recover, "check_at": check_at},
|
||||
)
|
||||
return celery_task_id
|
||||
|
||||
@@ -158,34 +307,81 @@ async def _reload_task(db: AsyncSession, task_id: str) -> ChatGenerationTask | N
|
||||
return result.scalar_one_or_none()
|
||||
|
||||
|
||||
async def _reschedule_not_due_retry(task: ChatGenerationTask, *, now: datetime) -> None:
|
||||
countdown = _countdown_until(task.download_next_retry_at)
|
||||
await _register_active_from_task(
|
||||
task,
|
||||
check_at=ensure_aware_utc(task.download_next_retry_at) or (now + timedelta(seconds=countdown)),
|
||||
priority=settings.DOWNLOAD_TASK_PRIORITY_RECOVER,
|
||||
reason="retry_waiting_not_due",
|
||||
)
|
||||
await _apply_download_async(
|
||||
task,
|
||||
priority=settings.DOWNLOAD_TASK_PRIORITY_RECOVER,
|
||||
countdown=countdown,
|
||||
reason="retry_waiting_not_due",
|
||||
event_type=ChatGenerationTaskEventType.DOWNLOAD_RETRY_ENQUEUE,
|
||||
failed_event_type=ChatGenerationTaskEventType.DOWNLOAD_RETRY_ENQUEUE_FAILED,
|
||||
)
|
||||
|
||||
|
||||
async def _claim_download_lease(db: AsyncSession, task: ChatGenerationTask) -> bool:
|
||||
now = _now()
|
||||
|
||||
if not task or task.generation_mode not in ALLOWED_GENERATION_MODES:
|
||||
if not task:
|
||||
return False
|
||||
if task.status != "generating":
|
||||
if task.generation_mode not in ALLOWED_GENERATION_MODES:
|
||||
await _log_download_event(task, event_type=ChatGenerationTaskEventType.DOWNLOAD_SKIP_INVALID_MODE, message="下载任务跳过:不支持的 generation_mode")
|
||||
return False
|
||||
if task.status != ChatGenerationTaskStatus.GENERATING.value:
|
||||
await _log_download_event(task, event_type=ChatGenerationTaskEventType.DOWNLOAD_SKIP_NOT_GENERATING, message="下载任务跳过:任务不是 generating 状态", detail={"status": task.status, "stage": task.pipeline_stage})
|
||||
return False
|
||||
if _is_already_completed(task):
|
||||
await _log_download_event(task, event_type=ChatGenerationTaskEventType.DOWNLOAD_SKIP_ALREADY_COMPLETED, message="下载任务跳过:任务已完成", detail={"status": task.status, "stage": task.pipeline_stage})
|
||||
return False
|
||||
if not task.remote_result_url:
|
||||
await _log_download_event(task, event_type=ChatGenerationTaskEventType.DOWNLOAD_SKIP_NO_REMOTE_RESULT_URL, message="下载任务跳过:缺少 remote_result_url")
|
||||
return False
|
||||
|
||||
stage = task.pipeline_stage
|
||||
|
||||
if stage == DOWNLOAD_STAGE_DOWNLOADING:
|
||||
if not _is_expired(task.download_lease_until, now):
|
||||
await _log_download_event(
|
||||
task,
|
||||
event_type=ChatGenerationTaskEventType.DOWNLOAD_SKIP_DOWNLOADING_LEASE_ALIVE,
|
||||
message="下载任务跳过:已有 downloading lease 且未过期",
|
||||
detail={"lease_until": task.download_lease_until, "download_celery_task_id": task.download_celery_task_id},
|
||||
)
|
||||
return False
|
||||
await log_task_event(
|
||||
await _log_download_event(
|
||||
task,
|
||||
event_type="DOWNLOAD_STUCK_RECOVER",
|
||||
event_type=ChatGenerationTaskEventType.DOWNLOAD_STUCK_RECOVER,
|
||||
message=f"downloading lease 已过期,重新抢占下载。lease_until={task.download_lease_until}",
|
||||
)
|
||||
elif stage == DOWNLOAD_STAGE_RETRY_WAITING:
|
||||
if not _is_expired(task.download_next_retry_at, now):
|
||||
await _log_download_event(
|
||||
task,
|
||||
event_type=ChatGenerationTaskEventType.DOWNLOAD_RETRY_NOT_DUE,
|
||||
message="下载重试提前触发,尚未到 next_retry_at,已重新投递延后重试",
|
||||
detail={
|
||||
"now": now,
|
||||
"download_next_retry_at": task.download_next_retry_at,
|
||||
"download_celery_task_id": task.download_celery_task_id,
|
||||
},
|
||||
)
|
||||
await _reschedule_not_due_retry(task, now=now)
|
||||
return False
|
||||
elif stage in (DOWNLOAD_STAGE_QUEUED, "result_ready"):
|
||||
elif stage in (DOWNLOAD_STAGE_QUEUED, RESULT_READY_STAGE):
|
||||
pass
|
||||
else:
|
||||
await _log_download_event(
|
||||
task,
|
||||
event_type=ChatGenerationTaskEventType.DOWNLOAD_SKIP_STAGE_NOT_ALLOWED,
|
||||
message="下载任务跳过:当前阶段不允许下载",
|
||||
detail={"stage": stage, "status": task.status, "download_celery_task_id": task.download_celery_task_id},
|
||||
)
|
||||
return False
|
||||
|
||||
old_stage = stage
|
||||
@@ -207,9 +403,9 @@ async def _claim_download_lease(db: AsyncSession, task: ChatGenerationTask) -> b
|
||||
reason="claim_download_lease",
|
||||
)
|
||||
|
||||
await log_task_event(
|
||||
await _log_download_event(
|
||||
task,
|
||||
event_type="DOWNLOAD_START",
|
||||
event_type=ChatGenerationTaskEventType.DOWNLOAD_START,
|
||||
from_stage=old_stage,
|
||||
to_stage=DOWNLOAD_STAGE_DOWNLOADING,
|
||||
detail={
|
||||
@@ -244,9 +440,9 @@ async def _mark_retry_waiting(db: AsyncSession, task: ChatGenerationTask, exc: E
|
||||
reason="download_retry_waiting",
|
||||
)
|
||||
|
||||
await log_task_event(
|
||||
await _log_download_event(
|
||||
task,
|
||||
event_type="DOWNLOAD_RETRY_WAITING",
|
||||
event_type=ChatGenerationTaskEventType.DOWNLOAD_RETRY_WAITING,
|
||||
message=error_message,
|
||||
to_stage=DOWNLOAD_STAGE_RETRY_WAITING,
|
||||
detail={
|
||||
@@ -262,16 +458,53 @@ def _should_final_fail(task: ChatGenerationTask) -> bool:
|
||||
return int(task.download_attempt_count or task.retry_count or 0) >= int(settings.DOWNLOAD_TASK_MAX_ATTEMPTS or 3)
|
||||
|
||||
|
||||
async def _mark_download_failed(
|
||||
db: AsyncSession,
|
||||
task: ChatGenerationTask,
|
||||
*,
|
||||
exc: Exception,
|
||||
non_retryable: bool = False,
|
||||
) -> None:
|
||||
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=DOWNLOAD_STAGE_FAILED,
|
||||
)
|
||||
task.download_last_error = error_message
|
||||
task.download_lease_until = None
|
||||
task.download_next_retry_at = None
|
||||
|
||||
from app.services.generation_module_hook_service import notify_chat_generation_task_finished
|
||||
|
||||
await notify_chat_generation_task_finished(db, task)
|
||||
await db.commit()
|
||||
await remove_download_active(task.id)
|
||||
|
||||
await _log_download_event(
|
||||
task,
|
||||
event_type=ChatGenerationTaskEventType.DOWNLOAD_FAILED_NON_RETRYABLE if non_retryable else ChatGenerationTaskEventType.DOWNLOAD_FAILED,
|
||||
message=task.error_message,
|
||||
detail={
|
||||
"download_attempt_count": task.download_attempt_count,
|
||||
"max_attempts": settings.DOWNLOAD_TASK_MAX_ATTEMPTS,
|
||||
"non_retryable": non_retryable,
|
||||
"download_last_error": task.download_last_error,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
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()
|
||||
task = await _reload_task(db, task_id)
|
||||
if not task:
|
||||
await _log_download_event(
|
||||
task_id=task_id,
|
||||
event_type=ChatGenerationTaskEventType.DOWNLOAD_SKIP_TASK_MISSING,
|
||||
message="下载任务跳过:ChatGenerationTask 不存在或已软删",
|
||||
)
|
||||
await remove_download_active(task_id)
|
||||
return
|
||||
|
||||
claimed = await _claim_download_lease(db, task)
|
||||
@@ -283,15 +516,21 @@ async def _run(task_id: str):
|
||||
|
||||
task = await _reload_task(db, task_id)
|
||||
if not task:
|
||||
await _log_download_event(
|
||||
task_id=task_id,
|
||||
event_type=ChatGenerationTaskEventType.DOWNLOAD_SKIP_TASK_MISSING,
|
||||
message="下载完成后任务不存在或已软删",
|
||||
)
|
||||
await remove_download_active(task_id)
|
||||
return
|
||||
|
||||
if task.gen_type == "image":
|
||||
if task.gen_type == GenerationType.IMAGE.value:
|
||||
task.image_url = downloaded.url
|
||||
else:
|
||||
task.video_url = downloaded.url
|
||||
task.video_cover_url = downloaded.cover_url
|
||||
|
||||
task.status = "completed"
|
||||
task.status = ChatGenerationTaskStatus.COMPLETED.value
|
||||
task.pipeline_stage = DOWNLOAD_STAGE_DONE
|
||||
task.generated_at = _now()
|
||||
task.retry_count = 0
|
||||
@@ -308,18 +547,18 @@ async def _run(task_id: str):
|
||||
remote_url=task.remote_result_url,
|
||||
generated_at=task.generated_at,
|
||||
)
|
||||
await sync_chat_generation_task_media_token_snapshot(db, task)
|
||||
|
||||
from app.services.generation_module_hook_service import notify_chat_generation_task_finished
|
||||
|
||||
await notify_chat_generation_task_finished(db, task)
|
||||
await db.commit()
|
||||
await remove_download_active(task.id)
|
||||
|
||||
from app.services.generation_module_hook_service import notify_chat_generation_task_finished
|
||||
await notify_chat_generation_task_finished(db, task)
|
||||
await db.commit()
|
||||
|
||||
await log_task_event(
|
||||
await _log_download_event(
|
||||
task,
|
||||
event_type="DOWNLOAD_SUCCESS",
|
||||
to_status="completed",
|
||||
event_type=ChatGenerationTaskEventType.DOWNLOAD_SUCCESS,
|
||||
to_status=ChatGenerationTaskStatus.COMPLETED.value,
|
||||
to_stage=DOWNLOAD_STAGE_DONE,
|
||||
detail={
|
||||
"resource_url": downloaded.url,
|
||||
@@ -337,54 +576,35 @@ async def _run(task_id: str):
|
||||
|
||||
task = await _reload_task(db, task_id)
|
||||
if not task:
|
||||
await remove_download_active(task_id)
|
||||
return
|
||||
|
||||
if _should_final_fail(task):
|
||||
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=DOWNLOAD_STAGE_FAILED,
|
||||
)
|
||||
task.download_last_error = error_message
|
||||
task.download_lease_until = None
|
||||
task.download_next_retry_at = None
|
||||
await db.commit()
|
||||
|
||||
await remove_download_active(task.id)
|
||||
|
||||
from app.services.generation_module_hook_service import notify_chat_generation_task_finished
|
||||
await notify_chat_generation_task_finished(db, task)
|
||||
await db.commit()
|
||||
|
||||
await log_task_event(
|
||||
task,
|
||||
event_type="DOWNLOAD_FAILED",
|
||||
message=task.error_message,
|
||||
detail={
|
||||
"download_attempt_count": task.download_attempt_count,
|
||||
"max_attempts": settings.DOWNLOAD_TASK_MAX_ATTEMPTS,
|
||||
},
|
||||
)
|
||||
non_retryable = _is_non_retryable_download_error(exc)
|
||||
if non_retryable or _should_final_fail(task):
|
||||
await _mark_download_failed(db, task, exc=exc, non_retryable=non_retryable)
|
||||
else:
|
||||
next_retry_at = await _mark_retry_waiting(db, task, exc)
|
||||
|
||||
if celery_app:
|
||||
delay_seconds = max(1, int((next_retry_at - _now()).total_seconds()))
|
||||
download_generation_result_task.apply_async(
|
||||
args=[task.id],
|
||||
queue=DOWNLOAD_QUEUE,
|
||||
priority=settings.DOWNLOAD_TASK_PRIORITY_RECOVER,
|
||||
countdown=delay_seconds,
|
||||
task_id=task.download_celery_task_id,
|
||||
)
|
||||
delay_seconds = _countdown_until(next_retry_at)
|
||||
await _apply_download_async(
|
||||
task,
|
||||
priority=settings.DOWNLOAD_TASK_PRIORITY_RECOVER,
|
||||
countdown=delay_seconds,
|
||||
reason="download_exception_retry",
|
||||
event_type=ChatGenerationTaskEventType.DOWNLOAD_RETRY_ENQUEUE,
|
||||
failed_event_type=ChatGenerationTaskEventType.DOWNLOAD_RETRY_ENQUEUE_FAILED,
|
||||
)
|
||||
|
||||
|
||||
if celery_app:
|
||||
@celery_app.task(name="generation.download_generation_result_task", bind=True, max_retries=3, default_retry_delay=30)
|
||||
def download_generation_result_task(self, task_id: str):
|
||||
return run_async(_run(task_id))
|
||||
try:
|
||||
return run_async(_run(task_id))
|
||||
except Exception as exc:
|
||||
# 只重试 run_async/连接池/worker 中断等基础设施异常;下载业务异常已在 _run 内写入 retry_waiting。
|
||||
retries = int(getattr(self.request, "retries", 0) or 0) + 1
|
||||
countdown = int(settings.DOWNLOAD_TASK_RETRY_BACKOFF_SECONDS or 30) * max(1, retries)
|
||||
raise self.retry(exc=exc, countdown=countdown)
|
||||
else:
|
||||
class _DisabledTask:
|
||||
def delay(self, *args, **kwargs):
|
||||
|
||||
@@ -12,6 +12,7 @@ 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.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,
|
||||
@@ -193,7 +194,8 @@ async def _run(task_id: str):
|
||||
await remove_poll_active(task.id)
|
||||
return
|
||||
|
||||
if _deadline_expired(task):
|
||||
final_poll_before_timeout = _deadline_expired(task)
|
||||
if final_poll_before_timeout and not (task.seedance_task_id or task.provider_task_id):
|
||||
await _mark_timeout(db, task, message="任务轮询超时")
|
||||
return
|
||||
|
||||
@@ -201,6 +203,14 @@ async def _run(task_id: str):
|
||||
await _mark_failed(db, task, message="缺少外部任务ID")
|
||||
return
|
||||
|
||||
if final_poll_before_timeout:
|
||||
await log_task_event(
|
||||
task,
|
||||
event_type="FINAL_POLL_BEFORE_TIMEOUT",
|
||||
message="任务已到 deadline,执行最后一次供应商查询后再判定超时",
|
||||
detail={"deadline_at": task.deadline_at, "stage": task.pipeline_stage},
|
||||
)
|
||||
|
||||
# 标记本次正在轮询,并登记 poll lease。
|
||||
# 如果 worker 在供应商接口调用过程中退出,启动恢复会在 lease 过期后重新投递。
|
||||
task.pipeline_stage = "polling"
|
||||
@@ -245,6 +255,7 @@ async def _run(task_id: str):
|
||||
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)
|
||||
@@ -272,6 +283,16 @@ async def _run(task_id: str):
|
||||
)
|
||||
return
|
||||
|
||||
if final_poll_before_timeout:
|
||||
await log_task_event(
|
||||
task,
|
||||
event_type="FINAL_POLL_BEFORE_TIMEOUT_PENDING",
|
||||
message=f"最终查询后供应商仍未完成,按超时处理。status={status}",
|
||||
detail=poll_result,
|
||||
)
|
||||
await _mark_timeout(db, task, message="任务轮询超时")
|
||||
return
|
||||
|
||||
# 供应商仍在 pending / running 时,把阶段从 polling 改回 waiting_remote。
|
||||
# 同时登记下一次 poll active,Celery countdown 丢失时可由恢复任务拉起。
|
||||
task.pipeline_stage = "waiting_remote"
|
||||
@@ -307,6 +328,15 @@ async def _run(task_id: str):
|
||||
await remove_poll_active(task_id)
|
||||
return
|
||||
|
||||
if final_poll_before_timeout:
|
||||
await log_task_event(
|
||||
task,
|
||||
event_type="FINAL_POLL_BEFORE_TIMEOUT_ERROR",
|
||||
message=str(exc),
|
||||
)
|
||||
await _mark_timeout(db, task, message="任务轮询超时")
|
||||
return
|
||||
|
||||
task.retry_count = (task.retry_count or 0) + 1
|
||||
|
||||
if task.retry_count > settings.CHATAPI_ASYNC_MAX_RETRIES:
|
||||
@@ -337,7 +367,13 @@ async def _run(task_id: str):
|
||||
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))
|
||||
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):
|
||||
|
||||
@@ -1,12 +1,20 @@
|
||||
# app/tasks/generation_recovery_tasks.py
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any, Dict
|
||||
import logging
|
||||
from typing import Any, Awaitable, Callable, Dict
|
||||
|
||||
from app.config import settings
|
||||
from app.models.base import async_session
|
||||
from app.services.redis_registry_service import get_registry_redis, redis_acquire_lock, redis_release_lock
|
||||
from app.tasks.async_runner import run_async
|
||||
from app.tasks.celery_app import celery_app
|
||||
|
||||
logger = logging.getLogger("video_gen")
|
||||
|
||||
RECOVERY_QUEUE = settings.CELERY_RECOVERY_QUEUE or "gen_recovery"
|
||||
RecoveryRunner = Callable[[], Awaitable[Dict[str, Any]]]
|
||||
|
||||
|
||||
async def _run_download_once() -> Dict[str, Any]:
|
||||
from app.services.generation_recovery_service import recover_download_tasks_once
|
||||
@@ -22,16 +30,195 @@ async def _run_generation_once() -> Dict[str, Any]:
|
||||
return await recover_generation_tasks_once(db)
|
||||
|
||||
|
||||
async def _run_module_async_once() -> Dict[str, Any]:
|
||||
from app.services.module_async_recovery_service import recover_module_async_tasks_once
|
||||
|
||||
async with async_session() as db:
|
||||
return await recover_module_async_tasks_once(db)
|
||||
|
||||
|
||||
async def _run_shot_split_once() -> Dict[str, Any]:
|
||||
from app.services.shot_replicate_recovery_service import recover_shot_split_tasks_once
|
||||
|
||||
async with async_session() as db:
|
||||
return await recover_shot_split_tasks_once(db)
|
||||
|
||||
|
||||
async def _run_with_execution_lock(
|
||||
*,
|
||||
lock_key: str,
|
||||
log_context: str,
|
||||
runner: RecoveryRunner,
|
||||
) -> Dict[str, Any]:
|
||||
"""恢复任务执行锁。
|
||||
|
||||
worker_ready 的启动锁只保证“只投递一次”;如果 broker 中残留旧消息,
|
||||
或者人工手动触发恢复任务,仍可能并发执行。这里再加执行锁,避免多个
|
||||
恢复扫描同时扫库、抢行锁、抢连接池。
|
||||
"""
|
||||
redis = await get_registry_redis()
|
||||
token: str | None = None
|
||||
if redis is not None:
|
||||
token = await redis_acquire_lock(
|
||||
lock_key=lock_key,
|
||||
ttl_seconds=int(settings.CELERY_RECOVERY_TASK_LOCK_TTL_SECONDS or 600),
|
||||
log_context=log_context,
|
||||
)
|
||||
if not token:
|
||||
return {"skipped": "lock_held", "lock_key": lock_key}
|
||||
else:
|
||||
# Redis 不可用时仍允许 DB fallback 执行一次,避免恢复能力彻底失效。
|
||||
logger.warning("恢复任务执行锁不可用,降级直接执行。context=%s", log_context)
|
||||
|
||||
try:
|
||||
result = await runner()
|
||||
result["execution_lock"] = "lock_acquired" if token else "redis_unavailable_run_db_fallback"
|
||||
return result
|
||||
finally:
|
||||
if token:
|
||||
await redis_release_lock(lock_key=lock_key, token=token, log_context=log_context)
|
||||
|
||||
|
||||
async def _acquire_download_recovery_loop_lock() -> tuple[bool, str]:
|
||||
"""下载恢复循环锁。
|
||||
|
||||
Redis 不可用时降级为直接执行 DB fallback,避免恢复能力彻底失效;
|
||||
Redis 可用但锁被其他 worker 持有时,本轮跳过,不再重复投递下一轮。
|
||||
"""
|
||||
redis = await get_registry_redis()
|
||||
if redis is None:
|
||||
return True, "redis_unavailable_run_db_fallback"
|
||||
token = await redis_acquire_lock(
|
||||
lock_key=settings.DOWNLOAD_RECOVERY_LOOP_LOCK_KEY,
|
||||
ttl_seconds=int(settings.DOWNLOAD_RECOVERY_LOOP_LOCK_TTL_SECONDS or 55),
|
||||
log_context="download_recovery_loop",
|
||||
)
|
||||
return (bool(token), "lock_acquired" if token else "lock_held")
|
||||
|
||||
|
||||
def _schedule_next_download_recovery_loop() -> None:
|
||||
if not celery_app or not bool(getattr(settings, "DOWNLOAD_RECOVERY_LOOP_ENABLED", False)):
|
||||
return
|
||||
try:
|
||||
recover_download_tasks_once.apply_async(
|
||||
countdown=max(1, int(settings.DOWNLOAD_RECOVERY_INTERVAL_SECONDS or 60)),
|
||||
queue=RECOVERY_QUEUE,
|
||||
priority=settings.DOWNLOAD_TASK_PRIORITY_RECOVER,
|
||||
)
|
||||
except Exception:
|
||||
logger.exception("下载恢复循环下一轮投递失败")
|
||||
|
||||
|
||||
async def _run_startup_recovery_once() -> Dict[str, Any]:
|
||||
"""启动容灾协调器:串行跑恢复扫描。
|
||||
|
||||
真实业务任务仍投递回原队列:
|
||||
- 创建/提词/视频分析 -> gen_chatapi_create
|
||||
- provider poll -> gen_provider_poll
|
||||
- 下载/ffmpeg 切片 -> gen_result_download
|
||||
恢复扫描本身只走 gen_recovery,避免堵住业务 worker。
|
||||
"""
|
||||
return await _run_with_execution_lock(
|
||||
lock_key=settings.CELERY_RECOVERY_STARTUP_TASK_LOCK_KEY,
|
||||
log_context="startup_recovery_once",
|
||||
runner=_run_startup_recovery_steps,
|
||||
)
|
||||
|
||||
|
||||
async def _run_startup_recovery_steps() -> Dict[str, Any]:
|
||||
results: Dict[str, Any] = {}
|
||||
|
||||
steps: list[tuple[str, str, str, RecoveryRunner]] = [
|
||||
(
|
||||
"module_async",
|
||||
settings.MODULE_ASYNC_RECOVERY_LOCK_KEY,
|
||||
"module_async_recovery",
|
||||
_run_module_async_once,
|
||||
),
|
||||
(
|
||||
"shot_split",
|
||||
settings.SHOT_SPLIT_RECOVERY_LOCK_KEY,
|
||||
"shot_split_recovery",
|
||||
_run_shot_split_once,
|
||||
),
|
||||
(
|
||||
"generation",
|
||||
settings.GENERATION_RECOVERY_LOCK_KEY,
|
||||
"generation_recovery",
|
||||
_run_generation_once,
|
||||
),
|
||||
(
|
||||
"download",
|
||||
settings.DOWNLOAD_RECOVERY_LOCK_KEY,
|
||||
"download_recovery",
|
||||
_run_download_once,
|
||||
),
|
||||
]
|
||||
|
||||
for name, lock_key, log_context, runner in steps:
|
||||
try:
|
||||
results[name] = await _run_with_execution_lock(
|
||||
lock_key=lock_key,
|
||||
log_context=log_context,
|
||||
runner=runner,
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.exception("启动容灾步骤执行失败。step=%s", name)
|
||||
results[name] = {"error": str(exc)}
|
||||
|
||||
return {"steps": results}
|
||||
|
||||
|
||||
if celery_app:
|
||||
|
||||
@celery_app.task(name="generation.recover_download_tasks_once")
|
||||
def recover_download_tasks_once() -> Dict[str, Any]:
|
||||
return run_async(_run_download_once())
|
||||
@celery_app.task(
|
||||
name="recovery.startup_recovery_once",
|
||||
bind=True,
|
||||
soft_time_limit=settings.CELERY_RECOVERY_SOFT_TIME_LIMIT_SECONDS,
|
||||
time_limit=settings.CELERY_RECOVERY_TIME_LIMIT_SECONDS,
|
||||
)
|
||||
def startup_recovery_once(self) -> Dict[str, Any]:
|
||||
return run_async(_run_startup_recovery_once())
|
||||
|
||||
|
||||
@celery_app.task(name="generation.recover_generation_tasks_once")
|
||||
def recover_generation_tasks_once() -> Dict[str, Any]:
|
||||
return run_async(_run_generation_once())
|
||||
@celery_app.task(
|
||||
name="generation.recover_download_tasks_once",
|
||||
bind=True,
|
||||
soft_time_limit=settings.CELERY_RECOVERY_SOFT_TIME_LIMIT_SECONDS,
|
||||
time_limit=settings.CELERY_RECOVERY_TIME_LIMIT_SECONDS,
|
||||
)
|
||||
def recover_download_tasks_once(self) -> Dict[str, Any]:
|
||||
acquired, reason = run_async(_acquire_download_recovery_loop_lock())
|
||||
if not acquired:
|
||||
return {"skipped": reason}
|
||||
try:
|
||||
result = run_async(
|
||||
_run_with_execution_lock(
|
||||
lock_key=settings.DOWNLOAD_RECOVERY_LOCK_KEY,
|
||||
log_context="download_recovery",
|
||||
runner=_run_download_once,
|
||||
)
|
||||
)
|
||||
result["loop_lock"] = reason
|
||||
return result
|
||||
finally:
|
||||
_schedule_next_download_recovery_loop()
|
||||
|
||||
|
||||
@celery_app.task(
|
||||
name="generation.recover_generation_tasks_once",
|
||||
bind=True,
|
||||
soft_time_limit=settings.CELERY_RECOVERY_SOFT_TIME_LIMIT_SECONDS,
|
||||
time_limit=settings.CELERY_RECOVERY_TIME_LIMIT_SECONDS,
|
||||
)
|
||||
def recover_generation_tasks_once(self) -> Dict[str, Any]:
|
||||
return run_async(
|
||||
_run_with_execution_lock(
|
||||
lock_key=settings.GENERATION_RECOVERY_LOCK_KEY,
|
||||
log_context="generation_recovery",
|
||||
runner=_run_generation_once,
|
||||
)
|
||||
)
|
||||
|
||||
else:
|
||||
|
||||
@@ -42,5 +229,6 @@ else:
|
||||
def apply_async(self, *args: Any, **kwargs: Any) -> None:
|
||||
raise RuntimeError("Celery is disabled")
|
||||
|
||||
startup_recovery_once = _DisabledTask()
|
||||
recover_download_tasks_once = _DisabledTask()
|
||||
recover_generation_tasks_once = _DisabledTask()
|
||||
|
||||
@@ -21,7 +21,7 @@ from app.tasks.celery_app import celery_app
|
||||
MODULE = ModuleCodeEnum.HOT_OPENING_REPLICATE.value
|
||||
|
||||
|
||||
async def _run_image_prompt(project_id: str, step_id: str | None = None):
|
||||
async def _run_image_prompt(project_id: str, step_id: str | None = None) -> None:
|
||||
lock_token: str | None = None
|
||||
if step_id:
|
||||
lock_token = await acquire_object_lock(object_type=OBJECT_MODULE_STEP, object_id=step_id)
|
||||
@@ -38,17 +38,17 @@ async def _run_image_prompt(project_id: str, step_id: str | None = None):
|
||||
await mark_active_started(object_type=OBJECT_MODULE_STEP, object_id=step_id)
|
||||
try:
|
||||
async with async_session() as db:
|
||||
result = await run_image_prompt_optimize(db, project_id=project_id, step_id=step_id)
|
||||
await run_image_prompt_optimize(db, project_id=project_id, step_id=step_id)
|
||||
await db.commit()
|
||||
if step_id:
|
||||
await cleanup_active_if_terminal(db, object_type=OBJECT_MODULE_STEP, object_id=step_id)
|
||||
return result
|
||||
return None
|
||||
finally:
|
||||
if step_id:
|
||||
await release_object_lock(object_type=OBJECT_MODULE_STEP, object_id=step_id, token=lock_token)
|
||||
|
||||
|
||||
async def _run_video_prompt(project_id: str, step_id: str | None = None):
|
||||
async def _run_video_prompt(project_id: str, step_id: str | None = None) -> None:
|
||||
lock_token: str | None = None
|
||||
if step_id:
|
||||
lock_token = await acquire_object_lock(object_type=OBJECT_MODULE_STEP, object_id=step_id)
|
||||
@@ -64,18 +64,18 @@ async def _run_video_prompt(project_id: str, step_id: str | None = None):
|
||||
await mark_active_started(object_type=OBJECT_MODULE_STEP, object_id=step_id)
|
||||
try:
|
||||
async with async_session() as db:
|
||||
result = await run_video_prompt_optimize(db, project_id=project_id, step_id=step_id)
|
||||
await run_video_prompt_optimize(db, project_id=project_id, step_id=step_id)
|
||||
await db.commit()
|
||||
if step_id:
|
||||
await cleanup_active_if_terminal(db, object_type=OBJECT_MODULE_STEP, object_id=step_id)
|
||||
return result
|
||||
return None
|
||||
finally:
|
||||
if step_id:
|
||||
await release_object_lock(object_type=OBJECT_MODULE_STEP, object_id=step_id, token=lock_token)
|
||||
|
||||
|
||||
if celery_app:
|
||||
@celery_app.task(name="hot_opening.start_image_prompt_optimize", bind=True, max_retries=3, default_retry_delay=30)
|
||||
@celery_app.task(name="hot_opening.start_image_prompt_optimize", bind=True, max_retries=3, default_retry_delay=30, ignore_result=True)
|
||||
def start_image_prompt_optimize(self, project_id: str, step_id: str | None = None):
|
||||
"""手动触发后的图片 AI 提词任务。
|
||||
|
||||
@@ -84,11 +84,12 @@ if celery_app:
|
||||
service 内部已经落库为业务失败的情况不会抛出异常,也不会重复 retry。
|
||||
"""
|
||||
try:
|
||||
return run_async(_run_image_prompt(project_id, step_id))
|
||||
run_async(_run_image_prompt(project_id, step_id))
|
||||
return None
|
||||
except Exception as exc:
|
||||
raise self.retry(exc=exc) from exc
|
||||
|
||||
@celery_app.task(name="hot_opening.start_video_prompt_optimize", bind=True, max_retries=3, default_retry_delay=30)
|
||||
@celery_app.task(name="hot_opening.start_video_prompt_optimize", bind=True, max_retries=3, default_retry_delay=30, ignore_result=True)
|
||||
def start_video_prompt_optimize(self, project_id: str, step_id: str | None = None):
|
||||
"""手动触发后的视频 AI 提词任务。
|
||||
|
||||
@@ -97,7 +98,8 @@ if celery_app:
|
||||
service 内部已经落库为业务失败的情况不会抛出异常,也不会重复 retry。
|
||||
"""
|
||||
try:
|
||||
return run_async(_run_video_prompt(project_id, step_id))
|
||||
run_async(_run_video_prompt(project_id, step_id))
|
||||
return None
|
||||
except Exception as exc:
|
||||
raise self.retry(exc=exc) from exc
|
||||
else:
|
||||
@@ -109,4 +111,4 @@ else:
|
||||
raise RuntimeError("Celery is disabled")
|
||||
|
||||
start_image_prompt_optimize = _DisabledTask()
|
||||
start_video_prompt_optimize = _DisabledTask()
|
||||
start_video_prompt_optimize = _DisabledTask()
|
||||
@@ -2,21 +2,49 @@ from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from app.config import settings
|
||||
from app.models.base import async_session
|
||||
from app.services.module_async_recovery_service import recover_module_async_tasks_once
|
||||
from app.services.redis_registry_service import get_registry_redis, redis_acquire_lock, redis_release_lock
|
||||
from app.tasks.async_runner import run_async
|
||||
from app.tasks.celery_app import celery_app
|
||||
|
||||
|
||||
async def _run_recover_module_async_tasks_once() -> dict[str, Any]:
|
||||
async with async_session() as db:
|
||||
return await recover_module_async_tasks_once(db)
|
||||
redis = await get_registry_redis()
|
||||
token: str | None = None
|
||||
if redis is not None:
|
||||
token = await redis_acquire_lock(
|
||||
lock_key=settings.MODULE_ASYNC_RECOVERY_LOCK_KEY,
|
||||
ttl_seconds=int(settings.CELERY_RECOVERY_TASK_LOCK_TTL_SECONDS or 600),
|
||||
log_context="module_async_recovery",
|
||||
)
|
||||
if not token:
|
||||
return {"skipped": "lock_held", "lock_key": settings.MODULE_ASYNC_RECOVERY_LOCK_KEY}
|
||||
|
||||
try:
|
||||
async with async_session() as db:
|
||||
result = await recover_module_async_tasks_once(db)
|
||||
result["execution_lock"] = "lock_acquired" if token else "redis_unavailable_run_db_fallback"
|
||||
return result
|
||||
finally:
|
||||
if token:
|
||||
await redis_release_lock(
|
||||
lock_key=settings.MODULE_ASYNC_RECOVERY_LOCK_KEY,
|
||||
token=token,
|
||||
log_context="module_async_recovery",
|
||||
)
|
||||
|
||||
|
||||
if celery_app:
|
||||
|
||||
@celery_app.task(name="module_async.recover_module_async_tasks_once")
|
||||
def recover_module_async_tasks_once_task() -> dict[str, Any]:
|
||||
@celery_app.task(
|
||||
name="module_async.recover_module_async_tasks_once",
|
||||
bind=True,
|
||||
soft_time_limit=settings.CELERY_RECOVERY_SOFT_TIME_LIMIT_SECONDS,
|
||||
time_limit=settings.CELERY_RECOVERY_TIME_LIMIT_SECONDS,
|
||||
)
|
||||
def recover_module_async_tasks_once_task(self) -> dict[str, Any]:
|
||||
return run_async(_run_recover_module_async_tasks_once())
|
||||
|
||||
else:
|
||||
|
||||
@@ -21,7 +21,7 @@ from app.tasks.celery_app import celery_app
|
||||
MODULE = ModuleCodeEnum.SHOT_REPLICATE.value
|
||||
|
||||
|
||||
async def _run_image_prompt(project_id: str, step_id: str | None = None):
|
||||
async def _run_image_prompt(project_id: str, step_id: str | None = None) -> None:
|
||||
lock_token: str | None = None
|
||||
if step_id:
|
||||
lock_token = await acquire_object_lock(object_type=OBJECT_MODULE_STEP, object_id=step_id)
|
||||
@@ -37,17 +37,17 @@ async def _run_image_prompt(project_id: str, step_id: str | None = None):
|
||||
await mark_active_started(object_type=OBJECT_MODULE_STEP, object_id=step_id)
|
||||
try:
|
||||
async with async_session() as db:
|
||||
result = await run_image_prompt_optimize(db, project_id=project_id, step_id=step_id)
|
||||
await run_image_prompt_optimize(db, project_id=project_id, step_id=step_id)
|
||||
await db.commit()
|
||||
if step_id:
|
||||
await cleanup_active_if_terminal(db, object_type=OBJECT_MODULE_STEP, object_id=step_id)
|
||||
return result
|
||||
return None
|
||||
finally:
|
||||
if step_id:
|
||||
await release_object_lock(object_type=OBJECT_MODULE_STEP, object_id=step_id, token=lock_token)
|
||||
|
||||
|
||||
async def _run_video_prompt(project_id: str, step_id: str | None = None):
|
||||
async def _run_video_prompt(project_id: str, step_id: str | None = None) -> None:
|
||||
lock_token: str | None = None
|
||||
if step_id:
|
||||
lock_token = await acquire_object_lock(object_type=OBJECT_MODULE_STEP, object_id=step_id)
|
||||
@@ -63,11 +63,11 @@ async def _run_video_prompt(project_id: str, step_id: str | None = None):
|
||||
await mark_active_started(object_type=OBJECT_MODULE_STEP, object_id=step_id)
|
||||
try:
|
||||
async with async_session() as db:
|
||||
result = await run_video_prompt_optimize(db, project_id=project_id, step_id=step_id)
|
||||
await run_video_prompt_optimize(db, project_id=project_id, step_id=step_id)
|
||||
await db.commit()
|
||||
if step_id:
|
||||
await cleanup_active_if_terminal(db, object_type=OBJECT_MODULE_STEP, object_id=step_id)
|
||||
return result
|
||||
return None
|
||||
finally:
|
||||
if step_id:
|
||||
await release_object_lock(object_type=OBJECT_MODULE_STEP, object_id=step_id, token=lock_token)
|
||||
@@ -75,7 +75,7 @@ async def _run_video_prompt(project_id: str, step_id: str | None = None):
|
||||
|
||||
if celery_app:
|
||||
|
||||
@celery_app.task(name="shot_replicate.start_image_prompt_optimize", bind=True, max_retries=3, default_retry_delay=30)
|
||||
@celery_app.task(name="shot_replicate.start_image_prompt_optimize", bind=True, max_retries=3, default_retry_delay=30, ignore_result=True)
|
||||
def start_image_prompt_optimize(self, project_id: str, step_id: str | None = None):
|
||||
"""手动触发后的图片 AI 提词任务。
|
||||
|
||||
@@ -84,12 +84,13 @@ if celery_app:
|
||||
service 内部已经落库为业务失败的情况不会抛出异常,也不会重复 retry。
|
||||
"""
|
||||
try:
|
||||
return run_async(_run_image_prompt(project_id, step_id))
|
||||
run_async(_run_image_prompt(project_id, step_id))
|
||||
return None
|
||||
except Exception as exc:
|
||||
raise self.retry(exc=exc) from exc
|
||||
|
||||
|
||||
@celery_app.task(name="shot_replicate.start_video_prompt_optimize", bind=True, max_retries=3, default_retry_delay=30)
|
||||
@celery_app.task(name="shot_replicate.start_video_prompt_optimize", bind=True, max_retries=3, default_retry_delay=30, ignore_result=True)
|
||||
def start_video_prompt_optimize(self, project_id: str, step_id: str | None = None):
|
||||
"""手动触发后的视频 AI 提词任务。
|
||||
|
||||
@@ -98,7 +99,8 @@ if celery_app:
|
||||
service 内部已经落库为业务失败的情况不会抛出异常,也不会重复 retry。
|
||||
"""
|
||||
try:
|
||||
return run_async(_run_video_prompt(project_id, step_id))
|
||||
run_async(_run_video_prompt(project_id, step_id))
|
||||
return None
|
||||
except Exception as exc:
|
||||
raise self.retry(exc=exc) from exc
|
||||
|
||||
@@ -112,4 +114,4 @@ else:
|
||||
raise RuntimeError("Celery is disabled")
|
||||
|
||||
start_image_prompt_optimize = _DisabledTask()
|
||||
start_video_prompt_optimize = _DisabledTask()
|
||||
start_video_prompt_optimize = _DisabledTask()
|
||||
@@ -20,7 +20,7 @@ from app.models.base import async_session
|
||||
from app.models.shot_replicate_segment import ShotReplicateSegment
|
||||
from app.models.shot_replicate_task_set import ShotReplicateTaskSet
|
||||
from app.services.module_generation_log_service import log_module_error, log_module_event_file, log_module_prompt_event
|
||||
from app.services.redis_registry_service import redis_acquire_lock, redis_release_lock
|
||||
from app.services.redis_registry_service import get_registry_redis, redis_acquire_lock, redis_release_lock
|
||||
from app.services.module_async_recovery_service import (
|
||||
OBJECT_SHOT_SEGMENT_ANALYSIS,
|
||||
OBJECT_SHOT_SPLIT_SEGMENT,
|
||||
@@ -561,8 +561,29 @@ async def _run_split_one_segment(segment_id: str) -> None:
|
||||
async def _run_recover_split_tasks_once() -> dict[str, Any]:
|
||||
from app.services.shot_replicate_recovery_service import recover_shot_split_tasks_once
|
||||
|
||||
async with async_session() as db:
|
||||
return await recover_shot_split_tasks_once(db)
|
||||
redis = await get_registry_redis()
|
||||
token: str | None = None
|
||||
if redis is not None:
|
||||
token = await redis_acquire_lock(
|
||||
lock_key=settings.SHOT_SPLIT_RECOVERY_LOCK_KEY,
|
||||
ttl_seconds=int(settings.CELERY_RECOVERY_TASK_LOCK_TTL_SECONDS or 600),
|
||||
log_context="shot_split_recovery",
|
||||
)
|
||||
if not token:
|
||||
return {"skipped": "lock_held", "lock_key": settings.SHOT_SPLIT_RECOVERY_LOCK_KEY}
|
||||
|
||||
try:
|
||||
async with async_session() as db:
|
||||
result = await recover_shot_split_tasks_once(db)
|
||||
result["execution_lock"] = "lock_acquired" if token else "redis_unavailable_run_db_fallback"
|
||||
return result
|
||||
finally:
|
||||
if token:
|
||||
await redis_release_lock(
|
||||
lock_key=settings.SHOT_SPLIT_RECOVERY_LOCK_KEY,
|
||||
token=token,
|
||||
log_context="shot_split_recovery",
|
||||
)
|
||||
|
||||
|
||||
if celery_app:
|
||||
@@ -582,8 +603,13 @@ if celery_app:
|
||||
return run_async(_run_analyze_custom_segment_video(segment_id))
|
||||
|
||||
|
||||
@celery_app.task(name="shot_replicate.recover_split_tasks_once")
|
||||
def recover_split_tasks_once() -> dict[str, Any]:
|
||||
@celery_app.task(
|
||||
name="shot_replicate.recover_split_tasks_once",
|
||||
bind=True,
|
||||
soft_time_limit=settings.CELERY_RECOVERY_SOFT_TIME_LIMIT_SECONDS,
|
||||
time_limit=settings.CELERY_RECOVERY_TIME_LIMIT_SECONDS,
|
||||
)
|
||||
def recover_split_tasks_once(self) -> dict[str, Any]:
|
||||
return run_async(_run_recover_split_tasks_once())
|
||||
|
||||
else:
|
||||
|
||||
Reference in New Issue
Block a user