AI创作批量生成任务 main V1 init
This commit is contained in:
@@ -0,0 +1 @@
|
||||
"""生成任务领域服务。"""
|
||||
@@ -0,0 +1 @@
|
||||
"""AI 创作生成编排服务。"""
|
||||
@@ -0,0 +1,122 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
|
||||
from fastapi import HTTPException
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.enums.common import MAX_GENERATION_COUNT, MIN_GENERATION_COUNT
|
||||
from app.enums.generation_provider import IMAGE_MULTI_OUTPUT_MAX, IMAGE_MULTI_REFERENCE_MAX
|
||||
from app.models.image_engine import ImageEngine
|
||||
from app.models.video_engine import VideoEngine
|
||||
|
||||
IMAGE_DEFAULT_SIZE = "2K"
|
||||
IMAGE_DEFAULT_PROPORTION = "1:1"
|
||||
IMAGE_DEFAULT_PX = "2048x2048"
|
||||
VIDEO_DEFAULT_DURATION = 4
|
||||
VIDEO_DEFAULT_RATIO = "16:9"
|
||||
VIDEO_DEFAULT_RESOLUTION = "480p"
|
||||
|
||||
|
||||
def normalize_px(value: str | None) -> str | None:
|
||||
if not value:
|
||||
return value
|
||||
return value.replace("×", "x").replace("X", "x").replace("×x", "x").replace("x×", "x")
|
||||
|
||||
|
||||
def parse_json_list(value: str | None, fallback: list):
|
||||
try:
|
||||
parsed = json.loads(value or "")
|
||||
return parsed if isinstance(parsed, list) else fallback
|
||||
except Exception:
|
||||
return fallback
|
||||
|
||||
|
||||
def image_supported_sizes(engine: ImageEngine) -> dict:
|
||||
try:
|
||||
data = json.loads(engine.supported_sizes or "{}")
|
||||
return data if isinstance(data, dict) else {}
|
||||
except Exception:
|
||||
return {}
|
||||
|
||||
|
||||
def normalize_generation_count(value: int | None) -> int:
|
||||
try:
|
||||
count = int(value or MIN_GENERATION_COUNT)
|
||||
except (TypeError, ValueError):
|
||||
count = MIN_GENERATION_COUNT
|
||||
return min(MAX_GENERATION_COUNT, max(MIN_GENERATION_COUNT, count))
|
||||
|
||||
|
||||
async def get_image_engine(db: AsyncSession, engine_id: str | None) -> ImageEngine:
|
||||
query = select(ImageEngine).where(ImageEngine.is_active == True)
|
||||
if engine_id:
|
||||
query = query.where(ImageEngine.id == engine_id)
|
||||
else:
|
||||
query = query.order_by(ImageEngine.priority.desc()).limit(1)
|
||||
result = await db.execute(query)
|
||||
engine = result.scalar_one_or_none()
|
||||
if not engine:
|
||||
raise HTTPException(status_code=400, detail="没有可用的图片引擎")
|
||||
return engine
|
||||
|
||||
|
||||
async def get_video_engine(db: AsyncSession, engine_id: str | None) -> VideoEngine:
|
||||
query = select(VideoEngine).where(VideoEngine.is_active == True)
|
||||
if engine_id:
|
||||
query = query.where(VideoEngine.id == engine_id)
|
||||
else:
|
||||
query = query.order_by(VideoEngine.priority.desc())
|
||||
result = await db.execute(query.limit(1))
|
||||
engine = result.scalar_one_or_none()
|
||||
if not engine:
|
||||
raise HTTPException(status_code=400, detail="没有可用的视频引擎")
|
||||
return engine
|
||||
|
||||
|
||||
def build_image_snapshot(engine: ImageEngine, size: str, proportion: str, px: str) -> dict:
|
||||
return {
|
||||
"engine_type": "image",
|
||||
"id": engine.id,
|
||||
"name": engine.name,
|
||||
"provider": engine.provider,
|
||||
"api_base": engine.api_base,
|
||||
"api_key_masked": "****" if engine.api_key else "",
|
||||
"model_name": engine.model_name,
|
||||
"generate_url": engine.generate_url,
|
||||
"supported_models": parse_json_list(engine.supported_models, []),
|
||||
"default_size": engine.default_size,
|
||||
"multi_generation_enabled": bool(getattr(engine, "multi_generation_enabled", False)),
|
||||
"max_generation_count": normalize_generation_count(getattr(engine, "max_generation_count", 1)),
|
||||
"multi_image_max_images": int(getattr(engine, "multi_image_max_images", IMAGE_MULTI_OUTPUT_MAX) or IMAGE_MULTI_OUTPUT_MAX),
|
||||
"max_reference_image_count": int(getattr(engine, "max_reference_image_count", IMAGE_MULTI_REFERENCE_MAX) or 0),
|
||||
"output_format": (getattr(engine, "output_format", "") or "").lower().strip(),
|
||||
"selected_size": size,
|
||||
"selected_proportion": proportion,
|
||||
"selected_px": px,
|
||||
}
|
||||
|
||||
|
||||
def build_video_snapshot(engine: VideoEngine, ratio: str, resolution: str, duration: int) -> dict:
|
||||
return {
|
||||
"engine_type": "video",
|
||||
"id": engine.id,
|
||||
"name": engine.name,
|
||||
"provider": engine.provider,
|
||||
"api_base": engine.api_base,
|
||||
"api_key_masked": "****" if engine.api_key else "",
|
||||
"model_name": engine.model_name,
|
||||
"generate_url": engine.generate_url,
|
||||
"query_url": engine.query_url,
|
||||
"supported_ratios": parse_json_list(engine.supported_ratios, []),
|
||||
"supported_resolutions": parse_json_list(engine.supported_resolutions, []),
|
||||
"supported_durations": parse_json_list(engine.supported_durations, []),
|
||||
"max_duration": engine.max_duration,
|
||||
"max_audio_count": engine.max_audio_count,
|
||||
"multi_generation_enabled": bool(getattr(engine, "multi_generation_enabled", False)),
|
||||
"max_generation_count": normalize_generation_count(getattr(engine, "max_generation_count", 1)),
|
||||
"selected_ratio": ratio,
|
||||
"selected_resolution": resolution,
|
||||
"selected_duration": duration,
|
||||
}
|
||||
@@ -0,0 +1,532 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from types import SimpleNamespace
|
||||
from uuid import uuid4
|
||||
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.enums.generation_provider import IMAGE_PROVIDER_CLAIM_LEASE_SECONDS
|
||||
from app.enums.generation_task import (
|
||||
ChatGenerationPipelineStage,
|
||||
ChatGenerationTaskEventType,
|
||||
ChatGenerationTaskStatus,
|
||||
GenerationMode,
|
||||
GenerationType,
|
||||
)
|
||||
from app.models.chat_generation_task import ChatGenerationTask
|
||||
from app.services.generation.ai.task_group_service import aggregate_main_task_status, load_children_map
|
||||
from app.services.generation.log_service import log_task_event
|
||||
from app.services.generation.provider_service import (
|
||||
create_image_sync_batch_result_with_engine,
|
||||
get_runtime_engine,
|
||||
)
|
||||
from app.services.generation.refund_service import mark_chat_generation_task_failed_and_refund_once
|
||||
from app.services.image_gen import ImageProviderError
|
||||
from app.services.operation_log_service import build_exception_detail, log_operation_event
|
||||
from app.utils.id_gen import generate_id
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class ImageBatchClaim:
|
||||
acquired: bool
|
||||
main_task_id: str
|
||||
claim_token: str | None = None
|
||||
task_snapshot: SimpleNamespace | None = None
|
||||
runtime_engine: SimpleNamespace | None = None
|
||||
existing_child_ids: list[str] | None = None
|
||||
reason: str | None = None
|
||||
|
||||
|
||||
def _now() -> datetime:
|
||||
return datetime.now(timezone.utc)
|
||||
|
||||
|
||||
def _json(value) -> str | None:
|
||||
if value is None:
|
||||
return None
|
||||
return json.dumps(value, ensure_ascii=False, default=str)
|
||||
|
||||
|
||||
def _aware(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 _lease_alive(task: ChatGenerationTask, now: datetime | None = None) -> bool:
|
||||
lease_until = _aware(task.provider_create_lease_until)
|
||||
return bool(task.provider_create_claim_token and lease_until and lease_until > (now or _now()))
|
||||
|
||||
|
||||
def _task_snapshot(main: ChatGenerationTask) -> SimpleNamespace:
|
||||
return SimpleNamespace(
|
||||
id=str(main.id),
|
||||
user_id=str(main.user_id),
|
||||
generation_mode=str(main.generation_mode),
|
||||
generation_count=int(main.generation_count or 1),
|
||||
original_prompt=main.original_prompt,
|
||||
optimized_prompt=main.optimized_prompt,
|
||||
media_references=main.media_references,
|
||||
gen_type=main.gen_type,
|
||||
duration=main.duration,
|
||||
aspect_ratio=main.aspect_ratio,
|
||||
resolution=main.resolution,
|
||||
image_size=main.image_size,
|
||||
image_proportion=main.image_proportion,
|
||||
image_px=main.image_px,
|
||||
engine_id=main.engine_id,
|
||||
)
|
||||
|
||||
|
||||
async def _claim_image_main_batch(db: AsyncSession, main_task_id: str) -> ImageBatchClaim:
|
||||
result = await db.execute(
|
||||
select(ChatGenerationTask)
|
||||
.where(
|
||||
ChatGenerationTask.id == main_task_id,
|
||||
ChatGenerationTask.generation_mode == GenerationMode.CHATAPI_MAIN.value,
|
||||
ChatGenerationTask.gen_type == GenerationType.IMAGE.value,
|
||||
ChatGenerationTask.deleted_at.is_(None),
|
||||
)
|
||||
.with_for_update()
|
||||
.limit(1)
|
||||
)
|
||||
main = result.scalar_one_or_none()
|
||||
if not main:
|
||||
await db.rollback()
|
||||
return ImageBatchClaim(False, main_task_id, reason="main_missing")
|
||||
|
||||
children_map = await load_children_map(db, [main.id], include_deleted=True)
|
||||
existing_children = children_map.get(main.id, [])
|
||||
if existing_children:
|
||||
child_ids = [str(child.id) for child in existing_children if child.deleted_at is None]
|
||||
main.provider_create_claim_token = None
|
||||
main.provider_create_lease_until = None
|
||||
await db.commit()
|
||||
return ImageBatchClaim(False, main_task_id, existing_child_ids=child_ids, reason="already_split")
|
||||
|
||||
if main.status != ChatGenerationTaskStatus.GENERATING.value:
|
||||
status = str(main.status)
|
||||
await db.rollback()
|
||||
return ImageBatchClaim(False, main_task_id, reason=f"status_{status}")
|
||||
|
||||
now = _now()
|
||||
if _lease_alive(main, now):
|
||||
user_id = str(main.user_id)
|
||||
group_id = str(main.id)
|
||||
lease_until = main.provider_create_lease_until
|
||||
await db.rollback()
|
||||
log_operation_event(
|
||||
domain="generation_ai_batch",
|
||||
event_type="IMAGE_MAIN_CLAIM_REJECTED",
|
||||
event_status="skipped",
|
||||
source="celery",
|
||||
user_id=user_id,
|
||||
group_id=group_id,
|
||||
task_id=group_id,
|
||||
detail={"reason": "lease_alive", "lease_until": lease_until},
|
||||
)
|
||||
return ImageBatchClaim(False, main_task_id, reason="lease_alive")
|
||||
|
||||
deadline = _aware(main.deadline_at)
|
||||
if deadline and deadline <= now:
|
||||
main.provider_create_claim_token = None
|
||||
main.provider_create_lease_until = None
|
||||
await mark_chat_generation_task_failed_and_refund_once(
|
||||
db,
|
||||
task=main,
|
||||
error_message="图片批量生成任务超时",
|
||||
pipeline_stage=ChatGenerationPipelineStage.TIMEOUT.value,
|
||||
)
|
||||
await db.commit()
|
||||
return ImageBatchClaim(False, main_task_id, reason="deadline_expired")
|
||||
|
||||
claim_token = uuid4().hex
|
||||
main.provider_create_claim_token = claim_token
|
||||
main.provider_create_started_at = now
|
||||
main.provider_create_lease_until = now + timedelta(seconds=IMAGE_PROVIDER_CLAIM_LEASE_SECONDS)
|
||||
main.pipeline_stage = ChatGenerationPipelineStage.CREATING_PROVIDER_TASK.value
|
||||
runtime_engine = await get_runtime_engine(db, main)
|
||||
snapshot = _task_snapshot(main)
|
||||
user_id = str(main.user_id)
|
||||
generation_count = int(main.generation_count or 1)
|
||||
lease_until = main.provider_create_lease_until
|
||||
await db.commit()
|
||||
|
||||
log_operation_event(
|
||||
domain="generation_ai_batch",
|
||||
event_type="IMAGE_MAIN_CLAIM_ACQUIRED",
|
||||
event_status="success",
|
||||
source="celery",
|
||||
user_id=user_id,
|
||||
group_id=main_task_id,
|
||||
task_id=main_task_id,
|
||||
detail={
|
||||
"generation_count": generation_count,
|
||||
"claim_token_suffix": claim_token[-8:],
|
||||
"lease_until": lease_until,
|
||||
},
|
||||
)
|
||||
return ImageBatchClaim(
|
||||
True,
|
||||
main_task_id,
|
||||
claim_token=claim_token,
|
||||
task_snapshot=snapshot,
|
||||
runtime_engine=runtime_engine,
|
||||
)
|
||||
|
||||
|
||||
def _validate_provider_batch(provider_result: dict, generation_count: int) -> list[dict]:
|
||||
items = provider_result.get("items") or []
|
||||
if not isinstance(items, list):
|
||||
raise RuntimeError("图片供应商返回 items 结构异常")
|
||||
|
||||
success_items: list[dict] = []
|
||||
errors: list[str] = []
|
||||
for position, item in enumerate(items, start=1):
|
||||
if not isinstance(item, dict):
|
||||
errors.append(f"第{position}项返回结构无效")
|
||||
continue
|
||||
if item.get("error_message") or item.get("error_code"):
|
||||
errors.append(
|
||||
f"第{position}项: {item.get('error_message') or item.get('error_code') or '生成失败'}"
|
||||
)
|
||||
continue
|
||||
remote_url = str(item.get("remote_result_url") or "").strip()
|
||||
if not remote_url:
|
||||
errors.append(f"第{position}项: 供应商未返回图片地址")
|
||||
continue
|
||||
normalized = dict(item)
|
||||
normalized["generation_index"] = position
|
||||
success_items.append(normalized)
|
||||
|
||||
generated_images = int(provider_result.get("generated_images") or 0)
|
||||
if generated_images and generated_images != len(success_items):
|
||||
errors.append(
|
||||
f"usage.generated_images={generated_images} 与有效图片数 {len(success_items)} 不一致"
|
||||
)
|
||||
if len(items) != generation_count:
|
||||
errors.append(f"返回条目数应为 {generation_count},实际 {len(items)}")
|
||||
if len(success_items) != generation_count:
|
||||
errors.append(f"成功图片数应为 {generation_count},实际 {len(success_items)}")
|
||||
if errors:
|
||||
raise RuntimeError("图片组图未全部成功;" + ";".join(errors))
|
||||
return success_items
|
||||
|
||||
|
||||
async def _fail_claimed_main(
|
||||
db: AsyncSession,
|
||||
*,
|
||||
main_task_id: str,
|
||||
claim_token: str,
|
||||
error_message: str,
|
||||
event_type: ChatGenerationTaskEventType,
|
||||
exception: Exception | None = None,
|
||||
) -> bool:
|
||||
try:
|
||||
await db.rollback()
|
||||
except Exception:
|
||||
pass
|
||||
result = await db.execute(
|
||||
select(ChatGenerationTask)
|
||||
.where(
|
||||
ChatGenerationTask.id == main_task_id,
|
||||
ChatGenerationTask.generation_mode == GenerationMode.CHATAPI_MAIN.value,
|
||||
ChatGenerationTask.deleted_at.is_(None),
|
||||
)
|
||||
.with_for_update()
|
||||
.limit(1)
|
||||
)
|
||||
main = result.scalar_one_or_none()
|
||||
if not main or main.provider_create_claim_token != claim_token:
|
||||
await db.rollback()
|
||||
return False
|
||||
|
||||
existing_map = await load_children_map(db, [main.id], include_deleted=True)
|
||||
if existing_map.get(main.id):
|
||||
# child 已经落库后不再允许图片生成退款。
|
||||
main.provider_create_claim_token = None
|
||||
main.provider_create_lease_until = None
|
||||
await db.commit()
|
||||
return False
|
||||
|
||||
main.provider_create_claim_token = None
|
||||
main.provider_create_lease_until = None
|
||||
await mark_chat_generation_task_failed_and_refund_once(
|
||||
db,
|
||||
task=main,
|
||||
error_message=error_message,
|
||||
pipeline_stage=ChatGenerationPipelineStage.FAILED.value,
|
||||
)
|
||||
task_id = str(main.id)
|
||||
user_id = str(main.user_id)
|
||||
await db.commit()
|
||||
await log_task_event(
|
||||
task_id=task_id,
|
||||
event_type=event_type.value,
|
||||
to_status=ChatGenerationTaskStatus.FAILED.value,
|
||||
to_stage=ChatGenerationPipelineStage.FAILED.value,
|
||||
message=error_message,
|
||||
)
|
||||
log_operation_event(
|
||||
domain="generation_ai_batch",
|
||||
event_type=event_type.value,
|
||||
event_status="failed",
|
||||
source="celery",
|
||||
user_id=user_id,
|
||||
group_id=task_id,
|
||||
task_id=task_id,
|
||||
message=error_message,
|
||||
detail=build_exception_detail(exception) if exception else {"message": error_message},
|
||||
error=error_message,
|
||||
)
|
||||
return True
|
||||
|
||||
|
||||
async def _split_children(
|
||||
db: AsyncSession,
|
||||
*,
|
||||
main_task_id: str,
|
||||
claim_token: str,
|
||||
provider_result: dict,
|
||||
provider_items: list[dict],
|
||||
) -> list[str]:
|
||||
result = await db.execute(
|
||||
select(ChatGenerationTask)
|
||||
.where(
|
||||
ChatGenerationTask.id == main_task_id,
|
||||
ChatGenerationTask.generation_mode == GenerationMode.CHATAPI_MAIN.value,
|
||||
ChatGenerationTask.gen_type == GenerationType.IMAGE.value,
|
||||
ChatGenerationTask.deleted_at.is_(None),
|
||||
)
|
||||
.with_for_update()
|
||||
.limit(1)
|
||||
)
|
||||
main = result.scalar_one_or_none()
|
||||
if not main:
|
||||
raise RuntimeError("图片主任务不存在或已删除")
|
||||
if main.provider_create_claim_token != claim_token:
|
||||
raise RuntimeError("图片主任务执行租约已失效,拒绝拆分子任务")
|
||||
if main.status != ChatGenerationTaskStatus.GENERATING.value:
|
||||
raise RuntimeError(f"图片主任务当前状态不允许拆分: {main.status}")
|
||||
|
||||
existing_map = await load_children_map(db, [main.id], include_deleted=True)
|
||||
existing = existing_map.get(main.id, [])
|
||||
if existing:
|
||||
main.provider_create_claim_token = None
|
||||
main.provider_create_lease_until = None
|
||||
await db.commit()
|
||||
return [str(child.id) for child in existing if child.deleted_at is None]
|
||||
|
||||
expected_count = max(1, int(main.generation_count or 1))
|
||||
if len(provider_items) != expected_count:
|
||||
raise RuntimeError(f"图片批量拆分数量不一致,期望 {expected_count},实际 {len(provider_items)}")
|
||||
|
||||
log_operation_event(
|
||||
domain="generation_ai_batch",
|
||||
event_type=ChatGenerationTaskEventType.IMAGE_BATCH_SPLIT_START.value,
|
||||
event_status="started",
|
||||
source="celery",
|
||||
user_id=main.user_id,
|
||||
group_id=main.id,
|
||||
task_id=main.id,
|
||||
detail={"generation_count": expected_count},
|
||||
)
|
||||
|
||||
children: list[ChatGenerationTask] = []
|
||||
for item in provider_items:
|
||||
index = int(item.get("generation_index") or 0)
|
||||
if index < 1 or index > expected_count:
|
||||
raise RuntimeError(f"无效的图片生成序号: {index}")
|
||||
child = ChatGenerationTask(
|
||||
id=generate_id(),
|
||||
user_id=main.user_id,
|
||||
original_prompt=main.original_prompt,
|
||||
optimized_prompt=main.optimized_prompt,
|
||||
gen_type=main.gen_type,
|
||||
image_size=main.image_size,
|
||||
image_proportion=main.image_proportion,
|
||||
image_px=main.image_px,
|
||||
status=ChatGenerationTaskStatus.GENERATING.value,
|
||||
pipeline_stage=ChatGenerationPipelineStage.RESULT_READY.value,
|
||||
generation_mode=GenerationMode.CHATAPI_CHILD.value,
|
||||
parent_task_id=main.id,
|
||||
generation_count=expected_count,
|
||||
generation_index=index,
|
||||
media_references=main.media_references,
|
||||
remote_result_url=item.get("remote_result_url"),
|
||||
engine_id=main.engine_id,
|
||||
engine_snapshot_json=main.engine_snapshot_json,
|
||||
provider_response_json=_json(item.get("response_data") or {}),
|
||||
# 图片生成计费和 token 都归属于 main;child 只负责下载和资源展示。
|
||||
credits_cost=0,
|
||||
image_tokens_used=0,
|
||||
deadline_at=main.deadline_at,
|
||||
)
|
||||
children.append(child)
|
||||
|
||||
children.sort(key=lambda child: int(child.generation_index or 0))
|
||||
db.add_all(children)
|
||||
main.provider_response_json = _json(provider_result.get("response_data") or provider_result)
|
||||
main.image_tokens_used = int(provider_result.get("image_tokens") or 0)
|
||||
main.provider_create_claim_token = None
|
||||
main.provider_create_lease_until = None
|
||||
await db.flush()
|
||||
child_ids = [str(child.id) for child in children]
|
||||
main_id = str(main.id)
|
||||
main_user_id = str(main.user_id)
|
||||
await aggregate_main_task_status(db, parent_task_id=main_id)
|
||||
await db.commit()
|
||||
|
||||
log_operation_event(
|
||||
domain="generation_ai_batch",
|
||||
event_type=ChatGenerationTaskEventType.IMAGE_BATCH_SPLIT_SUCCESS.value,
|
||||
event_status="success",
|
||||
source="celery",
|
||||
user_id=main_user_id,
|
||||
group_id=main_id,
|
||||
task_id=main_id,
|
||||
detail={"child_task_ids": child_ids},
|
||||
)
|
||||
return child_ids
|
||||
|
||||
|
||||
async def _enqueue_child_downloads(db: AsyncSession, child_ids: list[str]) -> dict[str, list[str]]:
|
||||
if not child_ids:
|
||||
return {"enqueued": [], "failed": []}
|
||||
result = await db.execute(
|
||||
select(ChatGenerationTask)
|
||||
.where(
|
||||
ChatGenerationTask.id.in_(child_ids),
|
||||
ChatGenerationTask.deleted_at.is_(None),
|
||||
)
|
||||
.order_by(ChatGenerationTask.generation_index.asc())
|
||||
)
|
||||
children = list(result.scalars().all())
|
||||
from app.tasks.generation_download_tasks import enqueue_download_task
|
||||
|
||||
enqueued: list[str] = []
|
||||
failed: list[str] = []
|
||||
for child in children:
|
||||
if child.status == ChatGenerationTaskStatus.COMPLETED.value:
|
||||
continue
|
||||
if child.pipeline_stage in {
|
||||
ChatGenerationPipelineStage.DOWNLOAD_QUEUED.value,
|
||||
ChatGenerationPipelineStage.DOWNLOADING.value,
|
||||
ChatGenerationPipelineStage.RETRY_WAITING.value,
|
||||
}:
|
||||
continue
|
||||
celery_task_id = await enqueue_download_task(db, child, reason="image_batch_split")
|
||||
if celery_task_id:
|
||||
enqueued.append(str(child.id))
|
||||
else:
|
||||
failed.append(str(child.id))
|
||||
|
||||
if children and children[0].parent_task_id:
|
||||
await aggregate_main_task_status(db, parent_task_id=str(children[0].parent_task_id))
|
||||
await db.commit()
|
||||
return {"enqueued": enqueued, "failed": failed}
|
||||
|
||||
|
||||
async def run_image_main_batch(db: AsyncSession, main_task: ChatGenerationTask) -> list[str]:
|
||||
"""单次同步组图,全部成功后原子拆分 child。
|
||||
|
||||
绝不在组图 API 失败后退化为 N 次单图请求。
|
||||
"""
|
||||
main_task_id = str(main_task.id)
|
||||
claim = await _claim_image_main_batch(db, main_task_id)
|
||||
if claim.existing_child_ids is not None:
|
||||
await _enqueue_child_downloads(db, claim.existing_child_ids)
|
||||
return claim.existing_child_ids
|
||||
if not claim.acquired or not claim.claim_token or not claim.task_snapshot or not claim.runtime_engine:
|
||||
return []
|
||||
|
||||
generation_count = max(1, int(claim.task_snapshot.generation_count or 1))
|
||||
try:
|
||||
log_operation_event(
|
||||
domain="generation_ai_batch",
|
||||
event_type=ChatGenerationTaskEventType.IMAGE_BATCH_PROVIDER_START.value,
|
||||
event_status="started",
|
||||
source="celery",
|
||||
user_id=claim.task_snapshot.user_id,
|
||||
group_id=main_task_id,
|
||||
task_id=main_task_id,
|
||||
detail={"generation_count": generation_count},
|
||||
)
|
||||
provider_result = await create_image_sync_batch_result_with_engine(
|
||||
claim.task_snapshot,
|
||||
claim.runtime_engine,
|
||||
generation_count=generation_count,
|
||||
)
|
||||
provider_items = _validate_provider_batch(provider_result, generation_count)
|
||||
log_operation_event(
|
||||
domain="generation_ai_batch",
|
||||
event_type=ChatGenerationTaskEventType.IMAGE_BATCH_PROVIDER_SUCCESS.value,
|
||||
event_status="success",
|
||||
source="celery",
|
||||
user_id=claim.task_snapshot.user_id,
|
||||
group_id=main_task_id,
|
||||
task_id=main_task_id,
|
||||
detail={
|
||||
"generation_count": generation_count,
|
||||
"result_count": len(provider_items),
|
||||
"image_tokens": int(provider_result.get("image_tokens") or 0),
|
||||
"single_provider_request": True,
|
||||
"fallback_to_single_requests": False,
|
||||
},
|
||||
)
|
||||
except Exception as exc:
|
||||
message = exc.safe_message if isinstance(exc, ImageProviderError) else str(exc)
|
||||
await _fail_claimed_main(
|
||||
db,
|
||||
main_task_id=main_task_id,
|
||||
claim_token=claim.claim_token,
|
||||
error_message=message or "图片批量生成失败",
|
||||
event_type=ChatGenerationTaskEventType.IMAGE_BATCH_PROVIDER_FAILED,
|
||||
exception=exc,
|
||||
)
|
||||
return []
|
||||
|
||||
try:
|
||||
child_ids = await _split_children(
|
||||
db,
|
||||
main_task_id=main_task_id,
|
||||
claim_token=claim.claim_token,
|
||||
provider_result=provider_result,
|
||||
provider_items=provider_items,
|
||||
)
|
||||
except Exception as exc:
|
||||
await _fail_claimed_main(
|
||||
db,
|
||||
main_task_id=main_task_id,
|
||||
claim_token=claim.claim_token,
|
||||
error_message=f"图片批量结果拆分失败: {exc}",
|
||||
event_type=ChatGenerationTaskEventType.IMAGE_BATCH_SPLIT_FAILED,
|
||||
exception=exc,
|
||||
)
|
||||
return []
|
||||
|
||||
# child 已提交后,下载投递失败不属于图片生成失败,不退款、不重新请求供应商。
|
||||
enqueue_result = await _enqueue_child_downloads(db, child_ids)
|
||||
if enqueue_result["failed"]:
|
||||
log_operation_event(
|
||||
domain="generation_ai_batch",
|
||||
event_type="DOWNLOAD_ENQUEUE_FAILED",
|
||||
event_status="failed",
|
||||
source="celery",
|
||||
user_id=claim.task_snapshot.user_id,
|
||||
group_id=main_task_id,
|
||||
task_id=main_task_id,
|
||||
detail={
|
||||
"failed_child_task_ids": enqueue_result["failed"],
|
||||
"enqueued_child_task_ids": enqueue_result["enqueued"],
|
||||
"provider_regenerated": False,
|
||||
"generation_refunded": False,
|
||||
},
|
||||
)
|
||||
return child_ids
|
||||
@@ -0,0 +1,987 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from datetime import datetime, date
|
||||
from typing import Any
|
||||
|
||||
from fastapi import HTTPException
|
||||
from sqlalchemy import and_, func, select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.models.chat_generation_task import ChatGenerationTask
|
||||
from app.models.generation_record import GenerationRecord
|
||||
from app.models.project import Project
|
||||
from app.models.image_engine import ImageEngine
|
||||
from app.models.user import User
|
||||
from app.models.video_engine import VideoEngine
|
||||
from app.enums.generation_task import CHAT_TOP_LEVEL_MODES, GenerationMode
|
||||
from app.enums.generation_history import (
|
||||
GenerationHistorySourceEnum,
|
||||
get_generation_history_source_label,
|
||||
get_generation_history_task_modes,
|
||||
normalize_generation_history_source,
|
||||
HISTORY_DAY_PAGE_SIZE_MAX,
|
||||
HISTORY_GROUP_ITEM_LIMIT,
|
||||
)
|
||||
from app.schemas.generation_ai import (
|
||||
GenerationAIEngineGroupOut,
|
||||
GenerationAIEngineOptionsOut,
|
||||
GenerationAIImageEngineOptionOut,
|
||||
GenerationAIRecordHistoryItemOut,
|
||||
GenerationAITaskOut,
|
||||
GenerationAIVideoEngineOptionOut,
|
||||
)
|
||||
from app.services.resource_accounting_service import (
|
||||
SOURCE_MODEL_CHAT_TASK,
|
||||
SOURCE_MODEL_GENERATION_RECORD,
|
||||
batch_get_generated_resource_info_map,
|
||||
)
|
||||
from app.services.resource_signed_url_service import build_resource_signed_url
|
||||
from app.services.generation.history_meta_service import (
|
||||
GenerationHistoryMeta,
|
||||
batch_load_generation_history_meta_map,
|
||||
build_empty_history_meta,
|
||||
)
|
||||
from app.services.generation.ai.task_group_service import get_display_status, load_children_map
|
||||
from app.services.generation.ai.engine_service import (
|
||||
image_supported_sizes,
|
||||
normalize_generation_count,
|
||||
parse_json_list,
|
||||
)
|
||||
from app.services.private_portrait.reference_resolver import batch_resolve_private_portrait_reference_display_urls
|
||||
|
||||
def _json(data: Any) -> str | None:
|
||||
if data is None:
|
||||
return None
|
||||
return json.dumps(data, ensure_ascii=False, default=str)
|
||||
|
||||
|
||||
def _parse_json(text: str | None):
|
||||
if not text:
|
||||
return None
|
||||
try:
|
||||
return json.loads(text)
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
|
||||
async def _resolve_task_reference_display_map(
|
||||
db: AsyncSession,
|
||||
tasks: list[ChatGenerationTask],
|
||||
*,
|
||||
user_id: str | None = None,
|
||||
) -> dict[str, list[dict] | None]:
|
||||
return await batch_resolve_private_portrait_reference_display_urls(
|
||||
db,
|
||||
{task.id: _parse_json(task.media_references) for task in tasks},
|
||||
user_id=user_id,
|
||||
)
|
||||
|
||||
|
||||
async def _resolve_generation_record_reference_display_map(
|
||||
db: AsyncSession,
|
||||
records: list[GenerationRecord],
|
||||
*,
|
||||
user_id: str | None = None,
|
||||
) -> dict[str, list[dict] | None]:
|
||||
return await batch_resolve_private_portrait_reference_display_urls(
|
||||
db,
|
||||
{record.id: _parse_json(record.media_references) for record in records},
|
||||
user_id=user_id,
|
||||
)
|
||||
|
||||
|
||||
async def list_generation_ai_engine_options(db: AsyncSession) -> GenerationAIEngineOptionsOut:
|
||||
"""获取当前启用的图片/视频生成引擎,供前端创建任务时选择 engine_id。"""
|
||||
image_result = await db.execute(
|
||||
select(ImageEngine)
|
||||
.where(ImageEngine.is_active == True)
|
||||
.order_by(ImageEngine.priority.desc())
|
||||
)
|
||||
video_result = await db.execute(
|
||||
select(VideoEngine)
|
||||
.where(VideoEngine.is_active == True)
|
||||
.order_by(VideoEngine.priority.desc())
|
||||
)
|
||||
|
||||
image_items = [
|
||||
GenerationAIImageEngineOptionOut(
|
||||
id=engine.id,
|
||||
name=engine.name,
|
||||
provider=engine.provider,
|
||||
model_name=engine.model_name,
|
||||
supported_models=parse_json_list(engine.supported_models, []),
|
||||
supported_sizes=image_supported_sizes(engine),
|
||||
default_size=engine.default_size,
|
||||
priority=engine.priority or 0,
|
||||
max_image_count=engine.max_image_count,
|
||||
multi_generation_enabled=bool(getattr(engine, "multi_generation_enabled", False)),
|
||||
max_generation_count=normalize_generation_count(getattr(engine, "max_generation_count", 1)),
|
||||
multi_image_max_images=int(getattr(engine, "multi_image_max_images", 15) or 15),
|
||||
max_reference_image_count=int(getattr(engine, "max_reference_image_count", 14) or 0),
|
||||
)
|
||||
for engine in image_result.scalars().all()
|
||||
]
|
||||
video_items = [
|
||||
GenerationAIVideoEngineOptionOut(
|
||||
id=engine.id,
|
||||
name=engine.name,
|
||||
provider=engine.provider,
|
||||
model_name=engine.model_name,
|
||||
supported_ratios=parse_json_list(engine.supported_ratios, []),
|
||||
supported_resolutions=parse_json_list(engine.supported_resolutions, []),
|
||||
supported_durations=parse_json_list(engine.supported_durations, []),
|
||||
max_duration=engine.max_duration,
|
||||
priority=engine.priority or 0,
|
||||
max_image_count=engine.max_image_count,
|
||||
max_video_count=engine.max_video_count,
|
||||
max_audio_count=engine.max_audio_count,
|
||||
supports_first_last_frame=engine.supports_first_last_frame,
|
||||
supports_universal_reference=engine.supports_universal_reference,
|
||||
multi_generation_enabled=bool(getattr(engine, "multi_generation_enabled", False)),
|
||||
max_generation_count=normalize_generation_count(getattr(engine, "max_generation_count", 1)),
|
||||
)
|
||||
for engine in video_result.scalars().all()
|
||||
]
|
||||
|
||||
return GenerationAIEngineOptionsOut(
|
||||
engine=GenerationAIEngineGroupOut(image=image_items, video=video_items)
|
||||
)
|
||||
|
||||
|
||||
def _resolve_error_message(error_message: str | None) -> str | None:
|
||||
"""匹配 ARK_ERRORS 字典,将原始错误码转换为友好提示。
|
||||
与 app/api/v1/generation.py 的 _record_to_out 保持一致。
|
||||
|
||||
注意:celery 任务中已调用 extract_error_message 将错误码转为中文提示后存入数据库,
|
||||
所以到达此函数的 message 可能是:
|
||||
1. 已翻译的中文提示(ARK_ERRORS 的 value)→ 直接返回
|
||||
2. 原始错误字符串(含 code='...' 或 JSON 格式)→ 匹配 ARK_ERRORS
|
||||
3. 未知内容 → 返回 "生成失败"
|
||||
"""
|
||||
if not error_message:
|
||||
return error_message
|
||||
|
||||
from app.services.error_codes import ARK_ERRORS
|
||||
|
||||
# 如果已经是 ARK_ERRORS 中已翻译的中文值,直接返回
|
||||
if error_message in ARK_ERRORS.values():
|
||||
return error_message
|
||||
|
||||
import re
|
||||
|
||||
# 匹配以下格式中的错误码:
|
||||
# 1. {'error': {'code': 'XXX', ...}} — str(error_obj) 的 Python dict 形式
|
||||
# 2. {"error": {"code": "XXX", ...}} — JSON 形式
|
||||
# 3. code='XXX' — 旧格式
|
||||
for pattern in [
|
||||
r"'code'\s*:\s*'([^']+)'", # 'code': 'XXX'
|
||||
r'"code"\s*:\s*"([^"]+)"', # "code": "XXX"
|
||||
r"code='([^']+)'", # code='XXX'
|
||||
]:
|
||||
match = re.search(pattern, error_message)
|
||||
if match:
|
||||
code = match.group(1)
|
||||
if code in ARK_ERRORS:
|
||||
return ARK_ERRORS[code]
|
||||
|
||||
# 兜底:按冒号分割,检查第二部分是否是已知错误码
|
||||
parts = error_message.split(":")
|
||||
if len(parts) >= 2 and parts[1].strip() in ARK_ERRORS:
|
||||
return ARK_ERRORS[parts[1].strip()]
|
||||
|
||||
# 没有匹配到已知错误码时,直接返回"生成失败"
|
||||
return "生成失败"
|
||||
|
||||
|
||||
def record_to_out(
|
||||
task: ChatGenerationTask,
|
||||
is_admin: bool = False,
|
||||
generated_resource_id: str | None = None,
|
||||
file_name: str | None = None,
|
||||
history_meta: GenerationHistoryMeta | None = None,
|
||||
media_references: list[dict] | None = None,
|
||||
child_items: list[GenerationAITaskOut] | None = None,
|
||||
) -> GenerationAITaskOut:
|
||||
refs = media_references if media_references is not None else _parse_json(task.media_references)
|
||||
snapshot = engine_snapshot_out(_parse_json(task.engine_snapshot_json))
|
||||
|
||||
source = GenerationHistorySourceEnum.CHAT_TASK
|
||||
try:
|
||||
if task.generation_mode in {
|
||||
GenerationMode.CHATAPI_ASYNC.value,
|
||||
GenerationMode.CHATAPI_MAIN.value,
|
||||
GenerationMode.CHATAPI_CHILD.value,
|
||||
}:
|
||||
source = GenerationHistorySourceEnum.CHAT_TASK
|
||||
else:
|
||||
source = GenerationHistorySourceEnum(str(task.generation_mode or "chat_task"))
|
||||
except ValueError:
|
||||
source = GenerationHistorySourceEnum.CHAT_TASK
|
||||
meta = history_meta or build_empty_history_meta(source)
|
||||
|
||||
is_deleted = task.deleted_at is not None
|
||||
is_main = task.generation_mode == GenerationMode.CHATAPI_MAIN.value
|
||||
hide_resource = is_deleted or is_main
|
||||
|
||||
return GenerationAITaskOut(
|
||||
id=task.id,
|
||||
user_id=task.user_id if is_admin else None,
|
||||
user_name=getattr(task, "username", None) if is_admin else None,
|
||||
project_id=None,
|
||||
generated_resource_id=None if hide_resource else generated_resource_id,
|
||||
file_name=None if hide_resource else file_name,
|
||||
history_source=meta.get("history_source"),
|
||||
history_source_label=meta.get("history_source_label"),
|
||||
module_project_id=meta.get("module_project_id"),
|
||||
module_project_title=meta.get("module_project_title"),
|
||||
module_step_id=meta.get("module_step_id"),
|
||||
module_step_code=meta.get("module_step_code"),
|
||||
hot_opening_project_id=meta.get("hot_opening_project_id"),
|
||||
hot_opening_project_title=meta.get("hot_opening_project_title"),
|
||||
shot_replicate_project_id=meta.get("shot_replicate_project_id"),
|
||||
shot_replicate_project_title=meta.get("shot_replicate_project_title"),
|
||||
shot_task_set_id=meta.get("shot_task_set_id"),
|
||||
shot_segment_id=meta.get("shot_segment_id"),
|
||||
shot_segment_index=meta.get("shot_segment_index"),
|
||||
shot_segment_label=meta.get("shot_segment_label"),
|
||||
gen_type=task.gen_type,
|
||||
generation_mode=task.generation_mode,
|
||||
parent_task_id=task.parent_task_id,
|
||||
generation_count=max(1, min(5, int(task.generation_count or 1))),
|
||||
generation_index=task.generation_index,
|
||||
display_status=get_display_status(task),
|
||||
pipeline_stage=task.pipeline_stage,
|
||||
status=task.status,
|
||||
original_prompt=task.original_prompt,
|
||||
optimized_prompt=task.optimized_prompt,
|
||||
duration=task.duration,
|
||||
aspect_ratio=task.aspect_ratio,
|
||||
resolution=task.resolution,
|
||||
image_size=task.image_size,
|
||||
image_proportion=task.image_proportion,
|
||||
image_px=task.image_px,
|
||||
media_references=refs,
|
||||
provider_task_id=task.provider_task_id,
|
||||
seedance_task_id=task.seedance_task_id,
|
||||
image_url="" if hide_resource else (build_resource_signed_url(task.image_url) if task.image_url else ""),
|
||||
video_url="" if hide_resource else (build_resource_signed_url(task.video_url) if task.video_url else ""),
|
||||
video_cover_url="" if hide_resource else (build_resource_signed_url(task.video_cover_url) if task.video_cover_url else ""),
|
||||
engine_id=task.engine_id,
|
||||
engine_snapshot=snapshot,
|
||||
credits_cost=task.credits_cost or 0.0,
|
||||
text_credits_cost=task.text_credits_cost or 0.0,
|
||||
text_tokens_used=task.text_tokens_used or 0,
|
||||
image_tokens_used=task.image_tokens_used or 0,
|
||||
video_tokens_used=task.video_tokens_used or 0,
|
||||
retry_count=task.retry_count or 0,
|
||||
poll_count=task.poll_count or 0,
|
||||
error_message=task.error_message if is_main else _resolve_error_message(task.error_message),
|
||||
created_at=task.created_at,
|
||||
generated_at=task.generated_at,
|
||||
child_items=child_items or [],
|
||||
)
|
||||
|
||||
def engine_snapshot_out(snapshot: dict) -> dict:
|
||||
"""从完整引擎快照中过滤前端允许展示的字段。"""
|
||||
if not snapshot:
|
||||
return {}
|
||||
keys = (
|
||||
"engine_type", "id", "name", "provider", "model_name",
|
||||
"supported_models", "default_size", "selected_size",
|
||||
"selected_proportion", "selected_px", "supported_ratios",
|
||||
"supported_resolutions", "supported_durations", "max_duration",
|
||||
"max_audio_count", "selected_ratio", "selected_resolution",
|
||||
"selected_duration", "generation_count", "multi_generation_enabled",
|
||||
"max_generation_count", "multi_image_max_images", "max_reference_image_count", "output_format",
|
||||
)
|
||||
result = {key: snapshot.get(key) for key in keys if key in snapshot}
|
||||
result.setdefault("generation_count", 1)
|
||||
return result
|
||||
|
||||
|
||||
async def build_task_out_list(
|
||||
db: AsyncSession,
|
||||
tasks: list[ChatGenerationTask],
|
||||
*,
|
||||
is_admin: bool = False,
|
||||
viewer_user_id: str | None = None,
|
||||
) -> list[GenerationAITaskOut]:
|
||||
"""批量回填主任务子项、资源账本和参考素材,避免列表 N+1。"""
|
||||
if not tasks:
|
||||
return []
|
||||
parent_ids = [
|
||||
task.id for task in tasks
|
||||
if task.generation_mode == GenerationMode.CHATAPI_MAIN.value
|
||||
]
|
||||
children_map = await load_children_map(db, parent_ids, include_deleted=True)
|
||||
children = [child for items in children_map.values() for child in items]
|
||||
resource_task_ids = [
|
||||
task.id for task in [*tasks, *children]
|
||||
if task.generation_mode != GenerationMode.CHATAPI_MAIN.value and task.deleted_at is None
|
||||
]
|
||||
resource_info_map = await batch_get_generated_resource_info_map(
|
||||
db,
|
||||
source_model=SOURCE_MODEL_CHAT_TASK,
|
||||
source_ids=resource_task_ids,
|
||||
)
|
||||
reference_display_map = await _resolve_task_reference_display_map(
|
||||
db,
|
||||
tasks,
|
||||
user_id=viewer_user_id,
|
||||
)
|
||||
|
||||
output: list[GenerationAITaskOut] = []
|
||||
for task in tasks:
|
||||
refs = reference_display_map.get(task.id)
|
||||
child_out: list[GenerationAITaskOut] = []
|
||||
for child in children_map.get(task.id, []):
|
||||
if is_admin:
|
||||
child.username = getattr(task, "username", None)
|
||||
resource = resource_info_map.get(child.id, {})
|
||||
child_out.append(
|
||||
record_to_out(
|
||||
child,
|
||||
is_admin=is_admin,
|
||||
generated_resource_id=resource.get("resource_id"),
|
||||
file_name=resource.get("file_name"),
|
||||
media_references=refs,
|
||||
)
|
||||
)
|
||||
resource = resource_info_map.get(task.id, {})
|
||||
output.append(
|
||||
record_to_out(
|
||||
task,
|
||||
is_admin=is_admin,
|
||||
generated_resource_id=resource.get("resource_id"),
|
||||
file_name=resource.get("file_name"),
|
||||
media_references=refs,
|
||||
child_items=child_out,
|
||||
)
|
||||
)
|
||||
return output
|
||||
|
||||
async def list_async_generation_tasks(
|
||||
db: AsyncSession,
|
||||
user_id: str | None,
|
||||
user_name: str | None,
|
||||
gen_type: str | None,
|
||||
status: str | None,
|
||||
page: int,
|
||||
page_size: int,
|
||||
is_admin: bool = False,
|
||||
engine_id: str | None = None,
|
||||
created_start: datetime | None = None,
|
||||
created_end: datetime | None = None,
|
||||
):
|
||||
if is_admin:
|
||||
query = (
|
||||
select(ChatGenerationTask, User.username)
|
||||
.join(User, ChatGenerationTask.user_id == User.id)
|
||||
)
|
||||
|
||||
if user_name:
|
||||
query = query.where(User.username.like(f"%{user_name}%"))
|
||||
else:
|
||||
query = select(ChatGenerationTask)
|
||||
|
||||
query = query.where(
|
||||
ChatGenerationTask.generation_mode.in_(list(CHAT_TOP_LEVEL_MODES)),
|
||||
ChatGenerationTask.deleted_at.is_(None),
|
||||
)
|
||||
|
||||
if user_id:
|
||||
query = query.where(ChatGenerationTask.user_id == user_id)
|
||||
|
||||
if gen_type:
|
||||
query = query.where(ChatGenerationTask.gen_type == gen_type)
|
||||
|
||||
if status:
|
||||
query = query.where(ChatGenerationTask.status == status)
|
||||
|
||||
if engine_id:
|
||||
query = query.where(ChatGenerationTask.engine_id == engine_id)
|
||||
|
||||
if created_start is not None or created_end is not None:
|
||||
range_filters = []
|
||||
if created_start is not None:
|
||||
range_filters.append(ChatGenerationTask.created_at >= created_start)
|
||||
if created_end is not None:
|
||||
range_filters.append(ChatGenerationTask.created_at <= created_end)
|
||||
if range_filters:
|
||||
query = query.where(and_(*range_filters))
|
||||
|
||||
count_query = select(func.count()).select_from(query.subquery())
|
||||
total = (await db.execute(count_query)).scalar_one()
|
||||
|
||||
result = await db.execute(
|
||||
query.order_by(ChatGenerationTask.created_at.desc(), ChatGenerationTask.id.desc())
|
||||
.offset((page - 1) * page_size)
|
||||
.limit(page_size)
|
||||
)
|
||||
|
||||
if is_admin:
|
||||
tasks = []
|
||||
for task, username in result.all():
|
||||
task.username = username
|
||||
tasks.append(task)
|
||||
return total, tasks
|
||||
|
||||
return total, list(result.scalars().all())
|
||||
|
||||
def _normalize_history_gen_type(gen_type: str | None) -> str:
|
||||
value = (gen_type or "").lower().strip()
|
||||
if value not in ("image", "video"):
|
||||
raise HTTPException(status_code=400, detail="gen_type 仅支持 image 或 video")
|
||||
return value
|
||||
|
||||
|
||||
|
||||
def _normalize_history_source(history_source: str | None) -> GenerationHistorySourceEnum:
|
||||
"""Normalize history_source query param.
|
||||
|
||||
默认保持原来的 chat_generation_tasks / chatapi_async 历史;
|
||||
显式传 hot_opening_replicate 或 shot_replicate 时查询对应模块素材;
|
||||
显式传 generation_record 时查询旧 generation_records 历史。
|
||||
"""
|
||||
try:
|
||||
return normalize_generation_history_source(history_source)
|
||||
except ValueError:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail="history_source 仅支持 chat_task、generation_record、hot_opening_replicate、shot_replicate",
|
||||
)
|
||||
|
||||
|
||||
def _history_day_to_str(value) -> str:
|
||||
if isinstance(value, datetime):
|
||||
return value.date().strftime("%Y-%m-%d")
|
||||
if isinstance(value, date):
|
||||
return value.strftime("%Y-%m-%d")
|
||||
return str(value)[:10]
|
||||
|
||||
|
||||
def _parse_history_date(value: str) -> date:
|
||||
try:
|
||||
return datetime.strptime(value, "%Y-%m-%d").date()
|
||||
except ValueError:
|
||||
raise HTTPException(status_code=400, detail="generated_date 格式必须是 YYYY-MM-DD")
|
||||
|
||||
|
||||
def _history_base_filters(user_id: str, gen_type: str, source: GenerationHistorySourceEnum):
|
||||
task_modes = get_generation_history_task_modes(source)
|
||||
if not task_modes:
|
||||
raise HTTPException(status_code=400, detail="history_source 不支持查询 ChatGenerationTask 历史")
|
||||
return [
|
||||
ChatGenerationTask.user_id == user_id,
|
||||
ChatGenerationTask.generation_mode.in_([mode.value for mode in task_modes]),
|
||||
ChatGenerationTask.deleted_at.is_(None),
|
||||
ChatGenerationTask.status == "completed",
|
||||
ChatGenerationTask.gen_type == gen_type,
|
||||
ChatGenerationTask.generated_at.is_not(None),
|
||||
]
|
||||
|
||||
|
||||
def _generation_record_history_base_filters(user_id: str, gen_type: str):
|
||||
return [
|
||||
GenerationRecord.user_id == user_id,
|
||||
GenerationRecord.deleted_at.is_(None),
|
||||
GenerationRecord.status == "completed",
|
||||
GenerationRecord.gen_type == gen_type,
|
||||
GenerationRecord.generated_at.is_not(None),
|
||||
]
|
||||
|
||||
|
||||
def generation_record_to_history_out(
|
||||
record: GenerationRecord,
|
||||
project_name: str | None = None,
|
||||
generated_resource_id: str | None = None,
|
||||
file_name: str | None = None,
|
||||
media_references: list[dict] | None = None,
|
||||
) -> GenerationAIRecordHistoryItemOut:
|
||||
refs = media_references if media_references is not None else _parse_json(record.media_references)
|
||||
return GenerationAIRecordHistoryItemOut(
|
||||
id=record.id,
|
||||
source_type="generation_record",
|
||||
project_id=record.project_id,
|
||||
project_name=project_name,
|
||||
generated_resource_id=generated_resource_id,
|
||||
file_name=file_name,
|
||||
history_source=GenerationHistorySourceEnum.GENERATION_RECORD.value,
|
||||
history_source_label=get_generation_history_source_label(GenerationHistorySourceEnum.GENERATION_RECORD),
|
||||
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,
|
||||
gen_type=record.gen_type,
|
||||
generation_mode="generation_record",
|
||||
pipeline_stage=None,
|
||||
status=record.status,
|
||||
original_prompt=record.original_prompt,
|
||||
# optimized_prompt=None,
|
||||
duration=record.duration,
|
||||
aspect_ratio=record.aspect_ratio,
|
||||
resolution=record.resolution,
|
||||
image_size=record.image_size,
|
||||
image_proportion=record.image_proportion,
|
||||
image_px=record.image_px,
|
||||
references=refs,
|
||||
media_references=refs,
|
||||
provider_task_id=record.seedance_task_id,
|
||||
seedance_task_id=record.seedance_task_id,
|
||||
remote_result_url=None,
|
||||
image_url=build_resource_signed_url(record.image_url) if record.image_url else '',
|
||||
video_url=build_resource_signed_url(record.video_url) if record.video_url else '',
|
||||
video_cover_url=build_resource_signed_url(record.video_cover_url) if record.video_cover_url else '',
|
||||
engine_id=None,
|
||||
engine_snapshot=None,
|
||||
credits_cost=record.credits_cost or 0.0,
|
||||
text_credits_cost=record.text_credits_cost or 0.0,
|
||||
text_tokens_used=record.text_tokens_used or 0,
|
||||
image_tokens_used=record.image_tokens_used or 0,
|
||||
video_tokens_used=record.video_tokens_used or 0,
|
||||
retry_count=0,
|
||||
poll_count=0,
|
||||
error_message=_resolve_error_message(record.error_message),
|
||||
created_at=record.created_at,
|
||||
generated_at=record.generated_at,
|
||||
)
|
||||
|
||||
|
||||
async def list_generation_record_history_grouped_days(
|
||||
db: AsyncSession,
|
||||
user_id: str,
|
||||
gen_type: str,
|
||||
page: int,
|
||||
page_size: int,
|
||||
keyword: str | None = None,
|
||||
):
|
||||
"""
|
||||
按生成日期倒序返回旧 generation_records 历史记录分组。
|
||||
|
||||
- 每页最多返回 10 个生成日期
|
||||
- 每个日期分组内最多返回倒序前 10 条旧记录
|
||||
- 只返回 completed 成功记录
|
||||
"""
|
||||
gen_type = _normalize_history_gen_type(gen_type)
|
||||
page = max(page, 1)
|
||||
page_size = min(max(page_size, 1), HISTORY_DAY_PAGE_SIZE_MAX)
|
||||
|
||||
filters = _generation_record_history_base_filters(user_id, gen_type)
|
||||
if keyword and keyword.strip():
|
||||
filters.append(GenerationRecord.original_prompt.ilike(f"%{keyword.strip()}%"))
|
||||
day_expr = func.date(GenerationRecord.generated_at).label("generated_date")
|
||||
|
||||
days_subquery = (
|
||||
select(day_expr)
|
||||
.where(*filters)
|
||||
.group_by(day_expr)
|
||||
.subquery()
|
||||
)
|
||||
|
||||
total_days = (
|
||||
await db.execute(select(func.count()).select_from(days_subquery))
|
||||
).scalar_one()
|
||||
|
||||
day_rows_result = await db.execute(
|
||||
select(
|
||||
day_expr,
|
||||
func.count(GenerationRecord.id).label("total"),
|
||||
)
|
||||
.where(*filters)
|
||||
.group_by(day_expr)
|
||||
.order_by(day_expr.desc())
|
||||
.offset((page - 1) * page_size)
|
||||
.limit(page_size)
|
||||
)
|
||||
day_rows = day_rows_result.all()
|
||||
|
||||
if not day_rows:
|
||||
return {
|
||||
"total_days": int(total_days or 0),
|
||||
"page": page,
|
||||
"page_size": page_size,
|
||||
"groups": [],
|
||||
}
|
||||
|
||||
day_list = [row[0] for row in day_rows]
|
||||
day_total_map = {row[0]: row[1] for row in day_rows}
|
||||
|
||||
min_day = min(day_list)
|
||||
max_day = max(day_list)
|
||||
|
||||
all_rows_result = await db.execute(
|
||||
select(GenerationRecord, Project.name.label("project_name"))
|
||||
.outerjoin(Project, (GenerationRecord.project_id == Project.id) & (Project.deleted_at.is_(None)))
|
||||
.where(
|
||||
*filters,
|
||||
func.date(GenerationRecord.generated_at).between(min_day, max_day),
|
||||
)
|
||||
.order_by(GenerationRecord.generated_at.desc(), GenerationRecord.created_at.desc())
|
||||
)
|
||||
all_rows = all_rows_result.all()
|
||||
|
||||
rows_by_day: dict[str, list[tuple[GenerationRecord, str]]] = {}
|
||||
for record, project_name in all_rows:
|
||||
if record.generated_at is None:
|
||||
continue
|
||||
day_key = record.generated_at.date()
|
||||
if day_key not in rows_by_day:
|
||||
rows_by_day[day_key] = []
|
||||
rows_by_day[day_key].append((record, project_name))
|
||||
|
||||
raw_groups = []
|
||||
all_record_ids: list[str] = []
|
||||
for generated_day in day_list:
|
||||
day_rows_list = rows_by_day.get(generated_day, [])
|
||||
limited_rows = day_rows_list[:HISTORY_GROUP_ITEM_LIMIT]
|
||||
day_total = day_total_map.get(generated_day, 0)
|
||||
raw_groups.append((generated_day, day_total, limited_rows))
|
||||
all_record_ids.extend(record.id for record, _project_name in limited_rows)
|
||||
|
||||
resource_info_map = await batch_get_generated_resource_info_map(
|
||||
db,
|
||||
source_model=SOURCE_MODEL_GENERATION_RECORD,
|
||||
source_ids=all_record_ids,
|
||||
resource_type=gen_type,
|
||||
)
|
||||
all_records = [record for _generated_day, _day_total, rows in raw_groups for record, _project_name in rows]
|
||||
reference_display_map = await _resolve_generation_record_reference_display_map(db, all_records, user_id=user_id)
|
||||
|
||||
groups = [
|
||||
{
|
||||
"generated_date": _history_day_to_str(generated_day),
|
||||
"total": int(day_total or 0),
|
||||
"items": [
|
||||
generation_record_to_history_out(
|
||||
record,
|
||||
project_name,
|
||||
generated_resource_id=resource_info_map.get(record.id, {}).get("resource_id"),
|
||||
file_name=resource_info_map.get(record.id, {}).get("file_name"),
|
||||
media_references=reference_display_map.get(record.id),
|
||||
)
|
||||
for record, project_name in rows
|
||||
],
|
||||
}
|
||||
for generated_day, day_total, rows in raw_groups
|
||||
]
|
||||
|
||||
return {
|
||||
"total_days": int(total_days or 0),
|
||||
"page": page,
|
||||
"page_size": page_size,
|
||||
"groups": groups,
|
||||
}
|
||||
|
||||
|
||||
async def list_generation_record_history_day_items(
|
||||
db: AsyncSession,
|
||||
user_id: str,
|
||||
gen_type: str,
|
||||
generated_date: str,
|
||||
page: int,
|
||||
page_size: int,
|
||||
keyword: str | None = None,
|
||||
):
|
||||
"""
|
||||
获取旧 generation_records 指定生成日期下的历史记录分页。
|
||||
"""
|
||||
gen_type = _normalize_history_gen_type(gen_type)
|
||||
target_day = _parse_history_date(generated_date)
|
||||
page = max(page, 1)
|
||||
page_size = min(max(page_size, 1), 100)
|
||||
|
||||
filters = _generation_record_history_base_filters(user_id, gen_type)
|
||||
if keyword and keyword.strip():
|
||||
filters.append(GenerationRecord.original_prompt.ilike(f"%{keyword.strip()}%"))
|
||||
day_expr = func.date(GenerationRecord.generated_at)
|
||||
|
||||
total = (
|
||||
await db.execute(
|
||||
select(func.count(GenerationRecord.id)).where(
|
||||
*filters,
|
||||
day_expr == target_day,
|
||||
)
|
||||
)
|
||||
).scalar_one()
|
||||
|
||||
result = await db.execute(
|
||||
select(GenerationRecord, Project.name.label("project_name"))
|
||||
.outerjoin(Project, (GenerationRecord.project_id == Project.id) & (Project.deleted_at.is_(None)))
|
||||
.where(
|
||||
*filters,
|
||||
day_expr == target_day,
|
||||
)
|
||||
.order_by(GenerationRecord.generated_at.desc(), GenerationRecord.created_at.desc())
|
||||
.offset((page - 1) * page_size)
|
||||
.limit(page_size)
|
||||
)
|
||||
|
||||
rows = result.all()
|
||||
resource_info_map = await batch_get_generated_resource_info_map(
|
||||
db,
|
||||
source_model=SOURCE_MODEL_GENERATION_RECORD,
|
||||
source_ids=[record.id for record, _project_name in rows],
|
||||
resource_type=gen_type,
|
||||
)
|
||||
reference_display_map = await _resolve_generation_record_reference_display_map(db, [record for record, _project_name in rows], user_id=user_id)
|
||||
|
||||
return {
|
||||
"generated_date": target_day.strftime("%Y-%m-%d"),
|
||||
"total": int(total or 0),
|
||||
"page": page,
|
||||
"page_size": page_size,
|
||||
"items": [
|
||||
generation_record_to_history_out(
|
||||
record,
|
||||
project_name,
|
||||
generated_resource_id=resource_info_map.get(record.id, {}).get("resource_id"),
|
||||
file_name=resource_info_map.get(record.id, {}).get("file_name"),
|
||||
media_references=reference_display_map.get(record.id),
|
||||
)
|
||||
for record, project_name in rows
|
||||
],
|
||||
}
|
||||
|
||||
|
||||
async def list_generation_history_grouped_days(
|
||||
db: AsyncSession,
|
||||
user_id: str,
|
||||
gen_type: str,
|
||||
page: int,
|
||||
page_size: int,
|
||||
history_source: str | None = None,
|
||||
keyword: str | None = None,
|
||||
):
|
||||
"""
|
||||
按生成日期倒序返回历史记录分组。
|
||||
|
||||
- 每页最多返回 10 个生成日期
|
||||
- 每个日期分组内最多返回倒序前 10 条任务
|
||||
- 只返回 completed 成功任务
|
||||
- history_source 支持 chat_task / generation_record / hot_opening_replicate / shot_replicate
|
||||
"""
|
||||
source = _normalize_history_source(history_source)
|
||||
if source == GenerationHistorySourceEnum.GENERATION_RECORD:
|
||||
return await list_generation_record_history_grouped_days(
|
||||
db=db,
|
||||
user_id=user_id,
|
||||
gen_type=gen_type,
|
||||
page=page,
|
||||
page_size=page_size,
|
||||
keyword=keyword,
|
||||
)
|
||||
|
||||
gen_type = _normalize_history_gen_type(gen_type)
|
||||
page = max(page, 1)
|
||||
page_size = min(max(page_size, 1), HISTORY_DAY_PAGE_SIZE_MAX)
|
||||
|
||||
filters = _history_base_filters(user_id, gen_type, source)
|
||||
if keyword and keyword.strip():
|
||||
filters.append(ChatGenerationTask.original_prompt.ilike(f"%{keyword.strip()}%"))
|
||||
day_expr = func.date(ChatGenerationTask.generated_at).label("generated_date")
|
||||
|
||||
days_subquery = (
|
||||
select(day_expr)
|
||||
.where(*filters)
|
||||
.group_by(day_expr)
|
||||
.subquery()
|
||||
)
|
||||
|
||||
total_days = (
|
||||
await db.execute(select(func.count()).select_from(days_subquery))
|
||||
).scalar_one()
|
||||
|
||||
day_rows_result = await db.execute(
|
||||
select(
|
||||
day_expr,
|
||||
func.count(ChatGenerationTask.id).label("total"),
|
||||
)
|
||||
.where(*filters)
|
||||
.group_by(day_expr)
|
||||
.order_by(day_expr.desc())
|
||||
.offset((page - 1) * page_size)
|
||||
.limit(page_size)
|
||||
)
|
||||
day_rows = day_rows_result.all()
|
||||
|
||||
if not day_rows:
|
||||
return {
|
||||
"total_days": int(total_days or 0),
|
||||
"page": page,
|
||||
"page_size": page_size,
|
||||
"groups": [],
|
||||
}
|
||||
|
||||
day_list = [row[0] for row in day_rows]
|
||||
day_total_map = {row[0]: row[1] for row in day_rows}
|
||||
|
||||
min_day = min(day_list)
|
||||
max_day = max(day_list)
|
||||
|
||||
all_tasks_result = await db.execute(
|
||||
select(ChatGenerationTask)
|
||||
.where(
|
||||
*filters,
|
||||
func.date(ChatGenerationTask.generated_at).between(min_day, max_day),
|
||||
)
|
||||
.order_by(ChatGenerationTask.generated_at.desc(), ChatGenerationTask.created_at.desc())
|
||||
)
|
||||
all_tasks = list(all_tasks_result.scalars().all())
|
||||
|
||||
tasks_by_day: dict[str, list[ChatGenerationTask]] = {}
|
||||
for task in all_tasks:
|
||||
if task.generated_at is None:
|
||||
continue
|
||||
day_key = task.generated_at.date()
|
||||
if day_key not in tasks_by_day:
|
||||
tasks_by_day[day_key] = []
|
||||
tasks_by_day[day_key].append(task)
|
||||
|
||||
raw_groups = []
|
||||
all_task_ids: list[str] = []
|
||||
for generated_day in day_list:
|
||||
day_tasks = tasks_by_day.get(generated_day, [])
|
||||
limited_tasks = day_tasks[:HISTORY_GROUP_ITEM_LIMIT]
|
||||
day_total = day_total_map.get(generated_day, 0)
|
||||
raw_groups.append((generated_day, day_total, limited_tasks))
|
||||
all_task_ids.extend(task.id for task in limited_tasks)
|
||||
|
||||
resource_info_map = await batch_get_generated_resource_info_map(
|
||||
db,
|
||||
source_model=SOURCE_MODEL_CHAT_TASK,
|
||||
source_ids=all_task_ids,
|
||||
resource_type=gen_type,
|
||||
)
|
||||
history_meta_map = await batch_load_generation_history_meta_map(
|
||||
db,
|
||||
source=source,
|
||||
chat_task_ids=all_task_ids,
|
||||
)
|
||||
all_tasks_for_refs = [task for _generated_day, _day_total, tasks in raw_groups for task in tasks]
|
||||
reference_display_map = await _resolve_task_reference_display_map(db, all_tasks_for_refs, user_id=user_id)
|
||||
|
||||
groups = [
|
||||
{
|
||||
"generated_date": _history_day_to_str(generated_day),
|
||||
"total": int(day_total or 0),
|
||||
"items": [
|
||||
record_to_out(
|
||||
task,
|
||||
generated_resource_id=resource_info_map.get(task.id, {}).get("resource_id"),
|
||||
file_name=resource_info_map.get(task.id, {}).get("file_name"),
|
||||
history_meta=history_meta_map.get(task.id),
|
||||
media_references=reference_display_map.get(task.id),
|
||||
)
|
||||
for task in tasks
|
||||
],
|
||||
}
|
||||
for generated_day, day_total, tasks in raw_groups
|
||||
]
|
||||
|
||||
return {
|
||||
"total_days": int(total_days or 0),
|
||||
"page": page,
|
||||
"page_size": page_size,
|
||||
"groups": groups,
|
||||
}
|
||||
|
||||
|
||||
async def list_generation_history_day_items(
|
||||
db: AsyncSession,
|
||||
user_id: str,
|
||||
gen_type: str,
|
||||
generated_date: str,
|
||||
page: int,
|
||||
page_size: int,
|
||||
history_source: str | None = None,
|
||||
keyword: str | None = None,
|
||||
):
|
||||
"""
|
||||
获取指定生成日期下的历史记录分页。
|
||||
|
||||
用于前端点击某一天后,继续加载该日期下的第 2 页、第 3 页数据。
|
||||
history_source 支持 chat_task / generation_record / hot_opening_replicate / shot_replicate。
|
||||
"""
|
||||
source = _normalize_history_source(history_source)
|
||||
if source == GenerationHistorySourceEnum.GENERATION_RECORD:
|
||||
return await list_generation_record_history_day_items(
|
||||
db=db,
|
||||
user_id=user_id,
|
||||
gen_type=gen_type,
|
||||
generated_date=generated_date,
|
||||
page=page,
|
||||
page_size=page_size,
|
||||
keyword=keyword,
|
||||
)
|
||||
|
||||
gen_type = _normalize_history_gen_type(gen_type)
|
||||
target_day = _parse_history_date(generated_date)
|
||||
page = max(page, 1)
|
||||
page_size = min(max(page_size, 1), 100)
|
||||
|
||||
filters = _history_base_filters(user_id, gen_type, source)
|
||||
if keyword and keyword.strip():
|
||||
filters.append(ChatGenerationTask.original_prompt.ilike(f"%{keyword.strip()}%"))
|
||||
day_expr = func.date(ChatGenerationTask.generated_at)
|
||||
|
||||
total = (
|
||||
await db.execute(
|
||||
select(func.count(ChatGenerationTask.id)).where(
|
||||
*filters,
|
||||
day_expr == target_day,
|
||||
)
|
||||
)
|
||||
).scalar_one()
|
||||
|
||||
result = await db.execute(
|
||||
select(ChatGenerationTask)
|
||||
.where(
|
||||
*filters,
|
||||
day_expr == target_day,
|
||||
)
|
||||
.order_by(ChatGenerationTask.generated_at.desc(), ChatGenerationTask.created_at.desc())
|
||||
.offset((page - 1) * page_size)
|
||||
.limit(page_size)
|
||||
)
|
||||
|
||||
tasks = list(result.scalars().all())
|
||||
task_ids = [task.id for task in tasks]
|
||||
resource_info_map = await batch_get_generated_resource_info_map(
|
||||
db,
|
||||
source_model=SOURCE_MODEL_CHAT_TASK,
|
||||
source_ids=task_ids,
|
||||
resource_type=gen_type,
|
||||
)
|
||||
history_meta_map = await batch_load_generation_history_meta_map(
|
||||
db,
|
||||
source=source,
|
||||
chat_task_ids=task_ids,
|
||||
)
|
||||
reference_display_map = await _resolve_task_reference_display_map(db, tasks, user_id=user_id)
|
||||
|
||||
return {
|
||||
"generated_date": target_day.strftime("%Y-%m-%d"),
|
||||
"total": int(total or 0),
|
||||
"page": page,
|
||||
"page_size": page_size,
|
||||
"items": [
|
||||
record_to_out(
|
||||
task,
|
||||
generated_resource_id=resource_info_map.get(task.id, {}).get("resource_id"),
|
||||
file_name=resource_info_map.get(task.id, {}).get("file_name"),
|
||||
history_meta=history_meta_map.get(task.id),
|
||||
media_references=reference_display_map.get(task.id),
|
||||
)
|
||||
for task in tasks
|
||||
],
|
||||
}
|
||||
|
||||
@@ -0,0 +1,587 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Any
|
||||
|
||||
from fastapi import HTTPException
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.config import settings
|
||||
from app.enums.audio_reference import (
|
||||
AUDIO_MAX_COUNT_LIMIT,
|
||||
AUDIO_MAX_DURATION_SECONDS,
|
||||
AUDIO_MAX_TOTAL_DURATION_SECONDS,
|
||||
AUDIO_MIN_DURATION_SECONDS,
|
||||
)
|
||||
from app.enums.generation_task import CHAT_TOP_LEVEL_MODES, GenerationMode, GenerationType
|
||||
from app.models.chat_generation_task import ChatGenerationTask
|
||||
from app.models.user import User
|
||||
from app.schemas.generation_ai import GenerationAITaskCreate
|
||||
from app.services.generation.ai.engine_service import (
|
||||
IMAGE_DEFAULT_PROPORTION,
|
||||
IMAGE_DEFAULT_PX,
|
||||
IMAGE_DEFAULT_SIZE,
|
||||
VIDEO_DEFAULT_DURATION,
|
||||
VIDEO_DEFAULT_RATIO,
|
||||
VIDEO_DEFAULT_RESOLUTION,
|
||||
build_image_snapshot,
|
||||
build_video_snapshot,
|
||||
get_image_engine,
|
||||
get_video_engine,
|
||||
image_supported_sizes,
|
||||
normalize_generation_count,
|
||||
normalize_px,
|
||||
parse_json_list,
|
||||
)
|
||||
from app.services.generation.billing_service import OWNER_CHAT_GENERATION_TASK, charge_generation_media_by_params
|
||||
from app.services.operation_log_service import log_operation_event
|
||||
from app.services.private_portrait.reference_resolver import resolve_private_portrait_references
|
||||
from app.services.resource_capacity_service import assert_user_resource_capacity_available
|
||||
from app.utils.id_gen import generate_id
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class GenerationTaskCreateResult:
|
||||
top_level_task_id: str
|
||||
enqueue_task_ids: list[str] = field(default_factory=list)
|
||||
child_task_ids: list[str] = field(default_factory=list)
|
||||
generation_count: int = 1
|
||||
gen_type: str = GenerationType.IMAGE.value
|
||||
created: bool = True
|
||||
|
||||
|
||||
def _json(data: Any) -> str | None:
|
||||
if data is None:
|
||||
return None
|
||||
return json.dumps(data, ensure_ascii=False, default=str)
|
||||
|
||||
|
||||
async def find_existing_top_level_task(
|
||||
db: AsyncSession,
|
||||
*,
|
||||
user_id: str,
|
||||
idempotency_key: str | None,
|
||||
) -> ChatGenerationTask | None:
|
||||
if not idempotency_key:
|
||||
return None
|
||||
result = await db.execute(
|
||||
select(ChatGenerationTask)
|
||||
.where(
|
||||
ChatGenerationTask.user_id == user_id,
|
||||
ChatGenerationTask.idempotency_key == idempotency_key,
|
||||
ChatGenerationTask.generation_mode.in_(list(CHAT_TOP_LEVEL_MODES)),
|
||||
ChatGenerationTask.deleted_at.is_(None),
|
||||
)
|
||||
.order_by(ChatGenerationTask.created_at.desc())
|
||||
.limit(1)
|
||||
)
|
||||
return result.scalar_one_or_none()
|
||||
|
||||
|
||||
def _validate_video_references(refs: list[dict], *, max_audio_count: int) -> float:
|
||||
input_video_duration = 0.0
|
||||
for ref in refs:
|
||||
if (ref.get("type") or "").lower() != GenerationType.VIDEO.value:
|
||||
continue
|
||||
try:
|
||||
ref_duration = float(ref.get("duration") or 0)
|
||||
except (TypeError, ValueError):
|
||||
ref_duration = 0.0
|
||||
if ref_duration < 2:
|
||||
raise HTTPException(status_code=400, detail="视频素材最短不能少于 2 秒")
|
||||
input_video_duration += ref_duration
|
||||
if input_video_duration > 15:
|
||||
raise HTTPException(status_code=400, detail=f"所有视频素材总时长不能超过 15 秒,当前 {input_video_duration:.1f} 秒")
|
||||
|
||||
audio_refs = [ref for ref in refs if (ref.get("type") or "").lower() == "audio"]
|
||||
if audio_refs:
|
||||
allowed_count = min(AUDIO_MAX_COUNT_LIMIT, max(0, int(max_audio_count or 0)))
|
||||
if allowed_count <= 0:
|
||||
raise HTTPException(status_code=400, detail="当前视频引擎不支持音频参考素材")
|
||||
if len(audio_refs) > allowed_count:
|
||||
raise HTTPException(status_code=400, detail=f"参考音频最多可传 {allowed_count} 段,当前 {len(audio_refs)} 段")
|
||||
|
||||
input_audio_duration = 0.0
|
||||
for ref in audio_refs:
|
||||
try:
|
||||
ref_duration = float(ref.get("duration") or 0)
|
||||
except (TypeError, ValueError):
|
||||
ref_duration = 0.0
|
||||
if ref_duration < AUDIO_MIN_DURATION_SECONDS or ref_duration > AUDIO_MAX_DURATION_SECONDS:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"单段参考音频时长必须在 {AUDIO_MIN_DURATION_SECONDS}-{AUDIO_MAX_DURATION_SECONDS} 秒之间",
|
||||
)
|
||||
input_audio_duration += ref_duration
|
||||
if input_audio_duration > AUDIO_MAX_TOTAL_DURATION_SECONDS:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"所有参考音频总时长不能超过 {AUDIO_MAX_TOTAL_DURATION_SECONDS} 秒,当前 {input_audio_duration:.1f} 秒",
|
||||
)
|
||||
return input_video_duration
|
||||
|
||||
|
||||
def _base_task_kwargs(
|
||||
*,
|
||||
task_id: str,
|
||||
user_id: str,
|
||||
req: GenerationAITaskCreate,
|
||||
gen_type: str,
|
||||
generation_mode: str,
|
||||
generation_count: int,
|
||||
engine_id: str,
|
||||
engine_snapshot_json: str,
|
||||
media_references_json: str | None,
|
||||
deadline_at: datetime,
|
||||
parent_task_id: str | None = None,
|
||||
generation_index: int | None = None,
|
||||
credits_cost: float = 0.0,
|
||||
idempotency_key: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
return {
|
||||
"id": task_id,
|
||||
"user_id": user_id,
|
||||
"original_prompt": req.original_prompt,
|
||||
"gen_type": gen_type,
|
||||
"status": "generating",
|
||||
"generation_mode": generation_mode,
|
||||
"pipeline_stage": "queued",
|
||||
"parent_task_id": parent_task_id,
|
||||
"generation_count": generation_count,
|
||||
"generation_index": generation_index,
|
||||
"engine_id": engine_id,
|
||||
"engine_snapshot_json": engine_snapshot_json,
|
||||
"media_references": media_references_json,
|
||||
"credits_cost": round(float(credits_cost or 0), 2),
|
||||
"idempotency_key": idempotency_key,
|
||||
"deadline_at": deadline_at,
|
||||
}
|
||||
|
||||
|
||||
async def create_generation_task_group(
|
||||
db: AsyncSession,
|
||||
current_user: User,
|
||||
req: GenerationAITaskCreate,
|
||||
) -> GenerationTaskCreateResult:
|
||||
"""创建单份 chatapi_async 或多份 chatapi_main/chatapi_child 任务组。
|
||||
|
||||
本函数只 flush,不主动 commit。调用方提交成功后才能投递 Celery。
|
||||
"""
|
||||
gen_type = (req.gen_type or "").lower().strip()
|
||||
if gen_type not in (GenerationType.IMAGE.value, GenerationType.VIDEO.value):
|
||||
raise HTTPException(status_code=400, detail="gen_type 仅支持 image 或 video")
|
||||
|
||||
existing = await find_existing_top_level_task(
|
||||
db,
|
||||
user_id=current_user.id,
|
||||
idempotency_key=req.idempotency_key,
|
||||
)
|
||||
if existing:
|
||||
return GenerationTaskCreateResult(
|
||||
top_level_task_id=existing.id,
|
||||
generation_count=int(existing.generation_count or 1),
|
||||
gen_type=existing.gen_type,
|
||||
created=False,
|
||||
)
|
||||
|
||||
refs = [item.model_dump(exclude_none=True) for item in (req.media_references or [])]
|
||||
refs = await resolve_private_portrait_references(
|
||||
db,
|
||||
user_id=current_user.id,
|
||||
media_references=refs,
|
||||
gen_type=gen_type,
|
||||
)
|
||||
media_references_json = _json(refs) if refs else None
|
||||
await assert_user_resource_capacity_available(db, current_user.id)
|
||||
|
||||
now = datetime.now(timezone.utc)
|
||||
main_id = generate_id()
|
||||
child_ids: list[str] = []
|
||||
enqueue_ids: list[str] = []
|
||||
total_billed_credits = 0.0
|
||||
|
||||
log_operation_event(
|
||||
domain="generation_ai_batch",
|
||||
event_type="BATCH_CREATE_START",
|
||||
event_status="started",
|
||||
source="service",
|
||||
user_id=current_user.id,
|
||||
group_id=main_id,
|
||||
detail={
|
||||
"gen_type": gen_type,
|
||||
"requested_generation_count": normalize_generation_count(req.generation_count),
|
||||
"idempotency_key_present": bool(req.idempotency_key),
|
||||
},
|
||||
)
|
||||
|
||||
if gen_type == GenerationType.IMAGE.value:
|
||||
if any((ref.get("type") or "").lower() == "audio" for ref in refs):
|
||||
raise HTTPException(status_code=400, detail="图片生成不支持音频参考素材")
|
||||
|
||||
engine = await get_image_engine(db, req.engine_id)
|
||||
generation_count = normalize_generation_count(req.generation_count)
|
||||
multi_generation_enabled = bool(getattr(engine, "multi_generation_enabled", False))
|
||||
max_generation_count = normalize_generation_count(getattr(engine, "max_generation_count", 1))
|
||||
if generation_count > 1 and not multi_generation_enabled:
|
||||
raise HTTPException(status_code=400, detail="当前图片引擎未开启多份生成,本次生成数量只能为 1")
|
||||
if generation_count > max_generation_count:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"当前图片引擎本次最多允许生成 {max_generation_count} 份",
|
||||
)
|
||||
|
||||
reference_image_count = sum(
|
||||
1 for ref in refs if (ref.get("type") or "").lower() == GenerationType.IMAGE.value
|
||||
)
|
||||
max_reference_count = max(0, int(getattr(engine, "max_reference_image_count", 14) or 0))
|
||||
multi_image_max_images = max(1, int(getattr(engine, "multi_image_max_images", 15) or 15))
|
||||
if reference_image_count > max_reference_count:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"当前图片引擎最多支持 {max_reference_count} 张参考图,当前 {reference_image_count} 张",
|
||||
)
|
||||
if generation_count > 1 and reference_image_count + generation_count > multi_image_max_images:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=(
|
||||
f"参考图数量与生成数量合计不能超过 {multi_image_max_images} 张,"
|
||||
f"当前参考图 {reference_image_count} 张、生成 {generation_count} 张"
|
||||
),
|
||||
)
|
||||
|
||||
sizes = image_supported_sizes(engine)
|
||||
size = req.image_size or engine.default_size or IMAGE_DEFAULT_SIZE
|
||||
proportion = req.image_proportion or IMAGE_DEFAULT_PROPORTION
|
||||
px = normalize_px(req.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
|
||||
|
||||
mode = GenerationMode.CHATAPI_ASYNC.value if generation_count == 1 else GenerationMode.CHATAPI_MAIN.value
|
||||
billing = await charge_generation_media_by_params(
|
||||
db,
|
||||
user_id=current_user.id,
|
||||
record_id=main_id,
|
||||
gen_type=GenerationType.IMAGE.value,
|
||||
image_size=size,
|
||||
engine_id=engine.id,
|
||||
project_name="AI生成任务",
|
||||
description_prefix="AI创作-",
|
||||
owner_type=OWNER_CHAT_GENERATION_TASK,
|
||||
attempt_no=1,
|
||||
quantity=generation_count,
|
||||
)
|
||||
image_snapshot = build_image_snapshot(engine, size, proportion, px)
|
||||
image_snapshot["generation_count"] = generation_count
|
||||
snapshot_json = _json(image_snapshot) or "{}"
|
||||
total_billed_credits = round(float(billing.total_charged or 0), 2)
|
||||
task = ChatGenerationTask(
|
||||
**_base_task_kwargs(
|
||||
task_id=main_id,
|
||||
user_id=current_user.id,
|
||||
req=req,
|
||||
gen_type=GenerationType.IMAGE.value,
|
||||
generation_mode=mode,
|
||||
generation_count=generation_count,
|
||||
engine_id=engine.id,
|
||||
engine_snapshot_json=snapshot_json,
|
||||
media_references_json=media_references_json,
|
||||
deadline_at=now + timedelta(minutes=settings.CHATAPI_ASYNC_IMAGE_DEADLINE_MINUTES),
|
||||
credits_cost=billing.total_charged,
|
||||
idempotency_key=req.idempotency_key,
|
||||
),
|
||||
image_size=size,
|
||||
image_proportion=proportion,
|
||||
image_px=px,
|
||||
)
|
||||
db.add(task)
|
||||
enqueue_ids.append(task.id)
|
||||
else:
|
||||
engine = await get_video_engine(db, req.engine_id)
|
||||
generation_count = normalize_generation_count(req.generation_count)
|
||||
multi_generation_enabled = bool(getattr(engine, "multi_generation_enabled", False))
|
||||
max_generation_count = normalize_generation_count(getattr(engine, "max_generation_count", 1))
|
||||
if generation_count > 1 and not multi_generation_enabled:
|
||||
raise HTTPException(status_code=400, detail="当前视频引擎未开启多份生成,本次生成数量只能为 1")
|
||||
if generation_count > max_generation_count:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"当前视频引擎本次最多允许生成 {max_generation_count} 份",
|
||||
)
|
||||
ratio = req.aspect_ratio or VIDEO_DEFAULT_RATIO
|
||||
resolution = req.resolution or VIDEO_DEFAULT_RESOLUTION
|
||||
duration = req.duration or VIDEO_DEFAULT_DURATION
|
||||
ratios = parse_json_list(engine.supported_ratios, [])
|
||||
resolutions = parse_json_list(engine.supported_resolutions, [])
|
||||
durations = parse_json_list(engine.supported_durations, [])
|
||||
if ratios and ratio not in ratios:
|
||||
raise HTTPException(status_code=400, detail=f"视频比例不支持: {ratio}")
|
||||
if resolutions and resolution not in resolutions:
|
||||
raise HTTPException(status_code=400, detail=f"视频分辨率不支持: {resolution}")
|
||||
if durations and duration not in durations:
|
||||
raise HTTPException(status_code=400, detail=f"视频时长不支持: {duration}")
|
||||
if engine.max_duration and duration > engine.max_duration:
|
||||
raise HTTPException(status_code=400, detail=f"视频时长不能超过 {engine.max_duration} 秒")
|
||||
|
||||
input_video_duration = _validate_video_references(refs, max_audio_count=engine.max_audio_count)
|
||||
video_snapshot = build_video_snapshot(engine, ratio, resolution, duration)
|
||||
video_snapshot["generation_count"] = generation_count
|
||||
snapshot_json = _json(video_snapshot) or "{}"
|
||||
deadline_at = now + timedelta(hours=settings.CHATAPI_ASYNC_VIDEO_FINAL_DEADLINE_HOURS)
|
||||
|
||||
if generation_count == 1:
|
||||
billing = await charge_generation_media_by_params(
|
||||
db,
|
||||
user_id=current_user.id,
|
||||
record_id=main_id,
|
||||
gen_type=GenerationType.VIDEO.value,
|
||||
duration=duration,
|
||||
resolution=resolution,
|
||||
engine_id=engine.id,
|
||||
input_video_duration=input_video_duration if input_video_duration > 0 else None,
|
||||
project_name="AI生成任务",
|
||||
description_prefix="AI创作-",
|
||||
owner_type=OWNER_CHAT_GENERATION_TASK,
|
||||
attempt_no=1,
|
||||
)
|
||||
total_billed_credits = round(float(billing.total_charged or 0), 2)
|
||||
task = ChatGenerationTask(
|
||||
**_base_task_kwargs(
|
||||
task_id=main_id,
|
||||
user_id=current_user.id,
|
||||
req=req,
|
||||
gen_type=GenerationType.VIDEO.value,
|
||||
generation_mode=GenerationMode.CHATAPI_ASYNC.value,
|
||||
generation_count=1,
|
||||
engine_id=engine.id,
|
||||
engine_snapshot_json=snapshot_json,
|
||||
media_references_json=media_references_json,
|
||||
deadline_at=deadline_at,
|
||||
credits_cost=billing.total_charged,
|
||||
idempotency_key=req.idempotency_key,
|
||||
),
|
||||
duration=duration,
|
||||
aspect_ratio=ratio,
|
||||
resolution=resolution,
|
||||
image_size=req.image_size or IMAGE_DEFAULT_SIZE,
|
||||
image_proportion=req.image_proportion or IMAGE_DEFAULT_PROPORTION,
|
||||
image_px=normalize_px(req.image_px) or IMAGE_DEFAULT_PX,
|
||||
)
|
||||
db.add(task)
|
||||
enqueue_ids.append(task.id)
|
||||
else:
|
||||
main_task = ChatGenerationTask(
|
||||
**_base_task_kwargs(
|
||||
task_id=main_id,
|
||||
user_id=current_user.id,
|
||||
req=req,
|
||||
gen_type=GenerationType.VIDEO.value,
|
||||
generation_mode=GenerationMode.CHATAPI_MAIN.value,
|
||||
generation_count=generation_count,
|
||||
engine_id=engine.id,
|
||||
engine_snapshot_json=snapshot_json,
|
||||
media_references_json=media_references_json,
|
||||
deadline_at=deadline_at,
|
||||
idempotency_key=req.idempotency_key,
|
||||
),
|
||||
duration=duration,
|
||||
aspect_ratio=ratio,
|
||||
resolution=resolution,
|
||||
image_size=req.image_size or IMAGE_DEFAULT_SIZE,
|
||||
image_proportion=req.image_proportion or IMAGE_DEFAULT_PROPORTION,
|
||||
image_px=normalize_px(req.image_px) or IMAGE_DEFAULT_PX,
|
||||
)
|
||||
db.add(main_task)
|
||||
await db.flush()
|
||||
|
||||
total_credits = 0.0
|
||||
children: list[ChatGenerationTask] = []
|
||||
for generation_index in range(1, generation_count + 1):
|
||||
child_id = generate_id()
|
||||
billing = await charge_generation_media_by_params(
|
||||
db,
|
||||
user_id=current_user.id,
|
||||
record_id=child_id,
|
||||
gen_type=GenerationType.VIDEO.value,
|
||||
duration=duration,
|
||||
resolution=resolution,
|
||||
engine_id=engine.id,
|
||||
input_video_duration=input_video_duration if input_video_duration > 0 else None,
|
||||
project_name="AI生成任务",
|
||||
description_prefix=f"AI创作-第{generation_index}份-",
|
||||
owner_type=OWNER_CHAT_GENERATION_TASK,
|
||||
attempt_no=1,
|
||||
)
|
||||
child = ChatGenerationTask(
|
||||
**_base_task_kwargs(
|
||||
task_id=child_id,
|
||||
user_id=current_user.id,
|
||||
req=req,
|
||||
gen_type=GenerationType.VIDEO.value,
|
||||
generation_mode=GenerationMode.CHATAPI_CHILD.value,
|
||||
generation_count=generation_count,
|
||||
generation_index=generation_index,
|
||||
parent_task_id=main_id,
|
||||
engine_id=engine.id,
|
||||
engine_snapshot_json=snapshot_json,
|
||||
media_references_json=media_references_json,
|
||||
deadline_at=deadline_at,
|
||||
credits_cost=billing.total_charged,
|
||||
),
|
||||
duration=duration,
|
||||
aspect_ratio=ratio,
|
||||
resolution=resolution,
|
||||
image_size=req.image_size or IMAGE_DEFAULT_SIZE,
|
||||
image_proportion=req.image_proportion or IMAGE_DEFAULT_PROPORTION,
|
||||
image_px=normalize_px(req.image_px) or IMAGE_DEFAULT_PX,
|
||||
)
|
||||
children.append(child)
|
||||
child_ids.append(child_id)
|
||||
enqueue_ids.append(child_id)
|
||||
total_credits = round(total_credits + billing.total_charged, 2)
|
||||
db.add_all(children)
|
||||
main_task.credits_cost = total_credits
|
||||
total_billed_credits = total_credits
|
||||
|
||||
await db.flush()
|
||||
log_operation_event(
|
||||
domain="generation_ai_batch",
|
||||
event_type="BATCH_BILLING_SUCCESS",
|
||||
event_status="success",
|
||||
source="service",
|
||||
user_id=current_user.id,
|
||||
group_id=main_id,
|
||||
detail={
|
||||
"gen_type": gen_type,
|
||||
"generation_count": generation_count,
|
||||
"total_billed_credits": total_billed_credits,
|
||||
},
|
||||
)
|
||||
log_operation_event(
|
||||
domain="generation_ai_batch",
|
||||
event_type="BATCH_CHILDREN_CREATED" if child_ids else "BATCH_MAIN_CREATED",
|
||||
event_status="success",
|
||||
source="service",
|
||||
user_id=current_user.id,
|
||||
group_id=main_id,
|
||||
detail={
|
||||
"gen_type": gen_type,
|
||||
"generation_count": generation_count,
|
||||
"child_task_ids": child_ids,
|
||||
"enqueue_task_ids": enqueue_ids,
|
||||
},
|
||||
)
|
||||
return GenerationTaskCreateResult(
|
||||
top_level_task_id=main_id,
|
||||
enqueue_task_ids=enqueue_ids,
|
||||
child_task_ids=child_ids,
|
||||
generation_count=generation_count,
|
||||
gen_type=gen_type,
|
||||
created=True,
|
||||
)
|
||||
|
||||
|
||||
async def enqueue_created_generation_tasks(
|
||||
db: AsyncSession,
|
||||
*,
|
||||
task_ids: list[str],
|
||||
) -> list[str]:
|
||||
"""在业务事务提交后投递任务;返回投递失败的任务ID。
|
||||
|
||||
投递失败会在补偿事务中将对应任务置为失败并幂等退款,视频子任务
|
||||
同时触发主任务状态汇总。调用方不应在初始事务提交前调用本函数。
|
||||
"""
|
||||
from app.services.generation.ai.task_group_service import aggregate_parent_for_child
|
||||
from app.services.generation.log_service import log_task_event
|
||||
from app.services.generation.refund_service import mark_chat_generation_task_failed_and_refund_once
|
||||
from app.tasks.generation_create_tasks import chatapi_create_generation_task
|
||||
|
||||
normalized_ids = list(dict.fromkeys(str(item) for item in task_ids if item))
|
||||
meta_result = await db.execute(
|
||||
select(
|
||||
ChatGenerationTask.id,
|
||||
ChatGenerationTask.user_id,
|
||||
ChatGenerationTask.parent_task_id,
|
||||
ChatGenerationTask.generation_index,
|
||||
).where(ChatGenerationTask.id.in_(normalized_ids))
|
||||
) if normalized_ids else None
|
||||
task_meta = {
|
||||
str(row.id): {
|
||||
"user_id": str(row.user_id),
|
||||
"parent_task_id": str(row.parent_task_id) if row.parent_task_id else None,
|
||||
"generation_index": row.generation_index,
|
||||
}
|
||||
for row in (meta_result.all() if meta_result is not None else [])
|
||||
}
|
||||
|
||||
failed_ids: list[str] = []
|
||||
for task_id in normalized_ids:
|
||||
meta = task_meta.get(task_id, {})
|
||||
log_operation_event(
|
||||
domain="generation_ai_batch",
|
||||
event_type="CHILD_ENQUEUE_START",
|
||||
event_status="started",
|
||||
source="api",
|
||||
user_id=meta.get("user_id"),
|
||||
group_id=meta.get("parent_task_id") or task_id,
|
||||
task_id=task_id,
|
||||
detail={"generation_index": meta.get("generation_index")},
|
||||
)
|
||||
try:
|
||||
chatapi_create_generation_task.delay(task_id)
|
||||
await log_task_event(
|
||||
task_id=task_id,
|
||||
event_type="CHILD_ENQUEUE_SUCCESS",
|
||||
to_status="generating",
|
||||
to_stage="queued",
|
||||
detail={"task_id": task_id},
|
||||
)
|
||||
log_operation_event(
|
||||
domain="generation_ai_batch",
|
||||
event_type="CHILD_ENQUEUE_SUCCESS",
|
||||
event_status="success",
|
||||
source="api",
|
||||
user_id=meta.get("user_id"),
|
||||
group_id=meta.get("parent_task_id") or task_id,
|
||||
task_id=task_id,
|
||||
detail={"generation_index": meta.get("generation_index")},
|
||||
)
|
||||
except Exception as exc:
|
||||
failed_ids.append(task_id)
|
||||
await db.rollback()
|
||||
failed_task = await mark_chat_generation_task_failed_and_refund_once(
|
||||
db,
|
||||
task_id=task_id,
|
||||
error_message=f"任务队列投递失败: {exc}",
|
||||
pipeline_stage="failed",
|
||||
)
|
||||
await aggregate_parent_for_child(db, failed_task)
|
||||
await db.commit()
|
||||
await log_task_event(
|
||||
task_id=task_id,
|
||||
event_type="CHILD_ENQUEUE_FAILED",
|
||||
to_status="failed",
|
||||
to_stage="failed",
|
||||
message=str(exc),
|
||||
detail={"task_id": task_id},
|
||||
)
|
||||
log_operation_event(
|
||||
domain="generation_ai_batch",
|
||||
event_type="CHILD_ENQUEUE_FAILED",
|
||||
event_status="failed",
|
||||
source="api",
|
||||
user_id=getattr(failed_task, "user_id", None),
|
||||
group_id=getattr(failed_task, "parent_task_id", None) or task_id,
|
||||
task_id=task_id,
|
||||
message=str(exc),
|
||||
detail={"physical_files_deleted": False},
|
||||
error=str(exc),
|
||||
)
|
||||
return failed_ids
|
||||
@@ -0,0 +1,426 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from collections import Counter, defaultdict
|
||||
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.generation_task import (
|
||||
ChatGenerationPipelineStage,
|
||||
ChatGenerationTaskStatus,
|
||||
GenerationMode,
|
||||
)
|
||||
from app.models.chat_generation_task import ChatGenerationTask
|
||||
from app.services.operation_log_service import log_operation_event
|
||||
from app.services.resource_accounting_service import (
|
||||
SOURCE_MODEL_CHAT_TASK,
|
||||
soft_delete_resources_by_source,
|
||||
)
|
||||
|
||||
|
||||
ACTIVE_STAGES = {
|
||||
ChatGenerationPipelineStage.QUEUED.value,
|
||||
ChatGenerationPipelineStage.PREPARING.value,
|
||||
ChatGenerationPipelineStage.CREATING_PROVIDER_TASK.value,
|
||||
ChatGenerationPipelineStage.WAITING_REMOTE.value,
|
||||
ChatGenerationPipelineStage.POLLING.value,
|
||||
ChatGenerationPipelineStage.RESULT_READY.value,
|
||||
ChatGenerationPipelineStage.DOWNLOAD_QUEUED.value,
|
||||
ChatGenerationPipelineStage.DOWNLOADING.value,
|
||||
ChatGenerationPipelineStage.RETRY_WAITING.value,
|
||||
}
|
||||
|
||||
|
||||
def is_task_active(task: ChatGenerationTask) -> bool:
|
||||
return task.deleted_at is None and (
|
||||
task.status == ChatGenerationTaskStatus.GENERATING.value
|
||||
or (task.pipeline_stage or "") in ACTIVE_STAGES
|
||||
)
|
||||
|
||||
|
||||
def get_display_status(task: ChatGenerationTask) -> str:
|
||||
if task.deleted_at is not None:
|
||||
return "deleted"
|
||||
if task.pipeline_stage == ChatGenerationPipelineStage.DOWNLOAD_FAILED.value:
|
||||
return "download_failed"
|
||||
return task.status or ChatGenerationTaskStatus.PENDING.value
|
||||
|
||||
|
||||
async def load_children_map(
|
||||
db: AsyncSession,
|
||||
parent_ids: Sequence[str] | Iterable[str],
|
||||
*,
|
||||
include_deleted: bool = True,
|
||||
) -> dict[str, list[ChatGenerationTask]]:
|
||||
ids = list(dict.fromkeys(str(item) for item in parent_ids if item))
|
||||
if not ids:
|
||||
return {}
|
||||
query = select(ChatGenerationTask).where(
|
||||
ChatGenerationTask.parent_task_id.in_(ids),
|
||||
ChatGenerationTask.generation_mode == GenerationMode.CHATAPI_CHILD.value,
|
||||
)
|
||||
if not include_deleted:
|
||||
query = query.where(ChatGenerationTask.deleted_at.is_(None))
|
||||
result = await db.execute(
|
||||
query.order_by(
|
||||
ChatGenerationTask.parent_task_id.asc(),
|
||||
ChatGenerationTask.generation_index.asc(),
|
||||
ChatGenerationTask.created_at.asc(),
|
||||
)
|
||||
)
|
||||
grouped: dict[str, list[ChatGenerationTask]] = defaultdict(list)
|
||||
for task in result.scalars().all():
|
||||
if task.parent_task_id:
|
||||
grouped[task.parent_task_id].append(task)
|
||||
return dict(grouped)
|
||||
|
||||
|
||||
async def load_task_and_children(
|
||||
db: AsyncSession,
|
||||
*,
|
||||
task_id: str,
|
||||
user_id: str | None = None,
|
||||
include_deleted_children: bool = True,
|
||||
) -> tuple[ChatGenerationTask | None, list[ChatGenerationTask]]:
|
||||
query = select(ChatGenerationTask).where(ChatGenerationTask.id == task_id)
|
||||
if user_id:
|
||||
query = query.where(ChatGenerationTask.user_id == user_id)
|
||||
result = await db.execute(query.limit(1))
|
||||
task = result.scalar_one_or_none()
|
||||
if not task:
|
||||
return None, []
|
||||
if task.generation_mode == GenerationMode.CHATAPI_CHILD.value and task.parent_task_id:
|
||||
parent_result = await db.execute(
|
||||
select(ChatGenerationTask).where(ChatGenerationTask.id == task.parent_task_id).limit(1)
|
||||
)
|
||||
parent = parent_result.scalar_one_or_none()
|
||||
return parent or task, [task]
|
||||
if task.generation_mode != GenerationMode.CHATAPI_MAIN.value:
|
||||
return task, []
|
||||
children_map = await load_children_map(
|
||||
db,
|
||||
[task.id],
|
||||
include_deleted=include_deleted_children,
|
||||
)
|
||||
return task, children_map.get(task.id, [])
|
||||
|
||||
|
||||
def _generation_result_status(task: ChatGenerationTask) -> str:
|
||||
"""返回任务真实生成结果,不受资源软删除影响。"""
|
||||
if task.pipeline_stage == ChatGenerationPipelineStage.DOWNLOAD_FAILED.value:
|
||||
return "download_failed"
|
||||
if task.status == ChatGenerationTaskStatus.FAILED.value or (task.pipeline_stage or "") in {
|
||||
ChatGenerationPipelineStage.FAILED.value,
|
||||
ChatGenerationPipelineStage.TIMEOUT.value,
|
||||
}:
|
||||
return "failed"
|
||||
if is_task_active(task):
|
||||
return "generating"
|
||||
if task.status == ChatGenerationTaskStatus.COMPLETED.value:
|
||||
return "completed"
|
||||
return task.status or "pending"
|
||||
|
||||
|
||||
def _build_summary(children: list[ChatGenerationTask]) -> str | None:
|
||||
if not children:
|
||||
return None
|
||||
result_counters: Counter[str] = Counter(_generation_result_status(child) for child in children)
|
||||
labels = {
|
||||
"completed": "完成",
|
||||
"failed": "生成失败",
|
||||
"download_failed": "下载失败",
|
||||
"generating": "生成中",
|
||||
"pending": "待处理",
|
||||
}
|
||||
parts = [f"{count}项{labels.get(status, status)}" for status, count in result_counters.items() if count]
|
||||
deleted_count = sum(1 for child in children if child.deleted_at is not None)
|
||||
if deleted_count:
|
||||
parts.append(f"{deleted_count}项资源已删除")
|
||||
return f"{len(children)}项中" + ",".join(parts)
|
||||
|
||||
|
||||
async def aggregate_main_task_status(
|
||||
db: AsyncSession,
|
||||
*,
|
||||
parent_task_id: str,
|
||||
) -> ChatGenerationTask | None:
|
||||
result = await db.execute(
|
||||
select(ChatGenerationTask)
|
||||
.where(
|
||||
ChatGenerationTask.id == parent_task_id,
|
||||
ChatGenerationTask.generation_mode == GenerationMode.CHATAPI_MAIN.value,
|
||||
)
|
||||
.with_for_update()
|
||||
.limit(1)
|
||||
)
|
||||
main = result.scalar_one_or_none()
|
||||
if not main or main.deleted_at is not None:
|
||||
return main
|
||||
|
||||
children_map = await load_children_map(db, [parent_task_id], include_deleted=True)
|
||||
children = children_map.get(parent_task_id, [])
|
||||
if not children:
|
||||
return main
|
||||
|
||||
previous_status = main.status
|
||||
previous_stage = main.pipeline_stage
|
||||
active_children = [child for child in children if is_task_active(child)]
|
||||
failed_children = [
|
||||
child
|
||||
for child in children
|
||||
if (
|
||||
child.status == ChatGenerationTaskStatus.FAILED.value
|
||||
or (child.pipeline_stage or "") in {
|
||||
ChatGenerationPipelineStage.FAILED.value,
|
||||
ChatGenerationPipelineStage.TIMEOUT.value,
|
||||
ChatGenerationPipelineStage.DOWNLOAD_FAILED.value,
|
||||
}
|
||||
)
|
||||
]
|
||||
completed_children = [
|
||||
child for child in children if child.status == ChatGenerationTaskStatus.COMPLETED.value
|
||||
]
|
||||
|
||||
if active_children:
|
||||
main.status = ChatGenerationTaskStatus.GENERATING.value
|
||||
main.pipeline_stage = active_children[0].pipeline_stage or ChatGenerationPipelineStage.QUEUED.value
|
||||
main.generated_at = None
|
||||
main.error_message = _build_summary(children)
|
||||
elif failed_children:
|
||||
main.status = ChatGenerationTaskStatus.FAILED.value
|
||||
main.pipeline_stage = (
|
||||
ChatGenerationPipelineStage.DOWNLOAD_FAILED.value
|
||||
if any(child.pipeline_stage == ChatGenerationPipelineStage.DOWNLOAD_FAILED.value for child in failed_children)
|
||||
else ChatGenerationPipelineStage.FAILED.value
|
||||
)
|
||||
main.generated_at = max(
|
||||
(child.generated_at for child in completed_children if child.generated_at),
|
||||
default=datetime.now(timezone.utc),
|
||||
)
|
||||
main.error_message = _build_summary(children)
|
||||
else:
|
||||
# 所有子任务真实生成结果均成功;资源是否软删除不改变生成历史终态。
|
||||
main.status = ChatGenerationTaskStatus.COMPLETED.value
|
||||
main.pipeline_stage = ChatGenerationPipelineStage.DONE.value
|
||||
main.generated_at = max(
|
||||
(child.generated_at for child in children if child.generated_at),
|
||||
default=main.generated_at or datetime.now(timezone.utc),
|
||||
)
|
||||
main.error_message = None
|
||||
|
||||
if main.gen_type == "video":
|
||||
main.credits_cost = round(sum(float(child.credits_cost or 0) for child in children), 2)
|
||||
main.text_credits_cost = round(sum(float(child.text_credits_cost or 0) for child in children), 2)
|
||||
main.text_tokens_used = sum(int(child.text_tokens_used or 0) for child in children)
|
||||
main.image_tokens_used = sum(int(child.image_tokens_used or 0) for child in children)
|
||||
main.video_tokens_used = sum(int(child.video_tokens_used or 0) for child in children)
|
||||
main.retry_count = sum(int(child.retry_count or 0) for child in children)
|
||||
main.poll_count = sum(int(child.poll_count or 0) for child in children)
|
||||
|
||||
await db.flush()
|
||||
log_operation_event(
|
||||
domain="generation_ai_batch",
|
||||
event_type="MAIN_STATUS_AGGREGATED",
|
||||
event_status="success",
|
||||
source="service",
|
||||
user_id=main.user_id,
|
||||
group_id=main.id,
|
||||
task_id=main.id,
|
||||
detail={
|
||||
"before_status": previous_status,
|
||||
"before_stage": previous_stage,
|
||||
"after_status": main.status,
|
||||
"after_stage": main.pipeline_stage,
|
||||
"summary": _build_summary(children),
|
||||
},
|
||||
)
|
||||
return main
|
||||
|
||||
|
||||
async def aggregate_parent_for_child(db: AsyncSession, child: ChatGenerationTask | None) -> ChatGenerationTask | None:
|
||||
if not child or child.generation_mode != GenerationMode.CHATAPI_CHILD.value or not child.parent_task_id:
|
||||
return None
|
||||
# 项目关闭了 autoflush,先显式 flush 子任务的终态,确保聚合查询读取到本事务最新状态。
|
||||
await db.flush()
|
||||
return await aggregate_main_task_status(db, parent_task_id=str(child.parent_task_id))
|
||||
|
||||
|
||||
async def soft_delete_child_tasks_batch(
|
||||
db: AsyncSession,
|
||||
*,
|
||||
child_task_ids: Sequence[str] | Iterable[str],
|
||||
user_id: str,
|
||||
deleted_at: datetime | None = None,
|
||||
require_completed: bool = False,
|
||||
) -> int:
|
||||
ids = list(dict.fromkeys(str(item) for item in child_task_ids if item))
|
||||
if not ids:
|
||||
return 0
|
||||
deleted_at = deleted_at or datetime.now(timezone.utc)
|
||||
result = await db.execute(
|
||||
select(ChatGenerationTask)
|
||||
.where(
|
||||
ChatGenerationTask.id.in_(ids),
|
||||
ChatGenerationTask.user_id == user_id,
|
||||
ChatGenerationTask.generation_mode == GenerationMode.CHATAPI_CHILD.value,
|
||||
)
|
||||
.with_for_update()
|
||||
)
|
||||
children = list(result.scalars().all())
|
||||
found_ids = {str(child.id) for child in children}
|
||||
missing_ids = [item for item in ids if item not in found_ids]
|
||||
if missing_ids:
|
||||
raise HTTPException(status_code=404, detail=f"子任务不存在: {','.join(missing_ids)}")
|
||||
|
||||
active_children = [child for child in children if child.deleted_at is None]
|
||||
running_ids = [child.id for child in active_children if is_task_active(child)]
|
||||
if running_ids:
|
||||
raise HTTPException(status_code=400, detail=f"仍有 {len(running_ids)} 个子任务生成中,暂不能删除")
|
||||
if require_completed:
|
||||
invalid_ids = [
|
||||
child.id for child in active_children
|
||||
if child.status != ChatGenerationTaskStatus.COMPLETED.value or child.generated_at is None
|
||||
]
|
||||
if invalid_ids:
|
||||
raise HTTPException(status_code=409, detail=f"只有生成完成的资源才能从素材云删除: {','.join(invalid_ids)}")
|
||||
|
||||
source_ids = [str(child.id) for child in active_children]
|
||||
freed_size = await soft_delete_resources_by_source(
|
||||
db,
|
||||
source_model=SOURCE_MODEL_CHAT_TASK,
|
||||
source_ids=source_ids,
|
||||
deleted_at=deleted_at,
|
||||
)
|
||||
parent_ids = list(dict.fromkeys(str(child.parent_task_id) for child in active_children if child.parent_task_id))
|
||||
for child in active_children:
|
||||
child.deleted_at = deleted_at
|
||||
await db.flush()
|
||||
for parent_id in parent_ids:
|
||||
await aggregate_main_task_status(db, parent_task_id=parent_id)
|
||||
return int(freed_size or 0)
|
||||
|
||||
|
||||
async def soft_delete_child_task(
|
||||
db: AsyncSession,
|
||||
*,
|
||||
child_task_id: str,
|
||||
user_id: str,
|
||||
deleted_at: datetime | None = None,
|
||||
) -> int:
|
||||
deleted_at = deleted_at or datetime.now(timezone.utc)
|
||||
detail_result = await db.execute(
|
||||
select(
|
||||
ChatGenerationTask.parent_task_id,
|
||||
ChatGenerationTask.generation_index,
|
||||
).where(
|
||||
ChatGenerationTask.id == child_task_id,
|
||||
ChatGenerationTask.user_id == user_id,
|
||||
ChatGenerationTask.generation_mode == GenerationMode.CHATAPI_CHILD.value,
|
||||
).limit(1)
|
||||
)
|
||||
detail = detail_result.one_or_none()
|
||||
if not detail:
|
||||
raise HTTPException(status_code=404, detail="子任务不存在")
|
||||
log_operation_event(
|
||||
domain="generation_ai_batch",
|
||||
event_type="CHILD_RESOURCE_DELETE_START",
|
||||
event_status="started",
|
||||
source="service",
|
||||
user_id=user_id,
|
||||
group_id=detail.parent_task_id,
|
||||
task_id=child_task_id,
|
||||
detail={"generation_index": detail.generation_index},
|
||||
)
|
||||
freed_size = await soft_delete_child_tasks_batch(
|
||||
db,
|
||||
child_task_ids=[child_task_id],
|
||||
user_id=user_id,
|
||||
deleted_at=deleted_at,
|
||||
)
|
||||
log_operation_event(
|
||||
domain="generation_ai_batch",
|
||||
event_type="CHILD_RESOURCE_DELETE_SUCCESS",
|
||||
event_status="success",
|
||||
source="service",
|
||||
user_id=user_id,
|
||||
group_id=detail.parent_task_id,
|
||||
task_id=child_task_id,
|
||||
detail={"generation_index": detail.generation_index, "freed_size_bytes": freed_size},
|
||||
)
|
||||
return freed_size
|
||||
|
||||
|
||||
async def soft_delete_top_level_task_group(
|
||||
db: AsyncSession,
|
||||
*,
|
||||
task_id: str,
|
||||
user_id: str,
|
||||
deleted_at: datetime | None = None,
|
||||
) -> int:
|
||||
deleted_at = deleted_at or datetime.now(timezone.utc)
|
||||
result = await db.execute(
|
||||
select(ChatGenerationTask)
|
||||
.where(
|
||||
ChatGenerationTask.id == task_id,
|
||||
ChatGenerationTask.user_id == user_id,
|
||||
ChatGenerationTask.generation_mode.in_(
|
||||
[GenerationMode.CHATAPI_ASYNC.value, GenerationMode.CHATAPI_MAIN.value]
|
||||
),
|
||||
ChatGenerationTask.deleted_at.is_(None),
|
||||
)
|
||||
.with_for_update()
|
||||
.limit(1)
|
||||
)
|
||||
task = result.scalar_one_or_none()
|
||||
if not task:
|
||||
raise HTTPException(status_code=404, detail="任务不存在")
|
||||
|
||||
if task.generation_mode == GenerationMode.CHATAPI_ASYNC.value:
|
||||
if is_task_active(task):
|
||||
raise HTTPException(status_code=400, detail="当前任务正在生成中,暂不能删除")
|
||||
freed_size = await soft_delete_resources_by_source(
|
||||
db,
|
||||
source_model=SOURCE_MODEL_CHAT_TASK,
|
||||
source_ids=[task.id],
|
||||
deleted_at=deleted_at,
|
||||
)
|
||||
task.deleted_at = deleted_at
|
||||
await db.flush()
|
||||
return int(freed_size or 0)
|
||||
|
||||
children_map = await load_children_map(db, [task.id], include_deleted=True)
|
||||
children = children_map.get(task.id, [])
|
||||
active_ids = [child.id for child in children if is_task_active(child)]
|
||||
if active_ids:
|
||||
raise HTTPException(status_code=400, detail=f"任务组仍有 {len(active_ids)} 个子任务生成中,暂不能删除")
|
||||
|
||||
active_children = [child for child in children if child.deleted_at is None]
|
||||
child_ids = [child.id for child in active_children]
|
||||
freed_size = await soft_delete_resources_by_source(
|
||||
db,
|
||||
source_model=SOURCE_MODEL_CHAT_TASK,
|
||||
source_ids=child_ids,
|
||||
deleted_at=deleted_at,
|
||||
)
|
||||
for child in active_children:
|
||||
child.deleted_at = deleted_at
|
||||
task.deleted_at = deleted_at
|
||||
await db.flush()
|
||||
log_operation_event(
|
||||
domain="generation_ai_batch",
|
||||
event_type="BATCH_GROUP_DELETE_SUCCESS",
|
||||
event_status="success",
|
||||
source="service",
|
||||
user_id=user_id,
|
||||
group_id=task.id,
|
||||
task_id=task.id,
|
||||
detail={
|
||||
"child_task_ids": child_ids,
|
||||
"freed_size_bytes": int(freed_size or 0),
|
||||
"physical_files_deleted": False,
|
||||
},
|
||||
)
|
||||
return int(freed_size or 0)
|
||||
@@ -0,0 +1,575 @@
|
||||
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,
|
||||
quantity: int = 1,
|
||||
) -> BillingSummary:
|
||||
"""图片/视频媒体生成扣费。
|
||||
|
||||
正式幂等由 owner_type + record_id + attempt_no + media + charge 组成。
|
||||
每次用户主动重试必须传入新的 attempt_no。
|
||||
"""
|
||||
project_name = project_name or "AI生成任务"
|
||||
gen_type = (gen_type or "").lower().strip()
|
||||
quantity = max(1, int(quantity or 1))
|
||||
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"
|
||||
unit_amount = await calc_image_credits(db, size, engine_id=engine_id)
|
||||
amount = round(unit_amount * quantity, 2)
|
||||
items.append(
|
||||
await deduct_credits_locked_once(
|
||||
db,
|
||||
user_id=user_id,
|
||||
amount=amount,
|
||||
description=f"{description_prefix}图片生成" + (f"×{quantity}" if quantity > 1 else ""),
|
||||
related_id=record_id,
|
||||
charge_key=CHARGE_MEDIA,
|
||||
biz_key=biz_key,
|
||||
attempt_no=attempt_no,
|
||||
record_meta=record_meta,
|
||||
)
|
||||
)
|
||||
elif gen_type == "video":
|
||||
unit_amount = await calc_video_credits(
|
||||
db, duration or 5, resolution or "720p",
|
||||
engine_id=engine_id,
|
||||
input_video_duration=input_video_duration,
|
||||
)
|
||||
amount = round(unit_amount * quantity, 2)
|
||||
items.append(
|
||||
await deduct_credits_locked_once(
|
||||
db,
|
||||
user_id=user_id,
|
||||
amount=amount,
|
||||
description=f"{description_prefix}视频生成" + (f"×{quantity}" if quantity > 1 else ""),
|
||||
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,
|
||||
)
|
||||
@@ -0,0 +1,143 @@
|
||||
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,
|
||||
)
|
||||
@@ -0,0 +1,554 @@
|
||||
from __future__ import annotations
|
||||
|
||||
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.generation.ai.task_group_service import soft_delete_child_tasks_batch
|
||||
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_service import log_operation_event
|
||||
# from app.services.operation_log import log_operation
|
||||
from app.services.resource_accounting_service import (
|
||||
SOURCE_MODEL_CHAT_TASK,
|
||||
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.in_([
|
||||
GenerationMode.CHATAPI_ASYNC.value,
|
||||
GenerationMode.CHATAPI_CHILD.value,
|
||||
]),
|
||||
)
|
||||
.with_for_update()
|
||||
)
|
||||
tasks = list(result.scalars().all())
|
||||
_raise_missing_if_any(ids=ids, found_ids=[task.id for task in tasks], message="AI 创作记录不存在或已删除")
|
||||
|
||||
already_deleted_ids = [str(task.id) for task in tasks if task.deleted_at is not None]
|
||||
_raise_invalid_if_any(invalid_ids=already_deleted_ids, message="AI 创作记录不存在或已删除", status_code=404)
|
||||
|
||||
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 创作记录只有生成完成后才能删除")
|
||||
|
||||
async_ids = [str(task.id) for task in tasks if task.generation_mode == GenerationMode.CHATAPI_ASYNC.value]
|
||||
child_ids = [str(task.id) for task in tasks if task.generation_mode == GenerationMode.CHATAPI_CHILD.value]
|
||||
|
||||
freed_size = 0
|
||||
if async_ids:
|
||||
freed_size += int(await soft_delete_resources_by_source(
|
||||
db,
|
||||
source_model=SOURCE_MODEL_CHAT_TASK,
|
||||
source_ids=async_ids,
|
||||
deleted_at=deleted_at,
|
||||
) or 0)
|
||||
async_id_set = set(async_ids)
|
||||
for task in tasks:
|
||||
if str(task.id) in async_id_set:
|
||||
task.deleted_at = deleted_at
|
||||
|
||||
if child_ids:
|
||||
freed_size += await soft_delete_child_tasks_batch(
|
||||
db,
|
||||
child_task_ids=child_ids,
|
||||
user_id=current_user.id,
|
||||
deleted_at=deleted_at,
|
||||
require_completed=True,
|
||||
)
|
||||
|
||||
parent_task_ids = list(dict.fromkeys(
|
||||
str(task.parent_task_id) for task in tasks if task.parent_task_id
|
||||
))
|
||||
log_operation_event(
|
||||
domain="generation_ai_batch",
|
||||
event_type="CHILD_RESOURCE_DELETE_SUCCESS",
|
||||
event_status="success",
|
||||
source="service",
|
||||
user_id=current_user.id,
|
||||
group_id=parent_task_ids[0] if len(parent_task_ids) == 1 else None,
|
||||
detail={
|
||||
"batch": True,
|
||||
"task_ids": [str(task.id) for task in tasks],
|
||||
"parent_task_ids": parent_task_ids,
|
||||
"freed_size_bytes": int(freed_size or 0),
|
||||
"physical_files_deleted": False,
|
||||
},
|
||||
)
|
||||
|
||||
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
|
||||
@@ -0,0 +1,283 @@
|
||||
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
|
||||
@@ -0,0 +1,134 @@
|
||||
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
|
||||
@@ -0,0 +1,41 @@
|
||||
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
|
||||
@@ -0,0 +1,153 @@
|
||||
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,
|
||||
)
|
||||
@@ -0,0 +1,192 @@
|
||||
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,
|
||||
}
|
||||
@@ -0,0 +1,226 @@
|
||||
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 ImageProviderError, 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
|
||||
from app.types.generation.provider import ImageProviderBatchResult
|
||||
|
||||
|
||||
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 {}
|
||||
|
||||
|
||||
def _try_json(value: Any) -> Any:
|
||||
if not isinstance(value, str):
|
||||
return value
|
||||
try:
|
||||
return json.loads(value)
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
|
||||
async def get_runtime_engine(db: AsyncSession, task: ChatGenerationTask) -> Any:
|
||||
"""使用任务快照冻结历史参数,只从当前引擎记录读取密钥。"""
|
||||
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"),
|
||||
multi_generation_enabled=bool(
|
||||
snapshot.get("multi_generation_enabled")
|
||||
if snapshot.get("multi_generation_enabled") is not None
|
||||
else getattr(engine, "multi_generation_enabled", False)
|
||||
),
|
||||
max_generation_count=int(
|
||||
snapshot.get("max_generation_count")
|
||||
or getattr(engine, "max_generation_count", 1)
|
||||
or 1
|
||||
),
|
||||
multi_image_max_images=int(
|
||||
snapshot.get("multi_image_max_images")
|
||||
or getattr(engine, "multi_image_max_images", 15)
|
||||
or 15
|
||||
),
|
||||
max_reference_image_count=int(
|
||||
snapshot.get("max_reference_image_count")
|
||||
if snapshot.get("max_reference_image_count") is not None
|
||||
else getattr(engine, "max_reference_image_count", 14)
|
||||
),
|
||||
output_format=(
|
||||
snapshot.get("output_format")
|
||||
if snapshot.get("output_format") is not None
|
||||
else getattr(engine, "output_format", "")
|
||||
) or "",
|
||||
)
|
||||
|
||||
|
||||
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_result(db, task)
|
||||
raise ValueError(f"不支持的生成类型: {task.gen_type}")
|
||||
|
||||
|
||||
async def _create_video_task(db: AsyncSession, task: ChatGenerationTask) -> dict:
|
||||
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(None, 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_batch_result(
|
||||
db: AsyncSession,
|
||||
task: ChatGenerationTask,
|
||||
*,
|
||||
generation_count: int,
|
||||
) -> ImageProviderBatchResult:
|
||||
engine = await get_runtime_engine(db, task)
|
||||
return await create_image_sync_batch_result_with_engine(
|
||||
task,
|
||||
engine,
|
||||
generation_count=generation_count,
|
||||
)
|
||||
|
||||
|
||||
async def create_image_sync_batch_result_with_engine(
|
||||
task: ChatGenerationTask,
|
||||
engine: Any,
|
||||
*,
|
||||
generation_count: int,
|
||||
) -> ImageProviderBatchResult:
|
||||
"""执行一次同步图片请求。
|
||||
|
||||
generation_count > 1 时是一次组图 API 调用;失败后绝不退化为多次单图调用。
|
||||
"""
|
||||
count = max(1, int(generation_count or 1))
|
||||
started = time.perf_counter()
|
||||
api_type = "image_sync_batch_create" if count > 1 else "image_sync_create"
|
||||
async with provider_limit("ark_image_sync_create", settings.ARK_IMAGE_CREATE_MAX_CONCURRENCY):
|
||||
try:
|
||||
result = await asyncio.to_thread(
|
||||
submit_image_task,
|
||||
None,
|
||||
engine,
|
||||
task,
|
||||
include_media_references=True,
|
||||
generation_count=count,
|
||||
)
|
||||
response_data = result.get("response_data") or result
|
||||
await log_provider_call(
|
||||
task,
|
||||
provider=engine.provider,
|
||||
api_type=api_type,
|
||||
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,
|
||||
total_tokens=int(result.get("image_tokens", 0) or 0),
|
||||
)
|
||||
return result
|
||||
except Exception as exc:
|
||||
error_message = exc.safe_message if isinstance(exc, ImageProviderError) else str(exc)
|
||||
await log_provider_call(
|
||||
task,
|
||||
provider=engine.provider,
|
||||
api_type=api_type,
|
||||
model=engine.model_name,
|
||||
engine_id=task.engine_id,
|
||||
status="failed",
|
||||
latency_ms=int((time.perf_counter() - started) * 1000),
|
||||
error_message=error_message,
|
||||
response_data=exc.as_dict() if isinstance(exc, ImageProviderError) else None,
|
||||
)
|
||||
raise
|
||||
|
||||
|
||||
async def create_image_sync_result(db: AsyncSession, task: ChatGenerationTask) -> dict:
|
||||
result = await create_image_sync_batch_result(db, task, generation_count=1)
|
||||
items = result.get("items") or []
|
||||
if len(items) != 1:
|
||||
raise RuntimeError(f"图片供应商单图返回数量异常,期望 1,实际 {len(items)}")
|
||||
item = items[0]
|
||||
if item.get("error_message"):
|
||||
raise RuntimeError(item.get("error_message") or "图片生成失败")
|
||||
image_url = item.get("remote_result_url")
|
||||
if not image_url:
|
||||
raise RuntimeError("图片供应商未返回有效图片地址")
|
||||
return {
|
||||
"task_id": None,
|
||||
"remote_result_url": image_url,
|
||||
"image_tokens": int(result.get("image_tokens", 0) or 0),
|
||||
"response_data": result.get("response_data") or {},
|
||||
}
|
||||
|
||||
|
||||
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)
|
||||
@@ -0,0 +1,943 @@
|
||||
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,
|
||||
GenerationMode,
|
||||
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)
|
||||
from app.services.generation.ai.task_group_service import aggregate_parent_for_child
|
||||
await aggregate_parent_for_child(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)
|
||||
from app.services.generation.ai.task_group_service import aggregate_parent_for_child
|
||||
await aggregate_parent_for_child(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、无供应商任务 ID:deadline 未过才恢复 create。
|
||||
4. 无结果 URL、无供应商任务 ID:deadline 已过直接超时失败,不再补救生成。
|
||||
"""
|
||||
from app.tasks.generation_create_tasks import chatapi_create_generation_task
|
||||
from app.tasks.generation_download_tasks import enqueue_download_task
|
||||
from app.tasks.generation_poll_tasks import poll_generation_task, register_poll_active
|
||||
|
||||
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] = {}
|
||||
|
||||
# 图片多份主任务只补投递,不在恢复服务内直接调用供应商。
|
||||
# 有效 claim 未过期时必须跳过,防止与正在运行的 Worker 重复调用组图 API。
|
||||
from app.tasks.generation_create_tasks import chatapi_create_generation_task
|
||||
image_main_cursor: str | None = None
|
||||
image_main_batch_size = max(1, int(settings.GENERATION_RECOVERY_BATCH_SIZE or 100))
|
||||
while True:
|
||||
image_main_query = select(ChatGenerationTask).where(
|
||||
ChatGenerationTask.deleted_at.is_(None),
|
||||
ChatGenerationTask.generation_mode == GenerationMode.CHATAPI_MAIN.value,
|
||||
ChatGenerationTask.gen_type == GenerationType.IMAGE.value,
|
||||
ChatGenerationTask.status == ChatGenerationTaskStatus.GENERATING.value,
|
||||
ChatGenerationTask.pipeline_stage.in_([
|
||||
ChatGenerationPipelineStage.QUEUED.value,
|
||||
ChatGenerationPipelineStage.PREPARING.value,
|
||||
ChatGenerationPipelineStage.CREATING_PROVIDER_TASK.value,
|
||||
]),
|
||||
)
|
||||
if image_main_cursor:
|
||||
image_main_query = image_main_query.where(ChatGenerationTask.id > image_main_cursor)
|
||||
image_main_result = await db.execute(
|
||||
image_main_query.order_by(ChatGenerationTask.id.asc())
|
||||
.limit(image_main_batch_size)
|
||||
.with_for_update()
|
||||
)
|
||||
image_mains = list(image_main_result.scalars().all())
|
||||
if not image_mains:
|
||||
break
|
||||
|
||||
for main in image_mains:
|
||||
main_id = str(main.id)
|
||||
image_main_cursor = main_id
|
||||
checked_ids.add(main_id)
|
||||
|
||||
child_result = await db.execute(
|
||||
select(ChatGenerationTask.id)
|
||||
.where(
|
||||
ChatGenerationTask.parent_task_id == main_id,
|
||||
ChatGenerationTask.generation_mode == GenerationMode.CHATAPI_CHILD.value,
|
||||
)
|
||||
.limit(1)
|
||||
)
|
||||
if child_result.scalar_one_or_none() is not None:
|
||||
main.provider_create_claim_token = None
|
||||
main.provider_create_lease_until = None
|
||||
await db.commit()
|
||||
results["image_main_already_split"] = results.get("image_main_already_split", 0) + 1
|
||||
continue
|
||||
|
||||
now = _now()
|
||||
lease_until = ensure_aware_utc(main.provider_create_lease_until)
|
||||
lease_alive = bool(main.provider_create_claim_token and lease_until and lease_until > now)
|
||||
if lease_alive:
|
||||
await db.commit()
|
||||
results["image_main_claim_alive"] = results.get("image_main_claim_alive", 0) + 1
|
||||
continue
|
||||
|
||||
if _is_expired(main.deadline_at, now):
|
||||
main.provider_create_claim_token = None
|
||||
main.provider_create_lease_until = None
|
||||
await mark_chat_generation_task_failed_and_refund_once(
|
||||
db,
|
||||
task=main,
|
||||
error_message="图片批量生成任务超时",
|
||||
pipeline_stage=ChatGenerationPipelineStage.TIMEOUT.value,
|
||||
)
|
||||
await db.commit()
|
||||
results["image_main_timeout"] = results.get("image_main_timeout", 0) + 1
|
||||
continue
|
||||
|
||||
if main.provider_create_claim_token or main.provider_create_lease_until:
|
||||
main.provider_create_claim_token = None
|
||||
main.provider_create_lease_until = None
|
||||
main.pipeline_stage = ChatGenerationPipelineStage.QUEUED.value
|
||||
await log_task_event(
|
||||
main,
|
||||
event_type=ChatGenerationTaskEventType.IMAGE_MAIN_CLAIM_EXPIRED.value,
|
||||
message="图片主任务供应商执行租约已过期,恢复重新投递",
|
||||
)
|
||||
await db.commit()
|
||||
try:
|
||||
chatapi_create_generation_task.apply_async(
|
||||
args=[main_id],
|
||||
queue=CeleryQueue.GEN_CHATAPI_CREATE.value,
|
||||
countdown=0,
|
||||
)
|
||||
results["recover_image_main_create"] = results.get("recover_image_main_create", 0) + 1
|
||||
except Exception as exc:
|
||||
logger.exception("恢复投递图片主任务失败 task_id=%s: %s", main_id, exc)
|
||||
results["recover_image_main_enqueue_failed"] = results.get("recover_image_main_enqueue_failed", 0) + 1
|
||||
|
||||
if len(image_mains) < image_main_batch_size:
|
||||
break
|
||||
|
||||
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
|
||||
|
||||
# 子任务可能在 worker 中断前已进入终态但主任务尚未汇总,按稳定游标完整重算全部主任务。
|
||||
from app.services.generation.ai.task_group_service import aggregate_main_task_status
|
||||
reconciled = 0
|
||||
main_cursor: str | None = None
|
||||
while True:
|
||||
main_query = select(ChatGenerationTask.id).where(
|
||||
ChatGenerationTask.deleted_at.is_(None),
|
||||
ChatGenerationTask.generation_mode == GenerationMode.CHATAPI_MAIN.value,
|
||||
)
|
||||
if main_cursor:
|
||||
main_query = main_query.where(ChatGenerationTask.id > main_cursor)
|
||||
main_result = await db.execute(main_query.order_by(ChatGenerationTask.id.asc()).limit(batch_size))
|
||||
parent_ids = list(main_result.scalars().all())
|
||||
if not parent_ids:
|
||||
break
|
||||
for parent_task_id in parent_ids:
|
||||
main_cursor = str(parent_task_id)
|
||||
await aggregate_main_task_status(db, parent_task_id=str(parent_task_id))
|
||||
await db.commit()
|
||||
reconciled += 1
|
||||
if len(parent_ids) < batch_size:
|
||||
break
|
||||
if reconciled:
|
||||
results["reconcile_main"] = reconciled
|
||||
|
||||
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,
|
||||
}
|
||||
@@ -0,0 +1,212 @@
|
||||
from __future__ import annotations
|
||||
|
||||
|
||||
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", "chatapi_main", "chatapi_child", "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
|
||||
@@ -0,0 +1,212 @@
|
||||
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.engine_service import (
|
||||
IMAGE_DEFAULT_PROPORTION,
|
||||
IMAGE_DEFAULT_PX,
|
||||
IMAGE_DEFAULT_SIZE,
|
||||
VIDEO_DEFAULT_RATIO,
|
||||
VIDEO_DEFAULT_RESOLUTION,
|
||||
build_image_snapshot as _build_image_snapshot,
|
||||
build_video_snapshot as _build_video_snapshot,
|
||||
get_image_engine as _get_image_engine,
|
||||
get_video_engine as _get_video_engine,
|
||||
image_supported_sizes as _image_supported_sizes,
|
||||
normalize_px,
|
||||
parse_json_list as _parse_list,
|
||||
)
|
||||
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)
|
||||
snapshot["generation_count"] = 1
|
||||
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)
|
||||
snapshot["generation_count"] = 1
|
||||
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
|
||||
@@ -1,14 +1,12 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from copy import deepcopy
|
||||
from datetime import datetime, timezone
|
||||
from datetime import datetime
|
||||
from typing import Any
|
||||
|
||||
from fastapi import HTTPException
|
||||
from sqlalchemy import String, cast, func, or_, select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from sqlalchemy.orm.attributes import flag_modified
|
||||
|
||||
from app.config import settings
|
||||
from app.enums.common import ModuleEventTypeEnum, ModuleProjectStatusEnum, ModulePromptTypeEnum, ModuleStepStatusEnum
|
||||
@@ -35,20 +33,19 @@ from app.schemas.hot_opening_replicate import (
|
||||
HotOpeningVideoGenerationOut,
|
||||
HotOpeningVideoPromptSchemaUpdateRequest,
|
||||
)
|
||||
from app.services.generation_ai_service import (
|
||||
from app.services.generation.ai.engine_service import (
|
||||
VIDEO_DEFAULT_DURATION,
|
||||
VIDEO_DEFAULT_RATIO,
|
||||
VIDEO_DEFAULT_RESOLUTION,
|
||||
_get_video_engine,
|
||||
_parse_list,
|
||||
get_video_engine,
|
||||
parse_json_list,
|
||||
)
|
||||
from app.services.generation_billing_service import charge_module_prompt_usage
|
||||
from app.services.generation_refund_service import mark_chat_generation_task_failed_and_refund_once
|
||||
from app.services.generation_task_factory_service import create_chat_generation_task_for_module
|
||||
from app.services.generation.billing_service import charge_module_prompt_usage
|
||||
from app.services.generation.refund_service import mark_chat_generation_task_failed_and_refund_once
|
||||
from app.services.generation.task_factory_service import create_chat_generation_task_for_module
|
||||
from app.services.hot_opening_video_prompt_service import build_final_video_prompt, optimize_hot_opening_video_prompt, patch_video_prompt_schema_from_client
|
||||
from app.services.module_generation_log_service import log_module_error, log_module_event_file, log_module_prompt_event
|
||||
from app.services.llm import optimize_prompt
|
||||
from app.services.resource_accounting_service import soft_delete_chat_task_resources
|
||||
from app.services.module_generation_flow_base_service import (
|
||||
chat_tasks_by_id as _base_chat_tasks_by_id,
|
||||
create_module_step as _base_create_step,
|
||||
@@ -1021,10 +1018,10 @@ async def generate_image_from_prompt(
|
||||
|
||||
|
||||
async def _resolve_video_prompt_config(db: AsyncSession, req: HotOpeningGenerateVideoPromptRequest) -> dict[str, Any]:
|
||||
engine = await _get_video_engine(db, req.engine_id)
|
||||
supported_ratios = _parse_list(engine.supported_ratios, [])
|
||||
supported_resolutions = _parse_list(engine.supported_resolutions, [])
|
||||
supported_durations = _parse_list(engine.supported_durations, [])
|
||||
engine = await get_video_engine(db, req.engine_id)
|
||||
supported_ratios = parse_json_list(engine.supported_ratios, [])
|
||||
supported_resolutions = parse_json_list(engine.supported_resolutions, [])
|
||||
supported_durations = parse_json_list(engine.supported_durations, [])
|
||||
|
||||
default_ratio = getattr(settings, "HOT_OPENING_DEFAULT_VIDEO_RATIO", None) or VIDEO_DEFAULT_RATIO
|
||||
default_resolution = getattr(settings, "HOT_OPENING_DEFAULT_VIDEO_RESOLUTION", None) or VIDEO_DEFAULT_RESOLUTION
|
||||
|
||||
@@ -1,9 +1,8 @@
|
||||
import base64
|
||||
import json
|
||||
import logging
|
||||
import mimetypes
|
||||
import os
|
||||
from datetime import datetime
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
from sqlalchemy import select
|
||||
@@ -11,10 +10,16 @@ from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from volcenginesdkarkruntime import AsyncArk
|
||||
|
||||
from app.config import settings
|
||||
from app.enums.generation_provider import (
|
||||
MULTI_IMAGE_PROMPT_TEMPLATE,
|
||||
ImageProviderErrorType,
|
||||
)
|
||||
from app.enums.private_portrait import PRIVATE_PORTRAIT_ASSET_URI_PREFIX
|
||||
from app.models.image_engine import ImageEngine
|
||||
from app.services.log_config import is_enabled, LOG_DIR, LOG_DATE_FORMAT, encrypt_data
|
||||
from app.services.generation_provider_types import (
|
||||
from app.services.log_config import LOG_DATE_FORMAT, LOG_DIR, encrypt_data, is_enabled
|
||||
from app.types.generation.provider import (
|
||||
ImageProviderBatchResult,
|
||||
ImageProviderItem,
|
||||
ProviderGenerationRecordLike,
|
||||
ProviderImageEngineLike,
|
||||
)
|
||||
@@ -22,8 +27,39 @@ from app.services.generation_provider_types import (
|
||||
logger = logging.getLogger("videogen")
|
||||
|
||||
|
||||
class ImageProviderError(RuntimeError):
|
||||
"""可被生成任务状态机安全收敛的图片供应商异常。"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
message: str,
|
||||
*,
|
||||
error_type: ImageProviderErrorType = ImageProviderErrorType.UNKNOWN,
|
||||
error_code: str | None = None,
|
||||
retryable: bool = False,
|
||||
http_status: int | None = None,
|
||||
provider_request_id: str | None = None,
|
||||
):
|
||||
super().__init__(message)
|
||||
self.safe_message = message
|
||||
self.error_type = error_type
|
||||
self.error_code = error_code
|
||||
self.retryable = retryable
|
||||
self.http_status = http_status
|
||||
self.provider_request_id = provider_request_id
|
||||
|
||||
def as_dict(self) -> dict[str, Any]:
|
||||
return {
|
||||
"error_type": self.error_type.value,
|
||||
"error_code": self.error_code,
|
||||
"message": self.safe_message,
|
||||
"retryable": self.retryable,
|
||||
"http_status": self.http_status,
|
||||
"provider_request_id": self.provider_request_id,
|
||||
}
|
||||
|
||||
|
||||
def _log_image_request(engine: ProviderImageEngineLike, record_id: str, request_data: dict):
|
||||
"""Log image generation request to log/AiModel/YYYY-MM-DD.log"""
|
||||
if not is_enabled():
|
||||
return
|
||||
try:
|
||||
@@ -41,14 +77,13 @@ def _log_image_request(engine: ProviderImageEngineLike, record_id: str, request_
|
||||
"request": request_encrypted,
|
||||
"request_length": len(request_str),
|
||||
}
|
||||
with open(log_file, "a", encoding="utf-8") as f:
|
||||
f.write(json.dumps(entry, ensure_ascii=False) + "\n")
|
||||
with open(log_file, "a", encoding="utf-8") as file:
|
||||
file.write(json.dumps(entry, ensure_ascii=False) + "\n")
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
def _log_image_response(record_id: str, response_data: dict, error: str | None = None):
|
||||
"""Log image generation response to log/AiModel/YYYY-MM-DD.log"""
|
||||
if not is_enabled():
|
||||
return
|
||||
try:
|
||||
@@ -63,17 +98,13 @@ def _log_image_response(record_id: str, response_data: dict, error: str | None =
|
||||
"response": response_encrypted,
|
||||
"error": error,
|
||||
}
|
||||
with open(log_file, "a", encoding="utf-8") as f:
|
||||
f.write(json.dumps(entry, ensure_ascii=False) + "\n")
|
||||
with open(log_file, "a", encoding="utf-8") as file:
|
||||
file.write(json.dumps(entry, ensure_ascii=False) + "\n")
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
async def get_active_image_engine(db: AsyncSession) -> ImageEngine:
|
||||
"""Get the active image engine with highest priority."""
|
||||
result = await db.execute(
|
||||
select(ImageEngine)
|
||||
.where(ImageEngine.is_active == True)
|
||||
@@ -103,105 +134,291 @@ def _resolve_url(url: str) -> str:
|
||||
return f"{settings.BASE_URL.rstrip('/')}/{url.lstrip('/')}"
|
||||
|
||||
|
||||
def _value(obj: Any, name: str, default: Any = None) -> Any:
|
||||
if obj is None:
|
||||
return default
|
||||
if isinstance(obj, dict):
|
||||
return obj.get(name, default)
|
||||
return getattr(obj, name, default)
|
||||
|
||||
|
||||
def _jsonable(value: Any) -> Any:
|
||||
if value is None or isinstance(value, (str, int, float, bool)):
|
||||
return value
|
||||
if isinstance(value, dict):
|
||||
return {str(key): _jsonable(item) for key, item in value.items()}
|
||||
if isinstance(value, (list, tuple)):
|
||||
return [_jsonable(item) for item in value]
|
||||
if hasattr(value, "model_dump"):
|
||||
try:
|
||||
return _jsonable(value.model_dump())
|
||||
except Exception:
|
||||
pass
|
||||
if hasattr(value, "to_dict"):
|
||||
try:
|
||||
return _jsonable(value.to_dict())
|
||||
except Exception:
|
||||
pass
|
||||
result: dict[str, Any] = {}
|
||||
for key in ("url", "b64_json", "size", "output_format", "error", "code", "message"):
|
||||
item = getattr(value, key, None)
|
||||
if item is not None:
|
||||
result[key] = _jsonable(item)
|
||||
return result or str(value)
|
||||
|
||||
|
||||
def _safe_text(value: Any, *, limit: int = 1000) -> str:
|
||||
text = str(value or "").strip()
|
||||
return text[:limit]
|
||||
|
||||
|
||||
def _classify_provider_exception(exc: Exception) -> ImageProviderError:
|
||||
if isinstance(exc, ImageProviderError):
|
||||
return exc
|
||||
if isinstance(exc, (httpx.TimeoutException, TimeoutError)):
|
||||
return ImageProviderError(
|
||||
"图片生成请求超时,请稍后重试",
|
||||
error_type=ImageProviderErrorType.TIMEOUT,
|
||||
retryable=True,
|
||||
)
|
||||
|
||||
status_code = getattr(exc, "status_code", None)
|
||||
request_id = getattr(exc, "request_id", None) or getattr(exc, "x_request_id", None)
|
||||
code = getattr(exc, "code", None)
|
||||
raw_message = _safe_text(getattr(exc, "message", None) or exc)
|
||||
lowered = raw_message.lower()
|
||||
|
||||
if status_code == 429 or "rate limit" in lowered or "限流" in raw_message:
|
||||
error_type = ImageProviderErrorType.RATE_LIMIT
|
||||
retryable = True
|
||||
message = "图片生成请求过于频繁,请稍后重试"
|
||||
elif status_code in {401, 403} or "api key" in lowered or "unauthorized" in lowered:
|
||||
error_type = ImageProviderErrorType.AUTH
|
||||
retryable = False
|
||||
message = "图片引擎鉴权失败,请联系管理员检查配置"
|
||||
elif status_code and int(status_code) >= 500:
|
||||
error_type = ImageProviderErrorType.PROVIDER_INTERNAL
|
||||
retryable = True
|
||||
message = "图片供应商服务异常,请稍后重试"
|
||||
elif "sequential_image_generation" in lowered or "not support" in lowered or "unsupported" in lowered:
|
||||
error_type = ImageProviderErrorType.CAPABILITY_MISMATCH
|
||||
retryable = False
|
||||
message = "图片引擎组图能力配置与供应商实际能力不匹配,请联系管理员"
|
||||
elif "content" in lowered and ("risk" in lowered or "moderation" in lowered or "policy" in lowered):
|
||||
error_type = ImageProviderErrorType.CONTENT_REJECTED
|
||||
retryable = False
|
||||
message = "图片内容未通过供应商审核,请调整提示词后重试"
|
||||
elif status_code and 400 <= int(status_code) < 500:
|
||||
error_type = ImageProviderErrorType.INVALID_REQUEST
|
||||
retryable = False
|
||||
message = "图片生成参数不被供应商支持,请联系管理员检查引擎配置"
|
||||
elif isinstance(exc, httpx.HTTPError):
|
||||
error_type = ImageProviderErrorType.NETWORK
|
||||
retryable = True
|
||||
message = "图片供应商网络连接异常,请稍后重试"
|
||||
else:
|
||||
error_type = ImageProviderErrorType.UNKNOWN
|
||||
retryable = False
|
||||
message = raw_message or "图片生成失败"
|
||||
|
||||
return ImageProviderError(
|
||||
message,
|
||||
error_type=error_type,
|
||||
error_code=_safe_text(code, limit=128) or None,
|
||||
retryable=retryable,
|
||||
http_status=int(status_code) if status_code is not None else None,
|
||||
provider_request_id=_safe_text(request_id, limit=128) or None,
|
||||
)
|
||||
|
||||
|
||||
def build_multi_image_provider_prompt(prompt: str, generation_count: int) -> str:
|
||||
base_prompt = (prompt or "").strip()
|
||||
if generation_count <= 1:
|
||||
return base_prompt
|
||||
suffix = MULTI_IMAGE_PROMPT_TEMPLATE.format(count=generation_count)
|
||||
return f"{base_prompt}\n\n{suffix}" if base_prompt else suffix
|
||||
|
||||
|
||||
def submit_image_task(
|
||||
db,
|
||||
engine: ProviderImageEngineLike,
|
||||
record: ProviderGenerationRecordLike,
|
||||
*,
|
||||
include_media_references: bool,
|
||||
) -> dict:
|
||||
"""Submit an image generation task via Ark SDK. Returns image_url."""
|
||||
from volcenginesdkarkruntime import Ark
|
||||
|
||||
client = Ark(
|
||||
base_url=engine.api_base,
|
||||
api_key=engine.api_key,
|
||||
timeout=300,
|
||||
)
|
||||
generation_count: int = 1,
|
||||
) -> ImageProviderBatchResult:
|
||||
"""通过 Ark 同步图片接口生成单图或单次组图。
|
||||
|
||||
prompt = record.optimized_prompt or record.original_prompt
|
||||
image_urls = []
|
||||
generation_count > 1 时只执行一次 sequential_auto 请求;任何失败都直接抛出,
|
||||
绝不退化为多次单图请求。
|
||||
"""
|
||||
from volcenginesdkarkruntime import Ark
|
||||
|
||||
count = max(1, int(generation_count or 1))
|
||||
multi_generation_enabled = bool(getattr(engine, "multi_generation_enabled", False))
|
||||
max_generation_count = max(1, min(5, int(getattr(engine, "max_generation_count", 1) or 1)))
|
||||
if count > 1 and not multi_generation_enabled:
|
||||
raise ImageProviderError(
|
||||
"当前图片引擎未开启多份生成",
|
||||
error_type=ImageProviderErrorType.CAPABILITY_MISMATCH,
|
||||
)
|
||||
if count > max_generation_count:
|
||||
raise ImageProviderError(
|
||||
f"当前图片引擎最多允许生成 {max_generation_count} 份",
|
||||
error_type=ImageProviderErrorType.CAPABILITY_MISMATCH,
|
||||
)
|
||||
|
||||
client = Ark(base_url=engine.api_base, api_key=engine.api_key, timeout=300)
|
||||
original_prompt = record.optimized_prompt or record.original_prompt
|
||||
provider_prompt = build_multi_image_provider_prompt(original_prompt, count)
|
||||
image_urls: list[str] = []
|
||||
|
||||
if include_media_references and record.media_references:
|
||||
try:
|
||||
refs = json.loads(record.media_references)
|
||||
for ref in refs:
|
||||
ref_type = ref.get("type")
|
||||
ref_url = ref.get("url", "")
|
||||
if ref_type == "image" and ref_url:
|
||||
resolved = _resolve_url(ref_url)
|
||||
image_urls.append(resolved)
|
||||
for ref in refs if isinstance(refs, list) else []:
|
||||
if (ref.get("type") or "").lower() == "image" and ref.get("url"):
|
||||
image_urls.append(_resolve_url(ref["url"]))
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
pass
|
||||
image_urls = []
|
||||
|
||||
request_payload = {
|
||||
request_log_payload: dict[str, Any] = {
|
||||
"model": engine.model_name,
|
||||
"prompt": prompt,
|
||||
"prompt": provider_prompt,
|
||||
"size": record.image_size or engine.default_size,
|
||||
"sequential_image_generation": "disabled",
|
||||
"output_format": "png",
|
||||
"response_format": "url",
|
||||
"watermark": False,
|
||||
"include_media_references": include_media_references,
|
||||
}
|
||||
|
||||
request_sdk_payload: dict[str, Any] = dict(request_log_payload)
|
||||
if image_urls:
|
||||
request_payload["image"] = image_urls
|
||||
request_log_payload["image"] = image_urls
|
||||
request_sdk_payload["image"] = image_urls
|
||||
output_format = (getattr(engine, "output_format", "") or "").lower().strip()
|
||||
if output_format:
|
||||
request_log_payload["output_format"] = output_format
|
||||
request_sdk_payload["output_format"] = output_format
|
||||
if count > 1:
|
||||
try:
|
||||
from volcenginesdkarkruntime.types.images import SequentialImageGenerationOptions
|
||||
except Exception:
|
||||
try:
|
||||
from volcenginesdkarkruntime.types.images.image_generate_params import (
|
||||
SequentialImageGenerationOptions,
|
||||
)
|
||||
except Exception as import_exc:
|
||||
raise ImageProviderError(
|
||||
"当前图片引擎运行依赖缺少组图参数对象,请升级火山 Ark SDK 后重试",
|
||||
error_type=ImageProviderErrorType.CAPABILITY_MISMATCH,
|
||||
) from import_exc
|
||||
|
||||
_log_image_request(engine, record.id, request_payload)
|
||||
request_log_payload["sequential_image_generation"] = "auto"
|
||||
request_log_payload["sequential_image_generation_options"] = {"max_images": count}
|
||||
request_log_payload["stream"] = False
|
||||
|
||||
request_sdk_payload["sequential_image_generation"] = "auto"
|
||||
request_sdk_payload["sequential_image_generation_options"] = SequentialImageGenerationOptions(
|
||||
max_images=count,
|
||||
)
|
||||
request_sdk_payload["stream"] = False
|
||||
|
||||
_log_image_request(engine, record.id, request_log_payload)
|
||||
|
||||
try:
|
||||
result = client.images.generate(
|
||||
model=engine.model_name,
|
||||
prompt=prompt,
|
||||
size=record.image_size or engine.default_size,
|
||||
output_format="png",
|
||||
response_format="url",
|
||||
watermark=False,
|
||||
image=image_urls if image_urls else None,
|
||||
)
|
||||
image_url = result.data[0].url
|
||||
|
||||
response_data = {
|
||||
"model": result.model,
|
||||
"created": result.created,
|
||||
"data": [{"url": item.url, "size": item.size} for item in result.data] if result.data else [],
|
||||
"usage": {
|
||||
"generated_images": result.usage.generated_images if hasattr(result.usage, 'generated_images') else 0,
|
||||
"output_tokens": result.usage.output_tokens if hasattr(result.usage, 'output_tokens') else 0,
|
||||
"total_tokens": result.usage.total_tokens if hasattr(result.usage, 'total_tokens') else 0,
|
||||
result = client.images.generate(**request_sdk_payload)
|
||||
top_error = _value(result, "error")
|
||||
if top_error:
|
||||
error_code = _value(top_error, "code")
|
||||
error_message = _value(top_error, "message") or str(top_error)
|
||||
raise ImageProviderError(
|
||||
_safe_text(error_message) or "图片供应商返回失败",
|
||||
error_type=ImageProviderErrorType.INVALID_REQUEST,
|
||||
error_code=_safe_text(error_code, limit=128) or None,
|
||||
)
|
||||
|
||||
raw_data = _value(result, "data", []) or []
|
||||
if not isinstance(raw_data, (list, tuple)):
|
||||
raise ImageProviderError(
|
||||
"图片供应商返回 data 结构异常",
|
||||
error_type=ImageProviderErrorType.INVALID_RESPONSE,
|
||||
)
|
||||
|
||||
items: list[ImageProviderItem] = []
|
||||
response_items: list[dict[str, Any]] = []
|
||||
for index, raw_item in enumerate(raw_data, start=1):
|
||||
item_error = _value(raw_item, "error")
|
||||
if item_error:
|
||||
error_code = _safe_text(_value(item_error, "code"), limit=128)
|
||||
error_message = _safe_text(_value(item_error, "message") or item_error)
|
||||
items.append({
|
||||
"generation_index": index,
|
||||
"error_code": error_code,
|
||||
"error_message": error_message or "单张图片生成失败",
|
||||
"response_data": _jsonable(raw_item),
|
||||
})
|
||||
response_items.append(_jsonable(raw_item))
|
||||
continue
|
||||
|
||||
url = _safe_text(_value(raw_item, "url"), limit=4000)
|
||||
b64_json = _safe_text(_value(raw_item, "b64_json"), limit=100) if not url else ""
|
||||
item: ImageProviderItem = {
|
||||
"generation_index": index,
|
||||
"remote_result_url": url,
|
||||
"size": _safe_text(_value(raw_item, "size"), limit=64),
|
||||
"output_format": _safe_text(_value(raw_item, "output_format"), limit=32),
|
||||
"response_data": _jsonable(raw_item),
|
||||
}
|
||||
if b64_json:
|
||||
item["b64_json"] = b64_json
|
||||
items.append(item)
|
||||
response_items.append(_jsonable(raw_item))
|
||||
|
||||
usage = _value(result, "usage")
|
||||
generated_images = int(_value(usage, "generated_images", 0) or 0)
|
||||
total_tokens = int(_value(usage, "total_tokens", 0) or 0)
|
||||
response_data = {
|
||||
"model": _value(result, "model", engine.model_name),
|
||||
"created": _value(result, "created"),
|
||||
"data": response_items,
|
||||
"usage": {
|
||||
"generated_images": generated_images,
|
||||
"input_images": int(_value(usage, "input_images", 0) or 0),
|
||||
"output_tokens": int(_value(usage, "output_tokens", 0) or 0),
|
||||
"total_tokens": total_tokens,
|
||||
},
|
||||
}
|
||||
except httpx.TimeoutException:
|
||||
error_msg = "图片生成超时,请稍后重试"
|
||||
logger.error(f"Image generation timeout for record {record.id}")
|
||||
_log_image_response(record.id, {}, error_msg)
|
||||
raise TimeoutError(error_msg)
|
||||
except Exception as e:
|
||||
error_msg = str(e)
|
||||
logger.error(f"Image generation failed for record {record.id}: {error_msg}")
|
||||
_log_image_response(record.id, {}, error_msg)
|
||||
raise
|
||||
_log_image_response(record.id, response_data)
|
||||
return {
|
||||
"items": items,
|
||||
"model": str(response_data["model"] or ""),
|
||||
"created": int(response_data["created"] or 0),
|
||||
"generated_images": generated_images,
|
||||
"image_tokens": total_tokens,
|
||||
"response_data": response_data,
|
||||
}
|
||||
except Exception as exc:
|
||||
provider_error = _classify_provider_exception(exc)
|
||||
logger.error(
|
||||
"Image generation failed for record %s: type=%s code=%s message=%s",
|
||||
record.id,
|
||||
provider_error.error_type.value,
|
||||
provider_error.error_code,
|
||||
provider_error.safe_message,
|
||||
)
|
||||
_log_image_response(record.id, provider_error.as_dict(), provider_error.safe_message)
|
||||
raise provider_error from exc
|
||||
finally:
|
||||
client.close()
|
||||
|
||||
return {
|
||||
"image_url": image_url,
|
||||
"image_tokens": getattr(result.usage, "total_tokens", 0),
|
||||
"response_data": json.dumps(response_data, ensure_ascii=False, default=str),
|
||||
"error": str(result.error) if result.error else "",
|
||||
}
|
||||
try:
|
||||
client.close()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
async def poll_image_task_status(engine: ImageEngine, task_id: str) -> dict:
|
||||
"""Query image task status via Ark SDK. Returns {status, image_url, response_data}."""
|
||||
client = AsyncArk(
|
||||
base_url=engine.api_base,
|
||||
api_key=engine.api_key,
|
||||
)
|
||||
|
||||
result = await client.image_generation.tasks.get(task_id=task_id)
|
||||
await client.close()
|
||||
client = AsyncArk(base_url=engine.api_base, api_key=engine.api_key)
|
||||
try:
|
||||
result = await client.image_generation.tasks.get(task_id=task_id)
|
||||
finally:
|
||||
await client.close()
|
||||
|
||||
response_dict = {
|
||||
"id": result.id,
|
||||
@@ -239,13 +456,11 @@ async def poll_image_task_status(engine: ImageEngine, task_id: str) -> dict:
|
||||
|
||||
|
||||
async def download_image(image_url: str, dest_path: str) -> str:
|
||||
"""Download image to local storage."""
|
||||
os.makedirs(os.path.dirname(dest_path), exist_ok=True)
|
||||
|
||||
async with httpx.AsyncClient(timeout=300) as client:
|
||||
async with client.stream("GET", image_url) as response:
|
||||
response.raise_for_status()
|
||||
with open(dest_path, "wb") as f:
|
||||
with open(dest_path, "wb") as file:
|
||||
async for chunk in response.aiter_bytes(chunk_size=8192):
|
||||
f.write(chunk)
|
||||
return dest_path
|
||||
file.write(chunk)
|
||||
return dest_path
|
||||
|
||||
@@ -4,7 +4,7 @@ from collections.abc import Iterable
|
||||
from datetime import datetime
|
||||
from typing import Any, TypedDict
|
||||
|
||||
from sqlalchemy import and_, func, or_, select
|
||||
from sqlalchemy import and_, case, func, or_, select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.enums.generation_task import GenerationType
|
||||
@@ -12,7 +12,7 @@ from app.enums.recent_generation import (
|
||||
RECENT_GENERATION_ALL_MODULES,
|
||||
RECENT_GENERATION_CHAT_TASK_MODULES,
|
||||
RECENT_GENERATION_COMPLETED_STATUS,
|
||||
RECENT_GENERATION_MODULE_TO_TASK_MODE,
|
||||
RECENT_GENERATION_MODULE_TO_TASK_MODES,
|
||||
RECENT_GENERATION_TASK_MODE_VALUE_TO_MODULE,
|
||||
RecentGenerationModuleEnum,
|
||||
RecentGenerationResourceTypeEnum,
|
||||
@@ -175,11 +175,12 @@ async def _list_chat_task_recent_rows(
|
||||
modules: list[RecentGenerationModuleEnum],
|
||||
limit: int,
|
||||
) -> list[dict[str, Any]]:
|
||||
task_mode_values = [
|
||||
RECENT_GENERATION_MODULE_TO_TASK_MODE[module].value
|
||||
task_mode_values = list(dict.fromkeys(
|
||||
task_mode.value
|
||||
for module in modules
|
||||
if module in RECENT_GENERATION_CHAT_TASK_MODULES
|
||||
]
|
||||
for task_mode in RECENT_GENERATION_MODULE_TO_TASK_MODES[module]
|
||||
))
|
||||
if not task_mode_values:
|
||||
return []
|
||||
|
||||
@@ -189,10 +190,19 @@ async def _list_chat_task_recent_rows(
|
||||
ChatGenerationTask.created_at,
|
||||
)
|
||||
|
||||
module_partition_expr = case(
|
||||
(
|
||||
ChatGenerationTask.generation_mode.in_(["chatapi_async", "chatapi_child"]),
|
||||
RecentGenerationModuleEnum.CHAT_AI.value,
|
||||
),
|
||||
else_=ChatGenerationTask.generation_mode,
|
||||
)
|
||||
|
||||
ranked_subquery = (
|
||||
select(
|
||||
ChatGenerationTask.id.label("generation_id"),
|
||||
ChatGenerationTask.generation_mode.label("generation_mode"),
|
||||
module_partition_expr.label("module_key"),
|
||||
ChatGenerationTask.gen_type.label("gen_type"),
|
||||
ChatGenerationTask.image_url.label("image_url"),
|
||||
ChatGenerationTask.video_url.label("video_url"),
|
||||
@@ -200,7 +210,7 @@ async def _list_chat_task_recent_rows(
|
||||
generated_time_expr.label("generated_time"),
|
||||
func.row_number()
|
||||
.over(
|
||||
partition_by=ChatGenerationTask.generation_mode,
|
||||
partition_by=module_partition_expr,
|
||||
order_by=(generated_time_expr.desc(), ChatGenerationTask.created_at.desc()),
|
||||
)
|
||||
.label("row_num"),
|
||||
@@ -218,7 +228,7 @@ async def _list_chat_task_recent_rows(
|
||||
stmt = (
|
||||
select(ranked_subquery)
|
||||
.where(ranked_subquery.c.row_num <= limit)
|
||||
.order_by(ranked_subquery.c.generation_mode.asc(), ranked_subquery.c.generated_time.desc())
|
||||
.order_by(ranked_subquery.c.module_key.asc(), ranked_subquery.c.generated_time.desc())
|
||||
)
|
||||
|
||||
return [dict(row) for row in (await db.execute(stmt)).mappings().all()]
|
||||
|
||||
@@ -1,14 +1,12 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from copy import deepcopy
|
||||
from datetime import datetime, timezone
|
||||
from datetime import datetime
|
||||
from typing import Any
|
||||
|
||||
from fastapi import HTTPException
|
||||
from sqlalchemy import func, select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from sqlalchemy.orm.attributes import flag_modified
|
||||
|
||||
from app.config import settings
|
||||
from app.enums.common import ModuleEventTypeEnum, ModuleProjectStatusEnum, ModulePromptTypeEnum, ModuleStepStatusEnum
|
||||
@@ -35,16 +33,16 @@ from app.schemas.shot_replicate import (
|
||||
ShotReplicateVideoGenerationOut,
|
||||
ShotReplicateVideoPromptSchemaUpdateRequest,
|
||||
)
|
||||
from app.services.generation_ai_service import (
|
||||
from app.services.generation.ai.engine_service import (
|
||||
VIDEO_DEFAULT_DURATION,
|
||||
VIDEO_DEFAULT_RATIO,
|
||||
VIDEO_DEFAULT_RESOLUTION,
|
||||
_get_video_engine,
|
||||
_parse_list,
|
||||
get_video_engine,
|
||||
parse_json_list,
|
||||
)
|
||||
from app.services.generation_billing_service import charge_module_prompt_usage
|
||||
from app.services.generation_refund_service import mark_chat_generation_task_failed_and_refund_once
|
||||
from app.services.generation_task_factory_service import create_chat_generation_task_for_module
|
||||
from app.services.generation.billing_service import charge_module_prompt_usage
|
||||
from app.services.generation.refund_service import mark_chat_generation_task_failed_and_refund_once
|
||||
from app.services.generation.task_factory_service import create_chat_generation_task_for_module
|
||||
from app.services.hot_opening_video_prompt_service import (
|
||||
build_final_video_prompt,
|
||||
optimize_hot_opening_video_prompt as optimize_shot_replicate_video_prompt,
|
||||
@@ -52,7 +50,6 @@ from app.services.hot_opening_video_prompt_service import (
|
||||
)
|
||||
from app.services.module_generation_log_service import log_module_error, log_module_event_file, log_module_prompt_event
|
||||
from app.services.llm import optimize_prompt
|
||||
from app.services.resource_accounting_service import soft_delete_chat_task_resources
|
||||
from app.services.module_generation_flow_base_service import (
|
||||
assert_project_has_no_active_chat_tasks as _base_assert_project_has_no_active_chat_tasks,
|
||||
chat_tasks_by_id as _base_chat_tasks_by_id,
|
||||
@@ -976,10 +973,10 @@ async def generate_image_from_prompt(
|
||||
|
||||
|
||||
async def _resolve_video_prompt_config(db: AsyncSession, req: ShotReplicateGenerateVideoPromptRequest) -> dict[str, Any]:
|
||||
engine = await _get_video_engine(db, req.engine_id)
|
||||
supported_ratios = _parse_list(engine.supported_ratios, [])
|
||||
supported_resolutions = _parse_list(engine.supported_resolutions, [])
|
||||
supported_durations = _parse_list(engine.supported_durations, [])
|
||||
engine = await get_video_engine(db, req.engine_id)
|
||||
supported_ratios = parse_json_list(engine.supported_ratios, [])
|
||||
supported_resolutions = parse_json_list(engine.supported_resolutions, [])
|
||||
supported_durations = parse_json_list(engine.supported_durations, [])
|
||||
|
||||
default_ratio = getattr(settings, "SHOT_REPLICATE_DEFAULT_VIDEO_RATIO", None) or VIDEO_DEFAULT_RATIO
|
||||
default_resolution = getattr(settings, "SHOT_REPLICATE_DEFAULT_VIDEO_RESOLUTION", None) or VIDEO_DEFAULT_RESOLUTION
|
||||
|
||||
@@ -14,7 +14,7 @@ from app.config import settings
|
||||
from app.enums.private_portrait import PRIVATE_PORTRAIT_ASSET_URI_PREFIX
|
||||
from app.models.video_engine import VideoEngine
|
||||
from app.services.log_config import is_enabled, LOG_DIR, LOG_DATE_FORMAT, encrypt_data
|
||||
from app.services.generation_provider_types import (
|
||||
from app.types.generation.provider import (
|
||||
ProviderGenerationRecordLike,
|
||||
ProviderVideoEngineLike,
|
||||
)
|
||||
|
||||
@@ -17,7 +17,7 @@ from app.services.resource_accounting_service import (
|
||||
)
|
||||
from app.services.video_cover_service import create_video_cover_for_local_video
|
||||
from app.config import settings
|
||||
from app.services.generation_refund_service import mark_generation_record_failed_and_refund_once
|
||||
from app.services.generation.refund_service import mark_generation_record_failed_and_refund_once
|
||||
|
||||
logger = logging.getLogger("videogen")
|
||||
|
||||
|
||||
Reference in New Issue
Block a user