Files
video-gen/video-gen-api/app/services/generation_ai_service.py
2026-07-11 12:51:48 +08:00

1206 lines
44 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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,
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
],
}
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)