From 313471c1dda4a05471a9963994e5e59ce92ceb2f Mon Sep 17 00:00:00 2001 From: 18610128193 <10574456+chenweiqiang-123@user.noreply.gitee.com> Date: Wed, 24 Jun 2026 14:30:05 +0800 Subject: [PATCH] =?UTF-8?q?=E6=96=B0=E5=A2=9E=E8=B5=84=E6=BA=90=E8=A1=A8?= =?UTF-8?q?=E6=96=87=E4=BB=B6=E5=90=8D=E5=AD=97=E6=AE=B5?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../5e2c124e5484_资源表新增文件名字段.py | 29 +++ video-gen-api/app/api/v1/upload_material.py | 188 ++++++++++++++++-- .../app/models/generated_resource.py | 1 + 3 files changed, 202 insertions(+), 16 deletions(-) create mode 100644 video-gen-api/alembic/versions/5e2c124e5484_资源表新增文件名字段.py diff --git a/video-gen-api/alembic/versions/5e2c124e5484_资源表新增文件名字段.py b/video-gen-api/alembic/versions/5e2c124e5484_资源表新增文件名字段.py new file mode 100644 index 00000000..4892e96c --- /dev/null +++ b/video-gen-api/alembic/versions/5e2c124e5484_资源表新增文件名字段.py @@ -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 ### diff --git a/video-gen-api/app/api/v1/upload_material.py b/video-gen-api/app/api/v1/upload_material.py index 30116a8d..eb0acc3e 100644 --- a/video-gen-api/app/api/v1/upload_material.py +++ b/video-gen-api/app/api/v1/upload_material.py @@ -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, } diff --git a/video-gen-api/app/models/generated_resource.py b/video-gen-api/app/models/generated_resource.py index 79dc6aee..17d0af25 100644 --- a/video-gen-api/app/models/generated_resource.py +++ b/video-gen-api/app/models/generated_resource.py @@ -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)