AI创作批量生成任务 main V2 build push

This commit is contained in:
2026-07-15 13:10:15 +08:00
parent 9756a86304
commit 0ffb35e2df
20 changed files with 668 additions and 5362 deletions
File diff suppressed because it is too large Load Diff
@@ -1,571 +0,0 @@
from __future__ import annotations
import re
from dataclasses import asdict, dataclass
from typing import Any, Mapping
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.enums.credit_record import CreditRecordBillingScene, CreditRecordChargeKind, CreditRecordOwnerType, CreditRecordSourceModule
from app.models.credit_record import CreditRecord
from app.models.generation_record import GenerationRecord
from app.models.module_generation_step import ModuleGenerationStep
from app.models.token_usage import TokenUsage
from app.models.system_config import SystemConfig
from app.services.credit_record_meta_service import (
CreditRecordMeta,
build_generation_media_meta,
build_generation_record_prompt_meta,
build_module_step_prompt_meta,
build_shot_video_analysis_meta,
)
from app.services.credits import calc_image_credits, calc_text_credits, calc_video_credits, deduct_credits
CHARGE_TEXT_PROMPT = CreditRecordChargeKind.TEXT_PROMPT.value
CHARGE_FILE_PARSE = CreditRecordChargeKind.FILE_PARSE.value
CHARGE_VISION_INPUT = CreditRecordChargeKind.VISION_INPUT.value
CHARGE_MEDIA = CreditRecordChargeKind.MEDIA.value
CHARGE_VIDEO_ANALYSIS = CreditRecordChargeKind.VIDEO_ANALYSIS.value
OWNER_GENERATION_RECORD = CreditRecordOwnerType.GENERATION_RECORD.value
OWNER_CHAT_GENERATION_TASK = CreditRecordOwnerType.CHAT_GENERATION_TASK.value
OWNER_MODULE_GENERATION_STEP = CreditRecordOwnerType.MODULE_GENERATION_STEP.value
OWNER_SHOT_REPLICATE_TASK_SET = CreditRecordOwnerType.SHOT_REPLICATE_TASK_SET.value
OWNER_SHOT_REPLICATE_SEGMENT = CreditRecordOwnerType.SHOT_REPLICATE_SEGMENT.value
_BIZ_KEY_PATTERN = re.compile(
r"^(?P<owner_type>[^:]+):(?P<owner_id>[^:]+):attempt:(?P<attempt_no>\d+):(?P<charge_kind>[^:]+):(?P<action>charge|refund)$"
)
@dataclass
class BillingItem:
charge_key: str
amount: float
charged: bool
skipped_reason: str | None = None
biz_key: str | None = None
attempt_no: int | None = None
@dataclass
class BillingSummary:
record_id: str
user_id: str
items: list[BillingItem]
@property
def total_charged(self) -> float:
return round(sum(item.amount for item in self.items if item.charged), 2)
def get_amount(self, charge_key: str) -> float:
return round(sum(item.amount for item in self.items if item.charge_key == charge_key and item.charged), 2)
def to_dict(self) -> dict[str, Any]:
data = asdict(self)
data["total_charged"] = self.total_charged
return data
def _round2(value: float | int | None) -> float:
return round(float(value or 0), 2)
def _safe_int(value: Any, default: int = 0) -> int:
try:
if value is None:
return default
return int(value)
except Exception:
return default
def build_credit_biz_key(
*,
owner_type: str,
owner_id: str,
attempt_no: int,
charge_kind: str,
action: str,
) -> str:
"""生成正式积分流水幂等键。
示例:generation_record:xxx:attempt:2:media:charge
"""
owner_type = owner_type.strip()
owner_id = owner_id.strip()
charge_kind = charge_kind.strip()
action = action.strip()
if action not in ("charge", "refund"):
raise ValueError("action 仅支持 charge/refund")
if attempt_no <= 0:
raise ValueError("attempt_no 必须大于 0")
return f"{owner_type}:{owner_id}:attempt:{attempt_no}:{charge_kind}:{action}"
def parse_credit_biz_key(biz_key: str | None) -> dict[str, Any] | None:
if not biz_key:
return None
match = _BIZ_KEY_PATTERN.match(biz_key)
if not match:
return None
data = match.groupdict()
data["attempt_no"] = int(data["attempt_no"])
return data
async def _get_config_float_or_none(db: AsyncSession, key: str) -> float | None:
result = await db.execute(select(SystemConfig).where(SystemConfig.key == key).limit(1))
config = result.scalar_one_or_none()
if not config:
return None
try:
return float(config.value)
except Exception:
return None
async def _calc_optional_token_credits(db: AsyncSession, tokens: int, config_key: str) -> float:
tokens = _safe_int(tokens)
if tokens <= 0:
return 0.0
rate = await _get_config_float_or_none(db, config_key)
if rate is None:
return 0.0
return round(tokens * rate / 1000, 2)
async def _find_existing_by_biz_key(db: AsyncSession, *, user_id: str, biz_key: str) -> CreditRecord | None:
result = await db.execute(
select(CreditRecord)
.where(CreditRecord.user_id == user_id, CreditRecord.biz_key == biz_key)
.limit(1)
)
return result.scalar_one_or_none()
async def get_next_credit_attempt_no(
db: AsyncSession,
*,
owner_type: str,
owner_id: str,
charge_kind: str = CHARGE_MEDIA,
) -> int:
"""根据已有正式 biz_key 计算下一轮扣费 attempt_no。
不依赖任务表 retry_count,防止 worker 崩溃/重复消息造成状态字段不可信。
"""
prefix = f"{owner_type}:{owner_id}:attempt:%:{charge_kind}:charge"
result = await db.execute(
select(CreditRecord.biz_key)
.where(CreditRecord.related_id == owner_id)
.where(CreditRecord.type == "consume")
.where(CreditRecord.biz_key.like(prefix))
)
max_attempt = 0
for (biz_key,) in result.all():
parsed = parse_credit_biz_key(biz_key)
if parsed and parsed.get("owner_type") == owner_type and parsed.get("owner_id") == owner_id:
max_attempt = max(max_attempt, int(parsed.get("attempt_no") or 0))
return max_attempt + 1
async def deduct_credits_locked_once(
db: AsyncSession,
*,
user_id: str,
amount: float,
description: str,
related_id: str,
charge_key: str,
biz_key: str | None = None,
attempt_no: int | None = None,
record_meta: CreditRecordMeta | dict | None = None,
) -> BillingItem:
"""按 biz_key 做幂等扣费。
charge_key 只保留为业务分类;正式幂等以 biz_key 为准。
record_meta 负责把业务归属、模块、步骤、token、模型快照写入 CreditRecord。
"""
amount = _round2(amount)
if amount <= 0:
return BillingItem(charge_key=charge_key, amount=0.0, charged=False, skipped_reason="amount_lte_zero", biz_key=biz_key, attempt_no=attempt_no)
if biz_key:
existing_charge = await _find_existing_by_biz_key(db, user_id=user_id, biz_key=biz_key)
if existing_charge:
return BillingItem(
charge_key=charge_key,
amount=abs(_round2(existing_charge.amount)),
charged=False,
skipped_reason="already_charged",
biz_key=biz_key,
attempt_no=attempt_no,
)
await deduct_credits(
db,
user_id=user_id,
amount=amount,
description=description,
related_id=related_id,
biz_key=biz_key,
record_meta=record_meta,
)
return BillingItem(charge_key=charge_key, amount=amount, charged=True, biz_key=biz_key, attempt_no=attempt_no)
async def charge_chatapi_prompt_usage(
db: AsyncSession,
*,
record: GenerationRecord,
usage: Mapping[str, Any],
project_name: str | None = None,
) -> BillingSummary:
"""项目记录提示词整理扣费;不参与生成失败媒体退款。"""
project_name = project_name or "AI生成任务"
items: list[BillingItem] = []
attempt_no = 1
owner_type = OWNER_GENERATION_RECORD
owner_id = record.id
input_tokens = _safe_int(usage.get("input_tokens"))
output_tokens = _safe_int(usage.get("output_tokens"))
text_credits = await calc_text_credits(db, input_tokens, output_tokens)
text_meta = await build_generation_record_prompt_meta(
db,
record_id=record.id,
attempt_no=attempt_no,
charge_kind=CHARGE_TEXT_PROMPT,
usage=usage,
)
items.append(
await deduct_credits_locked_once(
db,
user_id=record.user_id,
amount=text_credits,
description=f"提示词优化 - {project_name}",
related_id=record.id,
charge_key=CHARGE_TEXT_PROMPT,
biz_key=build_credit_biz_key(
owner_type=owner_type,
owner_id=owner_id,
attempt_no=attempt_no,
charge_kind=CHARGE_TEXT_PROMPT,
action="charge",
),
attempt_no=attempt_no,
record_meta=text_meta,
)
)
file_tokens = usage.get("file_parse_tokens") or usage.get("file_tokens") or usage.get("document_tokens") or 0
file_parse_credits = await _calc_optional_token_credits(db, _safe_int(file_tokens), "file_parse_credits_per_1000_tokens")
file_meta = await build_generation_record_prompt_meta(
db,
record_id=record.id,
attempt_no=attempt_no,
charge_kind=CHARGE_FILE_PARSE,
usage={**dict(usage), "total_tokens": _safe_int(file_tokens), "input_tokens": _safe_int(file_tokens), "output_tokens": 0},
)
items.append(
await deduct_credits_locked_once(
db,
user_id=record.user_id,
amount=file_parse_credits,
description="文件解析Token",
related_id=record.id,
charge_key=CHARGE_FILE_PARSE,
biz_key=build_credit_biz_key(
owner_type=owner_type,
owner_id=owner_id,
attempt_no=attempt_no,
charge_kind=CHARGE_FILE_PARSE,
action="charge",
),
attempt_no=attempt_no,
record_meta=file_meta,
)
)
vision_tokens = usage.get("vision_input_tokens") or usage.get("image_input_tokens") or usage.get("image_tokens") or 0
vision_input_credits = await _calc_optional_token_credits(db, _safe_int(vision_tokens), "vision_input_credits_per_1000_tokens")
vision_meta = await build_generation_record_prompt_meta(
db,
record_id=record.id,
attempt_no=attempt_no,
charge_kind=CHARGE_VISION_INPUT,
usage={**dict(usage), "total_tokens": _safe_int(vision_tokens), "input_tokens": _safe_int(vision_tokens), "output_tokens": 0},
)
items.append(
await deduct_credits_locked_once(
db,
user_id=record.user_id,
amount=vision_input_credits,
description="图片理解Token",
related_id=record.id,
charge_key=CHARGE_VISION_INPUT,
biz_key=build_credit_biz_key(
owner_type=owner_type,
owner_id=owner_id,
attempt_no=attempt_no,
charge_kind=CHARGE_VISION_INPUT,
action="charge",
),
attempt_no=attempt_no,
record_meta=vision_meta,
)
)
if hasattr(record, "text_credits_cost"):
record.text_credits_cost = round(text_credits + file_parse_credits + vision_input_credits, 2)
if hasattr(record, "text_tokens_used"):
record.text_tokens_used = _safe_int(usage.get("total_tokens"), input_tokens + output_tokens)
return BillingSummary(record_id=record.id, user_id=record.user_id, items=items)
async def charge_module_prompt_usage(
db: AsyncSession,
*,
user_id: str,
step_id: str,
usage: Mapping[str, Any],
description: str,
attempt_no: int = 1,
) -> BillingSummary:
"""爆款开头复刻/拆镜复刻模块图片/视频 AI 提词扣文本积分。
文本提词属于已经发生的 LLM 消费:
- 调用成功后按 input_tokens + output_tokens 扣费。
- 不参与后续图片/视频媒体生成失败退款。
- 通过 module_generation_step:{step_id}:attempt:1:text_prompt:charge 幂等。
"""
input_tokens = _safe_int(usage.get("input_tokens"))
output_tokens = _safe_int(usage.get("output_tokens"))
text_credits = await calc_text_credits(db, input_tokens, output_tokens)
biz_key = build_credit_biz_key(
owner_type=OWNER_MODULE_GENERATION_STEP,
owner_id=step_id,
attempt_no=attempt_no,
charge_kind=CHARGE_TEXT_PROMPT,
action="charge",
)
record_meta = await build_module_step_prompt_meta(db, step_id=step_id, attempt_no=attempt_no, usage=usage)
item = await deduct_credits_locked_once(
db,
user_id=user_id,
amount=text_credits,
description=description,
related_id=step_id,
charge_key=CHARGE_TEXT_PROMPT,
biz_key=biz_key,
attempt_no=attempt_no,
record_meta=record_meta,
)
result = await db.execute(select(ModuleGenerationStep).where(ModuleGenerationStep.id == step_id).limit(1))
step = result.scalar_one_or_none()
if step:
step.token_usage_id = record_meta.token_usage_id
step.model_config_id = usage.get("model_config_id")
step.input_tokens = record_meta.input_tokens
step.output_tokens = record_meta.output_tokens
step.total_tokens = record_meta.total_tokens
step.text_credits_cost = text_credits
return BillingSummary(record_id=step_id, user_id=user_id, items=[item])
async def charge_shot_video_analysis_usage(
db: AsyncSession,
*,
user_id: str,
owner_type: str,
owner_id: str,
usage: Mapping[str, Any],
description: str,
billing_scene: str,
source_project_id: str | None = None,
source_step_id: str | None = None,
attempt_no: int | None = None,
) -> BillingSummary:
"""拆镜复刻视频分析扣分析积分。
视频分析属于“文字提示词 + 视频素材”的模型调用类消费,
按 input_tokens + output_tokens 参考文本积分规则计费,
但账务归类为 analysis/video_analysis,避免混入提词优化统计。
"""
attempt_no = attempt_no or await get_next_credit_attempt_no(
db,
owner_type=owner_type,
owner_id=owner_id,
charge_kind=CHARGE_VIDEO_ANALYSIS,
)
input_tokens = _safe_int(usage.get("input_tokens"))
output_tokens = _safe_int(usage.get("output_tokens"))
amount = await calc_text_credits(db, input_tokens, output_tokens)
biz_key = build_credit_biz_key(
owner_type=owner_type,
owner_id=owner_id,
attempt_no=attempt_no,
charge_kind=CHARGE_VIDEO_ANALYSIS,
action="charge",
)
record_meta = await build_shot_video_analysis_meta(
db,
owner_type=owner_type,
owner_id=owner_id,
attempt_no=attempt_no,
usage=usage,
billing_scene=billing_scene,
source_project_id=source_project_id,
source_step_id=source_step_id,
)
item = await deduct_credits_locked_once(
db,
user_id=user_id,
amount=amount,
description=description,
related_id=owner_id,
charge_key=CHARGE_VIDEO_ANALYSIS,
biz_key=biz_key,
attempt_no=attempt_no,
record_meta=record_meta,
)
if record_meta.token_usage_id:
result = await db.execute(select(TokenUsage).where(TokenUsage.id == record_meta.token_usage_id).limit(1))
token_usage = result.scalar_one_or_none()
if token_usage:
token_usage.owner_type = token_usage.owner_type or owner_type
token_usage.owner_id = token_usage.owner_id or owner_id
token_usage.biz_key = token_usage.biz_key or biz_key
token_usage.source_module = token_usage.source_module or CreditRecordSourceModule.SHOT_REPLICATE.value
token_usage.source_step_code = token_usage.source_step_code or "video_analysis"
return BillingSummary(record_id=owner_id, user_id=user_id, items=[item])
async def charge_generation_media_by_params(
db: AsyncSession,
*,
user_id: str,
record_id: str,
gen_type: str,
image_size: str | None = None,
duration: int | None = None,
resolution: str | None = None,
engine_id: str | None = None,
input_video_duration: float | None = None,
project_name: str | None = None,
description_prefix: str = "AI创作-",
owner_type: str = OWNER_CHAT_GENERATION_TASK,
attempt_no: int | None = None,
source_module: str | None = None,
source_project_id: str | None = None,
source_step_id: str | None = None,
source_step_code: str | None = None,
billing_scene: str | None = None,
) -> BillingSummary:
"""图片/视频媒体生成扣费。
正式幂等由 owner_type + record_id + attempt_no + media + charge 组成。
每次用户主动重试必须传入新的 attempt_no。
"""
project_name = project_name or "AI生成任务"
gen_type = (gen_type or "").lower().strip()
attempt_no = attempt_no or await get_next_credit_attempt_no(
db,
owner_type=owner_type,
owner_id=record_id,
charge_kind=CHARGE_MEDIA,
)
biz_key = build_credit_biz_key(
owner_type=owner_type,
owner_id=record_id,
attempt_no=attempt_no,
charge_kind=CHARGE_MEDIA,
action="charge",
)
items: list[BillingItem] = []
record_meta = await build_generation_media_meta(
db,
owner_type=owner_type,
owner_id=record_id,
attempt_no=attempt_no,
gen_type=gen_type,
engine_id=engine_id,
source_module=source_module,
source_project_id=source_project_id,
source_step_id=source_step_id,
source_step_code=source_step_code,
billing_scene=billing_scene,
)
if gen_type == "image":
size = image_size or "2K"
amount = await calc_image_credits(db, size, engine_id=engine_id)
items.append(
await deduct_credits_locked_once(
db,
user_id=user_id,
amount=amount,
description=f"{description_prefix}图片生成",
related_id=record_id,
charge_key=CHARGE_MEDIA,
biz_key=biz_key,
attempt_no=attempt_no,
record_meta=record_meta,
)
)
elif gen_type == "video":
amount = await calc_video_credits(
db, duration or 5, resolution or "720p",
engine_id=engine_id,
input_video_duration=input_video_duration,
)
items.append(
await deduct_credits_locked_once(
db,
user_id=user_id,
amount=amount,
description=f"{description_prefix}视频生成",
related_id=record_id,
charge_key=CHARGE_MEDIA,
biz_key=biz_key,
attempt_no=attempt_no,
record_meta=record_meta,
)
)
else:
raise ValueError(f"不支持的生成类型: {gen_type}")
return BillingSummary(record_id=record_id, user_id=user_id, items=items)
async def charge_generation_media_for_record(
db: AsyncSession,
*,
record: GenerationRecord,
project_name: str | None = None,
description_prefix: str = "AI创作-",
attempt_no: int | None = None,
) -> BillingSummary:
return await charge_generation_media_by_params(
db,
user_id=record.user_id,
record_id=record.id,
gen_type=record.gen_type,
image_size=record.image_size,
duration=record.duration,
resolution=record.resolution,
project_name=project_name,
description_prefix=description_prefix,
owner_type=OWNER_GENERATION_RECORD,
attempt_no=attempt_no,
source_module=CreditRecordSourceModule.GENERATION_RECORD.value,
)
@@ -1,143 +0,0 @@
from __future__ import annotations
import os
import uuid
from dataclasses import dataclass
from datetime import datetime, timezone
from app.config import settings
from app.models.chat_generation_task import ChatGenerationTask
from app.services.image_gen import download_image
from app.services.provider_limit import provider_limit
from app.services.resource_accounting_service import safe_file_size
from app.services.video_cover_service import create_video_cover_for_local_video
from app.services.video_gen import download_video
@dataclass(slots=True)
class DownloadedGenerationResult:
url: str
storage_path: str | None
file_size_bytes: int
resource_type: str
storage_type: str = "local"
cover_url: str | None = None
cover_storage_path: str | None = None
def _to_aware_utc(value: datetime | None) -> datetime | None:
if value is None:
return None
if value.tzinfo is None:
return value.replace(tzinfo=timezone.utc)
return value.astimezone(timezone.utc)
def _build_storage_date_dir(record: ChatGenerationTask) -> str:
fixed = (getattr(record, "download_storage_date_dir", None) or "").strip().strip("/")
if fixed:
return fixed
created_at = _to_aware_utc(getattr(record, "created_at", None)) or datetime.now(timezone.utc)
return created_at.strftime("%Y/%m/%d")
def _make_part_path(final_path: str) -> str:
return f"{final_path}.{uuid.uuid4().hex}.part"
def _is_valid_file(path: str | None) -> bool:
if not path:
return False
try:
return os.path.isfile(path) and os.path.getsize(path) > 0
except OSError:
return False
def _safe_remove(path: str | None) -> None:
if not path:
return
try:
if os.path.exists(path):
os.remove(path)
except OSError:
pass
async def _download_image_atomically(remote_url: str, final_path: str) -> str:
if _is_valid_file(final_path):
return final_path
os.makedirs(os.path.dirname(final_path), exist_ok=True)
part_path = _make_part_path(final_path)
try:
await download_image(remote_url, part_path)
if not _is_valid_file(part_path):
raise RuntimeError("图片下载完成但临时文件为空")
os.replace(part_path, final_path)
return final_path
except Exception:
_safe_remove(part_path)
raise
async def _download_video_atomically(remote_url: str, final_path: str) -> str:
if _is_valid_file(final_path):
return final_path
os.makedirs(os.path.dirname(final_path), exist_ok=True)
part_path = _make_part_path(final_path)
try:
await download_video(remote_url, part_path)
if not _is_valid_file(part_path):
raise RuntimeError("视频下载完成但临时文件为空")
os.replace(part_path, final_path)
return final_path
except Exception:
_safe_remove(part_path)
raise
async def download_generation_result(record: ChatGenerationTask) -> DownloadedGenerationResult:
if not record.remote_result_url:
raise ValueError("缺少远程结果URL")
date_dir = _build_storage_date_dir(record)
if record.gen_type == "image":
dest_dir = os.path.join(settings.STORAGE_IMAGE_LOCAL_PATH, date_dir)
os.makedirs(dest_dir, exist_ok=True)
dest = os.path.join(dest_dir, f"{record.id}.png")
async with provider_limit("result_download", settings.RESULT_DOWNLOAD_MAX_CONCURRENCY):
await _download_image_atomically(record.remote_result_url if record.remote_result_url else "", dest)
return DownloadedGenerationResult(
url=f"/generate/images/{date_dir}/{record.id}.png",
storage_path=dest,
file_size_bytes=safe_file_size(dest),
resource_type="image",
)
dest_dir = os.path.join(settings.STORAGE_LOCAL_PATH, date_dir)
os.makedirs(dest_dir, exist_ok=True)
dest = os.path.join(dest_dir, f"{record.id}.mp4")
async with provider_limit("result_download", settings.RESULT_DOWNLOAD_MAX_CONCURRENCY):
await _download_video_atomically(record.remote_result_url if record.remote_result_url else "", dest)
cover_url, cover_storage_path = create_video_cover_for_local_video(
record_id=record.id,
video_path=dest,
date_dir=date_dir,
log_prefix=f"ChatGenerationTask视频封面生成 task_id={record.id}",
)
return DownloadedGenerationResult(
url=f"/generate/videos/{date_dir}/{record.id}.mp4",
storage_path=dest,
file_size_bytes=safe_file_size(dest),
resource_type="video",
cover_url=cover_url,
cover_storage_path=cover_storage_path,
)
@@ -1,514 +0,0 @@
from __future__ import annotations
import json
from datetime import datetime, timezone
from typing import Iterable, Sequence
from fastapi import HTTPException
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.enums.common import ModuleEventTypeEnum
from app.enums.generation_history import (
GenerationHistorySourceEnum,
get_generation_history_source_label,
normalize_generation_history_source,
MAX_BATCH_DELETE_COUNT,
)
from app.enums.generation_task import ChatGenerationTaskStatus, GenerationMode
from app.enums.shot_replicate import ShotSegmentAnalysisStatusEnum, ShotSegmentReplicateStatusEnum, ShotSplitStatusEnum
from app.models.chat_generation_task import ChatGenerationTask
from app.models.generation_record import GenerationRecord
from app.models.module_generation_project import ModuleGenerationProject
from app.models.module_generation_step import ModuleGenerationStep
from app.models.shot_replicate_segment import ShotReplicateSegment
from app.models.user import User
from app.schemas.generation_ai import GenerationAIHistoryBatchDeleteOut
from app.services.module_generation_flow_base_service import is_active_chat_generation_task
from app.services.module_generation_log_service import log_module_event_file
# from app.services.operation_log import log_operation
from app.services.resource_accounting_service import (
SOURCE_MODEL_CHAT_TASK,
SOURCE_MODEL_GENERATION_RECORD,
SOURCE_MODEL_SHOT_SEGMENT,
soft_delete_generation_record_resources,
soft_delete_resources_by_source,
)
COMPLETED_STATUS = ChatGenerationTaskStatus.COMPLETED.value
_ACTIVE_SPLIT_STATUSES = {
ShotSplitStatusEnum.PENDING.value,
ShotSplitStatusEnum.PROCESSING.value,
ShotSplitStatusEnum.RETRY_WAITING.value,
}
_ACTIVE_SEGMENT_ANALYSIS_STATUSES = {
ShotSegmentAnalysisStatusEnum.PENDING.value,
ShotSegmentAnalysisStatusEnum.PROCESSING.value,
}
_ACTIVE_SEGMENT_REPLICATE_STATUSES = {
ShotSegmentReplicateStatusEnum.PROCESSING.value,
}
def _normalize_ids(ids: Sequence[str] | Iterable[str]) -> list[str]:
normalized = [str(item).strip() for item in ids if str(item or "").strip()]
if not normalized:
raise HTTPException(status_code=400, detail="ids 不能为空")
if len(normalized) > MAX_BATCH_DELETE_COUNT:
raise HTTPException(status_code=400, detail=f"单次最多删除 {MAX_BATCH_DELETE_COUNT} 条记录")
if len(normalized) != len(set(normalized)):
raise HTTPException(status_code=400, detail="ids 不允许重复")
return normalized
def _missing_ids(request_ids: list[str], actual_ids: Iterable[str]) -> list[str]:
actual_set = {str(item) for item in actual_ids if item}
return [item for item in request_ids if item not in actual_set]
def _raise_missing_if_any(*, ids: list[str], found_ids: Iterable[str], message: str) -> None:
missing = _missing_ids(ids, found_ids)
if missing:
raise HTTPException(
status_code=404,
detail={
"message": message,
"missing_ids": missing,
"missing_count": len(missing),
},
)
def _raise_invalid_if_any(*, invalid_ids: list[str], message: str, status_code: int = 409) -> None:
if invalid_ids:
raise HTTPException(
status_code=status_code,
detail={
"message": message,
"invalid_ids": invalid_ids,
"invalid_count": len(invalid_ids),
},
)
def _task_id_list(steps: list[ModuleGenerationStep]) -> list[str]:
return list(dict.fromkeys(step.chat_task_id for step in steps if step.chat_task_id))
async def _load_chat_tasks_for_steps(
db: AsyncSession,
steps: list[ModuleGenerationStep],
) -> dict[str, ChatGenerationTask]:
task_ids = _task_id_list(steps)
if not task_ids:
return {}
result = await db.execute(
select(ChatGenerationTask)
.where(
ChatGenerationTask.id.in_(task_ids),
ChatGenerationTask.deleted_at.is_(None),
)
.with_for_update()
)
return {task.id: task for task in result.scalars().all()}
def _assert_no_active_chat_tasks(tasks: Iterable[ChatGenerationTask]) -> None:
active_task_ids = [task.id for task in tasks if is_active_chat_generation_task(task)]
if active_task_ids:
raise HTTPException(
status_code=409,
detail={
"message": "当前存在生成中任务,请等待生成完成或失败后再操作",
"active_count": len(active_task_ids),
},
)
# async def _log_history_batch_delete(
# db: AsyncSession,
# *,
# current_user: User,
# result: GenerationAIHistoryBatchDeleteOut,
# ) -> None:
# detail = {
# "history_source": result.history_source,
# "history_source_label": result.history_source_label,
# "requested_count": result.requested_count,
# "deleted_count": result.deleted_count,
# "requested_ids": result.requested_ids,
# "deleted_ids": result.deleted_ids,
# "generation_record_ids": result.generation_record_ids,
# "chat_task_ids": result.chat_task_ids,
# "module_project_ids": result.module_project_ids,
# "shot_segment_ids": result.shot_segment_ids,
# "freed_size_bytes": result.freed_size_bytes,
# }
# await log_operation(
# db,
# current_user.id,
# current_user.username,
# f"批量删除素材云历史-{result.history_source_label or result.history_source}",
# "DELETE",
# "/generation-ai/history/batch",
# detail=json.dumps(detail, ensure_ascii=False, default=str),
# )
def _build_out(
*,
source: GenerationHistorySourceEnum,
requested_ids: list[str],
deleted_ids: list[str],
generation_record_ids: list[str] | None = None,
chat_task_ids: list[str] | None = None,
module_project_ids: list[str] | None = None,
shot_segment_ids: list[str] | None = None,
freed_size_bytes: int = 0,
) -> GenerationAIHistoryBatchDeleteOut:
return GenerationAIHistoryBatchDeleteOut(
message="删除成功",
history_source=source.value,
history_source_label=get_generation_history_source_label(source),
requested_count=len(requested_ids),
deleted_count=len(deleted_ids),
requested_ids=requested_ids,
deleted_ids=deleted_ids,
generation_record_ids=generation_record_ids or [],
chat_task_ids=chat_task_ids or [],
module_project_ids=module_project_ids or [],
shot_segment_ids=shot_segment_ids or [],
deleted=True,
freed_size_bytes=int(freed_size_bytes or 0),
)
async def _delete_generation_records(
db: AsyncSession,
*,
current_user: User,
source: GenerationHistorySourceEnum,
ids: list[str],
deleted_at: datetime,
) -> GenerationAIHistoryBatchDeleteOut:
result = await db.execute(
select(GenerationRecord)
.where(
GenerationRecord.id.in_(ids),
GenerationRecord.user_id == current_user.id,
GenerationRecord.deleted_at.is_(None),
)
.with_for_update()
)
records = list(result.scalars().all())
_raise_missing_if_any(ids=ids, found_ids=[record.id for record in records], message="项目生成记录不存在或已删除")
invalid_ids = [
record.id
for record in records
if record.status != COMPLETED_STATUS or record.generated_at is None
]
_raise_invalid_if_any(invalid_ids=invalid_ids, message="项目生成记录只有生成完成后才能删除")
freed_size = await soft_delete_generation_record_resources(db, [record.id for record in records], deleted_at=deleted_at)
for record in records:
record.deleted_at = deleted_at
return _build_out(
source=source,
requested_ids=ids,
deleted_ids=ids,
generation_record_ids=ids,
freed_size_bytes=freed_size,
)
async def _delete_chat_tasks(
db: AsyncSession,
*,
current_user: User,
source: GenerationHistorySourceEnum,
ids: list[str],
deleted_at: datetime,
) -> GenerationAIHistoryBatchDeleteOut:
result = await db.execute(
select(ChatGenerationTask)
.where(
ChatGenerationTask.id.in_(ids),
ChatGenerationTask.user_id == current_user.id,
ChatGenerationTask.generation_mode == GenerationMode.CHATAPI_ASYNC.value,
ChatGenerationTask.deleted_at.is_(None),
)
.with_for_update()
)
tasks = list(result.scalars().all())
_raise_missing_if_any(ids=ids, found_ids=[task.id for task in tasks], message="AI 创作记录不存在或已删除")
invalid_ids = [
task.id
for task in tasks
if task.status != COMPLETED_STATUS or task.generated_at is None
]
_raise_invalid_if_any(invalid_ids=invalid_ids, message="AI 创作记录只有生成完成后才能删除")
freed_size = await soft_delete_resources_by_source(
db,
source_model=SOURCE_MODEL_CHAT_TASK,
source_ids=[task.id for task in tasks],
deleted_at=deleted_at,
)
for task in tasks:
task.deleted_at = deleted_at
return _build_out(
source=source,
requested_ids=ids,
deleted_ids=ids,
chat_task_ids=ids,
freed_size_bytes=freed_size,
)
async def _load_module_projects_by_ids(
db: AsyncSession,
*,
current_user: User,
source: GenerationHistorySourceEnum,
project_ids: list[str],
) -> list[ModuleGenerationProject]:
result = await db.execute(
select(ModuleGenerationProject)
.where(
ModuleGenerationProject.id.in_(project_ids),
ModuleGenerationProject.user_id == current_user.id,
ModuleGenerationProject.module == source.value,
ModuleGenerationProject.deleted_at.is_(None),
)
.with_for_update()
)
projects = list(result.scalars().all())
_raise_missing_if_any(ids=project_ids, found_ids=[project.id for project in projects], message="模块生成项目不存在或已删除")
return projects
async def _soft_delete_module_projects(
db: AsyncSession,
*,
source: GenerationHistorySourceEnum,
projects: list[ModuleGenerationProject],
deleted_at: datetime,
) -> tuple[list[str], list[str], int]:
project_ids = list(dict.fromkeys(project.id for project in projects))
if not project_ids:
return [], [], 0
step_result = await db.execute(
select(ModuleGenerationStep)
.where(
ModuleGenerationStep.project_id.in_(project_ids),
ModuleGenerationStep.module == source.value,
ModuleGenerationStep.deleted_at.is_(None),
ModuleGenerationStep.is_current == True,
)
.with_for_update()
)
steps = list(step_result.scalars().all())
task_map = await _load_chat_tasks_for_steps(db, steps)
_assert_no_active_chat_tasks(task_map.values())
chat_task_ids = list(task_map.keys())
freed_size = await soft_delete_resources_by_source(
db,
source_model=SOURCE_MODEL_CHAT_TASK,
source_ids=chat_task_ids,
deleted_at=deleted_at,
)
steps_by_project: dict[str, list[ModuleGenerationStep]] = {}
for step in steps:
steps_by_project.setdefault(step.project_id, []).append(step)
step.deleted_at = deleted_at
step.is_current = False
for task in task_map.values():
task.deleted_at = deleted_at
for project in projects:
project.deleted_at = deleted_at
project.final_image_url = None
project.final_video_url = None
project.final_video_cover_url = None
project.completed_at = None
project_steps = steps_by_project.get(project.id, [])
log_module_event_file(
module=source.value,
event_type=ModuleEventTypeEnum.PROJECT_DELETED.value,
project_id=project.id,
user_id=project.user_id,
message="素材云历史批量删除模块项目",
detail={
"source": "generation_history_batch_delete",
"step_ids": [step.id for step in project_steps],
"step_codes": [step.step_code for step in project_steps],
"chat_task_ids": [step.chat_task_id for step in project_steps if step.chat_task_id],
"refund_unfinished": False,
},
)
return project_ids, chat_task_ids, int(freed_size or 0)
async def _delete_hot_opening_projects(
db: AsyncSession,
*,
current_user: User,
source: GenerationHistorySourceEnum,
ids: list[str],
deleted_at: datetime,
) -> GenerationAIHistoryBatchDeleteOut:
projects = await _load_module_projects_by_ids(db, current_user=current_user, source=source, project_ids=ids)
module_project_ids, chat_task_ids, freed_size = await _soft_delete_module_projects(
db,
source=source,
projects=projects,
deleted_at=deleted_at,
)
return _build_out(
source=source,
requested_ids=ids,
deleted_ids=ids,
module_project_ids=module_project_ids,
chat_task_ids=chat_task_ids,
freed_size_bytes=freed_size,
)
def _assert_segments_not_active(segments: list[ShotReplicateSegment]) -> None:
invalid_ids = [
segment.id
for segment in segments
if segment.split_status in _ACTIVE_SPLIT_STATUSES
or segment.analysis_status in _ACTIVE_SEGMENT_ANALYSIS_STATUSES
or segment.replicate_status in _ACTIVE_SEGMENT_REPLICATE_STATUSES
]
_raise_invalid_if_any(invalid_ids=invalid_ids, message="拆镜片段仍在分割、分析或复刻处理中,暂不能删除")
async def _delete_shot_segments(
db: AsyncSession,
*,
current_user: User,
source: GenerationHistorySourceEnum,
ids: list[str],
deleted_at: datetime,
) -> GenerationAIHistoryBatchDeleteOut:
result = await db.execute(
select(ShotReplicateSegment)
.where(
ShotReplicateSegment.id.in_(ids),
ShotReplicateSegment.user_id == current_user.id,
ShotReplicateSegment.deleted_at.is_(None),
)
.with_for_update()
)
segments = list(result.scalars().all())
_raise_missing_if_any(ids=ids, found_ids=[segment.id for segment in segments], message="拆镜复刻片段不存在或已删除")
_assert_segments_not_active(segments)
missing_project_segment_ids = [segment.id for segment in segments if not segment.module_project_id]
_raise_invalid_if_any(invalid_ids=missing_project_segment_ids, message="拆镜复刻片段尚未关联复刻项目,不能按素材云历史删除")
module_project_ids = list(dict.fromkeys(segment.module_project_id for segment in segments if segment.module_project_id))
projects = await _load_module_projects_by_ids(
db,
current_user=current_user,
source=source,
project_ids=module_project_ids,
)
deleted_project_ids, chat_task_ids, project_freed_size = await _soft_delete_module_projects(
db,
source=source,
projects=projects,
deleted_at=deleted_at,
)
segment_freed_size = await soft_delete_resources_by_source(
db,
source_model=SOURCE_MODEL_SHOT_SEGMENT,
source_ids=[segment.id for segment in segments],
deleted_at=deleted_at,
)
for segment in segments:
segment.deleted_at = deleted_at
return _build_out(
source=source,
requested_ids=ids,
deleted_ids=ids,
module_project_ids=deleted_project_ids,
chat_task_ids=chat_task_ids,
shot_segment_ids=ids,
freed_size_bytes=int(project_freed_size or 0) + int(segment_freed_size or 0),
)
async def batch_delete_generation_history_items(
db: AsyncSession,
*,
current_user: User,
history_source: str,
ids: Sequence[str] | Iterable[str],
) -> GenerationAIHistoryBatchDeleteOut:
"""
按素材云 history_source 批量软删除历史记录。
事务由 get_db 统一提交/回滚;本服务只 flush,不主动 commit。
所有分支均为“先批量查询校验,再统一软删”,任何校验失败都会整体回滚。
"""
try:
source = normalize_generation_history_source(history_source)
except ValueError as exc:
raise HTTPException(status_code=400, detail="history_source 不支持") from exc
normalized_ids = _normalize_ids(ids)
deleted_at = datetime.now(timezone.utc)
if source == GenerationHistorySourceEnum.GENERATION_RECORD:
result = await _delete_generation_records(
db,
current_user=current_user,
source=source,
ids=normalized_ids,
deleted_at=deleted_at,
)
elif source == GenerationHistorySourceEnum.CHAT_TASK:
result = await _delete_chat_tasks(
db,
current_user=current_user,
source=source,
ids=normalized_ids,
deleted_at=deleted_at,
)
elif source == GenerationHistorySourceEnum.HOT_OPENING_REPLICATE:
result = await _delete_hot_opening_projects(
db,
current_user=current_user,
source=source,
ids=normalized_ids,
deleted_at=deleted_at,
)
elif source == GenerationHistorySourceEnum.SHOT_REPLICATE:
result = await _delete_shot_segments(
db,
current_user=current_user,
source=source,
ids=normalized_ids,
deleted_at=deleted_at,
)
else:
raise HTTPException(status_code=400, detail="history_source 不支持")
# await _log_history_batch_delete(db, current_user=current_user, result=result)
await db.flush()
return result
@@ -1,283 +0,0 @@
from __future__ import annotations
from typing import Any, TypedDict
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.enums.generation_history import (
GenerationHistorySourceEnum,
get_generation_history_source_label,
is_generation_history_module_source,
)
from app.models.module_generation_project import ModuleGenerationProject
from app.models.module_generation_step import ModuleGenerationStep
from app.models.shot_replicate_segment import ShotReplicateSegment
class GenerationHistoryMeta(TypedDict):
"""素材云历史项模块上下文回填字段。"""
history_source: str | None
history_source_label: str | None
module_project_id: str | None
module_project_title: str | None
module_step_id: str | None
module_step_code: str | None
hot_opening_project_id: str | None
hot_opening_project_title: str | None
shot_replicate_project_id: str | None
shot_replicate_project_title: str | None
shot_task_set_id: str | None
shot_segment_id: str | None
shot_segment_index: int | None
shot_segment_label: str | None
class _StepLinkInfo(TypedDict):
module_project_id: str | None
module_step_id: str | None
module_step_code: str | None
module: str | None
class _ProjectInfo(TypedDict):
module_project_id: str
module_project_title: str | None
module: str | None
class _ShotSegmentInfo(TypedDict):
shot_task_set_id: str | None
shot_segment_id: str | None
shot_segment_index: int | None
shot_segment_label: str | None
def _unique(values: list[str] | tuple[str, ...]) -> list[str]:
return list(dict.fromkeys(str(value) for value in values if value))
def _segment_label(segment_index: int | None) -> str | None:
if segment_index is None:
return None
return f"片段{segment_index}"
def build_empty_history_meta(source: GenerationHistorySourceEnum) -> GenerationHistoryMeta:
"""构造统一历史字段,非对应模块字段保持 null。"""
return {
"history_source": source.value,
"history_source_label": get_generation_history_source_label(source),
"module_project_id": None,
"module_project_title": None,
"module_step_id": None,
"module_step_code": None,
"hot_opening_project_id": None,
"hot_opening_project_title": None,
"shot_replicate_project_id": None,
"shot_replicate_project_title": None,
"shot_task_set_id": None,
"shot_segment_id": None,
"shot_segment_index": None,
"shot_segment_label": None,
}
async def _load_step_link_map(
db: AsyncSession,
*,
chat_task_ids: list[str],
source: GenerationHistorySourceEnum,
) -> dict[str, _StepLinkInfo]:
ids = _unique(chat_task_ids)
if not ids:
return {}
stmt = (
select(
ModuleGenerationStep.chat_task_id.label("chat_task_id"),
ModuleGenerationStep.id.label("module_step_id"),
ModuleGenerationStep.project_id.label("module_project_id"),
ModuleGenerationStep.step_code.label("module_step_code"),
ModuleGenerationStep.module.label("module"),
ModuleGenerationStep.is_current.label("is_current"),
ModuleGenerationStep.updated_at.label("updated_at"),
)
.where(
ModuleGenerationStep.deleted_at.is_(None),
ModuleGenerationStep.chat_task_id.in_(ids),
ModuleGenerationStep.module == source.value,
)
.order_by(
ModuleGenerationStep.chat_task_id.asc(),
ModuleGenerationStep.is_current.desc(),
ModuleGenerationStep.updated_at.desc(),
)
)
rows = (await db.execute(stmt)).mappings().all()
link_map: dict[str, _StepLinkInfo] = {}
for row in rows:
chat_task_id = row["chat_task_id"]
if not chat_task_id or chat_task_id in link_map:
continue
link_map[chat_task_id] = {
"module_project_id": row["module_project_id"],
"module_step_id": row["module_step_id"],
"module_step_code": row["module_step_code"],
"module": row["module"],
}
return link_map
async def _load_project_map(
db: AsyncSession,
*,
module_project_ids: list[str],
) -> dict[str, _ProjectInfo]:
ids = _unique(module_project_ids)
if not ids:
return {}
stmt = (
select(
ModuleGenerationProject.id.label("module_project_id"),
ModuleGenerationProject.title.label("module_project_title"),
ModuleGenerationProject.module.label("module"),
)
.where(
ModuleGenerationProject.deleted_at.is_(None),
ModuleGenerationProject.id.in_(ids),
)
)
return {
row["module_project_id"]: {
"module_project_id": row["module_project_id"],
"module_project_title": row["module_project_title"],
"module": row["module"],
}
for row in (await db.execute(stmt)).mappings().all()
if row["module_project_id"]
}
async def _load_shot_segment_map(
db: AsyncSession,
*,
module_project_ids: list[str],
) -> dict[str, _ShotSegmentInfo]:
ids = _unique(module_project_ids)
if not ids:
return {}
stmt = (
select(
ShotReplicateSegment.module_project_id.label("module_project_id"),
ShotReplicateSegment.id.label("shot_segment_id"),
ShotReplicateSegment.task_set_id.label("shot_task_set_id"),
ShotReplicateSegment.segment_index.label("shot_segment_index"),
ShotReplicateSegment.updated_at.label("updated_at"),
)
.where(
ShotReplicateSegment.deleted_at.is_(None),
ShotReplicateSegment.module_project_id.in_(ids),
)
.order_by(
ShotReplicateSegment.module_project_id.asc(),
ShotReplicateSegment.updated_at.desc(),
)
)
segment_map: dict[str, _ShotSegmentInfo] = {}
rows = (await db.execute(stmt)).mappings().all()
for row in rows:
module_project_id = row["module_project_id"]
if not module_project_id or module_project_id in segment_map:
continue
segment_index = row["shot_segment_index"]
segment_map[module_project_id] = {
"shot_task_set_id": row["shot_task_set_id"],
"shot_segment_id": row["shot_segment_id"],
"shot_segment_index": segment_index,
"shot_segment_label": _segment_label(segment_index),
}
return segment_map
def _merge_meta(
*,
source: GenerationHistorySourceEnum,
step_info: _StepLinkInfo | None,
project_info: _ProjectInfo | None,
shot_info: _ShotSegmentInfo | None,
) -> GenerationHistoryMeta:
meta = build_empty_history_meta(source)
module_project_id = step_info["module_project_id"] if step_info else None
module_step_id = step_info["module_step_id"] if step_info else None
module_step_code = step_info["module_step_code"] if step_info else None
module_project_title = project_info["module_project_title"] if project_info else None
meta["module_project_id"] = module_project_id
meta["module_project_title"] = module_project_title
meta["module_step_id"] = module_step_id
meta["module_step_code"] = module_step_code
if source == GenerationHistorySourceEnum.HOT_OPENING_REPLICATE:
meta["hot_opening_project_id"] = module_project_id
meta["hot_opening_project_title"] = module_project_title
elif source == GenerationHistorySourceEnum.SHOT_REPLICATE:
meta["shot_replicate_project_id"] = module_project_id
meta["shot_replicate_project_title"] = module_project_title
if shot_info:
meta["shot_task_set_id"] = shot_info["shot_task_set_id"]
meta["shot_segment_id"] = shot_info["shot_segment_id"]
meta["shot_segment_index"] = shot_info["shot_segment_index"]
meta["shot_segment_label"] = shot_info["shot_segment_label"]
return meta
async def batch_load_generation_history_meta_map(
db: AsyncSession,
*,
source: GenerationHistorySourceEnum,
chat_task_ids: list[str],
) -> dict[str, GenerationHistoryMeta]:
"""批量回填历史列表模块上下文,避免逐条链式查询。"""
ids = _unique(chat_task_ids)
if not ids:
return {}
if not is_generation_history_module_source(source):
return {chat_task_id: build_empty_history_meta(source) for chat_task_id in ids}
step_link_map = await _load_step_link_map(db, chat_task_ids=ids, source=source)
module_project_ids = [
step_info["module_project_id"]
for step_info in step_link_map.values()
if step_info.get("module_project_id")
]
project_map = await _load_project_map(db, module_project_ids=module_project_ids)
shot_segment_map: dict[str, _ShotSegmentInfo] = {}
if source == GenerationHistorySourceEnum.SHOT_REPLICATE:
shot_segment_map = await _load_shot_segment_map(db, module_project_ids=module_project_ids)
meta_map: dict[str, GenerationHistoryMeta] = {}
for chat_task_id in ids:
step_info = step_link_map.get(chat_task_id)
module_project_id = step_info.get("module_project_id") if step_info else None
project_info = project_map.get(module_project_id) if module_project_id else None
shot_info = shot_segment_map.get(module_project_id) if module_project_id else None
meta_map[chat_task_id] = _merge_meta(
source=source,
step_info=step_info,
project_info=project_info,
shot_info=shot_info,
)
return meta_map
@@ -1,134 +0,0 @@
from __future__ import annotations
import hashlib
import json
from typing import Any
from app.models.base import async_session
from app.models.chat_generation_task_event import ChatGenerationTaskEvent
from app.models.chat_provider_call_log import ChatProviderCallLog
from app.utils.id_gen import generate_id
MAX_EXCERPT_CHARS = 2000
def _safe_json(data: Any) -> str | None:
if data is None:
return None
try:
return json.dumps(data, ensure_ascii=False, default=str)
except Exception:
return str(data)
def _excerpt(data: Any, limit: int = MAX_EXCERPT_CHARS) -> str | None:
text = _safe_json(data)
if text is None:
return None
# Avoid storing secrets in logs.
text = text.replace("Authorization", "Authorization-REDACTED")
text = text.replace("api_key", "api_key_REDACTED")
if len(text) > limit:
return text[:limit] + "...[truncated]"
return text
def _hash(data: Any) -> str | None:
text = _safe_json(data)
if text is None:
return None
return hashlib.sha256(text.encode("utf-8")).hexdigest()
async def log_task_event(
task: Any | None = None,
*,
record: Any | None = None,
task_id: str | None = None,
record_id: str | None = None,
event_type: 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:
"""Write task event in a separate transaction; failure must not affect main flow."""
try:
obj = task or record
tid = task_id or record_id or (obj.id if obj else None)
if not tid:
return
async with async_session() as db:
db.add(ChatGenerationTaskEvent(
id=generate_id(),
task_id=tid,
generation_mode=getattr(obj, "generation_mode", "chatapi_async"),
event_type=event_type,
from_status=from_status,
to_status=to_status,
from_stage=from_stage,
to_stage=to_stage,
message=message,
detail_json=_excerpt(detail),
))
await db.commit()
except Exception:
return
async def log_provider_call(
task: Any | None = None,
*,
record: Any | None = None,
task_id: str | None = None,
record_id: str | None = None,
provider: str | None,
api_type: str,
model: str | None = None,
engine_id: str | None = None,
status: str,
latency_ms: int | None = None,
http_status: int | None = None,
provider_task_id: str | None = None,
request_data: Any = None,
response_data: Any = None,
prompt_tokens: int = 0,
completion_tokens: int = 0,
total_tokens: int = 0,
error_code: str | None = None,
error_message: str | None = None,
) -> None:
"""Write provider call log in a separate transaction; failure must not affect main flow."""
try:
obj = task or record
tid = task_id or record_id or (obj.id if obj else None)
if not tid:
return
async with async_session() as db:
db.add(ChatProviderCallLog(
id=generate_id(),
task_id=tid,
generation_mode=getattr(obj, "generation_mode", "chatapi_async"),
provider=provider,
api_type=api_type,
model=model,
engine_id=engine_id,
status=status,
latency_ms=latency_ms,
http_status=http_status,
provider_task_id=provider_task_id,
request_hash=_hash(request_data),
response_hash=_hash(response_data),
request_excerpt=_excerpt(request_data),
response_excerpt=_excerpt(response_data),
prompt_tokens=prompt_tokens or 0,
completion_tokens=completion_tokens or 0,
total_tokens=total_tokens or 0,
error_code=error_code,
error_message=error_message,
))
await db.commit()
except Exception:
return
@@ -1,41 +0,0 @@
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 已进入终态。
该方法必须幂等:下载恢复任务、重试任务、服务重启补偿都可能重复调用。
具体模块服务需要自行判断 step/project 是否已经完成或失败,避免重复推进。
"""
if not task:
return
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 status == ChatGenerationTaskStatus.COMPLETED.value:
await handle_chat_generation_task_completed(db, task)
elif status == ChatGenerationTaskStatus.FAILED.value:
await handle_chat_generation_task_failed(db, task)
return
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 status == ChatGenerationTaskStatus.COMPLETED.value:
await handle_chat_generation_task_completed(db, task)
elif status == ChatGenerationTaskStatus.FAILED.value:
await handle_chat_generation_task_failed(db, task)
return
@@ -1,153 +0,0 @@
from __future__ import annotations
from dataclasses import dataclass
from datetime import datetime, timedelta, timezone
from app.config import settings
from app.enums.generation_task import GenerationType
from app.models.chat_generation_task import ChatGenerationTask
from app.services.redis_registry_service import ensure_aware_utc
@dataclass(slots=True)
class PollScheduleDecision:
delay_seconds: int
next_poll_at: datetime
poll_interval_seconds: int
direct_countdown: bool
final_poll: bool
reason: str
def utc_now() -> datetime:
return datetime.now(timezone.utc)
def is_video_generation_task(task: ChatGenerationTask) -> bool:
return str(getattr(task, "gen_type", "") or "").lower() == GenerationType.VIDEO.value
def video_final_deadline_from(now: datetime | None = None) -> datetime:
current_time = now or utc_now()
hours = max(1, int(settings.CHATAPI_ASYNC_VIDEO_FINAL_DEADLINE_HOURS or 24))
return current_time + timedelta(hours=hours)
def ensure_video_poll_fields(task: ChatGenerationTask, *, now: datetime | None = None) -> None:
"""补齐视频轮询调度字段,兼容历史任务。"""
if not is_video_generation_task(task):
return
current_time = now or utc_now()
if ensure_aware_utc(getattr(task, "poll_started_at", None)) is None:
task.poll_started_at = current_time
if ensure_aware_utc(getattr(task, "deadline_at", None)) is None:
task.deadline_at = video_final_deadline_from(current_time)
if getattr(task, "poll_interval_seconds", None) is None:
task.poll_interval_seconds = 0
def is_final_poll_due(task: ChatGenerationTask, *, now: datetime | None = None) -> bool:
deadline_at = ensure_aware_utc(getattr(task, "deadline_at", None))
return bool(deadline_at and deadline_at <= (now or utc_now()))
def is_poll_not_due(task: ChatGenerationTask, *, now: datetime | None = None, tolerance_seconds: int = 1) -> bool:
"""判断当前 poll 任务是否早于 next_poll_at。只对视频降频轮询生效。"""
if not is_video_generation_task(task):
return False
if is_final_poll_due(task, now=now):
return False
next_poll_at = ensure_aware_utc(getattr(task, "next_poll_at", None))
if next_poll_at is None:
return False
return next_poll_at > (now or utc_now()) + timedelta(seconds=max(0, int(tolerance_seconds)))
def _clamp_positive_seconds(value: int | float | None, default: int) -> int:
try:
parsed = int(value if value is not None else default)
except (TypeError, ValueError):
parsed = default
return max(1, parsed)
def build_video_pending_poll_schedule(
task: ChatGenerationTask,
*,
now: datetime | None = None,
) -> PollScheduleDecision:
"""计算视频任务 pending/running 后的下一次轮询时间。"""
current_time = now or utc_now()
ensure_video_poll_fields(task, now=current_time)
deadline_at = ensure_aware_utc(task.deadline_at)
if deadline_at and deadline_at <= current_time:
return PollScheduleDecision(
delay_seconds=0,
next_poll_at=current_time,
poll_interval_seconds=int(task.poll_interval_seconds or 0),
direct_countdown=True,
final_poll=True,
reason="video_final_poll_due",
)
poll_started_at = ensure_aware_utc(task.poll_started_at) or current_time
elapsed_seconds = max(0, int((current_time - poll_started_at).total_seconds()))
high_freq_seconds = max(0, int(settings.CHATAPI_ASYNC_VIDEO_HIGH_FREQ_MINUTES or 10)) * 60
high_freq_poll_seconds = _clamp_positive_seconds(settings.CHATAPI_ASYNC_VIDEO_HIGH_FREQ_POLL_SECONDS, 30)
initial_backoff_seconds = _clamp_positive_seconds(settings.CHATAPI_ASYNC_VIDEO_BACKOFF_INITIAL_SECONDS, 60)
multiplier = max(1, int(settings.CHATAPI_ASYNC_VIDEO_BACKOFF_MULTIPLIER or 2))
max_backoff_seconds = _clamp_positive_seconds(settings.CHATAPI_ASYNC_VIDEO_BACKOFF_MAX_SECONDS, 3600)
if elapsed_seconds < high_freq_seconds:
delay_seconds = high_freq_poll_seconds
interval_seconds = int(task.poll_interval_seconds or 0)
reason = "video_high_freq_poll"
else:
previous_interval = int(task.poll_interval_seconds or 0)
if previous_interval < initial_backoff_seconds:
interval_seconds = initial_backoff_seconds
else:
interval_seconds = min(previous_interval * multiplier, max_backoff_seconds)
delay_seconds = interval_seconds
reason = "video_backoff_poll"
next_poll_at = current_time + timedelta(seconds=delay_seconds)
final_poll = False
if deadline_at and next_poll_at >= deadline_at:
next_poll_at = deadline_at
delay_seconds = max(0, int((deadline_at - current_time).total_seconds()))
final_poll = delay_seconds <= 0
reason = "video_schedule_to_final_deadline"
direct_max = max(0, int(settings.CHATAPI_ASYNC_VIDEO_DIRECT_COUNTDOWN_MAX_SECONDS or 300))
return PollScheduleDecision(
delay_seconds=delay_seconds,
next_poll_at=next_poll_at,
poll_interval_seconds=interval_seconds,
direct_countdown=delay_seconds <= direct_max,
final_poll=final_poll,
reason=reason,
)
def build_default_poll_schedule(
task: ChatGenerationTask,
*,
now: datetime | None = None,
delay_seconds: int | None = None,
reason: str = "default_poll",
) -> PollScheduleDecision:
"""图片和旧逻辑兼容用的固定间隔轮询计划。"""
current_time = now or utc_now()
delay = _clamp_positive_seconds(delay_seconds, int(settings.CHATAPI_ASYNC_POLL_INTERVAL_SECONDS or 30))
next_poll_at = current_time + timedelta(seconds=delay)
return PollScheduleDecision(
delay_seconds=delay,
next_poll_at=next_poll_at,
poll_interval_seconds=int(getattr(task, "poll_interval_seconds", 0) or 0),
direct_countdown=True,
final_poll=False,
reason=reason,
)
@@ -1,192 +0,0 @@
from __future__ import annotations
import json
import mimetypes
import os
import time
from typing import Any
import httpx
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.config import settings
from app.models.chat_generation_task import ChatGenerationTask
from app.models.model_config import ModelConfig
from app.models.token_usage import TokenUsage
from app.services.generation_log_service import log_provider_call
from app.services.provider_limit import provider_limit
from app.utils.id_gen import generate_id
def _absolute_url(url: str) -> str:
if url.startswith("http://") or url.startswith("https://") or url.startswith("data:"):
return url
base = settings.BASE_URL.rstrip("/")
return f"{base}/{url.lstrip('/')}"
def _load_refs(record: ChatGenerationTask) -> list[dict]:
if not record.media_references:
return []
try:
data = json.loads(record.media_references)
return data if isinstance(data, list) else []
except Exception:
return []
def _build_user_content(record: ChatGenerationTask) -> list[dict[str, Any]]:
if record.gen_type == "image":
params = f"图片参数:分辨率档位={record.image_size or '2K'},比例={record.image_proportion or '1:1'},像素={record.image_px or '2048x2048'}"
else:
params = f"视频参数:时长={record.duration or 4}秒,比例={record.aspect_ratio or '16:9'},分辨率={record.resolution or '480p'}"
text = (
f"生成类型:{record.gen_type}\n"
f"{params}\n"
f"用户描述:{record.original_prompt}\n\n"
"请只输出最终可直接用于图片/视频生成模型的 prompt,不要说你已经生成了图片或视频。"
)
parts: list[dict[str, Any]] = [{"type": "text", "text": text}]
for ref in _load_refs(record):
ref_type = ref.get("type")
ref_url = ref.get("url") or ""
if not ref_url:
continue
url = _absolute_url(ref_url)
if ref_type == "image":
parts.append({"type": "image_url", "image_url": {"url": url}})
elif ref_type == "video":
parts.append({"type": "video_url", "video_url": {"url": url, "fps": settings.CHATAPI_VIDEO_FPS}})
return parts
async def _get_model_config(db: AsyncSession) -> ModelConfig:
result = await db.execute(
select(ModelConfig)
.where(ModelConfig.is_active == True)
.order_by(ModelConfig.priority.desc())
.limit(1)
)
config = result.scalar_one_or_none()
if not config:
raise ValueError("没有可用的ChatAPI模型配置")
if config.provider == "mock":
return config
if not config.api_base or not config.api_key or not config.model_name:
raise ValueError("ChatAPI模型配置不完整")
return config
async def build_prompt_with_chatapi(db: AsyncSession, record: ChatGenerationTask) -> tuple[str, dict]:
"""Call ChatAPI once with current request params and attachments. No history context."""
config = await _get_model_config(db)
if config.provider == "mock":
return record.original_prompt, {"input_tokens": 0, "output_tokens": 0, "total_tokens": 0}
system_prompt = (
"你是图片/视频生成提示词整理助手。你的职责是根据用户文字、上传图片/视频和生成参数,"
"整理最终可直接用于生成模型的 prompt。不要声称你已经生成图片或视频,不要调用工具。"
"输出中文为主,内容具体、可执行,保留用户关键要求。"
)
request_data = {
"model": config.model_name,
"messages": [
{"role": "system", "content": system_prompt},
{"role": "user", "content": _build_user_content(record)},
],
"max_tokens": config.max_tokens,
"temperature": config.temperature,
}
started = time.perf_counter()
async with provider_limit("ark_chat_prompt", settings.ARK_CHAT_PROMPT_MAX_CONCURRENCY):
async with httpx.AsyncClient(timeout=settings.CHATAPI_REQUEST_TIMEOUT_SECONDS) as client:
try:
response = await client.post(
f"{config.api_base.rstrip('/')}/chat/completions",
headers={
"Authorization": f"Bearer {config.api_key}",
"Content-Type": "application/json",
},
json=request_data,
)
latency_ms = int((time.perf_counter() - started) * 1000)
if response.status_code >= 400:
await log_provider_call(
record,
provider=config.provider,
api_type="chat_prompt",
model=config.model_name,
engine_id=record.engine_id,
status="failed",
latency_ms=latency_ms,
http_status=response.status_code,
request_data=request_data,
response_data=response.text,
error_message=response.text[:1000],
)
raise RuntimeError(f"ChatAPI HTTP {response.status_code}: {response.text}")
data = response.json()
except Exception as exc:
latency_ms = int((time.perf_counter() - started) * 1000)
await log_provider_call(
record,
provider=config.provider,
api_type="chat_prompt",
model=config.model_name,
engine_id=record.engine_id,
status="failed",
latency_ms=latency_ms,
request_data=request_data,
response_data=None,
error_message=str(exc),
)
raise
usage = data.get("usage", {}) or {}
input_tokens = int(usage.get("prompt_tokens", 0) or 0)
output_tokens = int(usage.get("completion_tokens", 0) or 0)
total_tokens = int(usage.get("total_tokens", input_tokens + output_tokens) or 0)
content = data.get("choices", [{}])[0].get("message", {}).get("content", "").strip()
if not content:
raise RuntimeError("ChatAPI未返回有效prompt")
token_usage_id = generate_id()
db.add(TokenUsage(
id=token_usage_id,
model_config_id=config.id,
user_id=record.user_id,
owner_type="generation_record",
owner_id=record.id,
input_tokens=input_tokens,
output_tokens=output_tokens,
total_tokens=total_tokens,
))
await db.flush()
await log_provider_call(
record,
provider=config.provider,
api_type="chat_prompt",
model=config.model_name,
engine_id=record.engine_id,
status="success",
latency_ms=int((time.perf_counter() - started) * 1000),
http_status=200,
request_data=request_data,
response_data=data,
prompt_tokens=input_tokens,
completion_tokens=output_tokens,
total_tokens=total_tokens,
)
return content, {
"token_usage_id": token_usage_id,
"model_config_id": config.id,
"model_config_name": config.name,
"model_provider": config.provider,
"model_name": config.model_name,
"input_tokens": input_tokens,
"output_tokens": output_tokens,
"total_tokens": total_tokens,
}
@@ -1,172 +0,0 @@
from __future__ import annotations
import asyncio
import json
import time
from types import SimpleNamespace
from typing import Any
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.config import settings
from app.models.chat_generation_task import ChatGenerationTask
from app.models.image_engine import ImageEngine
from app.models.video_engine import VideoEngine
from app.services.generation_log_service import log_provider_call
from app.services.image_gen import poll_image_task_status, submit_image_task
from app.services.provider_limit import provider_limit
from app.services.video_gen import poll_task_status, submit_video_task
def _loads(data: str | None) -> dict:
if not data:
return {}
try:
obj = json.loads(data)
return obj if isinstance(obj, dict) else {}
except Exception:
return {}
async def get_runtime_engine(db: AsyncSession, task: ChatGenerationTask) -> Any:
"""Use frozen snapshot for historical params, current DB row only for secret api_key."""
snapshot = _loads(task.engine_snapshot_json)
if not task.engine_id:
raise ValueError("缺少 engine_id")
if task.gen_type == "image":
result = await db.execute(select(ImageEngine).where(ImageEngine.id == task.engine_id).limit(1))
else:
result = await db.execute(select(VideoEngine).where(VideoEngine.id == task.engine_id).limit(1))
engine = result.scalar_one_or_none()
if not engine:
raise ValueError("引擎不存在或已删除")
return SimpleNamespace(
id=task.engine_id,
name=snapshot.get("name") or engine.name,
provider=snapshot.get("provider") or engine.provider,
api_base=snapshot.get("api_base") or engine.api_base,
api_key=engine.api_key,
model_name=snapshot.get("model_name") or engine.model_name,
generate_url=snapshot.get("generate_url") or getattr(engine, "generate_url", ""),
query_url=snapshot.get("query_url") or getattr(engine, "query_url", ""),
default_size=snapshot.get("default_size") or getattr(engine, "default_size", "2K"),
)
async def create_provider_task(db: AsyncSession, task: ChatGenerationTask) -> dict:
if task.gen_type == "video":
return await _create_video_task(db, task)
if task.gen_type == "image":
return await _create_image_sync_task(db, task)
raise ValueError(f"不支持的生成类型: {task.gen_type}")
async def _create_video_task(db: AsyncSession, task: ChatGenerationTask) -> dict:
"""Create video provider task through the original Ark SDK async task API."""
engine = await get_runtime_engine(db, task)
started = time.perf_counter()
async with provider_limit("ark_video_create", settings.ARK_VIDEO_CREATE_MAX_CONCURRENCY):
try:
provider_task_id = await submit_video_task(
db,
engine,
task,
include_media_references=True,
)
response = {"task_id": provider_task_id}
await log_provider_call(
task,
provider=engine.provider,
api_type="video_create",
model=engine.model_name,
engine_id=task.engine_id,
status="success",
latency_ms=int((time.perf_counter() - started) * 1000),
provider_task_id=provider_task_id,
response_data=response,
)
return {"task_id": provider_task_id, "response_data": response}
except Exception as exc:
await log_provider_call(
task,
provider=engine.provider,
api_type="video_create",
model=engine.model_name,
engine_id=task.engine_id,
status="failed",
latency_ms=int((time.perf_counter() - started) * 1000),
error_message=str(exc),
)
raise
async def _create_image_sync_task(db: AsyncSession, task: ChatGenerationTask) -> dict:
"""Run the original synchronous image generation SDK under Celery control.
The legacy image SDK returns a final remote image URL immediately. We do
NOT use image_generation.tasks.create here, so image generation stays aligned
with the old working flow while no longer blocking the FastAPI request.
"""
engine = await get_runtime_engine(db, task)
started = time.perf_counter()
async with provider_limit("ark_image_sync_create", settings.ARK_IMAGE_CREATE_MAX_CONCURRENCY):
try:
result = await asyncio.to_thread(
submit_image_task,
db,
engine,
task,
include_media_references=True,
)
if result.get("error"):
raise RuntimeError(result.get("error"))
response_data = _try_json(result.get("response_data")) or result
await log_provider_call(
task,
provider=engine.provider,
api_type="image_sync_create",
model=engine.model_name,
engine_id=task.engine_id,
status="success",
latency_ms=int((time.perf_counter() - started) * 1000),
provider_task_id=None,
response_data=response_data,
)
return {
"task_id": None,
"remote_result_url": result.get("image_url"),
"image_tokens": result.get("image_tokens", 0) or 0,
"response_data": response_data,
}
except Exception as exc:
await log_provider_call(
task,
provider=engine.provider,
api_type="image_sync_create",
model=engine.model_name,
engine_id=task.engine_id,
status="failed",
latency_ms=int((time.perf_counter() - started) * 1000),
error_message=str(exc),
)
raise
def _try_json(text: Any) -> Any:
if not isinstance(text, str):
return text
try:
return json.loads(text)
except Exception:
return None
async def poll_provider_task(db: AsyncSession, task: ChatGenerationTask) -> dict:
engine = await get_runtime_engine(db, task)
task_id = task.seedance_task_id or task.provider_task_id
if task.gen_type == "video":
async with provider_limit("ark_video_poll", settings.ARK_VIDEO_POLL_MAX_CONCURRENCY):
return await poll_task_status(engine, task_id)
async with provider_limit("ark_image_poll", settings.ARK_IMAGE_POLL_MAX_CONCURRENCY):
return await poll_image_task_status(engine, task_id)
@@ -1,42 +0,0 @@
from __future__ import annotations
from typing import Protocol
class ProviderGenerationRecordLike(Protocol):
"""图片/视频供应商提交接口需要的任务字段协议。
GenerationRecord 与 ChatGenerationTask 都具备这些字段,但二者不是同一个 ORM 模型。
使用 Protocol 可以避免把 submit_image_task / submit_video_task 错误限制为某一个具体模型。
"""
id: str
original_prompt: str
optimized_prompt: str | None
media_references: str | None
gen_type: str
duration: int | None
aspect_ratio: str | None
resolution: str | None
image_size: str | None
image_proportion: str | None
image_px: str | None
class ProviderImageEngineLike(Protocol):
"""图片生成提交接口需要的引擎字段协议。"""
name: str
api_base: str
api_key: str
model_name: str
default_size: str | None
class ProviderVideoEngineLike(Protocol):
"""视频生成提交接口需要的引擎字段协议。"""
name: str
api_base: str
api_key: str
model_name: str
@@ -1,820 +0,0 @@
from __future__ import annotations
import json
import logging
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.celery_queue import CeleryQueue
from app.enums.generation_task import (
ALLOWED_GENERATION_MODES,
ChatGenerationPipelineStage,
ChatGenerationTaskEventType,
ChatGenerationTaskStatus,
GenerationType,
)
from app.models.chat_generation_task import ChatGenerationTask
from app.services.celery_download_recovery_service import (
ensure_aware_utc,
get_download_active_payloads,
get_due_download_record_ids,
postpone_download_active_check,
remove_download_active,
)
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_poll_schedule_service import ensure_video_poll_fields, is_poll_not_due, is_video_generation_task
from app.services.generation_refund_service import mark_chat_generation_task_failed_and_refund_once
from app.services.redis_registry_service import (
redis_get_due_registry_ids,
redis_get_registry_payloads,
redis_postpone_registry_item,
redis_remove_registry_item,
)
logger = logging.getLogger("video_gen")
POLL_QUEUE = CeleryQueue.GEN_PROVIDER_POLL.value
def _now() -> datetime:
return datetime.now(timezone.utc)
def _is_expired(value: datetime | None, now: datetime | None = None) -> bool:
checked = ensure_aware_utc(value)
if checked is None:
return True
return checked <= (now or _now())
def _queue_timeout_at(task: ChatGenerationTask, now: datetime | None = None) -> datetime:
current_time = now or _now()
enqueued_at = ensure_aware_utc(task.download_enqueued_at)
if enqueued_at is None:
return current_time
return enqueued_at + timedelta(seconds=int(settings.DOWNLOAD_TASK_QUEUE_TIMEOUT_SECONDS or 300))
def _is_queue_timeout(task: ChatGenerationTask, now: datetime | None = None) -> bool:
current_time = now or _now()
return _queue_timeout_at(task, current_time) <= current_time
def _is_final_task_state(task: ChatGenerationTask) -> bool:
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,
)
def _is_success(status: str | None) -> bool:
return str(status or "").lower() in ("succeeded", "success", "completed", "done")
def _is_failed(status: str | None) -> bool:
return str(status or "").lower() in ("failed", "error", "canceled", "cancelled")
def _engine_snapshot(task: ChatGenerationTask) -> dict[str, Any]:
try:
value = json.loads(task.engine_snapshot_json or "{}")
return value if isinstance(value, dict) else {}
except Exception:
return {}
def _poll_queue_timeout_at(now: datetime | None = None) -> datetime:
current_time = now or _now()
return current_time + timedelta(seconds=int(settings.POLL_TASK_QUEUE_TIMEOUT_SECONDS or 120))
async def _remove_poll_active(task_id: str) -> None:
await redis_remove_registry_item(
hash_key=settings.POLL_ACTIVE_REDIS_HASH_KEY,
zset_key=settings.POLL_ACTIVE_REDIS_ZSET_KEY,
item_id=task_id,
log_context="poll_active",
)
async def _postpone_poll_active(
*,
task_id: str,
payload: dict[str, Any] | None = None,
check_at: datetime | int | float | None = None,
) -> None:
await redis_postpone_registry_item(
hash_key=settings.POLL_ACTIVE_REDIS_HASH_KEY,
zset_key=settings.POLL_ACTIVE_REDIS_ZSET_KEY,
item_id=task_id,
payload=payload,
check_at=check_at or _poll_queue_timeout_at(),
log_context="poll_active",
)
async def recover_one_download_task(
db: AsyncSession,
task: ChatGenerationTask,
*,
payload: dict[str, Any] | None = None,
source: str = "startup_db",
) -> str:
from app.tasks.generation_download_tasks import (
DOWNLOAD_STAGE_DOWNLOADING,
DOWNLOAD_STAGE_QUEUED,
DOWNLOAD_STAGE_RETRY_WAITING,
enqueue_download_task,
)
current_time = _now()
if not task:
return "skip_missing_task"
if task.generation_mode not in ALLOWED_GENERATION_MODES:
await remove_download_active(task.id)
return "clean_invalid_mode"
if _is_final_task_state(task):
await remove_download_active(task.id)
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},
)
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 == ChatGenerationPipelineStage.RESULT_READY.value:
await log_task_event(
task,
event_type=ChatGenerationTaskEventType.DOWNLOAD_RECOVERY_ENQUEUE.value,
message=f"{source} 发现 result_ready 未完成下载,启动时恢复投递下载任务",
detail={"payload": redis_payload},
)
await enqueue_download_task(
db,
task,
recover=True,
reason=f"{source}_result_ready",
)
return "recover_result_ready"
if stage == DOWNLOAD_STAGE_QUEUED:
if _is_queue_timeout(task, current_time):
await log_task_event(
task,
event_type=ChatGenerationTaskEventType.DOWNLOAD_RECOVERY_ENQUEUE.value,
message=f"{source} 发现 download_queued 长时间未消费,启动时恢复投递下载任务",
detail={"payload": redis_payload},
)
await enqueue_download_task(
db,
task,
recover=True,
reason=f"{source}_download_queued_timeout",
)
return "recover_queued_timeout"
await postpone_download_active_check(
record_id=task.id,
payload=payload,
check_at=_queue_timeout_at(task, current_time),
)
return "skip_queued_not_timeout"
if stage == DOWNLOAD_STAGE_DOWNLOADING:
if _is_expired(task.download_lease_until, current_time):
await log_task_event(
task,
event_type=ChatGenerationTaskEventType.DOWNLOAD_RECOVERY_ENQUEUE.value,
message=f"{source} 发现 downloading lease 过期,启动时恢复投递下载任务",
detail={"payload": redis_payload},
)
await enqueue_download_task(
db,
task,
recover=True,
reason=f"{source}_downloading_lease_expired",
)
return "recover_downloading_expired"
await postpone_download_active_check(
record_id=task.id,
payload=payload,
check_at=task.download_lease_until,
)
return "skip_downloading_alive"
if stage == DOWNLOAD_STAGE_RETRY_WAITING:
if _is_expired(task.download_next_retry_at, current_time):
await log_task_event(
task,
event_type=ChatGenerationTaskEventType.DOWNLOAD_RECOVERY_ENQUEUE.value,
message=f"{source} 发现 retry_waiting 到期,启动时恢复投递下载任务",
detail={"payload": redis_payload},
)
await enqueue_download_task(
db,
task,
recover=True,
reason=f"{source}_retry_waiting_due",
)
return "recover_retry_due"
await postpone_download_active_check(
record_id=task.id,
payload=payload,
check_at=task.download_next_retry_at,
)
return "skip_retry_waiting_not_due"
return f"skip_stage_{stage}"
async def recover_download_tasks_once(db: AsyncSession) -> dict[str, Any]:
"""启动时下载容灾扫描。
先按 Redis active_index 找到到期下载任务;Redis 不可用或索引丢失时,
再通过 DB fallback 扫描 result_ready/download_* 状态,避免任务永久卡住。
"""
checked_ids: set[str] = set()
results: dict[str, int] = {}
due_ids = await get_due_download_record_ids(
limit=settings.DOWNLOAD_RECOVERY_BATCH_SIZE,
)
payloads = await get_download_active_payloads(due_ids)
for task_id in due_ids:
result = await db.execute(
select(ChatGenerationTask)
.where(
ChatGenerationTask.id == task_id,
ChatGenerationTask.deleted_at.is_(None),
)
.with_for_update()
.limit(1)
)
task = result.scalar_one_or_none()
if task is None:
await remove_download_active(task_id)
action = "clean_missing_task"
else:
checked_ids.add(task.id)
action = await recover_one_download_task(
db,
task,
payload=payloads.get(task_id),
source="startup_redis",
)
results[action] = results.get(action, 0) + 1
# DB fallback:不依赖 Redis active 注册表。
fallback_result = await db.execute(
select(ChatGenerationTask)
.where(
ChatGenerationTask.deleted_at.is_(None),
ChatGenerationTask.generation_mode.in_(list(ALLOWED_GENERATION_MODES)),
ChatGenerationTask.status == ChatGenerationTaskStatus.GENERATING.value,
ChatGenerationTask.remote_result_url.is_not(None),
ChatGenerationTask.pipeline_stage.in_(
[
ChatGenerationPipelineStage.RESULT_READY.value,
ChatGenerationPipelineStage.DOWNLOAD_QUEUED.value,
ChatGenerationPipelineStage.DOWNLOADING.value,
ChatGenerationPipelineStage.RETRY_WAITING.value,
]
),
)
.order_by(ChatGenerationTask.updated_at.asc())
.limit(int(settings.DOWNLOAD_RECOVERY_BATCH_SIZE or 100))
.with_for_update(skip_locked=True)
)
fallback_tasks = fallback_result.scalars().all()
for task in fallback_tasks:
if task.id in checked_ids:
continue
action = await recover_one_download_task(
db,
task,
payload=None,
source="startup_db",
)
results[action] = results.get(action, 0) + 1
checked_ids.add(task.id)
return {"checked": len(checked_ids), "results": results}
async def _mark_timeout(
db: AsyncSession,
task: ChatGenerationTask,
*,
error_message: str = "任务超时",
) -> str:
await mark_chat_generation_task_failed_and_refund_once(
db,
task=task,
error_message=error_message,
pipeline_stage=ChatGenerationPipelineStage.TIMEOUT.value,
)
await notify_chat_generation_task_finished(db, task)
await db.commit()
await _remove_poll_active(task.id)
await log_task_event(
task,
event_type=ChatGenerationTaskEventType.TASK_TIMEOUT.value,
to_status="failed",
to_stage=ChatGenerationPipelineStage.TIMEOUT.value,
)
return "mark_timeout"
async def _mark_failed(
db: AsyncSession,
task: ChatGenerationTask,
*,
error_message: str,
event_type: str = "POLL_FAILED",
detail: Any = None,
) -> str:
await mark_chat_generation_task_failed_and_refund_once(
db,
task=task,
error_message=error_message,
pipeline_stage=ChatGenerationPipelineStage.FAILED.value,
)
await notify_chat_generation_task_finished(db, task)
await db.commit()
await _remove_poll_active(task.id)
await log_task_event(task, event_type=event_type, message=task.error_message, detail=detail)
return "mark_failed"
async def recover_one_generation_task(
db: AsyncSession,
task: ChatGenerationTask,
*,
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
current_time = _now()
redis_payload = payload or {}
if not task:
return "skip_missing_task"
if task.generation_mode not in ALLOWED_GENERATION_MODES:
await _remove_poll_active(task.id)
return "clean_invalid_mode"
if _is_final_task_state(task):
await _remove_poll_active(task.id)
return "clean_final_state"
if task.status != ChatGenerationTaskStatus.GENERATING.value:
await _remove_poll_active(task.id)
return "clean_not_generating"
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))
# 最高优先级:只要远程结果 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=ChatGenerationTaskEventType.GENERATION_RECOVERY_ENQUEUE.value,
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=ChatGenerationTaskEventType.GENERATION_RECOVERY_ENQUEUE.value,
message=f"{source} 发现任务已到 deadline 且存在供应商任务ID,投递 poll 队列做最终查询",
detail={"pipeline_stage": task.pipeline_stage, "payload": redis_payload},
)
poll_generation_task.apply_async(
args=[task.id],
kwargs={"force_due": True},
queue=POLL_QUEUE,
countdown=0,
)
await register_poll_active(
task,
check_at=_poll_queue_timeout_at(),
reason=f"{source}_deadline_final_poll",
)
return "recover_deadline_final_poll"
await log_task_event(
task,
event_type=ChatGenerationTaskEventType.GENERATION_RECOVERY_TIMEOUT.value,
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 队列。
# 视频任务如果 next_poll_at 未到期,不提前 poll,只刷新 active 注册表等待 Beat dispatcher 到期投递。
if has_provider_task_id:
if is_video_generation_task(task):
ensure_video_poll_fields(task, now=current_time)
if is_poll_not_due(task, now=current_time):
await db.commit()
await register_poll_active(
task,
check_at=task.next_poll_at,
next_poll_at=task.next_poll_at,
reason=f"{source}_video_poll_not_due",
)
await log_task_event(
task,
event_type=ChatGenerationTaskEventType.POLL_SKIP_NOT_DUE.value,
message=f"{source} 发现视频任务尚未到下一次轮询时间,启动容灾不提前投递 poll",
detail={
"pipeline_stage": task.pipeline_stage,
"payload": redis_payload,
"next_poll_at": task.next_poll_at,
},
)
return "skip_video_poll_not_due"
original_next_poll_at = ensure_aware_utc(task.next_poll_at)
queue_hold_until = _poll_queue_timeout_at(current_time)
task.pipeline_stage = ChatGenerationPipelineStage.WAITING_REMOTE.value
# 这里仍复用 next_poll_at 做短暂队列保护,避免启动容灾重复投递。
# 真正消费时通过 force_due=True 跳过“未到期”校验,避免保护时间反向阻塞本次 poll。
task.next_poll_at = queue_hold_until
await db.commit()
await log_task_event(
task,
event_type=ChatGenerationTaskEventType.GENERATION_RECOVERY_ENQUEUE.value,
message=f"{source} 发现任务存在供应商任务ID,恢复投递轮询队列",
detail={
"pipeline_stage": task.pipeline_stage,
"payload": redis_payload,
"due_next_poll_at": original_next_poll_at,
"queue_hold_until": queue_hold_until,
},
)
poll_generation_task.apply_async(
args=[task.id],
kwargs={"force_due": True},
queue=POLL_QUEUE,
countdown=0,
)
await register_poll_active(
task,
check_at=task.next_poll_at,
next_poll_at=task.next_poll_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=ChatGenerationTaskEventType.GENERATION_RECOVERY_ENQUEUE.value,
message=f"{source} 发现任务未超时且缺少 remote_result_url/供应商任务ID,恢复投递创建队列",
detail={"pipeline_stage": task.pipeline_stage, "payload": redis_payload},
)
chatapi_create_generation_task.apply_async(
args=[task.id],
queue=CeleryQueue.GEN_CHATAPI_CREATE.value,
countdown=0,
)
return "recover_create_no_remote_no_provider_before_deadline"
# 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=ChatGenerationTaskEventType.GENERATION_RECOVERY_ENQUEUE.value,
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=CeleryQueue.GEN_CHATAPI_CREATE.value,
countdown=0,
)
return "recover_create_result_ready_no_url_before_deadline"
return f"skip_stage_{task.pipeline_stage}"
async def recover_generation_tasks_once(db: AsyncSession) -> dict[str, Any]:
"""启动时生成链路容灾扫描。
启动容灾由 worker_ready 触发,只跑一次完整恢复;周期性视频到期轮询由 Celery Beat 调度 dispatch_due_poll_tasks。
恢复顺序:
1. Redis poll active_index 到期任务;
2. DB fallback 扫描 queued/creating/waiting_remote/polling/result_ready
3. 下载阶段仍由 recover_download_tasks_once 兜底。
"""
checked_ids: set[str] = set()
results: dict[str, int] = {}
due_poll_ids = await redis_get_due_registry_ids(
zset_key=settings.POLL_ACTIVE_REDIS_ZSET_KEY,
limit=int(settings.POLL_RECOVERY_BATCH_SIZE or settings.GENERATION_RECOVERY_BATCH_SIZE or 100),
log_context="poll_active",
)
poll_payloads = await redis_get_registry_payloads(
hash_key=settings.POLL_ACTIVE_REDIS_HASH_KEY,
item_ids=due_poll_ids,
log_context="poll_active",
)
for task_id in due_poll_ids:
result = await db.execute(
select(ChatGenerationTask)
.where(
ChatGenerationTask.id == task_id,
ChatGenerationTask.deleted_at.is_(None),
)
.with_for_update()
.limit(1)
)
task = result.scalar_one_or_none()
if task is None:
await _remove_poll_active(task_id)
action = "clean_missing_poll_task"
else:
checked_ids.add(task.id)
action = await recover_one_generation_task(
db,
task,
payload=poll_payloads.get(task_id),
source="startup_poll_redis",
)
results[action] = results.get(action, 0) + 1
batch_size = int(settings.GENERATION_RECOVERY_BATCH_SIZE or settings.DOWNLOAD_RECOVERY_BATCH_SIZE or 100)
max_rounds = max(1, int(settings.GENERATION_RECOVERY_MAX_ROUNDS or 1))
total_db_checked = 0
for _round in range(max_rounds):
query_result = await db.execute(
select(ChatGenerationTask)
.where(
ChatGenerationTask.deleted_at.is_(None),
ChatGenerationTask.generation_mode.in_(list(ALLOWED_GENERATION_MODES)),
ChatGenerationTask.status == ChatGenerationTaskStatus.GENERATING.value,
ChatGenerationTask.pipeline_stage.in_(
[
"queued",
"preparing",
"creating_provider_task",
"waiting_remote",
"polling",
"result_ready",
]
),
)
.order_by(ChatGenerationTask.updated_at.asc())
.limit(batch_size)
.with_for_update(skip_locked=True)
)
tasks = query_result.scalars().all()
if not tasks:
break
progressed_this_round = 0
for task in tasks:
if task.id in checked_ids:
continue
action = await recover_one_generation_task(
db,
task,
payload=None,
source="startup_db",
)
results[action] = results.get(action, 0) + 1
checked_ids.add(task.id)
total_db_checked += 1
progressed_this_round += 1
if len(tasks) < batch_size or progressed_this_round <= 0:
break
return {
"checked": len(checked_ids),
"db_checked": total_db_checked,
"results": results,
}
async def dispatch_due_poll_tasks_once(db: AsyncSession) -> dict[str, Any]:
"""周期性轻量到期轮询调度。
只处理视频任务的 next_poll_at 到期记录,不替代启动容灾 recover_generation_tasks_once。
Beat 每分钟触发本任务,本任务把真实供应商轮询投递到 gen_provider_poll 队列。
"""
from app.tasks.generation_poll_tasks import poll_generation_task, register_poll_active
current_time = _now()
batch_size = max(1, int(settings.POLL_DUE_DISPATCH_BATCH_SIZE or 100))
poll_lease_expired_at = current_time - timedelta(seconds=int(settings.POLL_TASK_LEASE_SECONDS or 300))
query_result = await db.execute(
select(ChatGenerationTask)
.where(
ChatGenerationTask.deleted_at.is_(None),
ChatGenerationTask.generation_mode.in_(list(ALLOWED_GENERATION_MODES)),
ChatGenerationTask.status == ChatGenerationTaskStatus.GENERATING.value,
ChatGenerationTask.gen_type == GenerationType.VIDEO.value,
ChatGenerationTask.next_poll_at.is_not(None),
ChatGenerationTask.next_poll_at <= current_time,
ChatGenerationTask.pipeline_stage.in_(
[
ChatGenerationPipelineStage.WAITING_REMOTE.value,
ChatGenerationPipelineStage.POLLING.value,
]
),
)
.order_by(ChatGenerationTask.next_poll_at.asc(), ChatGenerationTask.updated_at.asc())
.limit(batch_size)
.with_for_update(skip_locked=True)
)
tasks = query_result.scalars().all()
results: dict[str, int] = {}
dispatched_task_ids: list[str] = []
dispatched_due_next_poll_at_by_id: dict[str, datetime | None] = {}
dispatched_queue_hold_until_by_id: dict[str, datetime] = {}
for task in tasks:
action = "skip_unknown"
try:
if task.pipeline_stage == ChatGenerationPipelineStage.POLLING.value:
last_poll_at = ensure_aware_utc(task.last_poll_at)
if last_poll_at and last_poll_at > poll_lease_expired_at:
await register_poll_active(
task,
check_at=_poll_queue_timeout_at(last_poll_at),
next_poll_at=task.next_poll_at,
reason="due_dispatch_polling_lease_alive",
)
action = "skip_polling_lease_alive"
continue
if not (task.seedance_task_id or task.provider_task_id):
# dispatcher 不负责重新 create;没有 provider id 的异常状态交给启动容灾或 create 任务处理。
await log_task_event(
task,
event_type=ChatGenerationTaskEventType.POLL_DISPATCH_SKIP.value,
message="视频到期轮询调度跳过:缺少外部任务ID",
detail={"pipeline_stage": task.pipeline_stage, "next_poll_at": task.next_poll_at},
)
action = "skip_no_provider_task_id"
continue
original_next_poll_at = ensure_aware_utc(task.next_poll_at)
queue_hold_until = _poll_queue_timeout_at(current_time)
task.pipeline_stage = ChatGenerationPipelineStage.WAITING_REMOTE.value
# 设置一个队列消费保护时间,避免 Beat 下一分钟看到旧 next_poll_at 又重复投递。
# poll worker 会通过 force_due=True 消费本次到期任务,避免该保护时间被误判为业务未到期。
task.next_poll_at = queue_hold_until
dispatched_due_next_poll_at_by_id[task.id] = original_next_poll_at
dispatched_queue_hold_until_by_id[task.id] = queue_hold_until
dispatched_task_ids.append(task.id)
action = "dispatch_poll"
except Exception as exc:
logger.exception("视频到期轮询调度单条处理失败。task_id=%s", getattr(task, "id", None))
action = "error"
await log_task_event(
task,
event_type=ChatGenerationTaskEventType.POLL_DISPATCH_SKIP.value,
message=f"视频到期轮询调度单条处理失败:{exc}",
)
finally:
results[action] = results.get(action, 0) + 1
await db.commit()
fresh_tasks = []
if dispatched_task_ids:
fresh_result = await db.execute(
select(ChatGenerationTask)
.where(
ChatGenerationTask.id.in_(dispatched_task_ids),
ChatGenerationTask.deleted_at.is_(None),
)
.execution_options(populate_existing=True)
)
fresh_tasks = fresh_result.scalars().all()
enqueued_count = 0
for task in fresh_tasks:
queue_hold_until = dispatched_queue_hold_until_by_id.get(task.id) or ensure_aware_utc(task.next_poll_at) or _poll_queue_timeout_at(current_time)
due_next_poll_at = dispatched_due_next_poll_at_by_id.get(task.id)
await register_poll_active(
task,
check_at=queue_hold_until,
next_poll_at=queue_hold_until,
reason="due_dispatch_poll_queued",
)
await log_task_event(
task,
event_type=ChatGenerationTaskEventType.POLL_DISPATCH_DUE.value,
message="视频 next_poll_at 到期,已投递 provider poll 队列",
detail={
"due_next_poll_at": due_next_poll_at,
"queue_hold_until": queue_hold_until,
"queue": POLL_QUEUE,
},
)
poll_generation_task.apply_async(
args=[task.id],
kwargs={"force_due": True},
queue=POLL_QUEUE,
countdown=0,
)
enqueued_count += 1
# 如果 log_task_event 内部不 commit,这里要提交一次
if fresh_tasks:
await db.commit()
return {
"checked": len(tasks),
"dispatched": len(dispatched_task_ids),
"enqueued": enqueued_count,
"results": results,
}
@@ -1,214 +0,0 @@
from __future__ import annotations
from datetime import datetime, timezone
from typing import Iterable
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.models.chat_generation_task import ChatGenerationTask
from app.models.credit_record import CreditRecord
from app.models.generation_record import GenerationRecord
from app.services.credits import refund_credits
from app.services.credit_record_meta_service import build_refund_meta_from_charge
from app.services.generation_billing_service import (
CHARGE_MEDIA,
OWNER_CHAT_GENERATION_TASK,
OWNER_GENERATION_RECORD,
build_credit_biz_key,
parse_credit_biz_key,
)
TERMINAL_FAILED_STAGES = {"failed", "timeout", "download_failed"}
def _round2(value: float | int | None) -> float:
return round(float(value or 0), 2)
async def _has_refund_for_biz_key(db: AsyncSession, *, user_id: str, charge_biz_key: str) -> bool:
result = await db.execute(
select(CreditRecord.id)
.where(
CreditRecord.user_id == user_id,
CreditRecord.type == "refund",
CreditRecord.refund_for_biz_key == charge_biz_key,
)
.limit(1)
)
return result.scalar_one_or_none() is not None
async def _find_unrefunded_media_charges(
db: AsyncSession,
*,
user_id: str,
owner_type: str,
owner_id: str,
) -> list[CreditRecord]:
"""查找当前任务下所有未退款的媒体生成扣费流水。"""
pattern = f"{owner_type}:{owner_id}:attempt:%:{CHARGE_MEDIA}:charge"
result = await db.execute(
select(CreditRecord)
.where(
CreditRecord.user_id == user_id,
CreditRecord.related_id == owner_id,
CreditRecord.type == "consume",
CreditRecord.biz_key.like(pattern),
)
.order_by(CreditRecord.created_at.asc())
)
charges = list(result.scalars().all())
unrefunded: list[CreditRecord] = []
for charge in charges:
if not charge.biz_key:
continue
if not await _has_refund_for_biz_key(db, user_id=user_id, charge_biz_key=charge.biz_key):
unrefunded.append(charge)
return unrefunded
async def refund_unrefunded_media_charges(
db: AsyncSession,
*,
user_id: str,
owner_type: str,
owner_id: str,
description_prefix: str,
) -> float:
"""回退当前任务所有未退款媒体扣费流水。
容灾考虑:
- 不依赖 retry_count 推断当前轮次。
- 如果扣费成功后 worker 崩溃,最终失败时会扫出未退款的 media charge 并补偿。
- refund_for_biz_key 保证同一轮扣费不会重复退款。
"""
total_refunded = 0.0
charges = await _find_unrefunded_media_charges(
db,
user_id=user_id,
owner_type=owner_type,
owner_id=owner_id,
)
for charge in charges:
parsed = parse_credit_biz_key(charge.biz_key)
if not parsed:
continue
attempt_no = int(parsed["attempt_no"])
refund_biz_key = build_credit_biz_key(
owner_type=owner_type,
owner_id=owner_id,
attempt_no=attempt_no,
charge_kind=CHARGE_MEDIA,
action="refund",
)
amount = abs(_round2(charge.amount))
if amount <= 0:
continue
await refund_credits(
db,
user_id=user_id,
amount=amount,
description=f"{description_prefix}失败积分回退",
related_id=owner_id,
biz_key=refund_biz_key,
refund_for_biz_key=charge.biz_key,
record_meta=build_refund_meta_from_charge(charge, attempt_no=attempt_no),
)
total_refunded = round(total_refunded + amount, 2)
return total_refunded
async def mark_generation_record_failed_and_refund_once(
db: AsyncSession,
*,
record_id: str | None = None,
record: GenerationRecord | None = None,
error_message: str | None = None,
) -> GenerationRecord | None:
"""把 GenerationRecord 标记为最终失败并幂等退回媒体生成积分。
调用方负责 commit;本函数只 flush,确保状态和退款在同一事务内提交。
"""
if record is None:
if not record_id:
return None
result = await db.execute(
select(GenerationRecord)
.where(GenerationRecord.id == record_id, GenerationRecord.deleted_at.is_(None))
.with_for_update()
.limit(1)
)
record = result.scalar_one_or_none()
if not record:
return None
if record.status == "completed":
return record
record.status = "failed"
if error_message:
record.error_message = error_message
refunded_amount = await refund_unrefunded_media_charges(
db,
user_id=record.user_id,
owner_type=OWNER_GENERATION_RECORD,
owner_id=record.id,
description_prefix="生成记录",
)
if refunded_amount > 0:
# GenerationRecord.credits_cost 只代表视频/图片生成媒体积分。
# 提示词优化积分在 text_credits_cost 中记录,优化成功后不参与生成失败退款。
record.credits_cost = max(0.0, round(float(record.credits_cost or 0) - refunded_amount, 2))
await db.flush()
return record
async def mark_chat_generation_task_failed_and_refund_once(
db: AsyncSession,
*,
task_id: str | None = None,
task: ChatGenerationTask | None = None,
error_message: str | None = None,
pipeline_stage: str = "failed",
) -> ChatGenerationTask | None:
"""把 ChatGenerationTask 标记为最终失败并幂等退回媒体生成积分。"""
if task is None:
if not task_id:
return None
result = await db.execute(
select(ChatGenerationTask)
.where(
ChatGenerationTask.id == task_id,
ChatGenerationTask.generation_mode.in_(["chatapi_async", "hot_opening_replicate", "shot_replicate"]),
ChatGenerationTask.deleted_at.is_(None),
)
.with_for_update()
.limit(1)
)
task = result.scalar_one_or_none()
if not task:
return None
if task.status == "completed":
return task
task.status = "failed"
task.pipeline_stage = pipeline_stage if pipeline_stage in TERMINAL_FAILED_STAGES else "failed"
if error_message:
task.error_message = error_message
refunded_amount = await refund_unrefunded_media_charges(
db,
user_id=task.user_id,
owner_type=OWNER_CHAT_GENERATION_TASK,
owner_id=task.id,
description_prefix="任务生成",
)
if refunded_amount > 0:
task.credits_cost = max(0.0, round(float(task.credits_cost or 0) - refunded_amount, 2))
await db.flush()
return task
@@ -1,210 +0,0 @@
from __future__ import annotations
import json
from datetime import datetime, timedelta, timezone
from typing import Any
from fastapi import HTTPException
from sqlalchemy.ext.asyncio import AsyncSession
from app.config import settings
from app.models.chat_generation_task import ChatGenerationTask
from app.models.user import User
from app.schemas.generation_ai import GenerationAIReference, GenerationAITaskCreate
from app.services.generation_ai_service import (
IMAGE_DEFAULT_PROPORTION,
IMAGE_DEFAULT_PX,
IMAGE_DEFAULT_SIZE,
VIDEO_DEFAULT_RATIO,
VIDEO_DEFAULT_RESOLUTION,
_build_image_snapshot,
_build_video_snapshot,
_get_image_engine,
_get_video_engine,
_image_supported_sizes,
_parse_list,
normalize_px,
)
from app.services.generation_billing_service import OWNER_CHAT_GENERATION_TASK, charge_generation_media_by_params
from app.services.resource_capacity_service import assert_user_resource_capacity_available
from app.services.private_portrait.reference_resolver import resolve_private_portrait_references
from app.utils.id_gen import generate_id
def _json(data: Any) -> str | None:
if data is None:
return None
return json.dumps(data, ensure_ascii=False, default=str)
def _build_backend_idempotency_key(*, generation_mode: str, gen_type: str, task_id: str) -> str:
"""模块生成关联 ChatGenerationTask 的幂等键由后端生成。
不再接收前端透传,避免 user_id + generation_mode + idempotency_key
唯一索引被前端固定 key 或重复 key 拦截。
"""
return f"{generation_mode}:{gen_type}:{task_id}"[:64]
async def create_chat_generation_task_for_module(
db: AsyncSession,
*,
current_user: User,
generation_mode: str,
gen_type: str,
original_prompt: str,
optimized_prompt: str | None = None,
engine_id: str | None = None,
media_references: list[dict[str, Any]] | None = None,
idempotency_key: str | None = None,
image_size: str | None = None,
image_proportion: str | None = None,
image_px: str | None = None,
duration: int | None = None,
aspect_ratio: str | None = None,
resolution: str | None = None,
billing_project_name: str = "模块生成任务",
billing_description_prefix: str = "模块生成-",
billing_source_module: str | None = None,
billing_source_project_id: str | None = None,
billing_source_step_id: str | None = None,
billing_source_step_code: str | None = None,
billing_scene: str | None = None,
) -> ChatGenerationTask:
"""创建可复用的 ChatGenerationTask 子任务。
和 /generation-ai 普通任务不同,generation_mode 由业务模块传入,
但仍复用同一套引擎校验、扣费、Celery 创建/轮询/下载逻辑。
"""
gen_type = gen_type.lower().strip()
if gen_type not in ("image", "video"):
raise HTTPException(status_code=400, detail="gen_type 仅支持 image 或 video")
task_id = generate_id()
now = datetime.now(timezone.utc)
refs = media_references or []
refs = await resolve_private_portrait_references(
db,
user_id=current_user.id,
media_references=refs,
gen_type=gen_type,
)
backend_idempotency_key = _build_backend_idempotency_key(
generation_mode=generation_mode,
gen_type=gen_type,
task_id=task_id,
)
await assert_user_resource_capacity_available(db, current_user.id)
if gen_type == "image":
engine = await _get_image_engine(db, engine_id)
sizes = _image_supported_sizes(engine)
size = image_size or engine.default_size or IMAGE_DEFAULT_SIZE
proportion = image_proportion or IMAGE_DEFAULT_PROPORTION
px = normalize_px(image_px)
if sizes:
if size not in sizes:
raise HTTPException(status_code=400, detail=f"图片分辨率档位不支持: {size}")
if proportion not in sizes.get(size, {}):
raise HTTPException(status_code=400, detail=f"图片比例不支持: {proportion}")
px = px or normalize_px((sizes.get(size) or {}).get(proportion))
px = px or IMAGE_DEFAULT_PX
media_billing = await charge_generation_media_by_params(
db,
user_id=current_user.id,
record_id=task_id,
gen_type="image",
image_size=size,
engine_id=engine.id,
project_name=billing_project_name,
description_prefix=billing_description_prefix,
owner_type=OWNER_CHAT_GENERATION_TASK,
attempt_no=1,
source_module=billing_source_module,
source_project_id=billing_source_project_id,
source_step_id=billing_source_step_id,
source_step_code=billing_source_step_code,
billing_scene=billing_scene,
)
snapshot = _build_image_snapshot(engine, size, proportion, px)
task = ChatGenerationTask(
id=task_id,
user_id=current_user.id,
original_prompt=original_prompt,
optimized_prompt=optimized_prompt,
gen_type="image",
image_size=size,
image_proportion=proportion,
image_px=px,
status="generating",
generation_mode=generation_mode,
pipeline_stage="queued",
engine_id=engine.id,
engine_snapshot_json=_json(snapshot),
media_references=_json(refs) if refs else None,
credits_cost=round(media_billing.total_charged, 2),
idempotency_key=backend_idempotency_key,
deadline_at=now + timedelta(minutes=settings.CHATAPI_ASYNC_IMAGE_DEADLINE_MINUTES),
)
else:
engine = await _get_video_engine(db, engine_id)
ratio = aspect_ratio or VIDEO_DEFAULT_RATIO
selected_resolution = resolution or VIDEO_DEFAULT_RESOLUTION
selected_duration = duration or 4
ratios = _parse_list(engine.supported_ratios, [])
resolutions = _parse_list(engine.supported_resolutions, [])
durations = _parse_list(engine.supported_durations, [])
if ratios and ratio not in ratios:
raise HTTPException(status_code=400, detail=f"视频比例不支持: {ratio}")
if resolutions and selected_resolution not in resolutions:
raise HTTPException(status_code=400, detail=f"视频分辨率不支持: {selected_resolution}")
if durations and selected_duration not in durations:
raise HTTPException(status_code=400, detail=f"视频时长不支持: {selected_duration}")
if engine.max_duration and selected_duration > engine.max_duration:
raise HTTPException(status_code=400, detail=f"视频时长不能超过 {engine.max_duration}")
media_billing = await charge_generation_media_by_params(
db,
user_id=current_user.id,
record_id=task_id,
gen_type="video",
duration=selected_duration,
resolution=selected_resolution,
engine_id=engine.id,
project_name=billing_project_name,
description_prefix=billing_description_prefix,
owner_type=OWNER_CHAT_GENERATION_TASK,
attempt_no=1,
source_module=billing_source_module,
source_project_id=billing_source_project_id,
source_step_id=billing_source_step_id,
source_step_code=billing_source_step_code,
billing_scene=billing_scene,
)
snapshot = _build_video_snapshot(engine, ratio, selected_resolution, selected_duration)
task = ChatGenerationTask(
id=task_id,
user_id=current_user.id,
original_prompt=original_prompt,
optimized_prompt=optimized_prompt,
gen_type="video",
duration=selected_duration,
aspect_ratio=ratio,
resolution=selected_resolution,
image_size=image_size or IMAGE_DEFAULT_SIZE,
image_proportion=image_proportion or IMAGE_DEFAULT_PROPORTION,
image_px=normalize_px(image_px) or IMAGE_DEFAULT_PX,
status="generating",
generation_mode=generation_mode,
pipeline_stage="queued",
engine_id=engine.id,
engine_snapshot_json=_json(snapshot),
media_references=_json(refs) if refs else None,
credits_cost=round(media_billing.total_charged, 2),
idempotency_key=backend_idempotency_key,
deadline_at=now + timedelta(hours=settings.CHATAPI_ASYNC_VIDEO_FINAL_DEADLINE_HOURS),
)
db.add(task)
await db.flush()
return task