diff --git a/video-gen-api/app/api/v1/upload_material.py b/video-gen-api/app/api/v1/upload_material.py index eb0acc3e..3d1f3f57 100644 --- a/video-gen-api/app/api/v1/upload_material.py +++ b/video-gen-api/app/api/v1/upload_material.py @@ -37,255 +37,12 @@ class UploadTaskRequest(BaseModel): class BatchUploadRequest(BaseModel): tasks: list[UploadTaskRequest] = Field(..., description="批量上传任务列表") +class UpdateFileName(BaseModel): + source_id: str = Field(..., description="资源id") + file_name: str = Field(..., description="文件名称,平台素材名称") -@router.post( - "/batch-upload", - summary="批量上传素材到平台", - description="支持批量上传多个授权账户下的资源到素材库,预留下前测功能", -) -async def batch_upload_material( - req: BatchUploadRequest, - current_user: User = Depends(get_current_user), - db: AsyncSession = Depends(get_db), -) -> Any | dict: - if not req.tasks: - return { - "code": 0, - "message": "批量上传完成", - "summary": { - "total_tasks": 0, - "total_success_count": 0, - "total_fail_count": 0, - "total_requested": 0, - }, - "details": [], - "error": "上传任务列表不能为空" - } - - all_results = [] - total_success = 0 - total_fail = 0 - - for task_index, task in enumerate(req.tasks, 1): - task_result = { - "task_index": task_index, - "oauth_id": task.oauth_id, - "advertiser_ids": task.advertiser_ids, - "resource_ids": task.resource_ids, - "is_pre_test": task.is_pre_test, - "pre_test_template": task.pre_test_template, - "result": None, - } - - try: - if not task.advertiser_ids: - task_result["result"] = { - "success": False, - "error": "广告主id数组不能为空", - "success_count": 0, - "fail_count": 0, - "total_count": 0, - "results": [], - } - all_results.append(task_result) - continue - - if not task.resource_ids: - task_result["result"] = { - "success": False, - "error": "资源id数组不能为空", - "success_count": 0, - "fail_count": 0, - "total_count": 0, - "results": [], - } - all_results.append(task_result) - continue - - if (task.is_pre_test == "1") and (not task.pre_test_template): - task_result["result"] = { - "success": False, - "error": "开启前测功能时,必须指定前测模板id", - "success_count": 0, - "fail_count": len(task.resource_ids) * len(task.advertiser_ids), - "total_count": len(task.resource_ids) * len(task.advertiser_ids), - "results": [{ - "resource_id": rid, - "advertiser_id": aid, - "filename": "", - "success": False, - "error": "开启前测功能时,必须指定前测模板id" - } for rid in task.resource_ids for aid in task.advertiser_ids], - } - total_fail += len(task.resource_ids) * len(task.advertiser_ids) - all_results.append(task_result) - continue - - if task.is_pre_test == "1": - template = await db.execute( - select(PreTestTemplate).where(PreTestTemplate.id == task.pre_test_template). - where(PreTestTemplate.deleted_at.is_(None)). - where(PreTestTemplate.user_id == current_user.id) - ) - template = template.scalar_one_or_none() - if not template: - task_result["result"] = { - "success": False, - "error": "前测模板id不存在", - "success_count": 0, - "fail_count": len(task.resource_ids) * len(task.advertiser_ids), - "total_count": len(task.resource_ids) * len(task.advertiser_ids), - "results": [{ - "resource_id": rid, - "advertiser_id": aid, - "filename": "", - "success": False, - "error": "前测模板id不存在" - } for rid in task.resource_ids for aid in task.advertiser_ids], - } - total_fail += len(task.resource_ids) * len(task.advertiser_ids) - all_results.append(task_result) - continue - - source_model_map = { - "generation_records": "GenerationRecord", - "generated_resources": None, - "chat_generation_tasks": "ChatGenerationTask", - } - - target_source_model = source_model_map.get(task.source_model) - - if target_source_model: - query = ( - select(GeneratedResource.id) - .where(GeneratedResource.source_model == target_source_model) - .where(GeneratedResource.source_id.in_(task.resource_ids)) - .where(GeneratedResource.user_id == current_user.id) - .where(GeneratedResource.deleted_at.is_(None)) - ) - result = await db.execute(query) - valid_resource_ids = [row[0] for row in result.all()] - - invalid_ids = set(task.resource_ids) - set(valid_resource_ids) - - if invalid_ids: - invalid_ids_str = ", ".join(invalid_ids) - task_result["result"] = { - "success": False, - "error": f"资源id [{invalid_ids_str}] 不可用或已删除", - "success_count": 0, - "fail_count": len(task.resource_ids) * len(task.advertiser_ids), - "total_count": len(task.resource_ids) * len(task.advertiser_ids), - "results": [{ - "resource_id": rid, - "advertiser_id": aid, - "filename": "", - "success": False, - "error": f"资源id {rid} 不可用或已删除" if rid in invalid_ids else "资源验证失败" - } for rid in task.resource_ids for aid in task.advertiser_ids], - } - total_fail += len(task.resource_ids) * len(task.advertiser_ids) - all_results.append(task_result) - continue - - resource_ids_to_upload = valid_resource_ids - else: - query = ( - select(GeneratedResource.id) - .where(GeneratedResource.id.in_(task.resource_ids)) - .where(GeneratedResource.user_id == current_user.id) - .where(GeneratedResource.deleted_at.is_(None)) - ) - result = await db.execute(query) - valid_resource_ids = [row[0] for row in result.all()] - - invalid_ids = set(task.resource_ids) - set(valid_resource_ids) - - if invalid_ids: - invalid_ids_str = ", ".join(invalid_ids) - task_result["result"] = { - "success": False, - "error": f"资源id [{invalid_ids_str}] 不可用或已删除", - "success_count": 0, - "fail_count": len(task.resource_ids) * len(task.advertiser_ids), - "total_count": len(task.resource_ids) * len(task.advertiser_ids), - "results": [{ - "resource_id": rid, - "advertiser_id": aid, - "filename": "", - "success": False, - "error": f"资源id {rid} 不可用或已删除" if rid in invalid_ids else "资源验证失败" - } for rid in task.resource_ids for aid in task.advertiser_ids], - } - total_fail += len(task.resource_ids) * len(task.advertiser_ids) - all_results.append(task_result) - continue - - resource_ids_to_upload = valid_resource_ids - - result = await upload_material_to_platform( - resource_ids_to_upload, - task.advertiser_ids, - task.oauth_id, - db, - current_user.id, - task.pre_test_template if task.is_pre_test == "1" else None, - ) - - task_result["result"] = { - "success": True, - **result, - } - total_success += result["success_count"] - total_fail += result["fail_count"] - - - except ValueError as e: - task_result["result"] = { - "success": False, - "error": str(e), - "success_count": 0, - "fail_count": len(task.resource_ids) * len(task.advertiser_ids), - "total_count": len(task.resource_ids) * len(task.advertiser_ids), - "results": [{ - "resource_id": rid, - "advertiser_id": aid, - "filename": "", - "success": False, - "error": str(e) - } for rid in task.resource_ids for aid in task.advertiser_ids], - } - total_fail += len(task.resource_ids) * len(task.advertiser_ids) - except Exception as e: - task_result["result"] = { - "success": False, - "error": f"上传失败: {str(e)}", - "success_count": 0, - "fail_count": len(task.resource_ids) * len(task.advertiser_ids), - "total_count": len(task.resource_ids) * len(task.advertiser_ids), - "results": [{ - "resource_id": rid, - "advertiser_id": aid, - "filename": "", - "success": False, - "error": f"上传失败: {str(e)}" - } for rid in task.resource_ids for aid in task.advertiser_ids], - } - total_fail += len(task.resource_ids) * len(task.advertiser_ids) - - all_results.append(task_result) - - return { - "code": 0, - "message": "批量上传完成", - "summary": { - "total_tasks": len(req.tasks), - "total_success_count": total_success, - "total_fail_count": total_fail, - "total_requested": sum(len(t.resource_ids) * len(t.advertiser_ids) for t in req.tasks), - }, - "details": all_results, - } +class FileNameUpdateRequest(BaseModel): + filenames: list[UpdateFileName] = Field(..., description="批量修改文件名列表,格式: [{\"source_id\":\"资源id\",\"file_name\":\"文件名称\"}]") @router.post( @@ -428,6 +185,160 @@ async def async_batch_upload_material( } +@router.post( + "/batch-update-filename", + summary="批量修改资源文件名", + description="批量修改generated_resources表中的file_name,支持自动处理重复文件名", +) +async def batch_update_filename( + req: FileNameUpdateRequest, + current_user: User = Depends(get_current_user), + db: AsyncSession = Depends(get_db), +) -> Any | dict: + try: + if not req.filenames: + return { + "code": 0, + "message": "修改列表不能为空", + "success_count": 0, + "fail_count": 0, + "results": [], + } + + success_count = 0 + fail_count = 0 + results = [] + valid_items = [] + + for item in req.filenames: + source_id = item.source_id + file_name = item.file_name + + if not source_id: + results.append({ + "source_id": source_id, + "file_name": file_name, + "success": False, + "error": "source_id不能为空", + }) + fail_count += 1 + continue + + if not file_name: + results.append({ + "source_id": source_id, + "file_name": file_name, + "success": False, + "error": "file_name不能为空", + }) + fail_count += 1 + continue + + resource = await db.execute( + select(GeneratedResource) + .where(GeneratedResource.id == source_id) + .where(GeneratedResource.user_id == current_user.id) + .where(GeneratedResource.deleted_at.is_(None)) + ) + resource = resource.scalar_one_or_none() + + if not resource: + results.append({ + "source_id": source_id, + "file_name": file_name, + "success": False, + "error": "资源不存在或不属于当前用户", + }) + fail_count += 1 + continue + + valid_items.append({ + "source_id": source_id, + "file_name": file_name, + "resource": resource, + }) + + query = ( + select(GeneratedResource.file_name) + .where(GeneratedResource.user_id == current_user.id) + .where(GeneratedResource.deleted_at.is_(None)) + .where(GeneratedResource.file_name.is_not(None)) + ) + result = await db.execute(query) + db_existing_names = set(row[0] for row in result.all()) + + name_counters = {} + + for item in valid_items: + file_name = item["file_name"] + resource = item["resource"] + base_name, ext = os.path.splitext(file_name) + + existing_names = db_existing_names.copy() + + if resource.file_name and resource.file_name in existing_names: + existing_names.remove(resource.file_name) + + if file_name not in name_counters: + counter = 1 + new_file_name = file_name + + while new_file_name in existing_names: + new_file_name = f"{base_name}{counter}{ext}" + counter += 1 + + name_counters[file_name] = { + "base_name": base_name, + "ext": ext, + "counter": counter, + } + existing_names.add(new_file_name) + db_existing_names.add(new_file_name) + else: + counter = name_counters[file_name]["counter"] + base_name = name_counters[file_name]["base_name"] + ext = name_counters[file_name]["ext"] + new_file_name = f"{base_name}{counter}{ext}" + + while new_file_name in existing_names: + counter += 1 + new_file_name = f"{base_name}{counter}{ext}" + + name_counters[file_name]["counter"] = counter + 1 + existing_names.add(new_file_name) + db_existing_names.add(new_file_name) + + item["resource"].file_name = new_file_name + db.add(item["resource"]) + + results.append({ + "source_id": item["source_id"], + "file_name": item["file_name"], + "new_file_name": new_file_name, + "success": True, + "error": None, + }) + success_count += 1 + + await db.commit() + + return { + "code": 0, + "message": f"批量修改完成,成功 {success_count} 条,失败 {fail_count} 条", + "success_count": success_count, + "fail_count": fail_count, + "results": results, + } + except Exception as e: + return { + "code": 0, + "message": f"批量修改文件名失败:{str(e)}", + "success_count": 0, + "fail_count": 0, + "results": [], + } + + @router.get( "/upload-history", summary="查询上传任务历史", diff --git a/video-gen-api/app/schemas/generation_ai.py b/video-gen-api/app/schemas/generation_ai.py index 6d7b167c..c7f2022d 100644 --- a/video-gen-api/app/schemas/generation_ai.py +++ b/video-gen-api/app/schemas/generation_ai.py @@ -321,6 +321,7 @@ class GenerationAITaskOut(BaseModel): None, description="关联的生成资源账本ID,来源于 generated_resources.id;历史脏数据可能为空", ) + file_name: str | None = Field(None, description="文件名,来源于 generated_resources.file_name") gen_type: str = Field(..., description="生成类型:image=图片,video=视频") generation_mode: str | None = Field( None, @@ -536,6 +537,7 @@ class GenerationAIRecordHistoryItemOut(BaseModel): None, description="关联的生成资源账本ID,来源于 generated_resources.id;历史脏数据可能为空", ) + file_name: str | None = Field(None, description="文件名,来源于 generated_resources.file_name") gen_type: str = Field(..., description="生成类型:image=图片,video=视频") generation_mode: str | None = Field( "generation_record", diff --git a/video-gen-api/app/services/generation_ai_service.py b/video-gen-api/app/services/generation_ai_service.py index cb158698..de9b4819 100644 --- a/video-gen-api/app/services/generation_ai_service.py +++ b/video-gen-api/app/services/generation_ai_service.py @@ -32,6 +32,7 @@ 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 @@ -510,6 +511,7 @@ 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( @@ -518,6 +520,7 @@ def generation_record_to_history_out( project_id=record.project_id, project_name=project_name, generated_resource_id=generated_resource_id, + file_name=file_name, gen_type=record.gen_type, generation_mode="generation_record", pipeline_stage=None, @@ -615,7 +618,7 @@ async def list_generation_record_history_grouped_days( raw_groups.append((generated_day, day_total, rows)) all_record_ids.extend(record.id for record, _project_name in rows) - 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_GENERATION_RECORD, source_ids=all_record_ids, @@ -630,7 +633,8 @@ async def list_generation_record_history_grouped_days( generation_record_to_history_out( record, project_name, - generated_resource_id=resource_id_map.get(record.id), + 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 ], @@ -687,7 +691,7 @@ async def list_generation_record_history_day_items( ) rows = result.all() - 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_GENERATION_RECORD, source_ids=[record.id for record, _project_name in rows], @@ -703,7 +707,8 @@ async def list_generation_record_history_day_items( generation_record_to_history_out( record, project_name, - generated_resource_id=resource_id_map.get(record.id), + 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 ], diff --git a/video-gen-api/app/services/resource_accounting_service.py b/video-gen-api/app/services/resource_accounting_service.py index cf798d59..723d321b 100644 --- a/video-gen-api/app/services/resource_accounting_service.py +++ b/video-gen-api/app/services/resource_accounting_service.py @@ -372,6 +372,50 @@ async def batch_get_generated_resource_id_map( return resource_id_map +async def batch_get_generated_resource_info_map( + db: AsyncSession, + *, + source_model: str, + source_ids: Sequence[str] | Iterable[str], + resource_type: str | None = None, +) -> dict[str, dict[str, str | None]]: + """批量查询来源记录对应的 GeneratedResource 信息(id 和 file_name)。 + + 用于历史列表接口批量回填资源账本信息,避免按记录一条条查询。 + 如果历史脏数据存在同一个 source_id 对应多条未软删资源账本,按 created_at 倒序取最新一条。 + + 返回格式: {source_id: {"resource_id": "...", "file_name": "..."}} + """ + ids = list(dict.fromkeys(str(item) for item in source_ids if item)) + if not ids: + return {} + + normalized_resource_type = (resource_type or "").lower().strip() + query = select(GeneratedResource).where( + GeneratedResource.source_model == source_model, + GeneratedResource.source_id.in_(ids), + GeneratedResource.deleted_at.is_(None), + ) + if normalized_resource_type in ("image", "video"): + query = query.where(GeneratedResource.resource_type == normalized_resource_type) + + result = await db.execute( + query.order_by( + GeneratedResource.source_id.asc(), + GeneratedResource.created_at.desc(), + ) + ) + + resource_info_map: dict[str, dict[str, str | None]] = {} + for resource in result.scalars().all(): + if resource.source_id not in resource_info_map: + resource_info_map[resource.source_id] = { + "resource_id": resource.id, + "file_name": resource.file_name, + } + return resource_info_map + + async def soft_delete_resources_by_source( db: AsyncSession, *,