1140 lines
42 KiB
Python
1140 lines
42 KiB
Python
from __future__ import annotations
|
||
|
||
import json
|
||
from datetime import datetime, timedelta, timezone, date
|
||
from typing import Any
|
||
|
||
from fastapi import HTTPException
|
||
from sqlalchemy import and_, func, select
|
||
from sqlalchemy.ext.asyncio import AsyncSession
|
||
|
||
from app.config import settings
|
||
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.audio_reference import (
|
||
AUDIO_ALLOWED_EXTENSIONS,
|
||
AUDIO_MAX_COUNT_LIMIT,
|
||
AUDIO_MAX_DURATION_SECONDS,
|
||
AUDIO_MAX_TOTAL_DURATION_SECONDS,
|
||
AUDIO_MIN_DURATION_SECONDS,
|
||
)
|
||
from app.enums.generation_history import (
|
||
GenerationHistorySourceEnum,
|
||
get_generation_history_source_label,
|
||
get_generation_history_task_mode,
|
||
normalize_generation_history_source,
|
||
)
|
||
from app.schemas.generation_ai import (
|
||
GenerationAIEngineGroupOut,
|
||
GenerationAIEngineOptionsOut,
|
||
GenerationAIImageEngineOptionOut,
|
||
GenerationAIRecordHistoryItemOut,
|
||
GenerationAITaskCreate,
|
||
GenerationAITaskOut,
|
||
GenerationAIVideoEngineOptionOut,
|
||
)
|
||
from app.services.generation_billing_service import (
|
||
OWNER_CHAT_GENERATION_TASK,
|
||
charge_generation_media_by_params,
|
||
)
|
||
from app.services.resource_accounting_service import (
|
||
SOURCE_MODEL_CHAT_TASK,
|
||
SOURCE_MODEL_GENERATION_RECORD,
|
||
batch_get_generated_resource_info_map,
|
||
soft_delete_chat_task_resources,
|
||
)
|
||
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.resource_capacity_service import assert_user_resource_capacity_available
|
||
from app.services.private_portrait.reference_resolver import batch_resolve_private_portrait_reference_display_urls, resolve_private_portrait_reference_display_urls, resolve_private_portrait_references
|
||
from app.utils.id_gen import generate_id
|
||
|
||
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"
|
||
|
||
HISTORY_DAY_PAGE_SIZE_MAX = 10
|
||
HISTORY_GROUP_ITEM_LIMIT = 10
|
||
|
||
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 _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 _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())
|
||
|
||
query = query.limit(1)
|
||
result = await db.execute(query)
|
||
engine = result.scalar_one_or_none()
|
||
if not engine:
|
||
raise HTTPException(status_code=400, detail="没有可用的视频引擎")
|
||
return engine
|
||
|
||
|
||
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 _parse_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 _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_list(engine.supported_models, []),
|
||
"default_size": engine.default_size,
|
||
"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_list(engine.supported_ratios, []),
|
||
"supported_resolutions": _parse_list(engine.supported_resolutions, []),
|
||
"supported_durations": _parse_list(engine.supported_durations, []),
|
||
"max_duration": engine.max_duration,
|
||
"max_audio_count": engine.max_audio_count,
|
||
"selected_ratio": ratio,
|
||
"selected_resolution": resolution,
|
||
"selected_duration": duration,
|
||
}
|
||
|
||
|
||
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_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,
|
||
)
|
||
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_list(engine.supported_ratios, []),
|
||
supported_resolutions=_parse_list(engine.supported_resolutions, []),
|
||
supported_durations=_parse_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,
|
||
)
|
||
for engine in video_result.scalars().all()
|
||
]
|
||
|
||
return GenerationAIEngineOptionsOut(
|
||
engine=GenerationAIEngineGroupOut(image=image_items, video=video_items)
|
||
)
|
||
|
||
|
||
async def create_async_generation_task(db: AsyncSession, current_user: User, req: GenerationAITaskCreate) -> ChatGenerationTask:
|
||
"""Create a project-independent chat generation task.
|
||
|
||
Important: this writes chat_generation_tasks, not generation_records, so chat
|
||
image/video generation no longer needs or validates a project_id.
|
||
"""
|
||
gen_type = req.gen_type.lower().strip()
|
||
if gen_type not in ("image", "video"):
|
||
raise HTTPException(status_code=400, detail="gen_type 仅支持 image 或 video")
|
||
|
||
if req.idempotency_key:
|
||
result = await db.execute(
|
||
select(ChatGenerationTask).where(
|
||
ChatGenerationTask.user_id == current_user.id,
|
||
ChatGenerationTask.idempotency_key == req.idempotency_key,
|
||
ChatGenerationTask.generation_mode == "chatapi_async",
|
||
ChatGenerationTask.deleted_at.is_(None),
|
||
).order_by(ChatGenerationTask.created_at.desc()).limit(1)
|
||
)
|
||
existing = result.scalar_one_or_none()
|
||
if existing:
|
||
return existing
|
||
|
||
refs = [r.model_dump(exclude_none=True) for r in (req.media_references or [])]
|
||
refs = await resolve_private_portrait_references(
|
||
db,
|
||
user_id=current_user.id,
|
||
media_references=refs,
|
||
gen_type=gen_type,
|
||
)
|
||
now = datetime.now(timezone.utc)
|
||
task_id = generate_id()
|
||
|
||
await assert_user_resource_capacity_available(db, current_user.id)
|
||
|
||
if gen_type == "image":
|
||
if any((r.get("type") or "").lower() == "audio" for r in refs):
|
||
raise HTTPException(status_code=400, detail="图片生成不支持音频参考素材")
|
||
engine = await _get_image_engine(db, req.engine_id)
|
||
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
|
||
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="AI生成任务",
|
||
description_prefix="AI创作-",
|
||
owner_type=OWNER_CHAT_GENERATION_TASK,
|
||
attempt_no=1,
|
||
)
|
||
snapshot = _build_image_snapshot(engine, size, proportion, px)
|
||
task = ChatGenerationTask(
|
||
id=task_id,
|
||
user_id=current_user.id,
|
||
original_prompt=req.original_prompt,
|
||
gen_type="image",
|
||
image_size=size,
|
||
image_proportion=proportion,
|
||
image_px=px,
|
||
status="generating",
|
||
generation_mode="chatapi_async",
|
||
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=req.idempotency_key,
|
||
deadline_at=now + timedelta(minutes=settings.CHATAPI_ASYNC_IMAGE_DEADLINE_MINUTES),
|
||
)
|
||
else:
|
||
engine = await _get_video_engine(db, req.engine_id)
|
||
ratio = req.aspect_ratio or VIDEO_DEFAULT_RATIO
|
||
resolution = req.resolution or VIDEO_DEFAULT_RESOLUTION
|
||
duration = req.duration or VIDEO_DEFAULT_DURATION
|
||
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 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 = 0.0
|
||
if refs:
|
||
video_refs = [r for r in refs if (r.get("type") or "").lower() == "video"]
|
||
for ref in video_refs:
|
||
ref_duration = float(ref.get("duration") or 0)
|
||
if ref_duration < 2:
|
||
raise HTTPException(status_code=400, detail=f"视频素材最短不能少于 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 = [r for r in refs if (r.get("type") or "").lower() == "audio"]
|
||
if audio_refs:
|
||
max_audio_count = int(engine.max_audio_count or 0)
|
||
if max_audio_count <= 0:
|
||
raise HTTPException(status_code=400, detail="当前视频引擎不支持音频参考素材")
|
||
if max_audio_count > AUDIO_MAX_COUNT_LIMIT:
|
||
max_audio_count = AUDIO_MAX_COUNT_LIMIT
|
||
if len(audio_refs) > max_audio_count:
|
||
raise HTTPException(
|
||
status_code=400,
|
||
detail=f"参考音频最多可传 {max_audio_count} 段,当前 {len(audio_refs)} 段",
|
||
)
|
||
|
||
input_audio_duration = 0.0
|
||
for ref in audio_refs:
|
||
raw_duration = ref.get("duration")
|
||
if raw_duration is None:
|
||
raw_duration = 0.0
|
||
try:
|
||
ref_duration = float(raw_duration)
|
||
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} 秒",
|
||
)
|
||
|
||
media_billing = await charge_generation_media_by_params(
|
||
db,
|
||
user_id=current_user.id,
|
||
record_id=task_id,
|
||
gen_type="video",
|
||
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,
|
||
)
|
||
snapshot = _build_video_snapshot(engine, ratio, resolution, duration)
|
||
task = ChatGenerationTask(
|
||
id=task_id,
|
||
user_id=current_user.id,
|
||
original_prompt=req.original_prompt,
|
||
gen_type="video",
|
||
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,
|
||
status="generating",
|
||
generation_mode="chatapi_async",
|
||
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=req.idempotency_key,
|
||
deadline_at=now + timedelta(hours=settings.CHATAPI_ASYNC_VIDEO_FINAL_DEADLINE_HOURS),
|
||
)
|
||
|
||
db.add(task)
|
||
await db.flush()
|
||
return task
|
||
|
||
|
||
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,
|
||
) -> 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:
|
||
source = GenerationHistorySourceEnum(
|
||
"chat_task" if task.generation_mode == "chatapi_async" else str(task.generation_mode or "chat_task")
|
||
)
|
||
except ValueError:
|
||
source = GenerationHistorySourceEnum.CHAT_TASK
|
||
meta = history_meta or build_empty_history_meta(source)
|
||
|
||
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=generated_resource_id,
|
||
file_name=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,
|
||
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,
|
||
# remote_result_url=task.remote_result_url,
|
||
image_url=build_resource_signed_url(task.image_url) if task.image_url else "",
|
||
video_url=build_resource_signed_url(task.video_url) if task.video_url else "",
|
||
video_cover_url=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=_resolve_error_message(task.error_message),
|
||
created_at=task.created_at,
|
||
generated_at=task.generated_at,
|
||
)
|
||
|
||
def engine_snapshot_out(snapshot: dict) -> dict:
|
||
"""
|
||
从完整的 engine_snapshot 中过滤出需要返回的字段
|
||
"""
|
||
if not snapshot:
|
||
return {}
|
||
|
||
return {
|
||
"engine_type": snapshot.get("engine_type"),
|
||
"id": snapshot.get("id"),
|
||
"name": snapshot.get("name"),
|
||
"provider": snapshot.get("provider"),
|
||
# "api_base": snapshot.get("api_base"),
|
||
# "api_key_masked": snapshot.get("api_key_masked"),
|
||
"model_name": snapshot.get("model_name"),
|
||
# "generate_url": snapshot.get("generate_url"),
|
||
"supported_models": snapshot.get("supported_models", []),
|
||
"default_size": snapshot.get("default_size"),
|
||
"selected_size": snapshot.get("selected_size"),
|
||
"selected_proportion": snapshot.get("selected_proportion"),
|
||
"selected_px": snapshot.get("selected_px")
|
||
}
|
||
|
||
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 == "chatapi_async",
|
||
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())
|
||
.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_mode = get_generation_history_task_mode(source)
|
||
if not task_mode:
|
||
raise HTTPException(status_code=400, detail="history_source 不支持查询 ChatGenerationTask 历史")
|
||
return [
|
||
ChatGenerationTask.user_id == user_id,
|
||
ChatGenerationTask.generation_mode == task_mode.value,
|
||
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,
|
||
):
|
||
"""
|
||
按生成日期倒序返回旧 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)
|
||
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()
|
||
|
||
raw_groups = []
|
||
all_record_ids: list[str] = []
|
||
for generated_day, day_total in day_rows:
|
||
item_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) == generated_day,
|
||
)
|
||
.order_by(GenerationRecord.generated_at.desc(), GenerationRecord.created_at.desc())
|
||
.limit(HISTORY_GROUP_ITEM_LIMIT)
|
||
)
|
||
rows = item_result.all()
|
||
raw_groups.append((generated_day, day_total, rows))
|
||
all_record_ids.extend(record.id for record, _project_name in 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,
|
||
):
|
||
"""
|
||
获取旧 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)
|
||
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,
|
||
):
|
||
"""
|
||
按生成日期倒序返回历史记录分组。
|
||
|
||
- 每页最多返回 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,
|
||
)
|
||
|
||
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)
|
||
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()
|
||
|
||
raw_groups = []
|
||
all_task_ids: list[str] = []
|
||
for generated_day, day_total in day_rows:
|
||
item_result = await db.execute(
|
||
select(ChatGenerationTask)
|
||
.where(
|
||
*filters,
|
||
func.date(ChatGenerationTask.generated_at) == generated_day,
|
||
)
|
||
.order_by(ChatGenerationTask.generated_at.desc(), ChatGenerationTask.created_at.desc())
|
||
.limit(HISTORY_GROUP_ITEM_LIMIT)
|
||
)
|
||
tasks = list(item_result.scalars().all())
|
||
raw_groups.append((generated_day, day_total, tasks))
|
||
all_task_ids.extend(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=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 = [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, 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,
|
||
):
|
||
"""
|
||
获取指定生成日期下的历史记录分页。
|
||
|
||
用于前端点击某一天后,继续加载该日期下的第 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,
|
||
)
|
||
|
||
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)
|
||
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
|
||
],
|
||
}
|
||
|
||
async def soft_delete_chat_generation_task(
|
||
db: AsyncSession,
|
||
*,
|
||
task: ChatGenerationTask,
|
||
deleted_at: datetime | None = None,
|
||
) -> int:
|
||
"""软删 ChatGenerationTask 并联动软删资源账本,返回释放的 active 空间字节数。"""
|
||
deleted_at = deleted_at or datetime.now(timezone.utc)
|
||
task.deleted_at = deleted_at
|
||
return await soft_delete_chat_task_resources(db, task.id, deleted_at=deleted_at)
|