988 lines
35 KiB
Python
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
|
|
],
|
|
}
|
|
|