Merge branch 'main' of https://gitee.com/wg123/video-gen
This commit is contained in:
@@ -0,0 +1,37 @@
|
||||
"""6idufv2q1c_add_credits_ratio_增加上传视频积分规则
|
||||
|
||||
Revision ID: a1b2c3d4e5f7
|
||||
Revises: 6idufv2q1c
|
||||
Create Date: 2026-07-01 00:00:00.000000
|
||||
|
||||
该文件包含 2026-07-01 的数据库迁移内容:
|
||||
1. 积分规则表增加传入视频计费字段(input_video_ratio, input_video_base_credits, input_video_per_second_credits)
|
||||
"""
|
||||
from typing import Sequence, Union
|
||||
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
|
||||
|
||||
revision: str = '6idufv2q1c'
|
||||
down_revision: Union[str, None] = 'a1b2c3d4e5f7'
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
# ========================================
|
||||
# 2026-07-01 - 积分规则表增加传入视频计费字段
|
||||
# ========================================
|
||||
op.add_column('credit_ratios', sa.Column('input_video_ratio', sa.Float(), server_default='1.0', nullable=False))
|
||||
op.add_column('credit_ratios', sa.Column('input_video_base_credits', sa.Float(), server_default='0.0', nullable=False))
|
||||
op.add_column('credit_ratios', sa.Column('input_video_per_second_credits', sa.Float(), server_default='0.5', nullable=False))
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
# ========================================
|
||||
# 2026-07-01 - 积分规则表增加传入视频计费字段(回滚)
|
||||
# ========================================
|
||||
op.drop_column('credit_ratios', 'input_video_per_second_credits')
|
||||
op.drop_column('credit_ratios', 'input_video_base_credits')
|
||||
op.drop_column('credit_ratios', 'input_video_ratio')
|
||||
@@ -0,0 +1,25 @@
|
||||
"""add max image count to image engines
|
||||
|
||||
Revision ID: a1b2c3d4e5f6
|
||||
Revises: f7a3b2c1d4e5
|
||||
Create Date: 2026-07-01 00:00:00.000000
|
||||
|
||||
"""
|
||||
from typing import Sequence, Union
|
||||
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
|
||||
|
||||
revision: str = 'a1b2c3d4e5f7'
|
||||
down_revision: Union[str, None] = 'f7a3b2c1d4e5'
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.add_column('image_engines', sa.Column('max_image_count', sa.Integer(), server_default='0', nullable=False))
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_column('image_engines', 'max_image_count')
|
||||
@@ -1400,9 +1400,10 @@ async def admin_list_generation_records(
|
||||
):
|
||||
"""List all generation records across all users, with optional filters."""
|
||||
query = (
|
||||
select(GenerationRecord, User.username, Project.name)
|
||||
select(GenerationRecord, User.username, Project.name, Project.industry, IndustryConfig.label)
|
||||
.join(User, GenerationRecord.user_id == User.id)
|
||||
.join(Project, GenerationRecord.project_id == Project.id)
|
||||
.outerjoin(IndustryConfig, Project.industry == IndustryConfig.key)
|
||||
.where(GenerationRecord.deleted_at.is_(None), Project.deleted_at.is_(None))
|
||||
.order_by(GenerationRecord.created_at.desc())
|
||||
)
|
||||
@@ -1427,7 +1428,7 @@ async def admin_list_generation_records(
|
||||
rows = result.all()
|
||||
|
||||
items = []
|
||||
for record, username, project_name in rows:
|
||||
for record, username, project_name, industry, industry_label in rows:
|
||||
refs = None
|
||||
if record.media_references:
|
||||
try:
|
||||
@@ -1440,6 +1441,7 @@ async def admin_list_generation_records(
|
||||
"username": username,
|
||||
"project_id": record.project_id,
|
||||
"project_name": project_name,
|
||||
"industry": industry_label or industry,
|
||||
"original_prompt": record.original_prompt,
|
||||
"optimized_prompt": record.optimized_prompt,
|
||||
"duration": record.duration,
|
||||
|
||||
@@ -48,70 +48,40 @@ async def get_credit_ratios(
|
||||
current_user: User = Depends(get_current_user),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
import json
|
||||
|
||||
async def get_engine_with_ratios(gen_type: str, engines: list):
|
||||
for engine in engines:
|
||||
async def get_ratios_for_engine_type(gen_type: str, engine_ids: list):
|
||||
for engine_id in engine_ids:
|
||||
result = await db.execute(
|
||||
select(CreditRatio)
|
||||
.where(CreditRatio.gen_type == gen_type)
|
||||
.where(CreditRatio.model_config_id == engine.id)
|
||||
.where(CreditRatio.model_config_id == engine_id)
|
||||
)
|
||||
ratios = result.scalars().all()
|
||||
if ratios:
|
||||
ratios_out = [CreditRatioOut.model_validate(r) for r in ratios]
|
||||
engine_info = {
|
||||
"id": engine.id,
|
||||
"name": engine.name,
|
||||
"provider": engine.provider,
|
||||
"ratios": ratios_out,
|
||||
}
|
||||
if gen_type == "video":
|
||||
try:
|
||||
supported_ratios = json.loads(engine.supported_ratios) if engine.supported_ratios else []
|
||||
except Exception:
|
||||
supported_ratios = []
|
||||
try:
|
||||
supported_resolutions = json.loads(engine.supported_resolutions) if engine.supported_resolutions else []
|
||||
except Exception:
|
||||
supported_resolutions = []
|
||||
try:
|
||||
supported_durations = json.loads(engine.supported_durations) if engine.supported_durations else []
|
||||
except Exception:
|
||||
supported_durations = []
|
||||
engine_info.update({
|
||||
"supported_ratios": supported_ratios,
|
||||
"supported_resolutions": supported_resolutions,
|
||||
"supported_durations": supported_durations,
|
||||
"max_duration": engine.max_duration,
|
||||
"max_image_count": engine.max_image_count,
|
||||
"max_video_count": engine.max_video_count,
|
||||
})
|
||||
return engine_info
|
||||
return None
|
||||
|
||||
return [CreditRatioOut.model_validate(r) for r in ratios]
|
||||
return []
|
||||
|
||||
video_engines_result = await db.execute(
|
||||
select(VideoEngine)
|
||||
select(VideoEngine.id)
|
||||
.where(VideoEngine.is_active == True)
|
||||
.order_by(VideoEngine.priority.desc())
|
||||
)
|
||||
video_engines = video_engines_result.scalars().all()
|
||||
|
||||
video_engine_ids = video_engines_result.scalars().all()
|
||||
|
||||
image_engines_result = await db.execute(
|
||||
select(ImageEngine)
|
||||
select(ImageEngine.id)
|
||||
.where(ImageEngine.is_active == True)
|
||||
.order_by(ImageEngine.priority.desc())
|
||||
)
|
||||
image_engines = image_engines_result.scalars().all()
|
||||
|
||||
image_engine_ids = image_engines_result.scalars().all()
|
||||
|
||||
grouped = {}
|
||||
|
||||
video_data = await get_engine_with_ratios("video", video_engines)
|
||||
if video_data:
|
||||
grouped["video"] = video_data
|
||||
|
||||
image_data = await get_engine_with_ratios("image", image_engines)
|
||||
if image_data:
|
||||
grouped["image"] = image_data
|
||||
|
||||
|
||||
video_ratios = await get_ratios_for_engine_type("video", video_engine_ids)
|
||||
if video_ratios:
|
||||
grouped["video"] = video_ratios
|
||||
|
||||
image_ratios = await get_ratios_for_engine_type("image", image_engine_ids)
|
||||
if image_ratios:
|
||||
grouped["image"] = image_ratios
|
||||
|
||||
return grouped
|
||||
|
||||
@@ -44,5 +44,6 @@ async def list_active_engines(
|
||||
"supported_models": models,
|
||||
"supported_sizes": sizes,
|
||||
"default_size": e.default_size,
|
||||
"max_image_count": e.max_image_count,
|
||||
})
|
||||
return {"items": items}
|
||||
@@ -1,11 +1,18 @@
|
||||
import logging
|
||||
from typing import Any, Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, Query
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from fastapi import APIRouter, Depends, Query, Body
|
||||
|
||||
from app.dependencies import get_db
|
||||
from app.schemas.resources_material import ResourcesMaterialListResponse
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from sqlalchemy import select
|
||||
from app.models.resources_material import ResourcesMaterial
|
||||
from app.models.pre_test_template import PreTestTemplate
|
||||
|
||||
from app.dependencies import get_db, get_current_user
|
||||
from app.schemas.resources_material import ResourcesMaterialListResponse, PreTestMaterialRequest
|
||||
from app.services.resources_material_service import get_resources_material_list
|
||||
from app.services.pre_test_queue import pre_test_queue
|
||||
from app.utils.id_gen import generate_id
|
||||
|
||||
router = APIRouter(prefix="/resources-material", tags=["resources-material"])
|
||||
|
||||
@@ -25,6 +32,7 @@ async def get_resources_material_list_api(
|
||||
page: int = Query(1, description="页码"),
|
||||
page_size: int = Query(20, description="每页数量"),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
current_user: dict = Depends(get_current_user),
|
||||
) -> Any | dict:
|
||||
items, total = await get_resources_material_list(
|
||||
db=db,
|
||||
@@ -35,6 +43,7 @@ async def get_resources_material_list_api(
|
||||
resource_type=resource_type,
|
||||
page=page,
|
||||
page_size=page_size,
|
||||
user_id=current_user.id,
|
||||
)
|
||||
|
||||
return {
|
||||
@@ -42,4 +51,96 @@ async def get_resources_material_list_api(
|
||||
"message": "查询成功",
|
||||
"data": items,
|
||||
"total": total,
|
||||
}
|
||||
}
|
||||
|
||||
#如果素材上传的时候没有指定前测,现在可以对已经上传好的素材进行前测
|
||||
@router.post(
|
||||
"/pre-commit",
|
||||
summary="素材列表提交前测",
|
||||
description="通过素材列表提交未前测的素材,异步处理,立即返回任务ID,结果稍后通过列表查询",
|
||||
)
|
||||
async def pre_test_material(
|
||||
req: PreTestMaterialRequest = Body(..., description="前测请求体"),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
current_user: dict = Depends(get_current_user),
|
||||
) -> Any | dict:
|
||||
|
||||
try:
|
||||
invalid_ids = []
|
||||
grouped_videos = {}
|
||||
|
||||
#检查前测模板是否有效
|
||||
pre_test_template = await db.execute(
|
||||
select(PreTestTemplate)
|
||||
.where(
|
||||
PreTestTemplate.id == req.pre_test_template_id,
|
||||
PreTestTemplate.user_id == current_user.id,
|
||||
PreTestTemplate.deleted_at.is_(None),
|
||||
)
|
||||
)
|
||||
pre_test_template = pre_test_template.scalar_one_or_none()
|
||||
if not pre_test_template:
|
||||
return {"code": 1, "message": f"前测模板{req.pre_test_template_id}不存在或不属于当前用户"}
|
||||
|
||||
for resource_material_id in req.resources_material_ids:
|
||||
resource_material = await db.execute(
|
||||
select(ResourcesMaterial)
|
||||
.where(
|
||||
ResourcesMaterial.id == resource_material_id,
|
||||
ResourcesMaterial.user_id == current_user.id,
|
||||
ResourcesMaterial.deleted_at.is_(None),
|
||||
)
|
||||
)
|
||||
resource_material = resource_material.scalar_one_or_none()
|
||||
|
||||
if not resource_material:
|
||||
invalid_ids.append(f"{resource_material_id}: 资源不存在或不属于当前用户")
|
||||
continue
|
||||
|
||||
if resource_material.status is not None:
|
||||
invalid_ids.append(f"{resource_material_id}: 该资源已进行过前测")
|
||||
continue
|
||||
|
||||
if not resource_material.upload_id:
|
||||
invalid_ids.append(f"{resource_material_id}: 该资源未上传到平台")
|
||||
continue
|
||||
|
||||
if resource_material.resource_type != "video":
|
||||
invalid_ids.append(f"{resource_material_id}: 前测仅支持视频类型")
|
||||
continue
|
||||
|
||||
key = f"{resource_material.oauth_id}|{resource_material.advertiser_id}"
|
||||
if key not in grouped_videos:
|
||||
grouped_videos[key] = []
|
||||
grouped_videos[key].append({
|
||||
"id": resource_material.id,
|
||||
"upload_id": resource_material.upload_id,
|
||||
})
|
||||
|
||||
if invalid_ids:
|
||||
return {"code": 1, "message": "; ".join(invalid_ids)}
|
||||
|
||||
if not grouped_videos:
|
||||
return {"code": 1, "message": "没有有效的视频资源"}
|
||||
|
||||
task_id = generate_id()
|
||||
|
||||
await pre_test_queue.enqueue({
|
||||
"task_id": task_id,
|
||||
"grouped_videos": grouped_videos,
|
||||
"pre_test_template_id": req.pre_test_template_id,
|
||||
})
|
||||
|
||||
return {
|
||||
"code": 0,
|
||||
"message": "任务已提交,正在处理中",
|
||||
"data": {
|
||||
"task_id": task_id,
|
||||
"total_groups": len(grouped_videos),
|
||||
},
|
||||
}
|
||||
|
||||
except ValueError as e:
|
||||
return {"code": 1, "message": str(e)}
|
||||
except Exception as e:
|
||||
return {"code": 1, "message": str(e)}
|
||||
@@ -10,6 +10,8 @@ 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
|
||||
@@ -355,51 +357,25 @@ async def batch_update_filename(
|
||||
.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 = {}
|
||||
existing_names = set(row[0] for row in result.all())
|
||||
|
||||
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)
|
||||
counter = 1
|
||||
new_file_name = file_name
|
||||
|
||||
while new_file_name in existing_names:
|
||||
new_file_name = f"{base_name}_{counter}{ext}"
|
||||
counter += 1
|
||||
|
||||
item["resource"].file_name = new_file_name
|
||||
db.add(item["resource"])
|
||||
existing_names.add(new_file_name)
|
||||
|
||||
results.append({
|
||||
"source_id": item["source_id"],
|
||||
@@ -421,7 +397,7 @@ async def batch_update_filename(
|
||||
}
|
||||
except Exception as e:
|
||||
return {
|
||||
"code": 0,
|
||||
"code": 1,
|
||||
"message": f"批量修改文件名失败:{str(e)}",
|
||||
"success_count": 0,
|
||||
"fail_count": 0,
|
||||
@@ -461,5 +437,3 @@ async def get_upload_history(
|
||||
"code": 0,
|
||||
"message": f"查询上传任务历史失败:{str(e)}",
|
||||
}
|
||||
|
||||
|
||||
@@ -80,6 +80,10 @@ async def lifespan(app: FastAPI):
|
||||
from app.tasks.pre_test_result_task import poll_pre_test_results
|
||||
pre_test_poll_task = asyncio.create_task(poll_pre_test_results())
|
||||
|
||||
# 启动前测任务队列
|
||||
from app.services.pre_test_queue import pre_test_queue
|
||||
pre_test_queue_task = asyncio.create_task(pre_test_queue.run())
|
||||
|
||||
# 启动时立即同步一次未支付订单
|
||||
asyncio.create_task(asyncio.sleep(5)) # 等待5秒后再同步,让系统完全启动
|
||||
async def startup_sync():
|
||||
@@ -104,6 +108,8 @@ async def lifespan(app: FastAPI):
|
||||
await queue_task
|
||||
upload_queue.stop()
|
||||
await upload_queue_task
|
||||
pre_test_queue.stop()
|
||||
await pre_test_queue_task
|
||||
material_consumption_queue.stop()
|
||||
await consumption_queue_task
|
||||
consumption_schedule_task.cancel()
|
||||
|
||||
@@ -22,3 +22,6 @@ class CreditRatio(Base, TimestampMixin):
|
||||
ratio: Mapped[float] = mapped_column(Float, nullable=False)
|
||||
base_credits: Mapped[float] = mapped_column(Float, default=80.0)
|
||||
per_second_credits: Mapped[float] = mapped_column(Float, default=2.0)
|
||||
input_video_ratio: Mapped[float] = mapped_column(Float, default=1.0)
|
||||
input_video_base_credits: Mapped[float] = mapped_column(Float, default=0.0)
|
||||
input_video_per_second_credits: Mapped[float] = mapped_column(Float, default=0.5)
|
||||
|
||||
@@ -17,6 +17,7 @@ class ImageEngine(Base, TimestampMixin):
|
||||
# {"2K":{"1:1":"2048×2048",...}, "4K":{"1:1":"4096×4096",...}}
|
||||
supported_sizes: Mapped[str] = mapped_column(Text, default='{}')
|
||||
default_size: Mapped[str] = mapped_column(String(32), default="2K")
|
||||
max_image_count: Mapped[int] = mapped_column(Integer, default=0)
|
||||
generate_url: Mapped[str | None] = mapped_column(String(512), nullable=True, default="")
|
||||
is_active: Mapped[bool] = mapped_column(Boolean, default=True)
|
||||
priority: Mapped[int] = mapped_column(Integer, default=0)
|
||||
|
||||
@@ -43,6 +43,24 @@ class CreditRatioCreate(BaseModel):
|
||||
description="视频每秒积分。图片规则通常为 0",
|
||||
examples=[2.0],
|
||||
)
|
||||
input_video_ratio: float = Field(
|
||||
default=1.0,
|
||||
ge=0,
|
||||
description="传入视频积分倍率。视频生成时,用户上传参考视频的额外积分倍率",
|
||||
examples=[1.0],
|
||||
)
|
||||
input_video_base_credits: float = Field(
|
||||
default=0.0,
|
||||
ge=0,
|
||||
description="传入视频基础积分。视频生成时,用户上传参考视频的基础积分",
|
||||
examples=[0.0],
|
||||
)
|
||||
input_video_per_second_credits: float = Field(
|
||||
default=0.5,
|
||||
ge=0,
|
||||
description="传入视频每秒积分。视频生成时,用户上传参考视频每秒消耗的积分",
|
||||
examples=[0.5],
|
||||
)
|
||||
|
||||
|
||||
class CreditRatioOut(CreditRatioCreate):
|
||||
|
||||
@@ -14,6 +14,7 @@ class GenerationAIReference(BaseModel):
|
||||
"url": "https://example.com/reference.png",
|
||||
"type": "image",
|
||||
"name": "参考图.png",
|
||||
"duration": 5.0,
|
||||
}
|
||||
}
|
||||
)
|
||||
@@ -33,6 +34,12 @@ class GenerationAIReference(BaseModel):
|
||||
description="参考素材名称,前端展示用,可为空",
|
||||
examples=["参考图.png"],
|
||||
)
|
||||
duration: float | None = Field(
|
||||
None,
|
||||
ge=0,
|
||||
description="视频素材时长(秒)。type=video 时使用,用于视频素材计费和时长校验",
|
||||
examples=[5.0],
|
||||
)
|
||||
|
||||
|
||||
class GenerationAITaskCreate(BaseModel):
|
||||
@@ -168,6 +175,7 @@ class GenerationAIImageEngineOptionOut(BaseModel):
|
||||
)
|
||||
default_size: str | None = Field(None, description="默认图片分辨率档位,例如 2K")
|
||||
priority: int = Field(0, description="引擎优先级,数值越大越优先")
|
||||
max_image_count: int = Field(0, description="最大图片数量")
|
||||
|
||||
|
||||
class GenerationAIVideoEngineOptionOut(BaseModel):
|
||||
@@ -182,6 +190,8 @@ class GenerationAIVideoEngineOptionOut(BaseModel):
|
||||
supported_durations: list[int] = Field(default_factory=list, description="支持的视频时长列表,单位秒")
|
||||
max_duration: int | None = Field(None, description="最大视频时长,单位秒")
|
||||
priority: int = Field(0, description="引擎优先级,数值越大越优先")
|
||||
max_image_count: int | None = Field(None, description="最大图片数量")
|
||||
max_video_count: int | None = Field(None, description="最大视频数量")
|
||||
|
||||
|
||||
class GenerationAIEngineGroupOut(BaseModel):
|
||||
@@ -219,6 +229,7 @@ class GenerationAIEngineOptionsOut(BaseModel):
|
||||
},
|
||||
"default_size": "2K",
|
||||
"priority": 10,
|
||||
"max_image_count": 0,
|
||||
}
|
||||
],
|
||||
"video": [
|
||||
@@ -232,6 +243,8 @@ class GenerationAIEngineOptionsOut(BaseModel):
|
||||
"supported_durations": [4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15],
|
||||
"max_duration": 15,
|
||||
"priority": 10,
|
||||
"max_image_count": 2,
|
||||
"max_video_count": 0,
|
||||
}
|
||||
],
|
||||
}
|
||||
|
||||
@@ -12,6 +12,7 @@ class ImageEngineCreate(BaseModel):
|
||||
supported_models: str = Field(default='["doubao-seedream-5-0-260128"]')
|
||||
supported_sizes: str = Field(default='{}')
|
||||
default_size: str = Field(default="2K", max_length=32)
|
||||
max_image_count: int = Field(default=0)
|
||||
generate_url: str = Field(default="", max_length=512)
|
||||
is_active: bool = True
|
||||
priority: int = 0
|
||||
@@ -31,6 +32,7 @@ class ImageEnginePublic(BaseModel):
|
||||
supported_models: list[str] = []
|
||||
supported_sizes: dict[str, dict[str, str]] = {}
|
||||
default_size: str = "2K"
|
||||
max_image_count: int = 0
|
||||
|
||||
|
||||
class ImageEngineListResponse(BaseModel):
|
||||
|
||||
@@ -56,4 +56,9 @@ class ResourcesMaterialListResponse(BaseModel):
|
||||
code: int = Field(0, description="返回码,0表示成功")
|
||||
message: str = Field("查询成功", description="返回消息")
|
||||
data: list[ResourcesMaterialOut] = Field(..., description="素材列表数据")
|
||||
total: int = Field(..., description="总记录数")
|
||||
total: int = Field(..., description="总记录数")
|
||||
|
||||
|
||||
class PreTestMaterialRequest(BaseModel):
|
||||
resources_material_ids: list[str] = Field(..., description="资源素材表id列表")
|
||||
pre_test_template_id: str = Field(..., description="前测模板id")
|
||||
@@ -65,6 +65,7 @@ async def calc_video_credits(
|
||||
duration: int,
|
||||
resolution: str,
|
||||
engine_id: str | None = None,
|
||||
input_video_duration: float | None = None,
|
||||
) -> float:
|
||||
"""Calculate video credits using CreditRatio table, with fallback to hardcoded.
|
||||
|
||||
@@ -72,8 +73,9 @@ async def calc_video_credits(
|
||||
1. gen_type=video + engine_id + resolution 精确规则;
|
||||
2. gen_type=video + resolution 下 base_credits/per_second_credits 最高规则;
|
||||
3. 原硬编码默认算法。
|
||||
|
||||
input_video_duration: 用户上传的参考视频总时长(秒),不为空时额外计费
|
||||
"""
|
||||
# 如果engine_id为空,默认查询权重最高的视频引擎积分规则
|
||||
if not engine_id:
|
||||
video_engines_result = await db.execute(
|
||||
select(VideoEngine.id)
|
||||
@@ -89,13 +91,21 @@ async def calc_video_credits(
|
||||
engine_id=engine_id,
|
||||
)
|
||||
if ratio:
|
||||
return round((ratio.base_credits + ratio.per_second_credits * duration) * ratio.ratio, 2)
|
||||
base_cost = (ratio.base_credits + ratio.per_second_credits * duration) * ratio.ratio
|
||||
if input_video_duration and input_video_duration > 0:
|
||||
input_video_cost = (
|
||||
ratio.input_video_base_credits + ratio.input_video_per_second_credits * input_video_duration
|
||||
) * ratio.input_video_ratio
|
||||
base_cost += input_video_cost
|
||||
return round(base_cost, 2)
|
||||
|
||||
# Fallback
|
||||
base = 60.0
|
||||
duration_cost = duration * 2.0
|
||||
multiplier = {"480p": 1, "1080p": 2, "720p": 1.5}.get(resolution, 1.0)
|
||||
return round((base + duration_cost) * multiplier, 2)
|
||||
total = (base + duration_cost) * multiplier
|
||||
if input_video_duration and input_video_duration > 0:
|
||||
total += input_video_duration * 0.5 * multiplier
|
||||
return round(total, 2)
|
||||
|
||||
|
||||
def calc_credits(duration: int, resolution: str) -> float:
|
||||
|
||||
@@ -186,6 +186,7 @@ async def list_generation_ai_engine_options(db: AsyncSession) -> GenerationAIEng
|
||||
supported_sizes=_image_supported_sizes(engine),
|
||||
default_size=engine.default_size,
|
||||
priority=engine.priority or 0,
|
||||
max_image_count=engine.max_image_count,
|
||||
)
|
||||
for engine in image_result.scalars().all()
|
||||
]
|
||||
@@ -200,6 +201,8 @@ async def list_generation_ai_engine_options(db: AsyncSession) -> GenerationAIEng
|
||||
supported_durations=_parse_list(engine.supported_durations, []),
|
||||
max_duration=engine.max_duration,
|
||||
priority=engine.priority or 0,
|
||||
max_image_count=engine.max_image_count,
|
||||
max_video_count=engine.max_video_count,
|
||||
)
|
||||
for engine in video_result.scalars().all()
|
||||
]
|
||||
@@ -298,6 +301,18 @@ async def create_async_generation_task(db: AsyncSession, current_user: User, req
|
||||
raise HTTPException(status_code=400, detail=f"视频时长不支持: {duration}")
|
||||
if engine.max_duration and duration > engine.max_duration:
|
||||
raise HTTPException(status_code=400, detail=f"视频时长不能超过 {engine.max_duration} 秒")
|
||||
|
||||
input_video_duration = 0.0
|
||||
if refs:
|
||||
video_refs = [r for r in refs if r.get("type") == "video"]
|
||||
for ref in video_refs:
|
||||
ref_duration = float(ref.get("duration") or 0)
|
||||
if ref_duration < 2:
|
||||
raise HTTPException(status_code=400, detail=f"视频素材最短不能少于 2 秒")
|
||||
input_video_duration += ref_duration
|
||||
if input_video_duration > 15:
|
||||
raise HTTPException(status_code=400, detail=f"所有视频素材总时长不能超过 15 秒,当前 {input_video_duration:.1f} 秒")
|
||||
|
||||
media_billing = await charge_generation_media_by_params(
|
||||
db,
|
||||
user_id=current_user.id,
|
||||
@@ -306,6 +321,7 @@ async def create_async_generation_task(db: AsyncSession, current_user: User, req
|
||||
duration=duration,
|
||||
resolution=resolution,
|
||||
engine_id=engine.id,
|
||||
input_video_duration=input_video_duration if input_video_duration > 0 else None,
|
||||
project_name="AI生成任务",
|
||||
description_prefix="AI创作-",
|
||||
owner_type=OWNER_CHAT_GENERATION_TASK,
|
||||
@@ -339,6 +355,51 @@ async def create_async_generation_task(db: AsyncSession, current_user: User, req
|
||||
return task
|
||||
|
||||
|
||||
def _resolve_error_message(error_message: str | None) -> str | None:
|
||||
"""匹配 ARK_ERRORS 字典,将原始错误码转换为友好提示。
|
||||
与 app/api/v1/generation.py 的 _record_to_out 保持一致。
|
||||
|
||||
注意:celery 任务中已调用 extract_error_message 将错误码转为中文提示后存入数据库,
|
||||
所以到达此函数的 message 可能是:
|
||||
1. 已翻译的中文提示(ARK_ERRORS 的 value)→ 直接返回
|
||||
2. 原始错误字符串(含 code='...' 或 JSON 格式)→ 匹配 ARK_ERRORS
|
||||
3. 未知内容 → 返回 "生成失败"
|
||||
"""
|
||||
if not error_message:
|
||||
return error_message
|
||||
|
||||
from app.services.error_codes import ARK_ERRORS
|
||||
|
||||
# 如果已经是 ARK_ERRORS 中已翻译的中文值,直接返回
|
||||
if error_message in ARK_ERRORS.values():
|
||||
return error_message
|
||||
|
||||
import re
|
||||
|
||||
# 匹配以下格式中的错误码:
|
||||
# 1. {'error': {'code': 'XXX', ...}} — str(error_obj) 的 Python dict 形式
|
||||
# 2. {"error": {"code": "XXX", ...}} — JSON 形式
|
||||
# 3. code='XXX' — 旧格式
|
||||
for pattern in [
|
||||
r"'code'\s*:\s*'([^']+)'", # 'code': 'XXX'
|
||||
r'"code"\s*:\s*"([^"]+)"', # "code": "XXX"
|
||||
r"code='([^']+)'", # code='XXX'
|
||||
]:
|
||||
match = re.search(pattern, error_message)
|
||||
if match:
|
||||
code = match.group(1)
|
||||
if code in ARK_ERRORS:
|
||||
return ARK_ERRORS[code]
|
||||
|
||||
# 兜底:按冒号分割,检查第二部分是否是已知错误码
|
||||
parts = error_message.split(":")
|
||||
if len(parts) >= 2 and parts[1].strip() in ARK_ERRORS:
|
||||
return ARK_ERRORS[parts[1].strip()]
|
||||
|
||||
# 没有匹配到已知错误码时,直接返回"生成失败"
|
||||
return "生成失败"
|
||||
|
||||
|
||||
def record_to_out(
|
||||
task: ChatGenerationTask,
|
||||
is_admin: bool = False,
|
||||
@@ -407,7 +468,7 @@ def record_to_out(
|
||||
video_tokens_used=task.video_tokens_used or 0,
|
||||
retry_count=task.retry_count or 0,
|
||||
poll_count=task.poll_count or 0,
|
||||
error_message=task.error_message,
|
||||
error_message=_resolve_error_message(task.error_message),
|
||||
created_at=task.created_at,
|
||||
generated_at=task.generated_at,
|
||||
)
|
||||
@@ -608,7 +669,7 @@ def generation_record_to_history_out(
|
||||
video_tokens_used=record.video_tokens_used or 0,
|
||||
retry_count=0,
|
||||
poll_count=0,
|
||||
error_message=record.error_message,
|
||||
error_message=_resolve_error_message(record.error_message),
|
||||
created_at=record.created_at,
|
||||
generated_at=record.generated_at,
|
||||
)
|
||||
|
||||
@@ -460,6 +460,7 @@ async def charge_generation_media_by_params(
|
||||
duration: int | None = None,
|
||||
resolution: str | None = None,
|
||||
engine_id: str | None = None,
|
||||
input_video_duration: float | None = None,
|
||||
project_name: str | None = None,
|
||||
description_prefix: str = "AI创作-",
|
||||
owner_type: str = OWNER_CHAT_GENERATION_TASK,
|
||||
@@ -522,7 +523,11 @@ async def charge_generation_media_by_params(
|
||||
)
|
||||
)
|
||||
elif gen_type == "video":
|
||||
amount = await calc_video_credits(db, duration or 5, resolution or "720p", engine_id=engine_id)
|
||||
amount = await calc_video_credits(
|
||||
db, duration or 5, resolution or "720p",
|
||||
engine_id=engine_id,
|
||||
input_video_duration=input_video_duration,
|
||||
)
|
||||
items.append(
|
||||
await deduct_credits_locked_once(
|
||||
db,
|
||||
|
||||
@@ -0,0 +1,246 @@
|
||||
import asyncio
|
||||
from sqlalchemy import select, update
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.models.base import async_session
|
||||
from app.models.resources_material import ResourcesMaterial
|
||||
from app.models.user_oauth import UserOAuth
|
||||
from app.models.pre_test_template import PreTestTemplate
|
||||
from app.utils.logger import get_logger
|
||||
from app.utils.douyinApi import DouyinApi
|
||||
import json
|
||||
|
||||
logger = get_logger("pre_test_queue", "pre_test_queue")
|
||||
douyin_api = DouyinApi()
|
||||
|
||||
|
||||
class PreTestQueue:
|
||||
def __init__(self):
|
||||
self.queue: asyncio.Queue[dict] = asyncio.Queue()
|
||||
self.running = False
|
||||
|
||||
async def enqueue(self, task_data: dict):
|
||||
"""Add a pre-test task to the queue."""
|
||||
await self.queue.put(task_data)
|
||||
logger.info(f"Enqueued pre-test task: {task_data.get('task_id')}")
|
||||
|
||||
async def run(self):
|
||||
"""Main processing loop."""
|
||||
self.running = True
|
||||
logger.info("Pre-test queue started")
|
||||
|
||||
while self.running:
|
||||
try:
|
||||
task_data = await asyncio.wait_for(self.queue.get(), timeout=5.0)
|
||||
except asyncio.TimeoutError:
|
||||
continue
|
||||
|
||||
try:
|
||||
await self._process(task_data)
|
||||
except Exception as e:
|
||||
logger.error(f"Error processing pre-test task: {e}", exc_info=True)
|
||||
finally:
|
||||
self.queue.task_done()
|
||||
|
||||
logger.info("Pre-test queue stopped")
|
||||
|
||||
async def _process(self, task_data: dict):
|
||||
"""Process a single pre-test task."""
|
||||
task_id = task_data.get("task_id")
|
||||
grouped_videos = task_data.get("grouped_videos", {})
|
||||
pre_test_template_id = task_data.get("pre_test_template_id")
|
||||
|
||||
logger.info(f"Processing pre-test task: {task_id}")
|
||||
|
||||
for key, videos in grouped_videos.items():
|
||||
oauth_id, advertiser_id = key.split("|")
|
||||
video_ids = [v["upload_id"] for v in videos]
|
||||
|
||||
logger.info(f"Processing {len(video_ids)} videos for oauth_id={oauth_id}, advertiser_id={advertiser_id}")
|
||||
|
||||
try:
|
||||
async with async_session() as db:
|
||||
await _pre_test_material_batch(
|
||||
oauth_id=oauth_id,
|
||||
advertiser_id=advertiser_id,
|
||||
video_ids=video_ids,
|
||||
pre_test_template_id=pre_test_template_id,
|
||||
db=db,
|
||||
resource_material_ids=[v["id"] for v in videos],
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(f"Failed to process pre-test for oauth_id={oauth_id}, advertiser_id={advertiser_id}: {e}", exc_info=True)
|
||||
|
||||
def stop(self):
|
||||
"""Stop the queue."""
|
||||
self.running = False
|
||||
|
||||
|
||||
async def _pre_test_material_batch(
|
||||
oauth_id: str,
|
||||
advertiser_id: str,
|
||||
video_ids: list[str],
|
||||
pre_test_template_id: str,
|
||||
db: AsyncSession,
|
||||
resource_material_ids: list[str],
|
||||
) -> 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, resource_material_id in zip(video_ids, resource_material_ids):
|
||||
await _update_material_pre_test_status(
|
||||
db, resource_material_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, resource_material_id in zip(video_ids, resource_material_ids):
|
||||
await _update_material_pre_test_status(
|
||||
db, resource_material_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:
|
||||
#记录日志
|
||||
logger.error(f"前测提交失败, oauth_id={oauth_id}, advertiser_id={advertiser_id}: {json.dumps(response)}")
|
||||
note = response.get("message", "未知错误")
|
||||
for video_id, resource_material_id in zip(video_ids, resource_material_ids):
|
||||
await _update_material_pre_test_status(
|
||||
db, resource_material_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, resource_material_id) in enumerate(zip(video_ids, resource_material_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, resource_material_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, resource_material_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,
|
||||
resource_material_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.id == resource_material_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,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
pre_test_queue = PreTestQueue()
|
||||
@@ -18,11 +18,14 @@ async def get_resources_material_list(
|
||||
resource_type: Optional[str] = None,
|
||||
page: int = 1,
|
||||
page_size: int = 20,
|
||||
user_id: Optional[str] = None,
|
||||
) -> Tuple[list[ResourcesMaterialOut], int]:
|
||||
resource_alias = aliased(GeneratedResource)
|
||||
|
||||
query = select(ResourcesMaterial).where(ResourcesMaterial.deleted_at.is_(None))
|
||||
|
||||
|
||||
if user_id:
|
||||
query = query.where(ResourcesMaterial.user_id == user_id)
|
||||
if advertiser_id:
|
||||
query = query.where(ResourcesMaterial.advertiser_id == advertiser_id)
|
||||
if material_id:
|
||||
|
||||
Reference in New Issue
Block a user