Files
video-gen/video-gen-api/app/api/v1/upload_material.py
T
2026-06-26 17:14:55 +08:00

465 lines
17 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.models.user_oauth import UserOAuth
from app.models.user_oauth_account import UserOAuthAccount
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(
# "/batch-upload",
# summary="批量上传素材到平台",
# description="支持批量上传多个授权账户下的资源到素材库,预留下前测功能",
# )
# async def batch_upload_material(
# current_user: User = Depends(get_current_user),
# db: AsyncSession = Depends(get_db),
# ) -> Any | dict:
# try:
# result = await upload_material_to_platform(
# ["0019ef7700c503991c9"],
# ["1856633793022992"],
# "0019f018b4b3612e184",
# db,
# current_user.id,
# None,
# )
# return result
# except Exception as e:
# return {
# "code": 0,
# "message": str(e),
# }
@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:
try:
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)
# 检查资源id是否存在,非资源id
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:
#用户提交的直接是资源id
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:
other_info = {}
if task.is_pre_test == "1":
other_info["is_pre_test"] = task.is_pre_test
other_info["pre_test_template"] = task.pre_test_template
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,
other_info=json.dumps(other_info) if other_info else 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,
}
except Exception as e:
return {
"code": 1,
"message": str(e),
"task_ids": [],
"errors": [],
}
@router.post(
"/oauth_account_list",
summary="获取授权账户列表",
description="获取用户授权账户列表",
)
async def async_get_oauth_account_list(
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
) -> Any | dict:
try:
result = await db.execute(
select(
UserOAuthAccount.advertiser_id,
UserOAuthAccount.advertiser_name,
UserOAuthAccount.oauth_id,
UserOAuth.open_type,
).join(
UserOAuth,
UserOAuth.id == UserOAuthAccount.oauth_id
).where(
UserOAuth.user_id == current_user.id,
UserOAuth.deleted_at.is_(None),
UserOAuthAccount.deleted_at.is_(None),
)
)
accounts = result.all()
return {
"code": 0,
"message": "查询成功",
"data": [
{
"advertiser_id": account.advertiser_id,
"advertiser_name": account.advertiser_name,
"oauth_id": account.oauth_id,
"open_type": account.open_type,
}
for account in accounts
],
}
except Exception as e:
return {
"code": 1,
"message": str(e),
}
@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)}",
}