修复生辰任务恢复机制BUG|修复生成视频/图片token快照不回落交易流水BUG

This commit is contained in:
2026-06-26 11:51:02 +08:00
parent f3675924dc
commit 06cdab9cbc
12 changed files with 773 additions and 125 deletions
+2
View File
@@ -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)
+9
View File
@@ -127,6 +127,15 @@ class Settings(BaseSettings):
DOWNLOAD_TASK_QUEUE_TIMEOUT_SECONDS: int = 5 * 60
DOWNLOAD_RECOVERY_BATCH_SIZE: int = 100
DOWNLOAD_RECOVERY_STARTUP_DELAY_SECONDS: int = 3
# 下载恢复自循环:不依赖 Celery beat,不新增 worker;由 gen_result_download 队列周期扫描 DB/Redis。
DOWNLOAD_RECOVERY_LOOP_ENABLED: bool = True
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"
+1
View File
@@ -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 *
+101
View File
@@ -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,
@@ -19,6 +25,7 @@ from app.services.celery_download_recovery_service import (
)
from app.services.generation_log_service import log_provider_call, 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 (
@@ -30,7 +37,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 +65,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 +143,21 @@ 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 +173,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 +196,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 +219,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 +284,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 +327,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 +336,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 +353,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()
@@ -397,6 +410,7 @@ async def _try_final_poll_before_timeout(db: AsyncSession, task: ChatGenerationT
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,
@@ -405,14 +419,14 @@ async def _try_final_poll_before_timeout(db: AsyncSession, task: ChatGenerationT
detail=poll_result,
)
task.pipeline_stage = "result_ready"
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="result_ready",
to_stage=ChatGenerationPipelineStage.RESULT_READY.value,
detail=poll_result,
)
await enqueue_download_task(db, task, recover=True, reason="final_poll_before_timeout_success")
@@ -463,13 +477,13 @@ async def recover_one_generation_task(
return "clean_not_generating"
if task.deadline_at and _is_expired(task.deadline_at, current_time):
if task.pipeline_stage in ("waiting_remote", "polling"):
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)
if task.pipeline_stage in ("queued", "preparing", "creating_provider_task"):
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:
task.pipeline_stage = "waiting_remote"
task.pipeline_stage = ChatGenerationPipelineStage.WAITING_REMOTE.value
await db.commit()
await log_task_event(
task,
@@ -498,7 +512,7 @@ async def recover_one_generation_task(
)
return "recover_create"
if task.pipeline_stage in ("waiting_remote", "polling"):
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(
@@ -516,7 +530,7 @@ async def recover_one_generation_task(
message=f"{source} 发现远程等待/轮询阶段任务未完成,恢复投递轮询队列",
detail={"pipeline_stage": task.pipeline_stage, "payload": redis_payload},
)
task.pipeline_stage = "waiting_remote"
task.pipeline_stage = ChatGenerationPipelineStage.WAITING_REMOTE.value
await db.commit()
poll_generation_task.apply_async(
args=[task.id],
@@ -536,7 +550,7 @@ async def recover_one_generation_task(
message=f"{source} 发现任务缺少供应商任务ID,恢复投递创建队列",
detail={"pipeline_stage": task.pipeline_stage, "payload": redis_payload},
)
task.pipeline_stage = "queued"
task.pipeline_stage = ChatGenerationPipelineStage.QUEUED.value
await db.commit()
await _remove_poll_active(task.id)
chatapi_create_generation_task.apply_async(
@@ -546,7 +560,7 @@ async def recover_one_generation_task(
)
return "recover_create_missing_provider_id"
if task.pipeline_stage == "result_ready":
if task.pipeline_stage == ChatGenerationPipelineStage.RESULT_READY.value:
await _remove_poll_active(task.id)
if task.remote_result_url:
await enqueue_download_task(
@@ -617,8 +631,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",
@@ -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:
@@ -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。
@@ -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,
await _apply_download_async(
task,
priority=priority,
countdown=countdown,
task_id=celery_task_id,
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):
return False
await log_task_event(
await _log_download_event(
task,
event_type="DOWNLOAD_STUCK_RECOVER",
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_download_event(
task,
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,47 +576,22 @@ 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,
delay_seconds = _countdown_until(next_retry_at)
await _apply_download_async(
task,
priority=settings.DOWNLOAD_TASK_PRIORITY_RECOVER,
countdown=delay_seconds,
task_id=task.download_celery_task_id,
reason="download_exception_retry",
event_type=ChatGenerationTaskEventType.DOWNLOAD_RETRY_ENQUEUE,
failed_event_type=ChatGenerationTaskEventType.DOWNLOAD_RETRY_ENQUEUE_FAILED,
)
@@ -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,
@@ -245,6 +246,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)
@@ -1,12 +1,17 @@
# app/tasks/generation_recovery_tasks.py
from __future__ import annotations
import logging
from typing import Any, 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
from app.tasks.async_runner import run_async
from app.tasks.celery_app import celery_app
logger = logging.getLogger("video_gen")
async def _run_download_once() -> Dict[str, Any]:
from app.services.generation_recovery_service import recover_download_tasks_once
@@ -22,11 +27,49 @@ async def _run_generation_once() -> Dict[str, Any]:
return await recover_generation_tasks_once(db)
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", True)):
return
try:
recover_download_tasks_once.apply_async(
countdown=max(1, int(settings.DOWNLOAD_RECOVERY_INTERVAL_SECONDS or 60)),
queue="gen_result_download",
priority=settings.DOWNLOAD_TASK_PRIORITY_RECOVER,
)
except Exception:
logger.exception("下载恢复循环下一轮投递失败")
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="generation.recover_download_tasks_once", bind=True)
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_download_once())
result["loop_lock"] = reason
return result
finally:
_schedule_next_download_recovery_loop()
@celery_app.task(name="generation.recover_generation_tasks_once")