新增资源表文件名字段

This commit is contained in:
18610128193
2026-06-24 14:30:05 +08:00
parent 143ae52878
commit 313471c1dd
3 changed files with 202 additions and 16 deletions
@@ -0,0 +1,29 @@
"""资源表新增文件名字段
Revision ID: 5e2c124e5484
Revises: 2bafebb4be14
Create Date: 2026-06-24 14:28:56.094105
"""
from typing import Sequence, Union
from alembic import op
import sqlalchemy as sa
# revision identifiers, used by Alembic.
revision: str = '5e2c124e5484'
down_revision: Union[str, None] = '2bafebb4be14'
branch_labels: Union[str, Sequence[str], None] = None
depends_on: Union[str, Sequence[str], None] = None
def upgrade() -> None:
# ### commands auto generated by Alembic - please adjust! ###
op.add_column('generated_resources', sa.Column('file_name', sa.String(length=255), nullable=True, comment='文件名,平台素材名称'))
# ### end Alembic commands ###
def downgrade() -> None:
# ### commands auto generated by Alembic - please adjust! ###
op.drop_column('generated_resources', 'file_name')
# ### end Alembic commands ###
+172 -16
View File
@@ -15,7 +15,7 @@ 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.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
@@ -31,11 +31,11 @@ class UploadTaskRequest(BaseModel):
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="批量上传任务列表")
tasks: list[UploadTaskRequest] = Field(..., description="批量上传任务列表")
@router.post(
@@ -147,8 +147,84 @@ async def batch_upload_material(
all_results.append(task_result)
continue
source_model_map = {
"generation_records": "GenerationRecord",
"generated_resources": None,
"chat_generation_tasks": "ChatGenerationTask",
}
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)
task_result["result"] = {
"success": False,
"error": f"资源id [{invalid_ids_str}] 不可用或已删除",
"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"资源id {rid} 不可用或已删除" if rid in invalid_ids else "资源验证失败"
} 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
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)
task_result["result"] = {
"success": False,
"error": f"资源id [{invalid_ids_str}] 不可用或已删除",
"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"资源id {rid} 不可用或已删除" if rid in invalid_ids else "资源验证失败"
} 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
resource_ids_to_upload = valid_resource_ids
result = await upload_material_to_platform(
task.resource_ids,
resource_ids_to_upload,
task.advertiser_ids,
task.oauth_id,
db,
@@ -222,35 +298,110 @@ async def async_batch_upload_material(
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": [],
"errors": [],
}
task_ids = []
errors = []
for task in req.tasks:
if not task.advertiser_ids or not task.resource_ids:
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 task.resource_ids:
for resource_id in resource_ids_to_upload:
task_id = generate_id()
upload_task = UploadTask(
id = task_id,
id=task_id,
user_id=current_user.id,
oauth_id=task.oauth_id,
advertiser_id=advertiser_id,
@@ -265,10 +416,15 @@ async def async_batch_upload_material(
await db.commit()
message = "上传任务已提交"
if errors:
message = f"部分任务提交成功,{len(errors)} 个任务失败"
return {
"code": 0,
"message": "上传任务已提交",
"message": message,
"task_ids": task_ids,
"errors": errors,
}
@@ -25,6 +25,7 @@ class GeneratedResource(Base, TimestampMixin, SoftDeleteMixin):
remote_url: Mapped[str | None] = mapped_column(Text, nullable=True)
storage_type: Mapped[str] = mapped_column(String(32), default="local", nullable=False)
storage_path: Mapped[str | None] = mapped_column(Text, nullable=True)
file_name: Mapped[str | None] = mapped_column(String(255), nullable=True,comment="文件名,平台素材名称")
file_size_bytes: Mapped[int] = mapped_column(BigInteger, default=0, nullable=False)
source_model: Mapped[str] = mapped_column(String(64), index=True, nullable=False)