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.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.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") class BatchUploadRequest(BaseModel): tasks: list[UploadTaskRequest] = Field(..., description="批量上传任务列表") 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: # param = { # "account_ids" : json.dumps([1836693172153543]), # } # account_info = await DouyinApi().get_account_info("0019eb9f130027c05b8", param) # return account_info 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 = 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() return { "code": 0, "message": "上传任务已提交", "task_ids": task_ids, } @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"], }