celery异步恢复任务独立队列|扩增celery子进程连接池上限配置

This commit is contained in:
2026-06-26 13:15:25 +08:00
parent 06cdab9cbc
commit ee03242e6c
10 changed files with 517 additions and 245 deletions
@@ -23,10 +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.media_token_usage_snapshot_service import sync_chat_generation_task_media_token_snapshot
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,
@@ -145,10 +143,20 @@ async def recover_one_download_task(
return "clean_final_state"
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})
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})
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
@@ -362,94 +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
await sync_chat_generation_task_media_token_snapshot(db, task, provider_response=response_data)
if not task.remote_result_url:
return await _mark_failed(
db,
task,
error_message="供应商任务成功但未返回结果URL",
detail=poll_result,
)
task.pipeline_stage = ChatGenerationPipelineStage.RESULT_READY.value
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=ChatGenerationPipelineStage.RESULT_READY.value,
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,
@@ -457,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、无供应商任务 IDdeadline 未过才恢复 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
@@ -472,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 (ChatGenerationPipelineStage.WAITING_REMOTE.value, ChatGenerationPipelineStage.POLLING.value):
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 (ChatGenerationPipelineStage.QUEUED.value, ChatGenerationPipelineStage.PREPARING.value, ChatGenerationPipelineStage.CREATING_PROVIDER_TASK.value):
if task.provider_task_id or task.seedance_task_id:
# 最高优先级:只要远程结果 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(
@@ -510,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 (ChatGenerationPipelineStage.WAITING_REMOTE.value, ChatGenerationPipelineStage.POLLING.value):
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 = ChatGenerationPipelineStage.WAITING_REMOTE.value
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},
)
# 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 == ChatGenerationPipelineStage.RESULT_READY.value:
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}"
@@ -670,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,
}
}