This commit is contained in:
2026-07-02 09:36:55 +08:00
47 changed files with 3381 additions and 2069 deletions
@@ -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')
+4 -2
View File
@@ -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,
+21 -51
View File
@@ -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}
+106 -5
View File
@@ -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)}
+11 -37
View File
@@ -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)}",
}
+6
View File
@@ -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()
+3
View File
@@ -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)
+1
View File
@@ -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)
+18
View File
@@ -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")
+14 -4
View File
@@ -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: