素材列表增加返回文件名称

This commit is contained in:
18610128193
2026-06-24 16:55:12 +08:00
parent 313471c1dd
commit 7f1b18314e
4 changed files with 214 additions and 252 deletions
+159 -248
View File
@@ -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="查询上传任务历史",