368 lines
13 KiB
Python
368 lines
13 KiB
Python
import os
|
|
import json
|
|
import uuid
|
|
import json
|
|
import uuid
|
|
from typing import Any, Optional, Dict
|
|
from fastapi import APIRouter, Depends, HTTPException, status, Query
|
|
from fastapi import APIRouter, Depends, HTTPException, status, Query
|
|
from pydantic import BaseModel, Field
|
|
from sqlalchemy.ext.asyncio import AsyncSession
|
|
from sqlalchemy import select, func
|
|
from sqlalchemy import select, func
|
|
|
|
from app.dependencies import get_current_user, get_db
|
|
from app.models.user import User
|
|
from app.models.pre_test_template import PreTestTemplate
|
|
from app.models.upload_task import UploadTask
|
|
from app.models.generated_resource import GeneratedResource
|
|
from app.services.upload_material_service import upload_material_to_platform
|
|
from app.services.upload_queue import upload_queue
|
|
from app.services.upload_queue import upload_queue
|
|
from app.utils.douyinApi import DouyinApi
|
|
from app.utils.id_gen import generate_id
|
|
|
|
|
|
router = APIRouter(prefix="/upload-material", tags=["上传素材"])
|
|
|
|
class UploadTaskRequest(BaseModel):
|
|
advertiser_ids: list[str] = Field(..., description="广告主id数组,支持多条")
|
|
resource_ids: list[str] = Field(..., description="资源id数组(generated_resources表主键)")
|
|
oauth_id: str = Field(..., description="授权表id")
|
|
is_pre_test: Optional[str] = Field(None, description="是否开启前测:1=是/2=否")
|
|
pre_test_template: Optional[str] = Field(None, description="前测模板id")
|
|
source_model: str = Field(default="generated_resources", description="上传素材来源模型,默认值为generated_resources,可选值:generation_records、generated_resources、chat_generation_tasks")
|
|
|
|
|
|
class BatchUploadRequest(BaseModel):
|
|
tasks: list[UploadTaskRequest] = Field(..., description="批量上传任务列表")
|
|
|
|
class UpdateFileName(BaseModel):
|
|
source_id: str = Field(..., description="资源id")
|
|
file_name: str = Field(..., description="文件名称,平台素材名称")
|
|
|
|
class FileNameUpdateRequest(BaseModel):
|
|
filenames: list[UpdateFileName] = Field(..., description="批量修改文件名列表,格式: [{\"source_id\":\"资源id\",\"file_name\":\"文件名称\"}]")
|
|
|
|
|
|
@router.post(
|
|
"/async-batch-upload",
|
|
summary="异步批量上传素材到平台",
|
|
description="支持批量上传多个授权账户下的资源到素材库,提交后立即返回,后台异步处理",
|
|
)
|
|
async def async_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": "上传任务列表不能为空",
|
|
"task_ids": [],
|
|
"errors": [],
|
|
}
|
|
|
|
task_ids = []
|
|
errors = []
|
|
|
|
source_model_map = {
|
|
"generation_records": "GenerationRecord",
|
|
"generated_resources": None,
|
|
"chat_generation_tasks": "ChatGenerationTask",
|
|
}
|
|
|
|
for task_index, task in enumerate(req.tasks, 1):
|
|
if not task.advertiser_ids:
|
|
errors.append({
|
|
"task_index": task_index,
|
|
"error": "广告主id数组不能为空",
|
|
})
|
|
continue
|
|
|
|
if not task.resource_ids:
|
|
errors.append({
|
|
"task_index": task_index,
|
|
"error": "资源id数组不能为空",
|
|
})
|
|
continue
|
|
|
|
if task.is_pre_test == "1" and not task.pre_test_template:
|
|
errors.append({
|
|
"task_index": task_index,
|
|
"error": "开启前测功能时,必须指定前测模板id",
|
|
})
|
|
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:
|
|
errors.append({
|
|
"task_index": task_index,
|
|
"error": "前测模板id不存在",
|
|
})
|
|
continue
|
|
|
|
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)
|
|
errors.append({
|
|
"task_index": task_index,
|
|
"error": f"资源id [{invalid_ids_str}] 不可用或已删除",
|
|
})
|
|
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)
|
|
errors.append({
|
|
"task_index": task_index,
|
|
"error": f"资源id [{invalid_ids_str}] 不可用或已删除",
|
|
})
|
|
continue
|
|
|
|
resource_ids_to_upload = valid_resource_ids
|
|
|
|
for advertiser_id in task.advertiser_ids:
|
|
for resource_id in resource_ids_to_upload:
|
|
task_id = generate_id()
|
|
upload_task = UploadTask(
|
|
id=task_id,
|
|
user_id=current_user.id,
|
|
oauth_id=task.oauth_id,
|
|
advertiser_id=advertiser_id,
|
|
resource_id=resource_id,
|
|
status=1,
|
|
note=None,
|
|
)
|
|
|
|
db.add(upload_task)
|
|
await upload_queue.enqueue(task_id)
|
|
task_ids.append(task_id)
|
|
|
|
await db.commit()
|
|
|
|
message = "上传任务已提交"
|
|
if errors:
|
|
message = f"部分任务提交成功,{len(errors)} 个任务失败"
|
|
|
|
return {
|
|
"code": 0,
|
|
"message": message,
|
|
"task_ids": task_ids,
|
|
"errors": errors,
|
|
}
|
|
|
|
|
|
@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="查询上传任务历史",
|
|
description="查询当前用户的上传任务历史列表",
|
|
)
|
|
async def get_upload_history(
|
|
page: int = Query(1, ge=1, description="页码"),
|
|
page_size: int = Query(20, ge=1, le=100, description="每页数量"),
|
|
status: Optional[int] = Query(None, description="上传状态筛选:1待上传,2上传中,3上传成功,4上传失败"),
|
|
current_user: User = Depends(get_current_user),
|
|
db: AsyncSession = Depends(get_db),
|
|
) -> Any | dict:
|
|
from app.services.upload_material_service import get_upload_history as get_upload_history_service
|
|
|
|
result = await get_upload_history_service(
|
|
user_id=current_user.id,
|
|
db=db,
|
|
page=page,
|
|
page_size=page_size,
|
|
status=status,
|
|
)
|
|
|
|
return {
|
|
"code": 0,
|
|
"data": result["data"],
|
|
"pagination": result["pagination"],
|
|
} |