531 lines
19 KiB
Python
531 lines
19 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.models.generated_resource import GeneratedResource
|
|
|
|
|
|
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)}",
|
|
}
|
|
|
|
#读取指定资源id的素材信息,包括size大小,尺寸,帧率,编码格式,码率,高宽比例
|
|
@router.get(
|
|
"/upload-material/{resource_id}",
|
|
summary="查询上传素材信息",
|
|
description="查询指定上传素材的详细信息",
|
|
)
|
|
async def get_upload_material_info(
|
|
resource_id: str,
|
|
db: AsyncSession = Depends(get_db),
|
|
) -> Any | dict:
|
|
try:
|
|
upload_material = await db.execute(
|
|
select(GeneratedResource)
|
|
.where(
|
|
GeneratedResource.id == resource_id,
|
|
)
|
|
)
|
|
upload_material = upload_material.scalar_one_or_none()
|
|
if not upload_material:
|
|
return {"code": 1, "message": f"素材{resource_id}不存在"}
|
|
storage_path = upload_material.storage_path
|
|
|
|
import ffmpeg
|
|
# 使用 ffmpeg.probe 获取视频的元数据[reference:20]
|
|
probe = ffmpeg.probe(storage_path)
|
|
|
|
# 从 'format' 中获取文件信息和码率[reference:21]
|
|
format_info = probe['format']
|
|
bit_rate = int(format_info.get('bit_rate', 0)) # 码率,单位 bps[reference:22]
|
|
file_size = int(format_info.get('size', 0)) # 文件大小,单位 bytes[reference:23]
|
|
duration = float(format_info.get('duration', 0)) # 时长,单位秒[reference:24]
|
|
|
|
# 从 'streams' 中查找视频流(通常是第一个视频流)
|
|
video_stream = next((stream for stream in probe['streams'] if stream['codec_type'] == 'video'), None)
|
|
if video_stream is None:
|
|
return None
|
|
|
|
width = int(video_stream['width'])
|
|
height = int(video_stream['height'])
|
|
# 帧率可能以分数形式表示,如 "30000/1001"[reference:25]
|
|
r_frame_rate = video_stream.get('r_frame_rate', '0/0')
|
|
if '/' in r_frame_rate:
|
|
num, den = map(int, r_frame_rate.split('/'))
|
|
fps = num / den if den != 0 else 0
|
|
else:
|
|
fps = float(r_frame_rate)
|
|
|
|
codec_name = video_stream.get('codec_name', 'unknown') # 编码格式名称,如 h264[reference:26]
|
|
|
|
return {
|
|
"width": width,
|
|
"height": height,
|
|
"fps": fps,
|
|
"codec": codec_name,
|
|
"bit_rate": bit_rate, # 单位 bps
|
|
"file_size": file_size, # 单位 bytes
|
|
"duration": duration,
|
|
"aspect_ratio": width / height
|
|
}
|
|
|
|
except Exception as e:
|
|
return {
|
|
"code": 0,
|
|
"message": f"查询上传素材信息失败:{str(e)}",
|
|
} |