Files
video-gen/video-gen-api/app/services/upload_queue.py
T

656 lines
23 KiB
Python

import asyncio
import logging
from datetime import datetime, timezone
from sqlalchemy import select, update
from sqlalchemy.orm import selectinload
from sqlalchemy.ext.asyncio import AsyncSession
from app.models.base import async_session
from app.models.upload_task import UploadTask
from app.models.generated_resource import GeneratedResource
from app.models.user_oauth import UserOAuth
from app.models.resources_material import ResourcesMaterial
from app.models.user_oauth_account import UserOAuthAccount
from app.models.pre_test_template import PreTestTemplate
from app.utils.id_gen import generate_id
from app.utils.douyinApi import DouyinApi
import os
import hashlib
import json
LOG_DIR = os.path.join(os.path.dirname(os.path.dirname(os.path.dirname(__file__))), "logs")
os.makedirs(LOG_DIR, exist_ok=True)
logger = logging.getLogger("upload_queue")
logger.setLevel(logging.INFO)
class DailyRotatingFileHandler(logging.FileHandler):
def __init__(self, directory, encoding=None):
self.directory = directory
filename = self._get_log_filename()
super().__init__(filename, encoding=encoding)
def _get_log_filename(self):
return os.path.join(self.directory, f"upload_queue-{datetime.now(timezone.utc).strftime('%Y-%m-%d')}.log")
def emit(self, record):
current_filename = self._get_log_filename()
if self.baseFilename != current_filename:
self.close()
self.baseFilename = current_filename
self.stream = self._open()
super().emit(record)
if not logger.handlers:
handler = DailyRotatingFileHandler(LOG_DIR, encoding="utf-8")
handler.setFormatter(logging.Formatter("%(asctime)s - %(levelname)s - %(message)s", "%Y-%m-%d %H:%M:%S"))
logger.addHandler(handler)
douyin_api = DouyinApi()
#上传素材队列,处理上传素材的任务
class UploadQueue:
def __init__(self):
self.queue: asyncio.Queue[str] = asyncio.Queue()
self.running = False
async def enqueue(self, task_id: str):
"""Add a task to the queue."""
await self.queue.put(task_id)
async def recover(self):
"""Recover pending tasks from DB on startup."""
async with async_session() as db:
result = await db.execute(
select(UploadTask).where(
UploadTask.status.in_([1, 2]),
UploadTask.deleted_at.is_(None),
)
)
tasks = result.scalars().all()
for task in tasks:
await self.queue.put(task.id)
logger.info(f"Recovered upload task: {task.id}")
async def run(self):
"""Main processing loop."""
self.running = True
logger.info("Upload queue started")
while self.running:
try:
task_id = await asyncio.wait_for(self.queue.get(), timeout=5.0)
except asyncio.TimeoutError:
continue
try:
await self._process(task_id)
except Exception as e:
logger.error(f"Error processing upload task {task_id}: {e}")
finally:
self.queue.task_done()
logger.info("Upload queue stopped")
async def _process(self, task_id: str):
"""Process a single upload task."""
async with async_session() as db:
result = await db.execute(
select(UploadTask).where(
UploadTask.id == task_id,
UploadTask.deleted_at.is_(None),
).with_for_update()
)
task = result.scalar_one_or_none()
if not task:
logger.warning(f"Upload task {task_id} not found or deleted")
return
if task.status == 2:
logger.warning(f"Upload task {task_id} is already running")
return
task.status = 2
await db.commit()
try:
result = await _upload_single_material(
task.id,
task.user_id,
task.oauth_id,
task.advertiser_id,
task.resource_id,
db=None,
other_info=json.loads(task.other_info) if task.other_info else None,
)
async with async_session() as db:
if result.get("success"):
await db.execute(
update(UploadTask).where(UploadTask.id == task_id).values(
status=3,
note=result.get("message", "上传成功"),
)
)
#请求巨量接口获取账户信息,如果存在则更新,否则插入新记录
param = {
"account_ids" : json.dumps([int(task.advertiser_id)]),
}
try:
account_info = await douyin_api.get_account_info(task.oauth_id, param)
code = account_info.get("code", 0)
try:
code = int(code)
except (ValueError, TypeError):
code = -1
if code != 0:
logger.error(f"Error getting account info: {json.dumps(account_info)}")
else:
existing_account = await db.execute(
select(UserOAuthAccount).where(
UserOAuthAccount.oauth_id == task.oauth_id,
UserOAuthAccount.advertiser_id == task.advertiser_id,
)
)
existing_account = existing_account.scalar_one_or_none()
if existing_account:
await db.execute(
update(UserOAuthAccount).where(
UserOAuthAccount.id == existing_account.id,
).values(
deleted_at=None,
advertiser_name= account_info.get("data", {}).get("account_detail_list", [{}])[0].get("advertiser_name", ""),
advertiser_role= "",
)
)
else:
db.add(UserOAuthAccount(
id=generate_id(),
oauth_id=task.oauth_id,
advertiser_id=task.advertiser_id,
advertiser_name= account_info.get("data", {}).get("account_detail_list", [{}])[0].get("advertiser_name", ""),
advertiser_role= "",
))
except Exception as e:
logger.error(f"Exception getting account info: {e}")
else:
await db.execute(
update(UploadTask).where(UploadTask.id == task_id).values(
status=4,
note=result.get("error", "上传失败"),
)
)
await db.commit()
message = result.get('message') or result.get('error', 'Unknown error')
logger.info(f"Upload task {task_id} completed: {'success' if result.get('success') else 'failed'}. Message: {message}")
except Exception as e:
async with async_session() as db:
await db.execute(
update(UploadTask).where(UploadTask.id == task_id).values(
status=4,
note=f"上传失败: {str(e)}",
)
)
await db.commit()
logger.error(f"Upload task {task_id} failed: {e}")
def stop(self):
"""Stop the queue."""
self.running = False
async def _upload_single_material(
task_id: str,
user_id: str,
oauth_id: str,
advertiser_id: str,
resource_id: str,
db=None,
other_info=None,
) -> dict:
own_db = False
if db is None:
from app.models.base import async_session
db = async_session()
own_db = True
try:
oauth = await db.execute(
select(UserOAuth).where(
UserOAuth.id == oauth_id,
UserOAuth.deleted_at.is_(None),
UserOAuth.user_id == user_id,
)
)
oauth = oauth.scalar_one_or_none()
if not oauth:
return {
"success": False,
"error": "授权记录不存在"
}
resource = await db.execute(
select(GeneratedResource).where(
GeneratedResource.id == resource_id,
GeneratedResource.user_id == user_id,
GeneratedResource.deleted_at.is_(None),
)
)
resource = resource.scalar_one_or_none()
if not resource:
return {
"success": False,
"error": f"资源 {resource_id} 不存在或不属于当前用户"
}
resource_type = resource.resource_type
storage_path = resource.storage_path
file_name = resource.file_name
if resource_type not in ["image", "video"]:
return {
"success": False,
"error": f"不支持的资源类型: {resource_type},仅支持image和video",
}
if not storage_path:
return {
"success": False,
"error": "资源本地存储路径为空",
}
result = await _upload_to_juliang(
oauth_id, storage_path, resource_type, advertiser_id, resource, db, user_id, file_name
)
if result.get("success") and resource_type == "video" and other_info and "is_pre_test" in other_info and other_info["is_pre_test"] == "1" and "pre_test_template" in other_info:
await _pre_test_material(
oauth_id,
advertiser_id,
[result.get("upload_id")],
other_info["pre_test_template"],
db,
)
return result
finally:
if own_db and db is not None:
await db.close()
async def _upload_to_juliang(
oauth_id: str,
storage_path: str,
resource_type: str,
advertiser_id: str,
resource: GeneratedResource,
db,
current_user_id: str,
file_name: str,
) -> dict:
#如果file_name不等于空,那么就是用file_name,否则用storage_path的文件名
if file_name:
filename = file_name
else:
filename = os.path.basename(storage_path)
if resource_type == "image":
if resource.file_size_bytes > 5 * 1024 * 1024:
return {
"resource_id": resource.id,
"advertiser_id": advertiser_id,
"filename": filename,
"success": False,
"error": "图片大小超过5M,不支持上传",
}
with open(storage_path, "rb") as f:
file_content = f.read()
image_signature = hashlib.md5(file_content).hexdigest()
data = {
"advertiser_id": advertiser_id,
"upload_type": "UPLOAD_BY_FILE",
"image_signature": image_signature,
"filename": filename,
}
files = {
"image_file": (filename, file_content, "image/png"),
}
response = await douyin_api.upload_image_material(oauth_id, data, files)
if response["code"] != 0:
return {
"resource_id": resource.id,
"advertiser_id": advertiser_id,
"filename": filename,
"success": False,
"error": response.get("message", "上传失败"),
}
image_id = response["data"]["id"]
material_id = response["data"]["material_id"]
existing = await db.execute(
select(ResourcesMaterial).where(
ResourcesMaterial.oauth_id == oauth_id,
ResourcesMaterial.advertiser_id == advertiser_id,
ResourcesMaterial.material_id == str(material_id),
ResourcesMaterial.upload_id == str(image_id),
ResourcesMaterial.deleted_at.is_(None),
)
)
existing = existing.scalar_one_or_none()
if existing:
return {
"resource_id": resource.id,
"advertiser_id": advertiser_id,
"filename": filename,
"success": True,
"material_id": str(material_id),
"upload_id": str(image_id),
"message": "素材已存在,跳过添加",
}
db.add(ResourcesMaterial(
id=generate_id(),
oauth_id=oauth_id,
advertiser_id=advertiser_id,
target_table="generated_resources",
target_id=resource.id,
material_id=str(material_id),
upload_id=str(image_id),
resource_type=resource_type,
user_id=current_user_id,
))
await db.commit()
return {
"resource_id": resource.id,
"advertiser_id": advertiser_id,
"filename": filename,
"success": True,
"material_id": str(material_id),
"upload_id": str(image_id),
}
elif resource_type == "video":
if resource.file_size_bytes > 500 * 1024 * 1024:
return {
"resource_id": resource.id,
"advertiser_id": advertiser_id,
"filename": filename,
"success": False,
"error": "视频大小超过500M,不支持上传",
}
with open(storage_path, "rb") as f:
file_content = f.read()
video_signature = hashlib.md5(file_content).hexdigest()
data = {
"advertiser_id": advertiser_id,
"upload_type": "UPLOAD_BY_FILE",
"video_signature": video_signature,
"filename": filename,
}
files = {
"video_file": (filename, file_content, "video/mp4"),
}
response = await douyin_api.upload_video_material(oauth_id, data, files)
if response["code"] != 0:
return {
"resource_id": resource.id,
"advertiser_id": advertiser_id,
"filename": filename,
"success": False,
"error": response.get("message", "上传失败"),
}
#上传成功,添加视频id,素材id到素材库
video_id = response["data"]["video_id"]
material_id = response["data"]["material_id"]
existing = await db.execute(
select(ResourcesMaterial).where(
ResourcesMaterial.oauth_id == oauth_id,
ResourcesMaterial.advertiser_id == advertiser_id,
ResourcesMaterial.material_id == str(material_id),
ResourcesMaterial.upload_id == str(video_id),
ResourcesMaterial.deleted_at.is_(None),
)
)
existing = existing.scalar_one_or_none()
if existing:
return {
"resource_id": resource.id,
"advertiser_id": advertiser_id,
"filename": filename,
"success": True,
"material_id": str(material_id),
"upload_id": str(video_id),
"message": "素材已存在,跳过添加",
}
db.add(ResourcesMaterial(
id=generate_id(),
oauth_id=oauth_id,
advertiser_id=advertiser_id,
target_table="generated_resources",
target_id=resource.id,
material_id=str(material_id),
upload_id=str(video_id),
resource_type=resource_type,
user_id=current_user_id,
))
await db.commit()
return {
"resource_id": resource.id,
"advertiser_id": advertiser_id,
"filename": filename,
"success": True,
"material_id": str(material_id),
"upload_id": str(video_id),
}
else:
return {
"resource_id": resource.id,
"advertiser_id": advertiser_id,
"filename": filename,
"success": False,
"error": f"不支持的资源类型: {resource_type}",
}
async def _pre_test_material(
oauth_id: str,
advertiser_id: str,
video_ids: list[str],
pre_test_template_id: str,
db: AsyncSession,
) -> any:
oauth = await db.execute(
select(UserOAuth).where(
UserOAuth.id == oauth_id,
UserOAuth.deleted_at.is_(None),
)
)
oauth = oauth.scalar_one_or_none()
if not oauth:
note = "授权记录不存在"
for video_id in video_ids:
await _update_material_pre_test_status(
db, oauth_id, advertiser_id, video_id,
task_id=None,
status="FAILED",
note=note,
pre_result=None,
pre_test_template_id=pre_test_template_id,
)
await db.commit()
return {"code": -1, "message": note, "data": {}}
pre_test_template = await db.execute(
select(PreTestTemplate).where(
PreTestTemplate.id == pre_test_template_id,
PreTestTemplate.deleted_at.is_(None),
PreTestTemplate.user_id == oauth.user_id,
)
)
pre_test_template = pre_test_template.scalar_one_or_none()
if not pre_test_template:
note = f"前测模板 {pre_test_template_id} 不存在"
for video_id in video_ids:
await _update_material_pre_test_status(
db, oauth_id, advertiser_id, video_id,
task_id=None,
status="FAILED",
note=note,
pre_result=None,
pre_test_template_id=pre_test_template_id,
)
await db.commit()
return {"code": -1, "message": note, "data": {}}
diagnose_config = {}
if pre_test_template.platform:
diagnose_config["platform"] = pre_test_template.platform
if pre_test_template.external_action:
diagnose_config["external_action"] = pre_test_template.external_action
if pre_test_template.cpa_bid:
diagnose_config["cpa_bid"] = pre_test_template.cpa_bid
if pre_test_template.audience_gender:
diagnose_config["audience_gender"] = pre_test_template.audience_gender
if pre_test_template.audience_age:
diagnose_config["audience_age"] = json.loads(pre_test_template.audience_age)
if pre_test_template.audience_region:
diagnose_config["audience_region"] = json.loads(pre_test_template.audience_region)
if pre_test_template.audience_network:
diagnose_config["audience_network"] = json.loads(pre_test_template.audience_network)
if pre_test_template.cus_name:
diagnose_config["cus_name"] = pre_test_template.cus_name
if pre_test_template.pricing_type:
diagnose_config["pricing_type"] = pre_test_template.pricing_type
if pre_test_template.cost_cap:
diagnose_config["cost_cap"] = pre_test_template.cost_cap
if pre_test_template.target_cost:
diagnose_config["target_cost"] = pre_test_template.target_cost
if pre_test_template.nobid:
diagnose_config["nobid"] = pre_test_template.nobid
if pre_test_template.cpc_bid:
diagnose_config["cpc_bid"] = pre_test_template.cpc_bid
if pre_test_template.budget:
diagnose_config["budget"] = pre_test_template.budget
params = {
"advertiser_id": int(advertiser_id),
"video_ids": video_ids,
"diagnose_config": diagnose_config,
}
response = await douyin_api.pre_test_material(oauth_id, params)
code = response.get("code", -1)
if code != 0:
note = response.get("message", "未知错误")
for video_id in video_ids:
await _update_material_pre_test_status(
db, oauth_id, advertiser_id, video_id,
task_id=None,
status="FAILED",
note=note,
pre_result=None,
pre_test_template_id=pre_test_template_id,
)
await db.commit()
return response
data = response.get("data", {})
task_ids = data.get("task_ids", [])
fail_video_ids = data.get("fail_video_ids", {})
success_count = 0
for i, video_id in enumerate(video_ids):
if video_id in fail_video_ids:
fail_info = fail_video_ids[video_id]
err_code = fail_info.get("err_code", "")
err_message = fail_info.get("err_message", "未知错误")
note = f"失败[{err_code}]: {err_message}"
await _update_material_pre_test_status(
db, oauth_id, advertiser_id, video_id,
task_id=None,
status="FAILED",
note=note,
pre_result=None,
pre_test_template_id=pre_test_template_id,
)
else:
task_id = str(task_ids[success_count]) if success_count < len(task_ids) else None
await _update_material_pre_test_status(
db, oauth_id, advertiser_id, video_id,
task_id=task_id,
status="PENDING",
note="",
pre_result=None,
pre_test_template_id=pre_test_template_id,
)
success_count += 1
await db.commit()
return response
async def _update_material_pre_test_status(
db: AsyncSession,
oauth_id: str,
advertiser_id: str,
upload_id: str,
task_id: str | None,
status: str,
note: str,
pre_result: str | None,
pre_test_template_id: str | None,
):
await db.execute(
update(ResourcesMaterial).where(
ResourcesMaterial.oauth_id == oauth_id,
ResourcesMaterial.advertiser_id == advertiser_id,
ResourcesMaterial.upload_id == upload_id,
ResourcesMaterial.deleted_at.is_(None),
).values(
task_id=task_id,
status=status,
note=note,
pre_result=pre_result,
pre_test_template_id=pre_test_template_id,
)
)
upload_queue = UploadQueue()