素材列表增加返回文件名称
This commit is contained in:
@@ -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="查询上传任务历史",
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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
|
||||
],
|
||||
|
||||
@@ -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,
|
||||
*,
|
||||
|
||||
Reference in New Issue
Block a user