diff --git a/video-gen-api/app/api/v1/generation_ai.py b/video-gen-api/app/api/v1/generation_ai.py index cddc0970..28db9b24 100644 --- a/video-gen-api/app/api/v1/generation_ai.py +++ b/video-gen-api/app/api/v1/generation_ai.py @@ -266,8 +266,11 @@ async def list_tasks( summary="获取AI生成历史日期分组", description=( "按生成完成日期倒序返回当前用户的AI生成历史记录。" - "默认查询 chat_generation_tasks / ChatGenerationTask 新任务历史。" - "当 history_source=generation_record 时,查询 generation_records / GenerationRecord 旧历史。" + "默认查询 chat_generation_tasks / ChatGenerationTask 的 AI创作历史。" + "history_source=chat_task 时查询 AI创作;" + "history_source=hot_opening_replicate 时查询爆款开头复刻生成素材;" + "history_source=shot_replicate 时查询拆镜复刻生成素材;" + "history_source=generation_record 时查询 generation_records / GenerationRecord 旧项目生成历史。" "该接口只返回生成成功的任务,即 status=completed 且 generated_at 不为空的数据。" "必须通过 gen_type 区分图片和视频。" "分页对象是生成日期,不是单条记录。" @@ -309,10 +312,12 @@ async def list_history_grouped_days( history_source: str | None = Query( None, description=( - "历史数据来源。默认不传或传 chat_task 查询 chat_generation_tasks / ChatGenerationTask;" - "传 generation_record 查询 generation_records / GenerationRecord 旧历史数据" + "历史数据来源。默认不传或传 chat_task 查询 AI创作;" + "传 generation_record 查询旧项目生成;" + "传 hot_opening_replicate 查询爆款开头复刻素材;" + "传 shot_replicate 查询拆镜复刻素材" ), - examples=["generation_record"], + examples=["shot_replicate"], ), current_user: User = Depends(get_current_user), db: AsyncSession = Depends(get_db), @@ -333,8 +338,11 @@ async def list_history_grouped_days( summary="获取指定日期下的AI生成历史分页", description=( "获取某一个生成日期下的生成成功记录分页。" - "默认查询 chat_generation_tasks / ChatGenerationTask 新任务历史。" - "当 history_source=generation_record 时,查询 generation_records / GenerationRecord 旧历史。" + "默认查询 chat_generation_tasks / ChatGenerationTask 的 AI创作历史。" + "history_source=chat_task 时查询 AI创作;" + "history_source=hot_opening_replicate 时查询爆款开头复刻生成素材;" + "history_source=shot_replicate 时查询拆镜复刻生成素材;" + "history_source=generation_record 时查询 generation_records / GenerationRecord 旧项目生成历史。" "该接口用于前端在历史分组列表中继续加载某一天的后续记录。" "例如 /history 接口中某一天 total=18,但 items 只返回前10条," "则前端可以调用本接口 page=2&page_size=10 获取该日期下剩余记录。" @@ -380,10 +388,12 @@ async def list_history_day_items( history_source: str | None = Query( None, description=( - "历史数据来源。默认不传或传 chat_task 查询 chat_generation_tasks / ChatGenerationTask;" - "传 generation_record 查询 generation_records / GenerationRecord 旧历史数据" + "历史数据来源。默认不传或传 chat_task 查询 AI创作;" + "传 generation_record 查询旧项目生成;" + "传 hot_opening_replicate 查询爆款开头复刻素材;" + "传 shot_replicate 查询拆镜复刻素材" ), - examples=["generation_record"], + examples=["hot_opening_replicate"], ), current_user: User = Depends(get_current_user), db: AsyncSession = Depends(get_db), diff --git a/video-gen-api/app/enums/__init__.py b/video-gen-api/app/enums/__init__.py index cf7cee1f..0c6efa60 100644 --- a/video-gen-api/app/enums/__init__.py +++ b/video-gen-api/app/enums/__init__.py @@ -7,6 +7,7 @@ from app.enums.user import * from app.enums.credit_record import * from app.enums.token_usage import * from app.enums.generation_task import * +from app.enums.generation_history import * from app.enums.generation_status import * from app.enums.sms import * from app.enums.notification import * diff --git a/video-gen-api/app/enums/generation_history.py b/video-gen-api/app/enums/generation_history.py new file mode 100644 index 00000000..974cd185 --- /dev/null +++ b/video-gen-api/app/enums/generation_history.py @@ -0,0 +1,120 @@ +from __future__ import annotations + +from enum import Enum + +from app.enums.generation_task import ChatGenerationTaskStatus, GenerationMode, GenerationType + + +class GenerationHistorySourceEnum(str, Enum): + """素材云历史接口支持的数据来源枚举。""" + + CHAT_TASK = "chat_task" + GENERATION_RECORD = "generation_record" + HOT_OPENING_REPLICATE = "hot_opening_replicate" + SHOT_REPLICATE = "shot_replicate" + + +class GenerationHistoryResourceTypeEnum(str, Enum): + """素材云历史接口支持的资源类型枚举。""" + + IMAGE = GenerationType.IMAGE.value + VIDEO = GenerationType.VIDEO.value + + +GENERATION_HISTORY_SOURCE_LABELS: dict[GenerationHistorySourceEnum, str] = { + GenerationHistorySourceEnum.CHAT_TASK: "AI创作", + GenerationHistorySourceEnum.GENERATION_RECORD: "项目生成", + GenerationHistorySourceEnum.HOT_OPENING_REPLICATE: "爆款开头复刻", + GenerationHistorySourceEnum.SHOT_REPLICATE: "拆镜复刻", +} +"""素材云历史来源中文名称。""" + + +GENERATION_HISTORY_CHAT_TASK_SOURCES: tuple[GenerationHistorySourceEnum, ...] = ( + GenerationHistorySourceEnum.CHAT_TASK, + GenerationHistorySourceEnum.HOT_OPENING_REPLICATE, + GenerationHistorySourceEnum.SHOT_REPLICATE, +) +"""来自 chat_generation_tasks 表的历史来源集合。""" + + +GENERATION_HISTORY_MODULE_SOURCES: tuple[GenerationHistorySourceEnum, ...] = ( + GenerationHistorySourceEnum.HOT_OPENING_REPLICATE, + GenerationHistorySourceEnum.SHOT_REPLICATE, +) +"""需要回填 module_generation_projects/module_generation_steps 的模块来源集合。""" + + +GENERATION_HISTORY_SOURCE_TO_TASK_MODE: dict[GenerationHistorySourceEnum, GenerationMode] = { + GenerationHistorySourceEnum.CHAT_TASK: GenerationMode.CHATAPI_ASYNC, + GenerationHistorySourceEnum.HOT_OPENING_REPLICATE: GenerationMode.HOT_OPENING_REPLICATE, + GenerationHistorySourceEnum.SHOT_REPLICATE: GenerationMode.SHOT_REPLICATE, +} +"""history_source 到 ChatGenerationTask.generation_mode 的映射。""" + + +GENERATION_HISTORY_TASK_MODE_VALUE_TO_SOURCE: dict[str, GenerationHistorySourceEnum] = { + task_mode.value: history_source + for history_source, task_mode in GENERATION_HISTORY_SOURCE_TO_TASK_MODE.items() +} +"""ChatGenerationTask.generation_mode 字符串值到 history_source 的映射。""" + + +GENERATION_HISTORY_COMPLETED_STATUS = ChatGenerationTaskStatus.COMPLETED.value +"""素材云历史只展示生成成功的数据。""" + + +GENERATION_HISTORY_SOURCE_ALIASES: dict[str, GenerationHistorySourceEnum] = { + "": GenerationHistorySourceEnum.CHAT_TASK, + "chat": GenerationHistorySourceEnum.CHAT_TASK, + "chat_ai": GenerationHistorySourceEnum.CHAT_TASK, + "chat_task": GenerationHistorySourceEnum.CHAT_TASK, + "chat_generation_task": GenerationHistorySourceEnum.CHAT_TASK, + "chat_generation_tasks": GenerationHistorySourceEnum.CHAT_TASK, + "record": GenerationHistorySourceEnum.GENERATION_RECORD, + "records": GenerationHistorySourceEnum.GENERATION_RECORD, + "project": GenerationHistorySourceEnum.GENERATION_RECORD, + "generation_record": GenerationHistorySourceEnum.GENERATION_RECORD, + "generation_records": GenerationHistorySourceEnum.GENERATION_RECORD, + "hot_opening": GenerationHistorySourceEnum.HOT_OPENING_REPLICATE, + "hot_opening_replicate": GenerationHistorySourceEnum.HOT_OPENING_REPLICATE, + "shot": GenerationHistorySourceEnum.SHOT_REPLICATE, + "shot_replicate": GenerationHistorySourceEnum.SHOT_REPLICATE, +} +"""history_source 兼容别名映射。""" + + +def normalize_generation_history_source(value: str | None) -> GenerationHistorySourceEnum: + """归一化素材云历史来源。""" + + key = (value or "chat_task").lower().strip() + if key in GENERATION_HISTORY_SOURCE_ALIASES: + return GENERATION_HISTORY_SOURCE_ALIASES[key] + return GenerationHistorySourceEnum(key) + + +def get_generation_history_source_label(source: GenerationHistorySourceEnum | str | None) -> str | None: + """获取素材云历史来源中文名称。""" + + if source is None: + return None + source_enum = source if isinstance(source, GenerationHistorySourceEnum) else GenerationHistorySourceEnum(str(source)) + return GENERATION_HISTORY_SOURCE_LABELS.get(source_enum) + + +def get_generation_history_task_mode(source: GenerationHistorySourceEnum) -> GenerationMode | None: + """获取 history_source 对应的 ChatGenerationTask.generation_mode。""" + + return GENERATION_HISTORY_SOURCE_TO_TASK_MODE.get(source) + + +def is_generation_history_chat_task_source(source: GenerationHistorySourceEnum) -> bool: + """判断当前来源是否走 chat_generation_tasks 表。""" + + return source in GENERATION_HISTORY_CHAT_TASK_SOURCES + + +def is_generation_history_module_source(source: GenerationHistorySourceEnum) -> bool: + """判断当前来源是否需要回填模块项目信息。""" + + return source in GENERATION_HISTORY_MODULE_SOURCES diff --git a/video-gen-api/app/schemas/generation_ai.py b/video-gen-api/app/schemas/generation_ai.py index c7f2022d..be5023c4 100644 --- a/video-gen-api/app/schemas/generation_ai.py +++ b/video-gen-api/app/schemas/generation_ai.py @@ -322,6 +322,26 @@ class GenerationAITaskOut(BaseModel): description="关联的生成资源账本ID,来源于 generated_resources.id;历史脏数据可能为空", ) file_name: str | None = Field(None, description="文件名,来源于 generated_resources.file_name") + history_source: str | None = Field( + None, + description=( + "素材云历史来源:chat_task=AI创作,generation_record=项目生成," + "hot_opening_replicate=爆款开头复刻,shot_replicate=拆镜复刻" + ), + ) + history_source_label: str | None = Field(None, description="素材云历史来源中文名称") + module_project_id: str | None = Field(None, description="模块生成项目ID;非模块生成历史返回 null") + module_project_title: str | None = Field(None, description="模块生成项目标题;非模块生成历史返回 null") + module_step_id: str | None = Field(None, description="模块生成步骤ID;非模块生成历史返回 null") + module_step_code: str | None = Field(None, description="模块生成步骤编码;非模块生成历史返回 null") + hot_opening_project_id: str | None = Field(None, description="爆款开头复刻项目ID;非爆款开头复刻返回 null") + hot_opening_project_title: str | None = Field(None, description="爆款开头复刻项目标题;非爆款开头复刻返回 null") + shot_replicate_project_id: str | None = Field(None, description="拆镜复刻项目ID;非拆镜复刻返回 null") + shot_replicate_project_title: str | None = Field(None, description="拆镜复刻项目标题;非拆镜复刻返回 null") + shot_task_set_id: str | None = Field(None, description="拆镜复刻总任务ID;非拆镜复刻返回 null") + shot_segment_id: str | None = Field(None, description="拆镜复刻片段ID;非拆镜复刻返回 null") + shot_segment_index: int | None = Field(None, description="拆镜复刻片段序号;非拆镜复刻返回 null") + shot_segment_label: str | None = Field(None, description="拆镜复刻片段展示名称,例如:拆镜复刻片段1;非拆镜复刻返回 null") gen_type: str = Field(..., description="生成类型:image=图片,video=视频") generation_mode: str | None = Field( None, @@ -538,6 +558,26 @@ class GenerationAIRecordHistoryItemOut(BaseModel): description="关联的生成资源账本ID,来源于 generated_resources.id;历史脏数据可能为空", ) file_name: str | None = Field(None, description="文件名,来源于 generated_resources.file_name") + history_source: str | None = Field( + None, + description=( + "素材云历史来源:chat_task=AI创作,generation_record=项目生成," + "hot_opening_replicate=爆款开头复刻,shot_replicate=拆镜复刻" + ), + ) + history_source_label: str | None = Field(None, description="素材云历史来源中文名称") + module_project_id: str | None = Field(None, description="模块生成项目ID;非模块生成历史返回 null") + module_project_title: str | None = Field(None, description="模块生成项目标题;非模块生成历史返回 null") + module_step_id: str | None = Field(None, description="模块生成步骤ID;非模块生成历史返回 null") + module_step_code: str | None = Field(None, description="模块生成步骤编码;非模块生成历史返回 null") + hot_opening_project_id: str | None = Field(None, description="爆款开头复刻项目ID;非爆款开头复刻返回 null") + hot_opening_project_title: str | None = Field(None, description="爆款开头复刻项目标题;非爆款开头复刻返回 null") + shot_replicate_project_id: str | None = Field(None, description="拆镜复刻项目ID;非拆镜复刻返回 null") + shot_replicate_project_title: str | None = Field(None, description="拆镜复刻项目标题;非拆镜复刻返回 null") + shot_task_set_id: str | None = Field(None, description="拆镜复刻总任务ID;非拆镜复刻返回 null") + shot_segment_id: str | None = Field(None, description="拆镜复刻片段ID;非拆镜复刻返回 null") + shot_segment_index: int | None = Field(None, description="拆镜复刻片段序号;非拆镜复刻返回 null") + shot_segment_label: str | None = Field(None, description="拆镜复刻片段展示名称,例如:拆镜复刻片段1;非拆镜复刻返回 null") gen_type: str = Field(..., description="生成类型:image=图片,video=视频") generation_mode: str | None = Field( "generation_record", @@ -624,7 +664,7 @@ class GenerationAIHistoryDayGroupOut(BaseModel): default_factory=list, description=( "当前生成日期下倒序前10条生成记录详情。" - "默认 history_source=chat_task 时 item 为 GenerationAITaskOut;" + "history_source=chat_task/hot_opening_replicate/shot_replicate 时 item 为 GenerationAITaskOut;" "history_source=generation_record 时 item 为 GenerationAIRecordHistoryItemOut" ), ) @@ -687,7 +727,7 @@ class GenerationAIHistoryDayItemsOut(BaseModel): default_factory=list, description=( "当前日期下的生成记录详情列表,按 generated_at 倒序排列。" - "默认 history_source=chat_task 时 item 为 GenerationAITaskOut;" + "history_source=chat_task/hot_opening_replicate/shot_replicate 时 item 为 GenerationAITaskOut;" "history_source=generation_record 时 item 为 GenerationAIRecordHistoryItemOut" ), ) \ No newline at end of file diff --git a/video-gen-api/app/services/generation_ai_service.py b/video-gen-api/app/services/generation_ai_service.py index 5253af1a..b1df03df 100644 --- a/video-gen-api/app/services/generation_ai_service.py +++ b/video-gen-api/app/services/generation_ai_service.py @@ -15,6 +15,12 @@ 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, @@ -31,11 +37,15 @@ from app.services.generation_billing_service import ( from app.services.resource_accounting_service import ( SOURCE_MODEL_CHAT_TASK, SOURCE_MODEL_GENERATION_RECORD, - batch_get_generated_resource_id_map, 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 @@ -333,15 +343,42 @@ 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, @@ -375,7 +412,6 @@ def record_to_out( generated_at=task.generated_at, ) - def engine_snapshot_out(snapshot: dict) -> dict: """ 从完整的 engine_snapshot 中过滤出需要返回的字段 @@ -460,18 +496,20 @@ def _normalize_history_gen_type(gen_type: str | None) -> str: -def _normalize_history_source(history_source: str | None) -> str: - """Normalize history source query param. +def _normalize_history_source(history_source: str | None) -> GenerationHistorySourceEnum: + """Normalize history_source query param. - 默认保持原来的 chat_generation_tasks 历史;只有显式传 generation_record - 才切换旧 generation_records 历史,避免影响现有前端。 + 默认保持原来的 chat_generation_tasks / chatapi_async 历史; + 显式传 hot_opening_replicate 或 shot_replicate 时查询对应模块素材; + 显式传 generation_record 时查询旧 generation_records 历史。 """ - value = (history_source or "chat_task").lower().strip() - if value in ("", "chat", "chat_task", "chat_generation_task", "chat_generation_tasks"): - return "chat_task" - if value in ("record", "records", "generation_record", "generation_records"): - return "generation_record" - raise HTTPException(status_code=400, detail="history_source 仅支持 chat_task 或 generation_record") + 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: @@ -489,10 +527,13 @@ def _parse_history_date(value: str) -> date: raise HTTPException(status_code=400, detail="generated_date 格式必须是 YYYY-MM-DD") -def _history_base_filters(user_id: str, gen_type: str): +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 == "chatapi_async", + ChatGenerationTask.generation_mode == task_mode.value, ChatGenerationTask.deleted_at.is_(None), ChatGenerationTask.status == "completed", ChatGenerationTask.gen_type == gen_type, @@ -524,6 +565,20 @@ def generation_record_to_history_out( 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, @@ -732,9 +787,10 @@ async def list_generation_history_grouped_days( - 每页最多返回 10 个生成日期 - 每个日期分组内最多返回倒序前 10 条任务 - 只返回 completed 成功任务 + - history_source 支持 chat_task / generation_record / hot_opening_replicate / shot_replicate """ source = _normalize_history_source(history_source) - if source == "generation_record": + if source == GenerationHistorySourceEnum.GENERATION_RECORD: return await list_generation_record_history_grouped_days( db=db, user_id=user_id, @@ -747,7 +803,7 @@ async def list_generation_history_grouped_days( page = max(page, 1) page_size = min(max(page_size, 1), HISTORY_DAY_PAGE_SIZE_MAX) - filters = _history_base_filters(user_id, gen_type) + filters = _history_base_filters(user_id, gen_type, source) day_expr = func.date(ChatGenerationTask.generated_at).label("generated_date") days_subquery = ( @@ -790,12 +846,17 @@ async def list_generation_history_grouped_days( raw_groups.append((generated_day, day_total, tasks)) all_task_ids.extend(task.id for task in tasks) - resource_id_map = await batch_get_generated_resource_id_map( + 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 = [ { @@ -804,7 +865,9 @@ async def list_generation_history_grouped_days( "items": [ record_to_out( task, - generated_resource_id=resource_id_map.get(task.id), + 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 ], @@ -833,9 +896,10 @@ async def list_generation_history_day_items( 获取指定生成日期下的历史记录分页。 用于前端点击某一天后,继续加载该日期下的第 2 页、第 3 页数据。 + history_source 支持 chat_task / generation_record / hot_opening_replicate / shot_replicate。 """ source = _normalize_history_source(history_source) - if source == "generation_record": + if source == GenerationHistorySourceEnum.GENERATION_RECORD: return await list_generation_record_history_day_items( db=db, user_id=user_id, @@ -850,7 +914,7 @@ async def list_generation_history_day_items( page = max(page, 1) page_size = min(max(page_size, 1), 100) - filters = _history_base_filters(user_id, gen_type) + filters = _history_base_filters(user_id, gen_type, source) day_expr = func.date(ChatGenerationTask.generated_at) total = ( @@ -874,12 +938,18 @@ async def list_generation_history_day_items( ) tasks = list(result.scalars().all()) - resource_id_map = await batch_get_generated_resource_id_map( + 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.id for task in tasks], + 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"), @@ -889,7 +959,9 @@ async def list_generation_history_day_items( "items": [ record_to_out( task, - generated_resource_id=resource_id_map.get(task.id), + 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 ], diff --git a/video-gen-api/app/services/generation_history_meta_service.py b/video-gen-api/app/services/generation_history_meta_service.py new file mode 100644 index 00000000..6326859e --- /dev/null +++ b/video-gen-api/app/services/generation_history_meta_service.py @@ -0,0 +1,283 @@ +from __future__ import annotations + +from typing import Any, TypedDict + +from sqlalchemy import select +from sqlalchemy.ext.asyncio import AsyncSession + +from app.enums.generation_history import ( + GenerationHistorySourceEnum, + get_generation_history_source_label, + is_generation_history_module_source, +) +from app.models.module_generation_project import ModuleGenerationProject +from app.models.module_generation_step import ModuleGenerationStep +from app.models.shot_replicate_segment import ShotReplicateSegment + + +class GenerationHistoryMeta(TypedDict): + """素材云历史项模块上下文回填字段。""" + + history_source: str | None + history_source_label: str | None + module_project_id: str | None + module_project_title: str | None + module_step_id: str | None + module_step_code: str | None + hot_opening_project_id: str | None + hot_opening_project_title: str | None + shot_replicate_project_id: str | None + shot_replicate_project_title: str | None + shot_task_set_id: str | None + shot_segment_id: str | None + shot_segment_index: int | None + shot_segment_label: str | None + + +class _StepLinkInfo(TypedDict): + module_project_id: str | None + module_step_id: str | None + module_step_code: str | None + module: str | None + + +class _ProjectInfo(TypedDict): + module_project_id: str + module_project_title: str | None + module: str | None + + +class _ShotSegmentInfo(TypedDict): + shot_task_set_id: str | None + shot_segment_id: str | None + shot_segment_index: int | None + shot_segment_label: str | None + + +def _unique(values: list[str] | tuple[str, ...]) -> list[str]: + return list(dict.fromkeys(str(value) for value in values if value)) + + +def _segment_label(segment_index: int | None) -> str | None: + if segment_index is None: + return None + return f"片段{segment_index}" + + +def build_empty_history_meta(source: GenerationHistorySourceEnum) -> GenerationHistoryMeta: + """构造统一历史字段,非对应模块字段保持 null。""" + + return { + "history_source": source.value, + "history_source_label": get_generation_history_source_label(source), + "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, + } + + +async def _load_step_link_map( + db: AsyncSession, + *, + chat_task_ids: list[str], + source: GenerationHistorySourceEnum, +) -> dict[str, _StepLinkInfo]: + ids = _unique(chat_task_ids) + if not ids: + return {} + + stmt = ( + select( + ModuleGenerationStep.chat_task_id.label("chat_task_id"), + ModuleGenerationStep.id.label("module_step_id"), + ModuleGenerationStep.project_id.label("module_project_id"), + ModuleGenerationStep.step_code.label("module_step_code"), + ModuleGenerationStep.module.label("module"), + ModuleGenerationStep.is_current.label("is_current"), + ModuleGenerationStep.updated_at.label("updated_at"), + ) + .where( + ModuleGenerationStep.deleted_at.is_(None), + ModuleGenerationStep.chat_task_id.in_(ids), + ModuleGenerationStep.module == source.value, + ) + .order_by( + ModuleGenerationStep.chat_task_id.asc(), + ModuleGenerationStep.is_current.desc(), + ModuleGenerationStep.updated_at.desc(), + ) + ) + + rows = (await db.execute(stmt)).mappings().all() + link_map: dict[str, _StepLinkInfo] = {} + for row in rows: + chat_task_id = row["chat_task_id"] + if not chat_task_id or chat_task_id in link_map: + continue + link_map[chat_task_id] = { + "module_project_id": row["module_project_id"], + "module_step_id": row["module_step_id"], + "module_step_code": row["module_step_code"], + "module": row["module"], + } + return link_map + + +async def _load_project_map( + db: AsyncSession, + *, + module_project_ids: list[str], +) -> dict[str, _ProjectInfo]: + ids = _unique(module_project_ids) + if not ids: + return {} + + stmt = ( + select( + ModuleGenerationProject.id.label("module_project_id"), + ModuleGenerationProject.title.label("module_project_title"), + ModuleGenerationProject.module.label("module"), + ) + .where( + ModuleGenerationProject.deleted_at.is_(None), + ModuleGenerationProject.id.in_(ids), + ) + ) + + return { + row["module_project_id"]: { + "module_project_id": row["module_project_id"], + "module_project_title": row["module_project_title"], + "module": row["module"], + } + for row in (await db.execute(stmt)).mappings().all() + if row["module_project_id"] + } + + +async def _load_shot_segment_map( + db: AsyncSession, + *, + module_project_ids: list[str], +) -> dict[str, _ShotSegmentInfo]: + ids = _unique(module_project_ids) + if not ids: + return {} + + stmt = ( + select( + ShotReplicateSegment.module_project_id.label("module_project_id"), + ShotReplicateSegment.id.label("shot_segment_id"), + ShotReplicateSegment.task_set_id.label("shot_task_set_id"), + ShotReplicateSegment.segment_index.label("shot_segment_index"), + ShotReplicateSegment.updated_at.label("updated_at"), + ) + .where( + ShotReplicateSegment.deleted_at.is_(None), + ShotReplicateSegment.module_project_id.in_(ids), + ) + .order_by( + ShotReplicateSegment.module_project_id.asc(), + ShotReplicateSegment.updated_at.desc(), + ) + ) + + segment_map: dict[str, _ShotSegmentInfo] = {} + rows = (await db.execute(stmt)).mappings().all() + for row in rows: + module_project_id = row["module_project_id"] + if not module_project_id or module_project_id in segment_map: + continue + segment_index = row["shot_segment_index"] + segment_map[module_project_id] = { + "shot_task_set_id": row["shot_task_set_id"], + "shot_segment_id": row["shot_segment_id"], + "shot_segment_index": segment_index, + "shot_segment_label": _segment_label(segment_index), + } + return segment_map + + +def _merge_meta( + *, + source: GenerationHistorySourceEnum, + step_info: _StepLinkInfo | None, + project_info: _ProjectInfo | None, + shot_info: _ShotSegmentInfo | None, +) -> GenerationHistoryMeta: + meta = build_empty_history_meta(source) + + module_project_id = step_info["module_project_id"] if step_info else None + module_step_id = step_info["module_step_id"] if step_info else None + module_step_code = step_info["module_step_code"] if step_info else None + module_project_title = project_info["module_project_title"] if project_info else None + + meta["module_project_id"] = module_project_id + meta["module_project_title"] = module_project_title + meta["module_step_id"] = module_step_id + meta["module_step_code"] = module_step_code + + if source == GenerationHistorySourceEnum.HOT_OPENING_REPLICATE: + meta["hot_opening_project_id"] = module_project_id + meta["hot_opening_project_title"] = module_project_title + elif source == GenerationHistorySourceEnum.SHOT_REPLICATE: + meta["shot_replicate_project_id"] = module_project_id + meta["shot_replicate_project_title"] = module_project_title + if shot_info: + meta["shot_task_set_id"] = shot_info["shot_task_set_id"] + meta["shot_segment_id"] = shot_info["shot_segment_id"] + meta["shot_segment_index"] = shot_info["shot_segment_index"] + meta["shot_segment_label"] = shot_info["shot_segment_label"] + + return meta + + +async def batch_load_generation_history_meta_map( + db: AsyncSession, + *, + source: GenerationHistorySourceEnum, + chat_task_ids: list[str], +) -> dict[str, GenerationHistoryMeta]: + """批量回填历史列表模块上下文,避免逐条链式查询。""" + + ids = _unique(chat_task_ids) + if not ids: + return {} + + if not is_generation_history_module_source(source): + return {chat_task_id: build_empty_history_meta(source) for chat_task_id in ids} + + step_link_map = await _load_step_link_map(db, chat_task_ids=ids, source=source) + module_project_ids = [ + step_info["module_project_id"] + for step_info in step_link_map.values() + if step_info.get("module_project_id") + ] + project_map = await _load_project_map(db, module_project_ids=module_project_ids) + + shot_segment_map: dict[str, _ShotSegmentInfo] = {} + if source == GenerationHistorySourceEnum.SHOT_REPLICATE: + shot_segment_map = await _load_shot_segment_map(db, module_project_ids=module_project_ids) + + meta_map: dict[str, GenerationHistoryMeta] = {} + for chat_task_id in ids: + step_info = step_link_map.get(chat_task_id) + module_project_id = step_info.get("module_project_id") if step_info else None + project_info = project_map.get(module_project_id) if module_project_id else None + shot_info = shot_segment_map.get(module_project_id) if module_project_id else None + meta_map[chat_task_id] = _merge_meta( + source=source, + step_info=step_info, + project_info=project_info, + shot_info=shot_info, + ) + return meta_map