Files
video-gen/video-gen-api/app/api/v1/upload_material.py
T

375 lines
13 KiB
Python

import os
import json
import uuid
from typing import Any, Optional, Dict
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 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.services.upload_material_service import upload_material_to_platform
from app.services.upload_queue import upload_queue
from app.utils.douyinApi import DouyinApi
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")
class BatchUploadRequest(BaseModel):
tasks: list[UploadTaskRequest] = 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
result = await upload_material_to_platform(
task.resource_ids,
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,
}
@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": [],
}
task_ids = []
for task in req.tasks:
if not task.advertiser_ids or not task.resource_ids:
continue
if task.is_pre_test == "1" and not task.pre_test_template:
continue
for advertiser_id in task.advertiser_ids:
for resource_id in task.resource_ids:
task_id = str(uuid.uuid4()).replace("-", "")[:32]
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()
return {
"code": 0,
"message": "上传任务已提交,请通过查询接口查看进度",
"task_ids": task_ids,
}
@router.get(
"/upload-status",
summary="查询上传任务状态",
description="查询单个上传任务的执行状态和结果",
)
async def get_upload_status(
task_id: str = Query(..., description="上传任务ID"),
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
) -> Any | dict:
result = await db.execute(
select(UploadTask).where(
UploadTask.id == task_id,
UploadTask.user_id == current_user.id,
UploadTask.deleted_at.is_(None),
)
)
task = result.scalar_one_or_none()
if not task:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail="任务不存在或无权限查看",
)
status_map = {
1: "待上传",
2: "上传中",
3: "上传成功",
4: "上传失败",
}
return {
"code": 0,
"data": {
"task_id": task.id,
"status": task.status,
"status_text": status_map.get(task.status, "未知"),
"advertiser_id": task.advertiser_id,
"resource_id": task.resource_id,
"note": task.note,
"created_at": task.created_at,
"updated_at": task.updated_at,
},
}
@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:
offset = (page - 1) * page_size
query = select(UploadTask).where(
UploadTask.user_id == current_user.id,
UploadTask.deleted_at.is_(None),
)
if status is not None:
query = query.where(UploadTask.status == status)
result = await db.execute(
query.order_by(UploadTask.created_at.desc())
.offset(offset)
.limit(page_size)
)
tasks = result.scalars().all()
count_query = select(func.count(UploadTask.id)).where(
UploadTask.user_id == current_user.id,
UploadTask.deleted_at.is_(None),
)
if status is not None:
count_query = count_query.where(UploadTask.status == status)
total_result = await db.execute(count_query)
total = total_result.scalar_one()
status_map = {
1: "待上传",
2: "上传中",
3: "上传成功",
4: "上传失败",
}
return {
"code": 0,
"data": [
{
"task_id": task.id,
"status": task.status,
"status_text": status_map.get(task.status, "未知"),
"advertiser_id": task.advertiser_id,
"resource_id": task.resource_id,
"note": task.note,
"created_at": task.created_at,
"updated_at": task.updated_at,
}
for task in tasks
],
"pagination": {
"page": page,
"page_size": page_size,
"total": total,
},
}