Files
video-gen/video-gen-api/app/services/generation/ai/service.py
T

988 lines
35 KiB
Python

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
],
}