Files
video-gen/video-gen-api/app/services/generation_ai_service.py
T
2026-07-01 11:56:30 +08:00

983 lines
34 KiB
Python
Raw 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 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.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.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 _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,
"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,
)
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() for r in (req.media_references or [])]
now = datetime.now(timezone.utc)
task_id = generate_id()
await assert_user_resource_capacity_available(db, current_user.id)
if gen_type == "image":
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} 秒")
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,
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(minutes=settings.CHATAPI_ASYNC_VIDEO_DEADLINE_MINUTES),
)
db.add(task)
await db.flush()
return task
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,
) -> GenerationAITaskOut:
refs = _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=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,
):
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)
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,
) -> GenerationAIRecordHistoryItemOut:
refs = _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=record.error_message,
created_at=record.created_at,
generated_at=record.generated_at,
)
async def list_generation_record_history_grouped_days(
db: AsyncSession,
user_id: str,
gen_type: str,
page: int,
page_size: int,
):
"""
按生成日期倒序返回旧 generation_records 历史记录分组。
- 每页最多返回 10 个生成日期
- 每个日期分组内最多返回倒序前 10 条旧记录
- 只返回 completed 成功记录
"""
gen_type = _normalize_history_gen_type(gen_type)
page = max(page, 1)
page_size = min(max(page_size, 1), HISTORY_DAY_PAGE_SIZE_MAX)
filters = _generation_record_history_base_filters(user_id, gen_type)
day_expr = func.date(GenerationRecord.generated_at).label("generated_date")
days_subquery = (
select(day_expr)
.where(*filters)
.group_by(day_expr)
.subquery()
)
total_days = (
await db.execute(select(func.count()).select_from(days_subquery))
).scalar_one()
day_rows_result = await db.execute(
select(
day_expr,
func.count(GenerationRecord.id).label("total"),
)
.where(*filters)
.group_by(day_expr)
.order_by(day_expr.desc())
.offset((page - 1) * page_size)
.limit(page_size)
)
day_rows = day_rows_result.all()
raw_groups = []
all_record_ids: list[str] = []
for generated_day, day_total in day_rows:
item_result = await db.execute(
select(GenerationRecord, Project.name.label("project_name"))
.outerjoin(Project, (GenerationRecord.project_id == Project.id) & (Project.deleted_at.is_(None)))
.where(
*filters,
func.date(GenerationRecord.generated_at) == generated_day,
)
.order_by(GenerationRecord.generated_at.desc(), GenerationRecord.created_at.desc())
.limit(HISTORY_GROUP_ITEM_LIMIT)
)
rows = item_result.all()
raw_groups.append((generated_day, day_total, rows))
all_record_ids.extend(record.id for record, _project_name in rows)
resource_info_map = await batch_get_generated_resource_info_map(
db,
source_model=SOURCE_MODEL_GENERATION_RECORD,
source_ids=all_record_ids,
resource_type=gen_type,
)
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"),
)
for record, project_name in rows
],
}
for generated_day, day_total, rows in raw_groups
]
return {
"total_days": int(total_days or 0),
"page": page,
"page_size": page_size,
"groups": groups,
}
async def list_generation_record_history_day_items(
db: AsyncSession,
user_id: str,
gen_type: str,
generated_date: str,
page: int,
page_size: int,
):
"""
获取旧 generation_records 指定生成日期下的历史记录分页。
"""
gen_type = _normalize_history_gen_type(gen_type)
target_day = _parse_history_date(generated_date)
page = max(page, 1)
page_size = min(max(page_size, 1), 100)
filters = _generation_record_history_base_filters(user_id, gen_type)
day_expr = func.date(GenerationRecord.generated_at)
total = (
await db.execute(
select(func.count(GenerationRecord.id)).where(
*filters,
day_expr == target_day,
)
)
).scalar_one()
result = await db.execute(
select(GenerationRecord, Project.name.label("project_name"))
.outerjoin(Project, (GenerationRecord.project_id == Project.id) & (Project.deleted_at.is_(None)))
.where(
*filters,
day_expr == target_day,
)
.order_by(GenerationRecord.generated_at.desc(), GenerationRecord.created_at.desc())
.offset((page - 1) * page_size)
.limit(page_size)
)
rows = result.all()
resource_info_map = await batch_get_generated_resource_info_map(
db,
source_model=SOURCE_MODEL_GENERATION_RECORD,
source_ids=[record.id for record, _project_name in rows],
resource_type=gen_type,
)
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"),
)
for record, project_name in rows
],
}
async def list_generation_history_grouped_days(
db: AsyncSession,
user_id: str,
gen_type: str,
page: int,
page_size: int,
history_source: str | None = None,
):
"""
按生成日期倒序返回历史记录分组。
- 每页最多返回 10 个生成日期
- 每个日期分组内最多返回倒序前 10 条任务
- 只返回 completed 成功任务
- history_source 支持 chat_task / generation_record / hot_opening_replicate / shot_replicate
"""
source = _normalize_history_source(history_source)
if source == GenerationHistorySourceEnum.GENERATION_RECORD:
return await list_generation_record_history_grouped_days(
db=db,
user_id=user_id,
gen_type=gen_type,
page=page,
page_size=page_size,
)
gen_type = _normalize_history_gen_type(gen_type)
page = max(page, 1)
page_size = min(max(page_size, 1), HISTORY_DAY_PAGE_SIZE_MAX)
filters = _history_base_filters(user_id, gen_type, source)
day_expr = func.date(ChatGenerationTask.generated_at).label("generated_date")
days_subquery = (
select(day_expr)
.where(*filters)
.group_by(day_expr)
.subquery()
)
total_days = (
await db.execute(select(func.count()).select_from(days_subquery))
).scalar_one()
day_rows_result = await db.execute(
select(
day_expr,
func.count(ChatGenerationTask.id).label("total"),
)
.where(*filters)
.group_by(day_expr)
.order_by(day_expr.desc())
.offset((page - 1) * page_size)
.limit(page_size)
)
day_rows = day_rows_result.all()
raw_groups = []
all_task_ids: list[str] = []
for generated_day, day_total in day_rows:
item_result = await db.execute(
select(ChatGenerationTask)
.where(
*filters,
func.date(ChatGenerationTask.generated_at) == generated_day,
)
.order_by(ChatGenerationTask.generated_at.desc(), ChatGenerationTask.created_at.desc())
.limit(HISTORY_GROUP_ITEM_LIMIT)
)
tasks = list(item_result.scalars().all())
raw_groups.append((generated_day, day_total, tasks))
all_task_ids.extend(task.id for task in tasks)
resource_info_map = await batch_get_generated_resource_info_map(
db,
source_model=SOURCE_MODEL_CHAT_TASK,
source_ids=all_task_ids,
resource_type=gen_type,
)
history_meta_map = await batch_load_generation_history_meta_map(
db,
source=source,
chat_task_ids=all_task_ids,
)
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),
)
for task in tasks
],
}
for generated_day, day_total, tasks in raw_groups
]
return {
"total_days": int(total_days or 0),
"page": page,
"page_size": page_size,
"groups": groups,
}
async def list_generation_history_day_items(
db: AsyncSession,
user_id: str,
gen_type: str,
generated_date: str,
page: int,
page_size: int,
history_source: str | None = None,
):
"""
获取指定生成日期下的历史记录分页。
用于前端点击某一天后,继续加载该日期下的第 2 页、第 3 页数据。
history_source 支持 chat_task / generation_record / hot_opening_replicate / shot_replicate。
"""
source = _normalize_history_source(history_source)
if source == GenerationHistorySourceEnum.GENERATION_RECORD:
return await list_generation_record_history_day_items(
db=db,
user_id=user_id,
gen_type=gen_type,
generated_date=generated_date,
page=page,
page_size=page_size,
)
gen_type = _normalize_history_gen_type(gen_type)
target_day = _parse_history_date(generated_date)
page = max(page, 1)
page_size = min(max(page_size, 1), 100)
filters = _history_base_filters(user_id, gen_type, source)
day_expr = func.date(ChatGenerationTask.generated_at)
total = (
await db.execute(
select(func.count(ChatGenerationTask.id)).where(
*filters,
day_expr == target_day,
)
)
).scalar_one()
result = await db.execute(
select(ChatGenerationTask)
.where(
*filters,
day_expr == target_day,
)
.order_by(ChatGenerationTask.generated_at.desc(), ChatGenerationTask.created_at.desc())
.offset((page - 1) * page_size)
.limit(page_size)
)
tasks = list(result.scalars().all())
task_ids = [task.id for task in tasks]
resource_info_map = await batch_get_generated_resource_info_map(
db,
source_model=SOURCE_MODEL_CHAT_TASK,
source_ids=task_ids,
resource_type=gen_type,
)
history_meta_map = await batch_load_generation_history_meta_map(
db,
source=source,
chat_task_ids=task_ids,
)
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),
)
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)