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: try: 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"], } except Exception as e: return { "code": 0, "message": f"查询上传任务历史失败:{str(e)}", }