Compare commits
55 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 3c782f89d1 | |||
| 2b4f11036f | |||
| 2881da65cf | |||
| 1037e218bb | |||
| f0514d7487 | |||
| bf8f62ec5b | |||
| cb911f5eba | |||
| f6fa3ea653 | |||
| 6a7709ba43 | |||
| 2c70dd4c29 | |||
| 5548e78eee | |||
| 44bd96b148 | |||
| dda67cd10f | |||
| 233d0272a9 | |||
| 631dd643c0 | |||
| 50489b05d6 | |||
| 252ea71d50 | |||
| d5f6e9499d | |||
| afd6b9c0b8 | |||
| 8a6ebb49be | |||
| 9e024caff5 | |||
| c7e9878f0d | |||
| b6243c8ab8 | |||
| 67e4b3fc6f | |||
| df654bce19 | |||
| bd99c161dd | |||
| 7fee203693 | |||
| 562ffc53cc | |||
| 9a67f727b3 | |||
| 5bb714f25b | |||
| 0f592add29 | |||
| 42dd5fabc6 | |||
| c7aed2c152 | |||
| 3428f4ef73 | |||
| 13c57771b9 | |||
| 5eedf190c1 | |||
| 7a9a97b31a | |||
| a6f89067e7 | |||
| 718deb32b9 | |||
| 054e81c3e7 | |||
| 07cf055a97 | |||
| 0e78f175fe | |||
| 371d1b9daf | |||
| 1a459d1ad2 | |||
| f097e16bc1 | |||
| 69eea54cd2 | |||
| b5cbc3b482 | |||
| 9106b4de2e | |||
| 7d639d5f9a | |||
| db6c237d68 | |||
| 7b68b94df6 | |||
| 8e19f24984 | |||
| 866d71a431 | |||
| 042512a527 | |||
| 3bb9c5dd4e |
+13
-4
@@ -211,15 +211,24 @@ COSYVOICE_CLONE_MODEL=voice-enrollment
|
||||
# 用于 AI 文案生成、智能剪辑等需要大模型能力的场景
|
||||
|
||||
DOUBAO_API_KEY=your-doubao-api-key
|
||||
DOUBAO_MODEL=doubao-seed-1-6-250615
|
||||
DOUBAO_FAST_MODEL=doubao-1-5-pro-32k-250115
|
||||
DOUBAO_MODEL=doubao-seed-2-1-pro-260915
|
||||
DOUBAO_FAST_MODEL=doubao-seed-2-1-lite-260915
|
||||
DOUBAO_BASE_URL=https://ark.cn-beijing.volces.com/api/v3
|
||||
DOUBAO_TIMEOUT=30
|
||||
DOUBAO_MAX_RETRIES=2
|
||||
# 视觉模型:pro 精度高,lite 速度快(viral-video 商品识别默认用 lite 提速)
|
||||
DOUBAO_VISION_MODEL=doubao-1-5-vision-pro-250328
|
||||
DOUBAO_VISION_LITE_MODEL=doubao-1-5-vision-lite-250315
|
||||
DOUBAO_VISION_MODEL=doubao-seed-2-1-pro-260915
|
||||
DOUBAO_VISION_LITE_MODEL=doubao-seed-2-1-lite-260915
|
||||
DOUBAO_VISION_USE_LITE=true
|
||||
# Embedding 向量化模型
|
||||
DOUBAO_EMBEDDING_MODEL=doubao-embedding-vision-251215
|
||||
# 视频模型(Seedance 2.5,统一走方舟;真人参考图通过信任链自动 AI 化)
|
||||
DOUBAO_VIDEO_MODEL=doubao-seedance-2-5-260628
|
||||
DOUBAO_VIDEO_TIMEOUT=480
|
||||
DOUBAO_VIDEO_POLL_INTERVAL=10
|
||||
# 图片模型(Seedream 5.0 Pro,用于信任链真人 AI 化 + 文生图)
|
||||
DOUBAO_IMAGE_MODEL=doubao-seedream-5-0-pro-260628
|
||||
DOUBAO_IMAGE_TIMEOUT=120
|
||||
|
||||
# ==================== 积分/会员系统 (#1895) ====================
|
||||
# 积分系统总开关:默认 false(暂停积分系统)。
|
||||
|
||||
@@ -1187,6 +1187,14 @@ jobs:
|
||||
DOUBAO_MODEL: "${{ secrets.DOUBAO_MODEL }}"
|
||||
DOUBAO_BASE_URL: "${{ secrets.DOUBAO_BASE_URL }}"
|
||||
DOUBAO_VISION_MODEL: "${{ secrets.DOUBAO_VISION_MODEL }}"
|
||||
DOUBAO_VISION_LITE_MODEL: "${{ secrets.DOUBAO_VISION_LITE_MODEL }}"
|
||||
DOUBAO_VISION_USE_LITE: "${{ secrets.DOUBAO_VISION_USE_LITE }}"
|
||||
DOUBAO_IMAGE_MODEL: "${{ secrets.DOUBAO_IMAGE_MODEL }}"
|
||||
DOUBAO_IMAGE_SIZE: "${{ secrets.DOUBAO_IMAGE_SIZE }}"
|
||||
DOUBAO_IMAGE_TIMEOUT: "${{ secrets.DOUBAO_IMAGE_TIMEOUT }}"
|
||||
DOUBAO_FAST_MODEL: "${{ secrets.DOUBAO_FAST_MODEL }}"
|
||||
DOUBAO_TIMEOUT: "${{ secrets.DOUBAO_TIMEOUT }}"
|
||||
DOUBAO_MAX_RETRIES: "${{ secrets.DOUBAO_MAX_RETRIES }}"
|
||||
WECHAT_APP_ID: "${{ secrets.WECHAT_APP_ID }}"
|
||||
WECHAT_APP_SECRET: "${{ secrets.WECHAT_APP_SECRET }}"
|
||||
TIKHUB_API_KEY: "${{ secrets.TIKHUB_API_KEY }}"
|
||||
|
||||
@@ -0,0 +1,31 @@
|
||||
"""viral_video_jobs 增加 pre_trusted_images 列(信任链Seedream预热结果)
|
||||
|
||||
Revision ID: 094_viral_video_pre_trusted
|
||||
Revises: 093_viral_video_pricing_points_float
|
||||
Create Date: 2026-10-04
|
||||
"""
|
||||
|
||||
import sqlalchemy as sa
|
||||
|
||||
from alembic import op
|
||||
|
||||
revision = "094_viral_video_pre_trusted"
|
||||
down_revision = "093"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
conn = op.get_bind()
|
||||
inspector = sa.inspect(conn)
|
||||
cols = {c["name"] for c in inspector.get_columns("viral_video_jobs")}
|
||||
if "pre_trusted_images" not in cols:
|
||||
op.add_column("viral_video_jobs", sa.Column("pre_trusted_images", sa.Text(), nullable=True))
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
conn = op.get_bind()
|
||||
inspector = sa.inspect(conn)
|
||||
cols = {c["name"] for c in inspector.get_columns("viral_video_jobs")}
|
||||
if "pre_trusted_images" in cols:
|
||||
op.drop_column("viral_video_jobs", "pre_trusted_images")
|
||||
@@ -0,0 +1,102 @@
|
||||
"""爆款视频 Prompt 模板配置表(#2040)。
|
||||
|
||||
086 曾预留同名旧表(id varchar / content / variables json),从未被业务使用;
|
||||
本迁移将其替换为 #2040 新结构。
|
||||
|
||||
Revision ID: 095_viral_video_prompt_templates
|
||||
Revises: 094_viral_video_pre_trusted
|
||||
Create Date: 2026-10-04
|
||||
"""
|
||||
|
||||
import sqlalchemy as sa
|
||||
|
||||
from alembic import op
|
||||
|
||||
revision = "095_viral_video_prompt_templates"
|
||||
down_revision = "094_viral_video_pre_trusted"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def _table_exists(conn, name: str) -> bool:
|
||||
return name in sa.inspect(conn).get_table_names()
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
conn = op.get_bind()
|
||||
# 086 预留的旧结构表:先删除(无业务数据、无任何引用)
|
||||
if _table_exists(conn, "viral_video_prompt_templates"):
|
||||
op.drop_table("viral_video_prompt_templates")
|
||||
|
||||
op.create_table(
|
||||
"viral_video_prompt_templates",
|
||||
sa.Column("id", sa.Integer, primary_key=True, autoincrement=True),
|
||||
sa.Column("name", sa.String(128), nullable=False),
|
||||
sa.Column("prompt_type", sa.String(32), nullable=False),
|
||||
sa.Column("version", sa.Integer, nullable=False, server_default="1"),
|
||||
sa.Column("system_prompt", sa.Text, nullable=False),
|
||||
sa.Column("user_prompt_template", sa.Text, nullable=False),
|
||||
sa.Column("example_output", sa.Text, nullable=True),
|
||||
sa.Column("is_active", sa.Boolean, nullable=False, server_default=sa.text("true")),
|
||||
sa.Column(
|
||||
"created_at",
|
||||
sa.DateTime(timezone=True),
|
||||
server_default=sa.func.now(),
|
||||
nullable=False,
|
||||
),
|
||||
sa.Column(
|
||||
"updated_at",
|
||||
sa.DateTime(timezone=True),
|
||||
server_default=sa.func.now(),
|
||||
nullable=False,
|
||||
),
|
||||
)
|
||||
op.create_index(
|
||||
"ix_vvpt_type_active",
|
||||
"viral_video_prompt_templates",
|
||||
["prompt_type", "is_active"],
|
||||
)
|
||||
op.create_index(
|
||||
"uq_vvpt_type_version",
|
||||
"viral_video_prompt_templates",
|
||||
["prompt_type", "version"],
|
||||
unique=True,
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
conn = op.get_bind()
|
||||
if _table_exists(conn, "viral_video_prompt_templates"):
|
||||
op.drop_index("uq_vvpt_type_version", table_name="viral_video_prompt_templates")
|
||||
op.drop_index("ix_vvpt_type_active", table_name="viral_video_prompt_templates")
|
||||
op.drop_table("viral_video_prompt_templates")
|
||||
|
||||
# 恢复 086 的旧预留结构
|
||||
op.create_table(
|
||||
"viral_video_prompt_templates",
|
||||
sa.Column("id", sa.String(36), primary_key=True),
|
||||
sa.Column("prompt_type", sa.String(50), nullable=False, index=True),
|
||||
sa.Column("name", sa.String(200), nullable=False),
|
||||
sa.Column("content", sa.Text, nullable=False, server_default=""),
|
||||
sa.Column("variables", sa.JSON, nullable=False, server_default="[]"),
|
||||
sa.Column("version", sa.Integer, nullable=False, server_default="1"),
|
||||
sa.Column(
|
||||
"is_active",
|
||||
sa.Boolean,
|
||||
nullable=False,
|
||||
server_default=sa.text("true"),
|
||||
index=True,
|
||||
),
|
||||
sa.Column(
|
||||
"created_at",
|
||||
sa.DateTime(timezone=True),
|
||||
nullable=False,
|
||||
server_default=sa.func.now(),
|
||||
),
|
||||
sa.Column(
|
||||
"updated_at",
|
||||
sa.DateTime(timezone=True),
|
||||
nullable=False,
|
||||
server_default=sa.func.now(),
|
||||
),
|
||||
)
|
||||
@@ -145,19 +145,22 @@ def get_rules(
|
||||
def get_packages(
|
||||
current_user: AuthenticatedUser = Depends(get_current_user),
|
||||
):
|
||||
"""查询可购买的积分包列表。"""
|
||||
packages = []
|
||||
for code, pkg in POINTS_PACKAGES.items():
|
||||
unit_price = f"¥{pkg['price_cents'] / 100 / pkg['points']:.3f}/积分"
|
||||
packages.append(
|
||||
PointsPackageItem(
|
||||
code=code,
|
||||
name=pkg["name"],
|
||||
points=pkg["points"],
|
||||
price_cents=pkg["price_cents"],
|
||||
unit_price=unit_price,
|
||||
)
|
||||
"""查询可购买的积分包列表(读管理后台 credit_packages 表真实数据)。
|
||||
|
||||
仅返回 is_active=true;后台改价/启停后最多 30 秒生效。
|
||||
"""
|
||||
from packages.application.catalog.admin_catalog import get_points_packages
|
||||
|
||||
packages = [
|
||||
PointsPackageItem(
|
||||
code=row["code"],
|
||||
name=row["name"],
|
||||
points=row["points"],
|
||||
price_cents=row["price_cents"],
|
||||
unit_price=row["unit_price"],
|
||||
)
|
||||
for row in get_points_packages()
|
||||
]
|
||||
mt = _member_type(current_user)
|
||||
discount = MEMBER_DISCOUNT.get(mt) if mt else None
|
||||
return PointsPackagesResponse(packages=packages, user_discount=discount)
|
||||
|
||||
@@ -86,33 +86,13 @@ async def get_current_subscription(
|
||||
def list_membership_plans(
|
||||
current_user: AuthenticatedUser = Depends(get_current_user),
|
||||
) -> dict[str, list[dict[str, Any]]]:
|
||||
"""查询所有会员档位(供前端会员购买页展示)。
|
||||
"""查询可购买的会员套餐(读管理后台 plans 表真实数据)。
|
||||
|
||||
返回 points 积分体系下的会员档位(月卡/季卡/年卡),含价格、时长、积分折扣等信息。
|
||||
仅返回 is_enabled=true 的套餐;后台启停/改价后最多 30 秒生效。
|
||||
"""
|
||||
from packages.domain.points_rules import MEMBER_DISCOUNT, MEMBERSHIP_PRICES
|
||||
from packages.application.catalog.admin_catalog import get_membership_plans
|
||||
|
||||
plans: list[dict[str, Any]] = []
|
||||
for plan_id, info in MEMBERSHIP_PRICES.items():
|
||||
days = info["duration_days"]
|
||||
monthly_cents = round(info["price_cents"] * 30 / days)
|
||||
features: dict[str, Any] = {"max_resolution": "1080p"}
|
||||
if plan_id == MembershipType.MONTHLY:
|
||||
features.update({"free_clips_daily": 2})
|
||||
elif plan_id == MembershipType.QUARTERLY:
|
||||
features.update({"free_clips_daily": 5})
|
||||
elif plan_id == MembershipType.YEARLY:
|
||||
features.update({"free_clips_daily": "unlimited"})
|
||||
plans.append({
|
||||
"plan_id": plan_id,
|
||||
"name": info["name"],
|
||||
"price_cents": info["price_cents"],
|
||||
"monthly_price_cents": monthly_cents,
|
||||
"duration_days": days,
|
||||
"points_discount": MEMBER_DISCOUNT.get(plan_id, 1.0),
|
||||
"features": features,
|
||||
})
|
||||
return {"plans": plans}
|
||||
return {"plans": get_membership_plans()}
|
||||
|
||||
|
||||
@router.get("/billing-records", response_model=list[BillingRecord])
|
||||
|
||||
@@ -146,6 +146,7 @@ def _to_response(job) -> ViralVideoJobResponse:
|
||||
video_model=getattr(job, "video_model", "") or "",
|
||||
intent_result=job.intent_result,
|
||||
result_video_url=job.result_video_url,
|
||||
pre_trusted_images=getattr(job, "pre_trusted_images", None) or None,
|
||||
video_resolution=getattr(job, "video_resolution", "720p") or "720p",
|
||||
credits_prepaid=float(getattr(job, "credits_prepaid", 0) or 0),
|
||||
credits_cost=float(getattr(job, "credits_cost", 0) or 0),
|
||||
|
||||
@@ -222,6 +222,7 @@ class ViralVideoJobResponse(BaseModel):
|
||||
video_model: str = ""
|
||||
intent_result: dict | None = None
|
||||
result_video_url: str = ""
|
||||
pre_trusted_images: list[str] | None = None
|
||||
video_resolution: str = "720p"
|
||||
credits_prepaid: float = 0.0
|
||||
credits_cost: float = 0.0
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -241,19 +241,35 @@ DOUBAO_API_KEY=${DOUBAO_API_KEY}
|
||||
|
||||
# 模型 Endpoint ID(在 ARK 控制台创建推理接入点后获得)
|
||||
DOUBAO_MODEL=${DOUBAO_MODEL}
|
||||
DOUBAO_FAST_MODEL=${DOUBAO_FAST_MODEL}
|
||||
|
||||
# API Base URL
|
||||
DOUBAO_BASE_URL=${DOUBAO_BASE_URL}
|
||||
|
||||
# 请求超时(秒)
|
||||
DOUBAO_TIMEOUT=60
|
||||
DOUBAO_TIMEOUT=${DOUBAO_TIMEOUT}
|
||||
|
||||
# 最大重试次数
|
||||
DOUBAO_MAX_RETRIES=2
|
||||
DOUBAO_MAX_RETRIES=${DOUBAO_MAX_RETRIES}
|
||||
|
||||
# 视觉模型 Endpoint ID(支持图片/视频理解的模型)
|
||||
# 视觉模型(支持图片/视频理解的模型,model name 格式)
|
||||
DOUBAO_VISION_MODEL=${DOUBAO_VISION_MODEL}
|
||||
|
||||
# 快速视觉模型(viral-video 图片分析 lite 路径)
|
||||
DOUBAO_VISION_LITE_MODEL=${DOUBAO_VISION_LITE_MODEL}
|
||||
|
||||
# 是否启用 lite 视觉路径(true/false)
|
||||
DOUBAO_VISION_USE_LITE=${DOUBAO_VISION_USE_LITE}
|
||||
|
||||
# 信任链文生图模型(Seedream)
|
||||
DOUBAO_IMAGE_MODEL=${DOUBAO_IMAGE_MODEL}
|
||||
|
||||
# 文生图尺寸
|
||||
DOUBAO_IMAGE_SIZE=${DOUBAO_IMAGE_SIZE}
|
||||
|
||||
# 文生图超时(秒)
|
||||
DOUBAO_IMAGE_TIMEOUT=${DOUBAO_IMAGE_TIMEOUT}
|
||||
|
||||
|
||||
# ==================== 微信开放平台 OAuth(网页扫码登录)====================
|
||||
# 回调域名:xiaoxiajianji.com(微信开放平台已配置)
|
||||
|
||||
@@ -928,6 +928,7 @@ class ViralVideoJobModel(Base):
|
||||
id = Column(String(36), primary_key=True)
|
||||
user_id = Column(String(36), nullable=False, index=True)
|
||||
images = Column(JSON, nullable=False, default=list) # 产品图片 URL 列表
|
||||
pre_trusted_images = Column(JSON, nullable=True) # #2172 信任链预热结果(Seedream AI 化 URL 列表)
|
||||
industry = Column(String(100), nullable=False, default="")
|
||||
target_customer = Column(String(500), nullable=False, default="")
|
||||
persona_id = Column(String(36), nullable=False, default="")
|
||||
@@ -990,16 +991,17 @@ class ViralVideoStyleTemplateModel(Base):
|
||||
|
||||
|
||||
class ViralVideoPromptTemplateModel(Base):
|
||||
"""爆款视频 Prompt 模板表(由 #2040 seed)"""
|
||||
"""爆款视频 Prompt 模板表(#2040:纯文本 XML 标签模板,运营可直接编辑)"""
|
||||
|
||||
__tablename__ = "viral_video_prompt_templates"
|
||||
|
||||
id = Column(String(36), primary_key=True)
|
||||
prompt_type = Column(String(50), nullable=False, index=True)
|
||||
name = Column(String(200), nullable=False)
|
||||
content = Column(Text, nullable=False, default="")
|
||||
variables = Column(JSON, nullable=False, default=list)
|
||||
id = Column(Integer, primary_key=True, autoincrement=True)
|
||||
name = Column(String(128), nullable=False)
|
||||
prompt_type = Column(String(32), nullable=False)
|
||||
version = Column(Integer, nullable=False, default=1)
|
||||
is_active = Column(Boolean, nullable=False, default=True, index=True)
|
||||
system_prompt = Column(Text, nullable=False)
|
||||
user_prompt_template = Column(Text, nullable=False)
|
||||
example_output = Column(Text, nullable=True)
|
||||
is_active = Column(Boolean, nullable=False, default=True)
|
||||
created_at = Column(DateTime(timezone=True), nullable=False, default=lambda: datetime.now(UTC))
|
||||
updated_at = Column(DateTime(timezone=True), nullable=False, default=lambda: datetime.now(UTC))
|
||||
|
||||
@@ -6,18 +6,33 @@ from sqlalchemy.orm import Session
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.models import (
|
||||
ViralVideoJobModel,
|
||||
ViralVideoPromptTemplateModel,
|
||||
ViralVideoStyleTemplateModel,
|
||||
)
|
||||
from packages.domain.viral_video import ViralVideoJob, ViralVideoStatus
|
||||
|
||||
|
||||
def _to_domain(model: ViralVideoJobModel) -> ViralVideoJob:
|
||||
"""ORM → 领域实体。"""
|
||||
"""ORM → 领域实体。pre_trusted_images 兼容脏数据:双序列化字符串/字符数组/list[str]。"""
|
||||
import json as _pti_json
|
||||
|
||||
_raw_pti = getattr(model, "pre_trusted_images", None)
|
||||
_pti: list[str] | None = None
|
||||
if _raw_pti is not None:
|
||||
if isinstance(_raw_pti, str):
|
||||
try:
|
||||
_p = _pti_json.loads(_raw_pti)
|
||||
if isinstance(_p, list):
|
||||
_pti = [u for u in _p if isinstance(u, str) and u] or None
|
||||
except Exception:
|
||||
_pti = None
|
||||
elif isinstance(_raw_pti, list):
|
||||
_f = [u for u in _raw_pti if isinstance(u, str) and len(u) > 5]
|
||||
_pti = _f if _f else None
|
||||
return ViralVideoJob(
|
||||
id=model.id,
|
||||
user_id=model.user_id,
|
||||
images=list(model.images or []),
|
||||
pre_trusted_images=_pti,
|
||||
industry=model.industry or "",
|
||||
target_customer=model.target_customer or "",
|
||||
persona_id=model.persona_id or "",
|
||||
@@ -70,6 +85,7 @@ class SQLAlchemyViralVideoJobRepository:
|
||||
id=job.id,
|
||||
user_id=job.user_id,
|
||||
images=job.images,
|
||||
pre_trusted_images=job.pre_trusted_images,
|
||||
industry=job.industry,
|
||||
target_customer=job.target_customer,
|
||||
persona_id=job.persona_id,
|
||||
@@ -126,6 +142,7 @@ class SQLAlchemyViralVideoJobRepository:
|
||||
model.storyboard = job.storyboard
|
||||
model.generated_copy_text = job.generated_copy_text or ""
|
||||
model.copy_result = job.copy_result
|
||||
model.pre_trusted_images = job.pre_trusted_images
|
||||
model.result_video_url = job.result_video_url
|
||||
model.video_resolution = getattr(job, "video_resolution", "720p") or "720p"
|
||||
model.credits_prepaid = float(getattr(job, "credits_prepaid", 0) or 0)
|
||||
@@ -227,31 +244,3 @@ class SQLAlchemyViralVideoStyleTemplateRepository:
|
||||
"style_config": dict(model.style_config) if model.style_config else {},
|
||||
"is_system": model.is_system,
|
||||
}
|
||||
|
||||
|
||||
class SQLAlchemyViralVideoPromptTemplateRepository:
|
||||
"""Prompt 模板仓储(由 #2040 seed,这里只读取)。"""
|
||||
|
||||
def __init__(self, session: Session):
|
||||
self.session = session
|
||||
|
||||
def get_active_by_type(self, prompt_type: str) -> dict | None:
|
||||
model = (
|
||||
self.session.query(ViralVideoPromptTemplateModel)
|
||||
.filter(
|
||||
ViralVideoPromptTemplateModel.prompt_type == prompt_type,
|
||||
ViralVideoPromptTemplateModel.is_active.is_(True),
|
||||
)
|
||||
.order_by(ViralVideoPromptTemplateModel.version.desc())
|
||||
.first()
|
||||
)
|
||||
if model is None:
|
||||
return None
|
||||
return {
|
||||
"id": model.id,
|
||||
"prompt_type": model.prompt_type,
|
||||
"name": model.name,
|
||||
"content": model.content,
|
||||
"variables": list(model.variables or []),
|
||||
"version": model.version,
|
||||
}
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
"""应用层:对外展示目录(套餐/积分包)。"""
|
||||
@@ -0,0 +1,152 @@
|
||||
"""读取管理后台配置的会员套餐 / 积分充值包(共享库真实数据)。
|
||||
|
||||
替代旧的硬编码 MEMBERSHIP_PRICES / POINTS_PACKAGES。
|
||||
短 TTL 缓存(30 秒),后台改价/启停后用户端最多 30 秒可见。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import threading
|
||||
import time
|
||||
from typing import Any
|
||||
|
||||
_CACHE_TTL = 30.0
|
||||
_lock = threading.Lock()
|
||||
_cache: dict[str, tuple[float, Any]] = {}
|
||||
|
||||
_QUOTA_LABELS = {
|
||||
"4k": "4K 超清分辨率",
|
||||
"batch_render": "批量渲染",
|
||||
"priority_queue": "优先处理队列",
|
||||
"ai_matting": "AI 智能抠像",
|
||||
"remove_watermark": "去水印",
|
||||
}
|
||||
|
||||
|
||||
def _cached(key: str, loader):
|
||||
now = time.time()
|
||||
hit = _cache.get(key)
|
||||
if hit and now - hit[0] < _CACHE_TTL:
|
||||
return hit[1]
|
||||
with _lock:
|
||||
hit = _cache.get(key)
|
||||
if hit and time.time() - hit[0] < _CACHE_TTL:
|
||||
return hit[1]
|
||||
value = loader()
|
||||
_cache[key] = (time.time(), value)
|
||||
return value
|
||||
|
||||
|
||||
def _quota_features(quotas: dict[str, Any] | None) -> dict[str, Any]:
|
||||
quotas = quotas or {}
|
||||
features: dict[str, Any] = {}
|
||||
for k, v in quotas.items():
|
||||
if k == "credits_per_month":
|
||||
features["credits_per_month"] = v
|
||||
elif k in _QUOTA_LABELS:
|
||||
features[_QUOTA_LABELS[k]] = v
|
||||
else:
|
||||
features[k] = v
|
||||
return features
|
||||
|
||||
|
||||
def get_membership_plans() -> list[dict[str, Any]]:
|
||||
"""读取 is_enabled=true 的套餐,按年/月周期展开为用户端档位。"""
|
||||
|
||||
def _load() -> list[dict[str, Any]]:
|
||||
from sqlalchemy import text
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.session import SessionLocal
|
||||
|
||||
if SessionLocal is None:
|
||||
return []
|
||||
|
||||
session = SessionLocal()
|
||||
try:
|
||||
rows = session.execute(text("""
|
||||
SELECT plan_key, name, description, monthly_price, yearly_price,
|
||||
quotas, display_order
|
||||
FROM plans
|
||||
WHERE is_enabled = TRUE
|
||||
ORDER BY display_order NULLS LAST, created_at
|
||||
""")).fetchall()
|
||||
finally:
|
||||
session.close()
|
||||
|
||||
plans: list[dict[str, Any]] = []
|
||||
for r in rows:
|
||||
base_features = _quota_features(r.quotas if isinstance(r.quotas, dict) else None)
|
||||
if r.yearly_price and float(r.yearly_price) > 0:
|
||||
plans.append(
|
||||
{
|
||||
"plan_id": r.plan_key,
|
||||
"billing_cycle": "yearly",
|
||||
"name": r.name,
|
||||
"description": r.description,
|
||||
"price_cents": int(round(float(r.yearly_price) * 100)),
|
||||
"monthly_price_cents": int(round(float(r.yearly_price) * 100 / 12)),
|
||||
"duration_days": 365,
|
||||
"features": dict(base_features),
|
||||
}
|
||||
)
|
||||
if r.monthly_price and float(r.monthly_price) > 0:
|
||||
plans.append(
|
||||
{
|
||||
"plan_id": r.plan_key,
|
||||
"billing_cycle": "monthly",
|
||||
"name": r.name,
|
||||
"description": r.description,
|
||||
"price_cents": int(round(float(r.monthly_price) * 100)),
|
||||
"monthly_price_cents": int(round(float(r.monthly_price) * 100)),
|
||||
"duration_days": 30,
|
||||
"features": dict(base_features),
|
||||
}
|
||||
)
|
||||
return plans
|
||||
|
||||
return _cached("membership_plans", _load)
|
||||
|
||||
|
||||
def get_points_packages() -> list[dict[str, Any]]:
|
||||
"""读取 is_active=true 的积分充值包。"""
|
||||
|
||||
def _load() -> list[dict[str, Any]]:
|
||||
from sqlalchemy import text
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.session import SessionLocal
|
||||
|
||||
if SessionLocal is None:
|
||||
return []
|
||||
|
||||
session = SessionLocal()
|
||||
try:
|
||||
rows = session.execute(text("""
|
||||
SELECT package_key, name, price, credits, bonus_credits,
|
||||
is_recommended, description, sort_order
|
||||
FROM credit_packages
|
||||
WHERE is_active = TRUE
|
||||
ORDER BY sort_order NULLS LAST, price
|
||||
""")).fetchall()
|
||||
finally:
|
||||
session.close()
|
||||
|
||||
packages: list[dict[str, Any]] = []
|
||||
for r in rows:
|
||||
total_points = int(r.credits or 0) + int(r.bonus_credits or 0)
|
||||
price_cents = int(round(float(r.price) * 100))
|
||||
unit = (price_cents / 100 / total_points) if total_points else 0
|
||||
packages.append(
|
||||
{
|
||||
"code": r.package_key,
|
||||
"name": r.name,
|
||||
"points": total_points,
|
||||
"bonus_credits": int(r.bonus_credits or 0),
|
||||
"price_cents": price_cents,
|
||||
"unit_price": f"¥{unit:.3f}/积分",
|
||||
"is_recommended": bool(r.is_recommended),
|
||||
"description": r.description,
|
||||
}
|
||||
)
|
||||
return packages
|
||||
|
||||
return _cached("points_packages", _load)
|
||||
@@ -0,0 +1 @@
|
||||
"""应用层:爆款视频 Prompt 模板系统(#2040)。"""
|
||||
@@ -0,0 +1,427 @@
|
||||
"""爆款视频 5 步编排:图片分析 → 意图解析 → 文案融合 → 分镜 → 审核重写。
|
||||
|
||||
所有 LLM 调用走 DoubaoClient,单测通过 client 参数注入 mock,不真调 API。
|
||||
任何一步解析失败都走规则 fallback,不抛异常阻断。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from typing import Optional
|
||||
|
||||
from packages.application.viral_video import xml_parser as xp
|
||||
from packages.application.viral_video.prompt_loader import (
|
||||
PromptTemplate,
|
||||
get_template,
|
||||
render_system_prompt,
|
||||
render_user_prompt,
|
||||
)
|
||||
from packages.application.viral_video.prompts import (
|
||||
FUSION_INSTRUCTIONS,
|
||||
GLOBAL_CONSTRAINTS,
|
||||
NEGATIVE_RULES,
|
||||
)
|
||||
from packages.application.viral_video.reviewer import Reviewer
|
||||
from packages.application.viral_video.schemas import (
|
||||
BodyPoint,
|
||||
Clip,
|
||||
ColorItem,
|
||||
CoreMessage,
|
||||
FusionResult,
|
||||
ImageAnalysis,
|
||||
IntentResult,
|
||||
KenBurns,
|
||||
PersonalBrand,
|
||||
ProductItem,
|
||||
ReviewResult,
|
||||
ScriptSegment,
|
||||
Storyboard,
|
||||
TextItem,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class CopyGenerator:
|
||||
"""5 步 Prompt 编排器。"""
|
||||
|
||||
def __init__(self, client=None, reviewer: Optional[Reviewer] = None):
|
||||
if client is None:
|
||||
from packages.shared.ai_client import get_doubao_client
|
||||
|
||||
client = get_doubao_client()
|
||||
self.client = client
|
||||
self.reviewer = reviewer or Reviewer(client)
|
||||
|
||||
# ── 底层调用 ────────────────────────────────────────────────────────
|
||||
def _chat(self, template: PromptTemplate, system_kwargs: dict | None, **user_kwargs) -> str:
|
||||
system = render_system_prompt(template, **(system_kwargs or {}))
|
||||
user = render_user_prompt(template, **user_kwargs)
|
||||
result = self.client.chat_completion(
|
||||
[
|
||||
{"role": "system", "content": system},
|
||||
{"role": "user", "content": user},
|
||||
],
|
||||
temperature=0.7,
|
||||
max_tokens=2048,
|
||||
)
|
||||
return result or ""
|
||||
|
||||
# ── 步骤1:图片多模态分析 ───────────────────────────────────────────
|
||||
def analyze_images(self, images: list[str], industry: str = "") -> ImageAnalysis:
|
||||
template = get_template("image_analysis")
|
||||
image_urls = "\n".join(f"第{i + 1}张:{url}" for i, url in enumerate(images))
|
||||
system = render_system_prompt(template)
|
||||
user = render_user_prompt(template, image_count=len(images), industry=industry or "通用", image_urls=image_urls)
|
||||
raw = self.client.vision_completion(
|
||||
[
|
||||
{"role": "system", "content": system},
|
||||
{"role": "user", "content": user},
|
||||
],
|
||||
images=images,
|
||||
max_tokens=2048,
|
||||
temperature=0.3,
|
||||
)
|
||||
analysis = self._parse_image_analysis(raw or "")
|
||||
if not analysis.products and not analysis.key_selling_points:
|
||||
logger.warning("图片分析标签解析失败,走规则 fallback")
|
||||
return self._fallback_image_analysis(images, raw or "")
|
||||
return analysis
|
||||
|
||||
def _parse_image_analysis(self, raw: str) -> ImageAnalysis:
|
||||
products = [
|
||||
ProductItem(
|
||||
name=n["attrs"].get("name", "无法判断"),
|
||||
features=n["attrs"].get("features", "无法判断"),
|
||||
position=n["attrs"].get("position", "secondary"),
|
||||
image_index=xp.attr_int(n["attrs"].get("image_index"), 0),
|
||||
)
|
||||
for n in xp.find_all(raw, "product")
|
||||
]
|
||||
colors = [
|
||||
ColorItem(
|
||||
hex=c["attrs"].get("hex", "#000000"),
|
||||
name=c["attrs"].get("name", "无法判断"),
|
||||
coverage=xp.attr_float(c["attrs"].get("coverage"), 0.0),
|
||||
)
|
||||
for c in xp.find_all(raw, "color")
|
||||
]
|
||||
people = xp.find_first(raw, "people")
|
||||
visible_text = [
|
||||
TextItem(text=t["attrs"].get("text", ""), position=t["attrs"].get("position", ""))
|
||||
for t in xp.find_all(raw, "text_item")
|
||||
]
|
||||
quality_node = xp.find_first(raw, "quality")
|
||||
selling_points = [n["text"] or n["attrs"].get("text", "") for n in xp.find_all(raw, "point")]
|
||||
return ImageAnalysis(
|
||||
products=products,
|
||||
colors=colors,
|
||||
has_person=xp.attr_bool(people["attrs"].get("has_person")) if people else False,
|
||||
person_count=xp.attr_int(people["attrs"].get("count"), 0) if people else 0,
|
||||
people=people["attrs"] if people else {},
|
||||
mood=xp.text_of(raw, "mood"),
|
||||
visible_text=visible_text,
|
||||
scene=xp.text_of(raw, "scene"),
|
||||
quality=quality_node["attrs"] if quality_node else {},
|
||||
key_selling_points=[p for p in selling_points if p],
|
||||
raw=raw,
|
||||
)
|
||||
|
||||
def _fallback_image_analysis(self, images: list[str], raw: str) -> ImageAnalysis:
|
||||
return ImageAnalysis(
|
||||
products=[ProductItem(name="无法判断(视觉分析不可用)", image_index=0)],
|
||||
scene="无法判断",
|
||||
raw=raw,
|
||||
)
|
||||
|
||||
# ── 步骤2:意图解析 ─────────────────────────────────────────────────
|
||||
def parse_intent(self, user_copy_text: str, image_analysis: ImageAnalysis, industry: str = "") -> IntentResult:
|
||||
template = get_template("intent_parsing")
|
||||
raw = self._chat(
|
||||
template,
|
||||
None,
|
||||
user_copy_text=user_copy_text or "(用户没有提供文案)",
|
||||
industry=industry or "通用",
|
||||
image_analysis=self._image_brief(image_analysis),
|
||||
)
|
||||
intent = self._parse_intent(raw)
|
||||
if not intent.intent_summary and not intent.core_messages:
|
||||
logger.warning("意图解析标签解析失败,走规则 fallback")
|
||||
return self._fallback_intent(user_copy_text, raw)
|
||||
return intent
|
||||
|
||||
def _parse_intent(self, raw: str) -> IntentResult:
|
||||
messages = [
|
||||
CoreMessage(
|
||||
text=n["text"],
|
||||
must_keep=xp.attr_bool(n["attrs"].get("must_keep"), default=False),
|
||||
confidence=xp.attr_float(n["attrs"].get("confidence"), 0.0),
|
||||
)
|
||||
for n in xp.find_all(raw, "message")
|
||||
if n["text"]
|
||||
]
|
||||
brands = [
|
||||
PersonalBrand(text=n["text"], category=n["attrs"].get("category", "brand"))
|
||||
for n in xp.find_all(raw, "brand")
|
||||
if n["text"]
|
||||
]
|
||||
missing = [n["text"] for n in xp.find_all(raw, "info") if n["text"]]
|
||||
return IntentResult(
|
||||
intent_summary=xp.text_of(raw, "intent_summary"),
|
||||
core_messages=messages,
|
||||
personal_brands=brands,
|
||||
emotion_tone=xp.text_of(raw, "emotion_tone"),
|
||||
missing_info=missing,
|
||||
raw=raw,
|
||||
)
|
||||
|
||||
def _fallback_intent(self, user_copy_text: str, raw: str) -> IntentResult:
|
||||
text = (user_copy_text or "").strip()
|
||||
messages = [CoreMessage(text=text[:80], must_keep=True, confidence=1.0)] if text else []
|
||||
return IntentResult(
|
||||
intent_summary=text[:30] or "未提供文案,按产品图片自由创作",
|
||||
core_messages=messages,
|
||||
personal_brands=[],
|
||||
raw=raw,
|
||||
)
|
||||
|
||||
# ── 步骤3:文案融合生成(三档)──────────────────────────────────────
|
||||
def fuse(
|
||||
self,
|
||||
fusion_level: str,
|
||||
image_analysis: ImageAnalysis,
|
||||
intent: IntentResult,
|
||||
industry: str = "",
|
||||
target_customer: str = "",
|
||||
marketing_purpose: str = "",
|
||||
duration: int = 15,
|
||||
) -> FusionResult:
|
||||
template = get_template("copy_fusion")
|
||||
system_kwargs = {
|
||||
"fusion_instruction": FUSION_INSTRUCTIONS.get(fusion_level, FUSION_INSTRUCTIONS["ai_polish"]),
|
||||
"global_constraints": GLOBAL_CONSTRAINTS,
|
||||
"negative_rules": NEGATIVE_RULES,
|
||||
}
|
||||
raw = self._chat(
|
||||
template,
|
||||
system_kwargs,
|
||||
industry=industry or "通用",
|
||||
target_customer=target_customer or "通用消费者",
|
||||
marketing_purpose=marketing_purpose or "产品种草",
|
||||
duration=duration,
|
||||
image_analysis=self._image_brief(image_analysis),
|
||||
intent_result=self._intent_brief(intent),
|
||||
)
|
||||
result = self._parse_fusion(raw)
|
||||
if not result.title and not result.script_segments:
|
||||
logger.warning("文案融合标签解析失败(fusion=%s),走规则 fallback", fusion_level)
|
||||
return self._fallback_fusion(fusion_level, image_analysis, intent, duration, raw)
|
||||
return result
|
||||
|
||||
def _parse_fusion(self, raw: str) -> FusionResult:
|
||||
body_points = [
|
||||
BodyPoint(
|
||||
text=n["text"],
|
||||
elaboration=n["attrs"].get("elaboration", ""),
|
||||
image_index=xp.attr_int(n["attrs"].get("image_index"), 0),
|
||||
)
|
||||
for n in xp.find_all(raw, "point")
|
||||
if n["text"]
|
||||
]
|
||||
segments = [
|
||||
ScriptSegment(
|
||||
text=n["text"],
|
||||
duration_sec=xp.attr_float(n["attrs"].get("duration_sec"), 0.0),
|
||||
image_index=xp.attr_int(n["attrs"].get("image_index"), 0),
|
||||
)
|
||||
for n in xp.find_all(raw, "segment")
|
||||
if n["text"]
|
||||
]
|
||||
return FusionResult(
|
||||
title=xp.text_of(raw, "title"),
|
||||
hook=xp.text_of(raw, "hook"),
|
||||
body_points=body_points,
|
||||
cta=xp.text_of(raw, "cta"),
|
||||
script_segments=segments,
|
||||
word_count=xp.attr_int(xp.text_of(raw, "word_count"), 0),
|
||||
estimated_duration=xp.attr_int(xp.text_of(raw, "estimated_duration"), 0),
|
||||
raw=raw,
|
||||
)
|
||||
|
||||
def _fallback_fusion(
|
||||
self,
|
||||
fusion_level: str,
|
||||
image_analysis: ImageAnalysis,
|
||||
intent: IntentResult,
|
||||
duration: int,
|
||||
raw: str,
|
||||
) -> FusionResult:
|
||||
product_name = image_analysis.products[0].name if image_analysis.products else "这款产品"
|
||||
selling = image_analysis.key_selling_points[:2]
|
||||
if fusion_level == "ai_full":
|
||||
title = f"{product_name},很多人用完都回购了"
|
||||
hook = f"这个{product_name},我想认真说说"
|
||||
body = selling or ["图片可见的产品卖点"]
|
||||
cta = "感兴趣的可以了解一下"
|
||||
elif fusion_level == "user_primary":
|
||||
user_text = intent.intent_summary or product_name
|
||||
title = user_text[:20]
|
||||
hook = user_text[:15]
|
||||
body = [m.text for m in intent.core_messages] or [user_text]
|
||||
cta = "想了解的可以看看"
|
||||
else:
|
||||
title = intent.intent_summary[:20] or product_name
|
||||
hook = intent.core_messages[0].text[:15] if intent.core_messages else product_name
|
||||
body = [m.text for m in intent.core_messages] or selling or [product_name]
|
||||
cta = "有需要的可以了解一下"
|
||||
|
||||
brand_texts = [b.text for b in intent.personal_brands]
|
||||
points = [BodyPoint(text=b) for b in body]
|
||||
lines = [hook] + body + brand_texts[:2] + [cta]
|
||||
joined = ",".join(lines)
|
||||
per = max(3, duration // max(1, len(lines)))
|
||||
segments = [ScriptSegment(text=line, duration_sec=per, image_index=0) for line in lines]
|
||||
return FusionResult(
|
||||
title=title,
|
||||
hook=hook,
|
||||
body_points=points,
|
||||
cta=cta,
|
||||
script_segments=segments,
|
||||
word_count=len(joined),
|
||||
estimated_duration=duration,
|
||||
raw=raw,
|
||||
)
|
||||
|
||||
# ── 步骤4:编导级分镜 ───────────────────────────────────────────────
|
||||
def storyboard(
|
||||
self, fusion: FusionResult, image_analysis: ImageAnalysis, images: list[str], duration: int
|
||||
) -> Storyboard:
|
||||
template = get_template("storyboard")
|
||||
raw = self._chat(
|
||||
template,
|
||||
None,
|
||||
duration=duration,
|
||||
image_count=len(images),
|
||||
fusion_result=self._fusion_brief(fusion),
|
||||
image_analysis=self._image_brief(image_analysis),
|
||||
)
|
||||
board = self._parse_storyboard(raw)
|
||||
if not board.clips:
|
||||
logger.warning("分镜标签解析失败,走规则 fallback")
|
||||
return self._fallback_storyboard(fusion, duration, raw)
|
||||
return board
|
||||
|
||||
def _parse_storyboard(self, raw: str) -> Storyboard:
|
||||
clips: list[Clip] = []
|
||||
for node in xp.find_all(raw, "clip"):
|
||||
attrs = node["attrs"]
|
||||
body = node["text"]
|
||||
kb = xp.find_first(node["text"] and f"<root>{node['text']}</root>", "ken_burns")
|
||||
clips.append(
|
||||
Clip(
|
||||
image_index=xp.attr_int(attrs.get("image_index"), 0),
|
||||
transition=attrs.get("transition", "cut"),
|
||||
zoom=(None if attrs.get("zoom") in (None, "null", "None", "") else attrs.get("zoom")),
|
||||
duration_sec=xp.attr_float(attrs.get("duration_sec"), 0.0),
|
||||
bgm_note=attrs.get("bgm_note", ""),
|
||||
voice_text=xp.text_of(body and f"<root>{body}</root>", "voice_text"),
|
||||
subtitle_text=xp.text_of(body and f"<root>{body}</root>", "subtitle_text"),
|
||||
ken_burns=KenBurns(
|
||||
start=kb["attrs"].get("start", "0,0") if kb else "0,0",
|
||||
end=kb["attrs"].get("end", "0,0") if kb else "0,0",
|
||||
ease=kb["attrs"].get("ease", "linear") if kb else "linear",
|
||||
),
|
||||
)
|
||||
)
|
||||
return Storyboard(clips=clips, raw=raw)
|
||||
|
||||
def _fallback_storyboard(self, fusion: FusionResult, duration: int, raw: str) -> Storyboard:
|
||||
segments = fusion.script_segments or [ScriptSegment(text=fusion.hook or fusion.title, duration_sec=duration)]
|
||||
total = sum(s.duration_sec for s in segments) or duration
|
||||
clips = [
|
||||
Clip(
|
||||
image_index=min(s.image_index, 0),
|
||||
transition="cut",
|
||||
duration_sec=max(2.0, s.duration_sec * duration / total if total else duration / len(segments)),
|
||||
voice_text=s.text,
|
||||
subtitle_text=s.text[:20],
|
||||
)
|
||||
for s in segments
|
||||
]
|
||||
return Storyboard(clips=clips, raw=raw)
|
||||
|
||||
# ── 步骤5:审核(不通过自动重写1次)─────────────────────────────────
|
||||
def review_and_rewrite(
|
||||
self, fusion: FusionResult, intent: IntentResult, fusion_level: str
|
||||
) -> tuple[FusionResult, ReviewResult, int]:
|
||||
"""返回最终文案、最后一次审核结果、重写次数(0或1)。"""
|
||||
review = self.reviewer.review(fusion, intent, fusion_level)
|
||||
if review.passed:
|
||||
return fusion, review, 0
|
||||
|
||||
logger.info("文案审核不通过,自动重写 1 次:%s", [i.text for i in review.issues])
|
||||
rewritten = self.reviewer.rewrite(fusion, review, intent, fusion_level)
|
||||
second = self.reviewer.review(rewritten, intent, fusion_level)
|
||||
if second.passed:
|
||||
return rewritten, second, 1
|
||||
# 二次仍不通过:带上重写结果和问题返回,由上游决定是否交给前端
|
||||
return rewritten, second, 1
|
||||
|
||||
# ── 全流程编排 ──────────────────────────────────────────────────────
|
||||
def generate(
|
||||
self,
|
||||
images: list[str],
|
||||
*,
|
||||
industry: str = "",
|
||||
target_customer: str = "",
|
||||
marketing_purpose: str = "",
|
||||
duration: int = 15,
|
||||
user_copy_text: str = "",
|
||||
fusion_level: str = "ai_polish",
|
||||
) -> dict:
|
||||
image_analysis = self.analyze_images(images, industry)
|
||||
intent = self.parse_intent(user_copy_text, image_analysis, industry)
|
||||
fusion = self.fuse(
|
||||
fusion_level,
|
||||
image_analysis,
|
||||
intent,
|
||||
industry=industry,
|
||||
target_customer=target_customer,
|
||||
marketing_purpose=marketing_purpose,
|
||||
duration=duration,
|
||||
)
|
||||
fusion, review, rewrites = self.review_and_rewrite(fusion, intent, fusion_level)
|
||||
board = self.storyboard(fusion, image_analysis, images, duration)
|
||||
return {
|
||||
"image_analysis": image_analysis,
|
||||
"intent_result": intent,
|
||||
"fusion_result": fusion,
|
||||
"review_result": review,
|
||||
"storyboard": board,
|
||||
"rewrite_count": rewrites,
|
||||
}
|
||||
|
||||
# ── 简报工具 ────────────────────────────────────────────────────────
|
||||
@staticmethod
|
||||
def _image_brief(a) -> str:
|
||||
if a is None:
|
||||
return "无图片分析信息"
|
||||
lines = [f"产品:{p.name}({p.features})" for p in a.products]
|
||||
lines += [f"卖点:{s}" for s in a.key_selling_points]
|
||||
lines.append(f"场景:{a.scene}")
|
||||
return "\n".join(lines) or "无图片分析信息"
|
||||
|
||||
@staticmethod
|
||||
def _intent_brief(i: IntentResult) -> str:
|
||||
lines = [f"意图:{i.intent_summary}"]
|
||||
lines += [f"核心信息[must_keep={m.must_keep}]:{m.text}" for m in i.core_messages]
|
||||
lines += [f"事实({b.category}):{b.text}" for b in i.personal_brands]
|
||||
return "\n".join(lines)
|
||||
|
||||
@staticmethod
|
||||
def _fusion_brief(f: FusionResult) -> str:
|
||||
lines = [f"标题:{f.title}", f"钩子:{f.hook}"]
|
||||
lines += [f"要点:{p.text}" for p in f.body_points]
|
||||
lines += [f"配音:{s.text}" for s in f.script_segments]
|
||||
lines.append(f"行动号召:{f.cta}")
|
||||
return "\n".join(lines)
|
||||
@@ -0,0 +1,160 @@
|
||||
"""Prompt 模板加载器:从 viral_video_prompt_templates 读模板,30 秒 TTL 热加载。
|
||||
|
||||
DB 不可用或没有数据时自动回落到 prompts.DEFAULT_TEMPLATES,保证流程不阻断。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import threading
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
from typing import Optional
|
||||
|
||||
import sqlalchemy as sa
|
||||
|
||||
from packages.adapters.sqlalchemy_impl import session as _session_mod
|
||||
from packages.application.viral_video.prompts import DEFAULT_TEMPLATES
|
||||
|
||||
CACHE_TTL_SECONDS = 30.0
|
||||
|
||||
_VALID_TYPES = {"image_analysis", "intent_parsing", "copy_fusion", "storyboard", "review"}
|
||||
|
||||
|
||||
@dataclass
|
||||
class PromptTemplate:
|
||||
name: str
|
||||
prompt_type: str
|
||||
version: int
|
||||
system_prompt: str
|
||||
user_prompt_template: str
|
||||
example_output: str = ""
|
||||
is_active: bool = True
|
||||
|
||||
|
||||
_lock = threading.Lock()
|
||||
_cache: dict[str, tuple[float, PromptTemplate]] = {}
|
||||
|
||||
|
||||
def _fallback(prompt_type: str) -> Optional[PromptTemplate]:
|
||||
for item in DEFAULT_TEMPLATES:
|
||||
if item["prompt_type"] == prompt_type:
|
||||
return PromptTemplate(
|
||||
name=item["name"],
|
||||
prompt_type=item["prompt_type"],
|
||||
version=item["version"],
|
||||
system_prompt=item["system_prompt"],
|
||||
user_prompt_template=item["user_prompt_template"],
|
||||
example_output=item["example_output"] or "",
|
||||
is_active=bool(item["is_active"]),
|
||||
)
|
||||
return None
|
||||
|
||||
|
||||
_lazy_session = None
|
||||
|
||||
|
||||
def _get_session():
|
||||
"""优先用全局 SessionLocal(worker);否则按应用配置懒建同步引擎(api)。"""
|
||||
global _lazy_session
|
||||
if _session_mod.SessionLocal is not None:
|
||||
return _session_mod.SessionLocal()
|
||||
if _lazy_session is not None:
|
||||
return _lazy_session()
|
||||
try:
|
||||
from packages.config import get_shared_settings
|
||||
|
||||
url = str(get_shared_settings().database_url)
|
||||
except Exception: # noqa: BLE001
|
||||
return None
|
||||
if not url:
|
||||
return None
|
||||
url = url.replace("postgresql+asyncpg://", "postgresql+psycopg://")
|
||||
url = url.replace("postgresql://", "postgresql+psycopg://") if url.startswith("postgresql://") else url
|
||||
engine = sa.create_engine(url, pool_pre_ping=True, pool_size=2, max_overflow=2)
|
||||
from sqlalchemy.orm import sessionmaker
|
||||
|
||||
_lazy_session = sessionmaker(bind=engine)
|
||||
return _lazy_session()
|
||||
|
||||
|
||||
def _load_from_db(prompt_type: str) -> Optional[PromptTemplate]:
|
||||
session = None
|
||||
try:
|
||||
session = _get_session()
|
||||
if session is None:
|
||||
return None
|
||||
sql = sa.text("""
|
||||
SELECT name, prompt_type, version, system_prompt,
|
||||
user_prompt_template, COALESCE(example_output, '') AS example_output,
|
||||
is_active
|
||||
FROM viral_video_prompt_templates
|
||||
WHERE prompt_type = :pt AND is_active = TRUE
|
||||
ORDER BY version DESC
|
||||
LIMIT 1
|
||||
""")
|
||||
row = session.execute(sql, {"pt": prompt_type}).first()
|
||||
if row is None:
|
||||
return None
|
||||
return PromptTemplate(
|
||||
name=row[0],
|
||||
prompt_type=row[1],
|
||||
version=int(row[2]),
|
||||
system_prompt=row[3],
|
||||
user_prompt_template=row[4],
|
||||
example_output=row[5] or "",
|
||||
is_active=bool(row[6]),
|
||||
)
|
||||
except Exception: # noqa: BLE001 - 表不存在/DB 不可用时静默回落
|
||||
return None
|
||||
finally:
|
||||
if session is not None:
|
||||
try:
|
||||
session.close()
|
||||
except Exception: # noqa: BLE001
|
||||
pass
|
||||
|
||||
|
||||
def get_template(prompt_type: str, *, force_refresh: bool = False) -> Optional[PromptTemplate]:
|
||||
"""取某类型当前启用模板,30 秒缓存;DB 无数据则回落到代码默认模板。"""
|
||||
if prompt_type not in _VALID_TYPES:
|
||||
raise ValueError(f"未知 prompt_type: {prompt_type}")
|
||||
|
||||
now = time.monotonic()
|
||||
with _lock:
|
||||
cached = _cache.get(prompt_type)
|
||||
if not force_refresh and cached and now - cached[0] < CACHE_TTL_SECONDS:
|
||||
return cached[1]
|
||||
|
||||
template = _load_from_db(prompt_type) or _fallback(prompt_type)
|
||||
if template is not None:
|
||||
with _lock:
|
||||
_cache[prompt_type] = (now, template)
|
||||
return template
|
||||
|
||||
|
||||
def invalidate() -> None:
|
||||
"""清空缓存(测试用)。"""
|
||||
with _lock:
|
||||
_cache.clear()
|
||||
|
||||
|
||||
class _SafeDict(dict):
|
||||
def __missing__(self, key: str) -> str:
|
||||
return "{" + key + "}"
|
||||
|
||||
|
||||
def _safe_format(text: str, kwargs: dict) -> str:
|
||||
try:
|
||||
return text.format_map(_SafeDict(kwargs))
|
||||
except Exception: # noqa: BLE001
|
||||
return text
|
||||
|
||||
|
||||
def render_user_prompt(template: PromptTemplate, **kwargs) -> str:
|
||||
"""填充 user_prompt_template 占位符,缺键原样保留不报错。"""
|
||||
return _safe_format(template.user_prompt_template, kwargs)
|
||||
|
||||
|
||||
def render_system_prompt(template: PromptTemplate, **kwargs) -> str:
|
||||
"""copy_fusion 等 system_prompt 含运行时变量时填充。"""
|
||||
return _safe_format(template.system_prompt, kwargs)
|
||||
@@ -0,0 +1,338 @@
|
||||
"""爆款视频 5 套 Prompt 模板默认值(#2040 核心资产)。
|
||||
|
||||
重要约定(用户明确要求):
|
||||
- 所有 system_prompt / user_prompt_template / example_output 都是**纯文本自然语言 + XML 标签**,
|
||||
运营可直接看懂和编辑,禁止 JSON、禁止 ```json 代码块。
|
||||
- LLM 按 XML 标签输出字段,程序用正则解析(见 xml_parser.py)。
|
||||
- user_prompt_template 中花括号占位符(如 {user_copy_text})在运行时填充。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
TEMPLATE_VERSION = 1
|
||||
|
||||
# 所有文案类 Prompt 自动注入的硬约束
|
||||
GLOBAL_CONSTRAINTS = """【必须遵守的硬约束】
|
||||
1. 不编造时间:不写“今年最新”“2024 爆款”等会过时的时间表述。
|
||||
2. 不承诺效果:不写“保证”“一定”“100%有效”“包治百病”等绝对化用语。
|
||||
3. 不编造价格、销量、认证、奖项:除非用户在文案中明确给出,否则一律不写。
|
||||
4. 符合广告法及平台社区规范。
|
||||
5. 只描述图片中真实可见的内容,看不到的不瞎猜。"""
|
||||
|
||||
# 反套路化要求
|
||||
NEGATIVE_RULES = """【反套路化要求】
|
||||
禁止使用“家人们谁懂啊”“绝绝子”“宝子们”“家人们”“太绝了”“yyds”等烂大街网络词;
|
||||
禁止固定模板化开头;语言要像真人朋友之间的分享,自然、具体、有信息量。"""
|
||||
|
||||
# 输出禁用套路词(测试会检查)
|
||||
BANNED_PHRASES = ["家人们谁懂啊", "绝绝子", "宝子们", "yyds", "太绝了"]
|
||||
|
||||
# 文案融合三档独立指令段
|
||||
FUSION_INSTRUCTIONS = {
|
||||
"ai_full": """【本次创作模式:AI 全权创作】
|
||||
你是资深短视频编导。用户只提供了产品图片,没有给出具体文案方向。请根据图片内容和营销参数,自由发挥创作完整的爆款短视频文案。充分挖掘产品真实可见的卖点,使用爆款结构,抓人眼球。""",
|
||||
"ai_polish": """【本次创作模式:AI 辅助润色】
|
||||
你是用户的文案助理。用户已经写了草稿/关键词/碎碎念,表达了他想讲的核心意思,但表达不完整、不够吸引人。你的任务是:以用户的意思为主,保留他想表达的所有核心信息点,在此基础上润色扩写、调整语序、增加衔接、优化表达,让文案更流畅更有吸引力。绝对不能改变用户想表达的核心意思,不能把用户的观点换成相反的,不能添加用户没提到的产品卖点。用户提到的品牌名、价格、人名、具体事实必须原样保留。""",
|
||||
"user_primary": """【本次创作模式:以用户原文为主】
|
||||
你是文案润色助手。用户已经写好了明确的文案,这是他最终想表达的内容。你的任务是最小化修改:只做必要的错别字修正、标点调整、语句通顺度优化,以及添加必要的衔接词让口播更自然。用户的核心句子、关键表述、事实信息一律不改。如果用户文案本身已经很好,直接返回,不要为了改而改。personal_brands 中的事实信息必须逐字保留。""",
|
||||
}
|
||||
|
||||
# ── 模板1:图片多模态分析(VLM)────────────────────────────────────────
|
||||
_IMAGE_ANALYSIS_SYSTEM = f"""你是电商商品视觉分析师,负责从商品图片中提取真实可见的商品信息。
|
||||
|
||||
工作方式(分步骤看,不要跳步):
|
||||
1. 先看整体:有哪些产品、什么场景、有没有人物。
|
||||
2. 再看细节:包装文字、颜色构成、人物状态、画面质感。
|
||||
3. 最后提炼卖点:只总结图片里能看到的卖点。
|
||||
|
||||
{GLOBAL_CONSTRAINTS}
|
||||
|
||||
请严格按下面的标签格式输出,标签名一个都不能改,不要输出任何解释,不要用代码块:
|
||||
<products> 下面每个产品用一个 <product> 标签,属性 name 是产品名、features 是外观特征、position 是 main 或 secondary、image_index 是第几张图(从0开始)。
|
||||
<colors> 下面每个主要颜色用一个 <color> 标签,属性 hex 是色值、name 是颜色名、coverage 是占比小数。
|
||||
<people> 用一个标签,属性 has_person、count、gender、age_range、hair(发型发色)、skin_tone(肤色)、face_shape(脸型)、outfit(穿着)、pose(姿态)、expression(表情)分别描述人物外貌。有人物时属性尽量具体(如hair="黑色长直发"、outfit="白色衬衫"),无人像时除has_person=false外其他填"无法判断"。
|
||||
<mood> 标签写画面整体情绪氛围。
|
||||
<visible_text> 下面每处可见文字用一个 <text_item> 标签,属性 text 是文字内容、position 是位置。
|
||||
<scene> 标签写场景描述。
|
||||
<quality> 用一个标签,属性 resolution、lighting、composition、blur 描述画质。
|
||||
<key_selling_points> 下面每个卖点用一个 <point> 标签。
|
||||
|
||||
【人物属性硬性要求(has_person=true时必须遵守)】
|
||||
hair/skin_tone/face_shape/outfit四项绝对禁止填“无法判断”,必须基于图片可见特征给出具体中文描述:
|
||||
- hair:必须描述发型+发色,如“黑色齐肩直发”“棕色微卷中长发”“深棕色短发”
|
||||
- skin_tone:必须描述肤色,如“暖调自然肤色”“白皙肤色”“小麦色”
|
||||
- face_shape:必须描述脸型,如“鹅蛋脸”“圆脸”“瓜子脸”“方脸”
|
||||
- outfit:必须描述可见穿着,如“米色翻领衬衫”“白色T恤”“黑色连衣裙”
|
||||
即使局部被遮挡也要根据可见部分合理推断;确实看不清时按最接近的直观印象描述。
|
||||
|
||||
其他非人物属性看不到或无法判断时填“无法判断”,布尔值填false,不要留空标签。
|
||||
|
||||
【有人物场景输出参考(女性手持商品示例,必须写全10个属性,禁止省略)】
|
||||
<people has_person="true" count="1" gender="女" age_range="青年" hair="黑色齐肩直发" skin_tone="暖调自然肤色" face_shape="鹅蛋脸" outfit="米色翻领衬衫" pose="正面半身,手持商品" expression="面带微笑"/>"""
|
||||
|
||||
_IMAGE_ANALYSIS_USER = """请分析以下商品图片,共 {image_count} 张。
|
||||
所属行业:{industry}
|
||||
图片地址:
|
||||
{image_urls}
|
||||
|
||||
按约定的标签格式输出分析结果。"""
|
||||
|
||||
_IMAGE_ANALYSIS_EXAMPLE = """<products>
|
||||
<product name="大公鸡头 多功能油污净 625ml" features="红色瓶盖白色瓶身,鸡头图案Logo" position="main" image_index="0"/>
|
||||
</products>
|
||||
<colors>
|
||||
<color hex="#D32F2F" name="红色" coverage="0.4"/>
|
||||
<color hex="#FFFFFF" name="白色" coverage="0.5"/>
|
||||
</colors>
|
||||
<people has_person="false" count="0" gender="无法判断" age_range="无法判断" hair="无法判断" skin_tone="无法判断" face_shape="无法判断" outfit="无法判断" pose="无法判断" expression="无法判断"/>
|
||||
<mood>干净、实用</mood>
|
||||
<visible_text>
|
||||
<text_item text="多功能油污净" position="瓶身正面"/>
|
||||
</visible_text>
|
||||
<scene>白底棚拍产品图</scene>
|
||||
<quality resolution="高清" lighting="均匀柔和" composition="主体居中" blur="false"/>
|
||||
<key_selling_points>
|
||||
<point>针对重油污设计</point>
|
||||
<point>大容量625ml</point>
|
||||
</key_selling_points>"""
|
||||
|
||||
# ── 模板2:用户文案意图解析(LLM)──────────────────────────────────────
|
||||
_INTENT_SYSTEM = f"""你负责理解用户的营销意图。用户给的文案可能只是几个关键词、碎碎念或者不完整的短句,你要读懂他真正想讲什么。
|
||||
|
||||
{GLOBAL_CONSTRAINTS}
|
||||
|
||||
请严格按下面的标签格式输出,不要解释,不要用代码块:
|
||||
<intent_summary> 用用户的语言风格,一句话、30字以内概括核心意图。
|
||||
<core_messages> 下面每个核心信息点用一个 <message> 标签,属性 must_keep 为 true 或 false、confidence 为 0 到 1 的小数,标签内容写信息点。
|
||||
<personal_brands> 把用户提到的具体事实——品牌名、价格、人名、地名、时间、产品名——每条用一个 <brand> 标签,属性 category 取 brand、price、person、place、time、product 之一。这些事实必须原样引用,一个字都不能改。
|
||||
<emotion_tone> 写文案的情绪调性。
|
||||
<missing_info> 把你认为缺失、后续生成时需要合理推断的信息,每条用一个 <info> 标签;没有就输出空标签。"""
|
||||
|
||||
_INTENT_USER = """用户原始文案:{user_copy_text}
|
||||
所属行业:{industry}
|
||||
图片分析结果(供参考):
|
||||
{image_analysis}
|
||||
|
||||
请理解用户意图,按标签格式输出。"""
|
||||
|
||||
_INTENT_EXAMPLE = """<intent_summary>一款厨房去油污神器,喷一喷油污就掉</intent_summary>
|
||||
<core_messages>
|
||||
<message must_keep="true" confidence="0.97">去油污效果好,喷上等几分钟再擦</message>
|
||||
<message must_keep="false" confidence="0.7">适合厨房重油污场景</message>
|
||||
</core_messages>
|
||||
<personal_brands>
|
||||
<brand category="product">大公鸡头多功能油污净</brand>
|
||||
<brand category="price">39块钱一瓶</brand>
|
||||
</personal_brands>
|
||||
<emotion_tone>亲切、真实、带分享感</emotion_tone>
|
||||
<missing_info>
|
||||
<info>没有说明具体容量,按图片读出的625ml处理</info>
|
||||
</missing_info>"""
|
||||
|
||||
# ── 模板3:文案融合生成(LLM)──────────────────────────────────────────
|
||||
_FUSION_SYSTEM = """你负责为短视频生成营销文案。请按思维链分步完成:先定人设和目标客户,再找卖点,再搭结构,再安排情绪,最后写行动号召,不要一步到位乱写。
|
||||
|
||||
{fusion_instruction}
|
||||
|
||||
{global_constraints}
|
||||
|
||||
{negative_rules}
|
||||
|
||||
请严格按下面的标签格式输出,不要解释,不要用代码块:
|
||||
<title> 视频标题。
|
||||
<hook> 开头3秒钩子,5到15字。
|
||||
<body_points> 每个要点用一个 <point> 标签,属性 elaboration 是展开说明、image_index 是对应第几张图(从0开始),标签内容写要点。
|
||||
<cta> 口语化的行动号召。
|
||||
<script_segments> 每段配音用一个 <segment> 标签,属性 duration_sec 是秒数、image_index 是对应图片,标签内容写配音文案(纯口播文本,不加旁白标注、不加镜头标注、不加"主播:"之类前缀)。
|
||||
<voiceover_script> 把所有 segment 的配音文案按顺序自然拼接成一段完整的纯口播文本(无标记、无括号、无前缀),长度要适配 {duration} 秒,约 {approx_chars} 字。
|
||||
<overview_theme> 视频主题(一句话概括)。
|
||||
<scene_and_lighting> 整体场景描述+光线设定(100-200字,要具体:在哪拍、什么光线、什么色调、什么氛围)。
|
||||
<word_count> 配音总字数,只写数字。
|
||||
<estimated_duration> 预计时长秒数,只写数字。
|
||||
|
||||
用户在 personal_brands 中提到的品牌名、价格、人名、地名、时间、产品名等事实信息,必须原样出现在文案里,一个字都不能改。"""
|
||||
|
||||
_FUSION_USER = """所属行业:{industry}
|
||||
目标客户:{target_customer}
|
||||
营销目的:{marketing_purpose}
|
||||
视频时长:{duration}秒
|
||||
图片分析结果:
|
||||
{image_analysis}
|
||||
用户意图解析结果:
|
||||
{intent_result}
|
||||
|
||||
请按标签格式生成文案。"""
|
||||
|
||||
_FUSION_EXAMPLE = """<title>厨房重油污,别再用洗洁精硬擦了</title>
|
||||
<hook>这油污,我真的忍很久了</hook>
|
||||
<body_points>
|
||||
<point elaboration="喷在油污上等几分钟,一擦就干净" image_index="0">大公鸡头油污净去油快</point>
|
||||
<point elaboration="39块钱625ml,能用很久" image_index="0">39块钱一瓶,性价比高</point>
|
||||
</body_points>
|
||||
<cta>厨房油污重的,真的可以试一瓶</cta>
|
||||
<script_segments>
|
||||
<segment duration_sec="3" image_index="0">这油污我真的忍很久了,用洗洁精擦半天都没用</segment>
|
||||
<segment duration_sec="6" image_index="0">后来换了这个大公鸡头油污净,喷上等几分钟,一擦就干净</segment>
|
||||
<segment duration_sec="4" image_index="0">39块钱625ml,厨房重油污的可以试一瓶</segment>
|
||||
</script_segments>
|
||||
<voiceover_script>这油污我真的忍很久了,用洗洁精擦半天都没用。后来换了这个大公鸡头油污净,喷上等几分钟,一擦就干净。39块钱625ml,厨房重油污的可以试一瓶。</voiceover_script>
|
||||
<overview_theme>厨房油污清洁好物分享</overview_theme>
|
||||
<scene_and_lighting>简洁明亮的厨房台面场景,自然光从窗户洒入,色调温暖柔和,突出产品白色瓶身与去油污对比效果。</scene_and_lighting>
|
||||
<word_count>58</word_count>
|
||||
<estimated_duration>13</estimated_duration>"""
|
||||
|
||||
# ── 模板4:编导级分镜(LLM)────────────────────────────────────────────
|
||||
_STORYBOARD_SYSTEM = """你是短视频编导,负责把文案拆成可拍摄的分镜,为 Seedance 2.5 视频模型写编导分镜脚本。脚本将整体作为 prompt 一次性传给视频模型,必须让模型在连贯镜头流中清楚每段时间拍什么、画面如何、人物说什么。
|
||||
|
||||
工作方式:
|
||||
1. 按文案的 script_segments 顺序分配镜头。
|
||||
2. 每个镜头确定景别/角度/运镜、画面场景与对白、人物动作细节、音效/BGM、转场。
|
||||
3. 检查所有镜头时长加起来接近目标时长,误差不超过2秒。
|
||||
4. image_index 必须在已上传图片范围内,第一张主图必须用在第一个镜头。
|
||||
|
||||
{fusion_instruction}
|
||||
|
||||
{global_constraints}
|
||||
|
||||
{negative_rules}
|
||||
|
||||
请严格按下面的标签格式输出,不要解释,不要用代码块:
|
||||
<clips> 下面每个镜头用一个 <clip> 标签,属性 image_index 是图片序号(从0开始)、transition 取 fade/cut/zoom_in/slide_left/dissolve/wipe 之一、zoom 取 in/out/null、duration_sec 是该镜头秒数、bgm_note 是该段BGM情绪。每个 <clip> 里面包含:
|
||||
<voice_text> 该镜头配音文本(纯口播文本,不加旁白标注);
|
||||
<subtitle_text> 字幕文本,可与配音一致或更精简;
|
||||
<shot_type_angle_movement> 景别+角度+运镜(例:近景俯拍45度,缓慢推镜;中景平视,固定镜头;特写平视,快速拉镜);
|
||||
<scene_and_dialogue> 画面场景描述 + 人物口播台词(对白要自然口语化,像朋友聊天,不要硬广推销腔);
|
||||
<action_details> 人物动作、表情、物品操作细节(手怎么动、表情变化、产品怎么展示);
|
||||
<audio_bgm> 环境音+BGM提示(例:轻快流行BGM,环境嘈杂咖啡店背景音);
|
||||
<transition> 硬切/淡入淡出/叠化(最后一镜写『结束』即可);
|
||||
<reference_image_index> 参考图片索引(0-based,对应第几张产品图,无则空);
|
||||
<ken_burns> 用一个空标签,属性 start、end 写"x,y"坐标、ease 写缓动方式;不需要运镜时坐标相同。"""
|
||||
|
||||
_STORYBOARD_USER = """目标时长:{duration}秒
|
||||
上传图片数量:{image_count}张(第1张是主图/封面)
|
||||
文案内容:
|
||||
{fusion_result}
|
||||
图片分析结果:
|
||||
{image_analysis}
|
||||
|
||||
请按标签格式输出分镜。"""
|
||||
|
||||
_STORYBOARD_EXAMPLE = """<clips>
|
||||
<clip image_index="0" transition="cut" zoom="null" duration_sec="3" bgm_note="日常、轻微烦躁">
|
||||
<voice_text>这油污我真的忍很久了</voice_text>
|
||||
<subtitle_text>这油污忍很久了</subtitle_text>
|
||||
<shot_type_angle_movement>近景俯拍45度,缓慢推镜</shot_type_angle_movement>
|
||||
<scene_and_dialogue>厨房台面,主妇皱眉看着灶台油污。对白:这油污我真的忍很久了</scene_and_dialogue>
|
||||
<action_details>右手拿着脏抹布,无奈摇头</action_details>
|
||||
<audio_bgm>轻快日常BGM,带一点烦躁感</audio_bgm>
|
||||
<transition>硬切</transition>
|
||||
<reference_image_index>0</reference_image_index>
|
||||
<ken_burns start="0,0" end="0,0" ease="linear"/>
|
||||
</clip>
|
||||
<clip image_index="0" transition="zoom_in" zoom="in" duration_sec="6" bgm_note="轻快、出现转机">
|
||||
<voice_text>后来换了大公鸡头油污净,喷上等几分钟,一擦就干净</voice_text>
|
||||
<subtitle_text>喷上等几分钟,一擦就干净</subtitle_text>
|
||||
<shot_type_angle_movement>特写平视,固定镜头</shot_type_angle_movement>
|
||||
<scene_and_dialogue>手部特写,喷油污净在油污处。对白:后来换了这个大公鸡头油污净,喷上等几分钟,一擦就干净</scene_and_dialogue>
|
||||
<action_details>左手拿产品瓶身,右手按压喷头,等待片刻后用抹布轻擦</action_details>
|
||||
<audio_bgm>轻快转折BGM,带清爽感</audio_bgm>
|
||||
<transition>淡入淡出</transition>
|
||||
<reference_image_index>0</reference_image_index>
|
||||
<ken_burns start="20,20" end="80,80" ease="ease-in-out"/>
|
||||
</clip>
|
||||
<clip image_index="0" transition="fade" zoom="null" duration_sec="4" bgm_note="温暖、推荐">
|
||||
<voice_text>39块钱625ml,厨房重油污的可以试一瓶</voice_text>
|
||||
<subtitle_text>39元625ml,可以试一瓶</subtitle_text>
|
||||
<shot_type_angle_movement>中景平视,缓慢拉镜</shot_type_angle_movement>
|
||||
<scene_and_dialogue>产品正面展示,明亮背景。对白:39块钱625ml,厨房重油污的可以试一瓶</scene_and_dialogue>
|
||||
<action_details>产品置于画面中央,轻微转动展示瓶身</action_details>
|
||||
<audio_bgm>温暖收尾BGM</audio_bgm>
|
||||
<transition>结束</transition>
|
||||
<reference_image_index>0</reference_image_index>
|
||||
<ken_burns start="50,50" end="20,20" ease="ease-in-out"/>
|
||||
</clip>
|
||||
</clips>"""
|
||||
|
||||
# ── 模板5:文案审核(LLM)──────────────────────────────────────────────
|
||||
_REVIEW_SYSTEM = f"""你是短视频文案合规审核员,从6个维度逐条检查文案:
|
||||
1. 违规词:有没有平台禁用词、敏感词。
|
||||
2. 夸大承诺:有没有“包治百病”“100%有效”“保证赚钱”等绝对化、夸大表述。
|
||||
3. 事实一致性:有没有编造价格、数据、认证,或者用户没提到的产品特性。
|
||||
4. 用户意图保留:在 ai_polish 和 user_primary 模式下,core_messages 中 must_keep=true 的点是否都保留了。
|
||||
5. 结构完整性:标题、钩子、正文、行动号召是否齐全。
|
||||
6. 语气人设:是否符合选定的人设语气,有没有“家人们谁懂啊”“绝绝子”“宝子们”等套路词。
|
||||
|
||||
{GLOBAL_CONSTRAINTS}
|
||||
|
||||
请严格按下面的标签格式输出,不要解释,不要用代码块:
|
||||
<passed> 整体是否通过,只写 true 或 false。
|
||||
<issues> 每个问题用一个 <issue> 标签,属性 dimension 是维度名、severity 取 error 或 warning、location 是问题所在(如 hook、body_points、cta),标签内容写问题描述;没有问题就输出空标签。
|
||||
<rewrite_suggestions> 每条具体修改建议用一个 <suggestion> 标签;没有就输出空标签。"""
|
||||
|
||||
_REVIEW_USER = """本次创作模式:{fusion_level}
|
||||
待审核文案:
|
||||
{fusion_result}
|
||||
用户意图解析(用于核对核心信息是否保留):
|
||||
{intent_result}
|
||||
|
||||
请按6个维度审核,按标签格式输出。"""
|
||||
|
||||
_REVIEW_EXAMPLE = """<passed>false</passed>
|
||||
<issues>
|
||||
<issue dimension="夸大承诺" severity="error" location="body_points">出现了“一喷100%掉光”的绝对化表述,违反广告法</issue>
|
||||
<issue dimension="用户意图保留" severity="warning" location="cta">用户强调的“39块钱”没有保留</issue>
|
||||
</issues>
|
||||
<rewrite_suggestions>
|
||||
<suggestion>把“一喷100%掉光”改为“喷上等几分钟,大部分油污能擦掉”</suggestion>
|
||||
<suggestion>在结尾补回“39块钱625ml”</suggestion>
|
||||
</rewrite_suggestions>"""
|
||||
|
||||
|
||||
# 5 套模板默认数据(seed 数据源与 loader 的兜底)
|
||||
DEFAULT_TEMPLATES: list[dict] = [
|
||||
{
|
||||
"name": "图片多模态分析",
|
||||
"prompt_type": "image_analysis",
|
||||
"version": TEMPLATE_VERSION,
|
||||
"system_prompt": _IMAGE_ANALYSIS_SYSTEM,
|
||||
"user_prompt_template": _IMAGE_ANALYSIS_USER,
|
||||
"example_output": _IMAGE_ANALYSIS_EXAMPLE,
|
||||
"is_active": True,
|
||||
},
|
||||
{
|
||||
"name": "用户文案意图解析",
|
||||
"prompt_type": "intent_parsing",
|
||||
"version": TEMPLATE_VERSION,
|
||||
"system_prompt": _INTENT_SYSTEM,
|
||||
"user_prompt_template": _INTENT_USER,
|
||||
"example_output": _INTENT_EXAMPLE,
|
||||
"is_active": True,
|
||||
},
|
||||
{
|
||||
"name": "文案融合生成",
|
||||
"prompt_type": "copy_fusion",
|
||||
"version": TEMPLATE_VERSION,
|
||||
"system_prompt": _FUSION_SYSTEM,
|
||||
"user_prompt_template": _FUSION_USER,
|
||||
"example_output": _FUSION_EXAMPLE,
|
||||
"is_active": True,
|
||||
},
|
||||
{
|
||||
"name": "编导级分镜",
|
||||
"prompt_type": "storyboard",
|
||||
"version": TEMPLATE_VERSION,
|
||||
"system_prompt": _STORYBOARD_SYSTEM,
|
||||
"user_prompt_template": _STORYBOARD_USER,
|
||||
"example_output": _STORYBOARD_EXAMPLE,
|
||||
"is_active": True,
|
||||
},
|
||||
{
|
||||
"name": "文案审核",
|
||||
"prompt_type": "review",
|
||||
"version": TEMPLATE_VERSION,
|
||||
"system_prompt": _REVIEW_SYSTEM,
|
||||
"user_prompt_template": _REVIEW_USER,
|
||||
"example_output": _REVIEW_EXAMPLE,
|
||||
"is_active": True,
|
||||
},
|
||||
]
|
||||
@@ -0,0 +1,312 @@
|
||||
"""文案审核 + 自动重写(#2040 第5套 Prompt)。
|
||||
|
||||
6 维度:违规词 / 夸大承诺 / 事实一致性 / 用户意图保留 / 结构完整性 / 语气人设。
|
||||
LLM 审核之外叠加本地规则预检(保证即使 LLM 不可用也能兜住广告法红线)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import re
|
||||
from typing import Optional
|
||||
|
||||
from packages.application.viral_video import xml_parser as xp
|
||||
from packages.application.viral_video.prompt_loader import (
|
||||
get_template,
|
||||
render_system_prompt,
|
||||
render_user_prompt,
|
||||
)
|
||||
from packages.application.viral_video.schemas import (
|
||||
FusionResult,
|
||||
IntentResult,
|
||||
ReviewIssue,
|
||||
ReviewResult,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# 本地规则:绝对化/夸大词
|
||||
_EXAGGERATION_PATTERNS = [
|
||||
r"100\s*%",
|
||||
r"百分百",
|
||||
r"包治百病",
|
||||
r"保证.{0,8}(有效|赚钱|瘦|好)",
|
||||
r"绝对(有效|安全|靠谱)",
|
||||
r"全网第一",
|
||||
r"国家级",
|
||||
r"特效",
|
||||
r"立刻见效",
|
||||
r"一喷(就|全|100)",
|
||||
]
|
||||
|
||||
# 本地规则:平台违规/套路词
|
||||
_VIOLATION_PHRASES = [
|
||||
"家人们谁懂啊",
|
||||
"绝绝子",
|
||||
"宝子们",
|
||||
"yyds",
|
||||
"最(好|强|牛|便宜)", # 广告法极限词
|
||||
"第一(名|品牌)?",
|
||||
]
|
||||
|
||||
_LOCATIONS = ["title", "hook", "body_points", "cta", "script_segments"]
|
||||
|
||||
|
||||
class Reviewer:
|
||||
def __init__(self, client=None):
|
||||
if client is None:
|
||||
from packages.shared.ai_client import get_doubao_client
|
||||
|
||||
client = get_doubao_client()
|
||||
self.client = client
|
||||
|
||||
# ── 审核 ────────────────────────────────────────────────────────────
|
||||
def review(self, fusion: FusionResult, intent: IntentResult, fusion_level: str) -> ReviewResult:
|
||||
local = self._rule_check(fusion, intent, fusion_level)
|
||||
llm_result = self._llm_review(fusion, intent, fusion_level)
|
||||
if llm_result is None:
|
||||
return ReviewResult(
|
||||
passed=not local,
|
||||
issues=local,
|
||||
rewrite_suggestions=[],
|
||||
raw="",
|
||||
)
|
||||
# LLM 与本地规则合并去重
|
||||
issues = self._merge_issues(llm_result.issues, local)
|
||||
return ReviewResult(
|
||||
passed=llm_result.passed and not local,
|
||||
issues=issues,
|
||||
rewrite_suggestions=llm_result.rewrite_suggestions,
|
||||
raw=llm_result.raw,
|
||||
)
|
||||
|
||||
def _llm_review(self, fusion: FusionResult, intent: IntentResult, fusion_level: str) -> Optional[ReviewResult]:
|
||||
template = get_template("review")
|
||||
system = render_system_prompt(template)
|
||||
user = render_user_prompt(
|
||||
template,
|
||||
fusion_level=fusion_level,
|
||||
fusion_result=self._fusion_text(fusion),
|
||||
intent_result=self._intent_text(intent),
|
||||
)
|
||||
raw = self.client.chat_completion(
|
||||
[
|
||||
{"role": "system", "content": system},
|
||||
{"role": "user", "content": user},
|
||||
],
|
||||
temperature=0.2,
|
||||
max_tokens=1024,
|
||||
)
|
||||
if not raw:
|
||||
return None
|
||||
passed = xp.text_of(raw, "passed").strip().lower()
|
||||
issues = [
|
||||
ReviewIssue(
|
||||
dimension=n["attrs"].get("dimension", "未知维度"),
|
||||
severity=n["attrs"].get("severity", "warning"),
|
||||
location=n["attrs"].get("location", ""),
|
||||
text=n["text"],
|
||||
)
|
||||
for n in xp.find_all(raw, "issue")
|
||||
if n["text"]
|
||||
]
|
||||
suggestions = [n["text"] for n in xp.find_all(raw, "suggestion") if n["text"]]
|
||||
parsed = ReviewResult(
|
||||
passed=passed == "true" and not issues,
|
||||
issues=issues,
|
||||
rewrite_suggestions=suggestions,
|
||||
raw=raw,
|
||||
)
|
||||
return parsed
|
||||
|
||||
# ── 本地规则预检 ────────────────────────────────────────────────────
|
||||
def _rule_check(self, fusion: FusionResult, intent, fusion_level: str) -> list[ReviewIssue]:
|
||||
issues: list[ReviewIssue] = []
|
||||
for location, text in self._segments(fusion):
|
||||
for pattern in _EXAGGERATION_PATTERNS:
|
||||
if re.search(pattern, text):
|
||||
issues.append(
|
||||
ReviewIssue(
|
||||
dimension="夸大承诺",
|
||||
severity="error",
|
||||
location=location,
|
||||
text=f"出现夸大/绝对化表述:{self._hit(text, pattern)}",
|
||||
)
|
||||
)
|
||||
for phrase in _VIOLATION_PHRASES:
|
||||
if re.search(phrase, text, flags=re.IGNORECASE):
|
||||
issues.append(
|
||||
ReviewIssue(
|
||||
dimension="违规词",
|
||||
severity="error",
|
||||
location=location,
|
||||
text=f"出现违规或套路词:{self._hit(text, phrase)}",
|
||||
)
|
||||
)
|
||||
|
||||
# 结构完整性
|
||||
if not fusion.title:
|
||||
issues.append(ReviewIssue(dimension="结构完整性", severity="warning", location="title", text="缺少标题"))
|
||||
if not fusion.hook:
|
||||
issues.append(ReviewIssue(dimension="结构完整性", severity="warning", location="hook", text="缺少开头钩子"))
|
||||
if not fusion.cta:
|
||||
issues.append(ReviewIssue(dimension="结构完整性", severity="warning", location="cta", text="缺少行动号召"))
|
||||
|
||||
# 用户意图保留(must_keep)
|
||||
full_text = self._fusion_text(fusion)
|
||||
if fusion_level in {"ai_polish", "user_primary"} and intent is not None:
|
||||
for message in intent.core_messages:
|
||||
if message.must_keep:
|
||||
key = self._compact(message.text)
|
||||
if key and key[:10] not in self._compact(full_text):
|
||||
issues.append(
|
||||
ReviewIssue(
|
||||
dimension="用户意图保留",
|
||||
severity="warning",
|
||||
location="script_segments",
|
||||
text=f"用户核心信息被丢失:{message.text[:30]}",
|
||||
)
|
||||
)
|
||||
for brand in intent.personal_brands:
|
||||
if brand.text and brand.text not in full_text:
|
||||
issues.append(
|
||||
ReviewIssue(
|
||||
dimension="事实一致性",
|
||||
severity="error",
|
||||
location="script_segments",
|
||||
text=f"personal_brands 事实信息未原样保留:{brand.text[:30]}",
|
||||
)
|
||||
)
|
||||
return issues
|
||||
|
||||
@staticmethod
|
||||
def _hit(text: str, pattern: str) -> str:
|
||||
match = re.search(pattern, text, flags=re.IGNORECASE)
|
||||
return match.group(0) if match else pattern
|
||||
|
||||
@staticmethod
|
||||
def _compact(text: str) -> str:
|
||||
return re.sub(r"[\s,。!?、,.!?;;::\"'“”‘’()()【】\[\]]", "", text)
|
||||
|
||||
@staticmethod
|
||||
def _merge_issues(llm_issues: list[ReviewIssue], local: list[ReviewIssue]) -> list[ReviewIssue]:
|
||||
merged = list(local)
|
||||
seen = {(i.dimension, Reviewer._compact(i.text)[:20]) for i in local}
|
||||
for issue in llm_issues:
|
||||
key = (issue.dimension, Reviewer._compact(issue.text)[:20])
|
||||
if key not in seen:
|
||||
merged.append(issue)
|
||||
seen.add(key)
|
||||
return merged
|
||||
|
||||
# ── 自动重写(1 次)─────────────────────────────────────────────────
|
||||
def rewrite(
|
||||
self,
|
||||
fusion: FusionResult,
|
||||
review: ReviewResult,
|
||||
intent: IntentResult,
|
||||
fusion_level: str,
|
||||
) -> FusionResult:
|
||||
from packages.application.viral_video.generator import CopyGenerator
|
||||
|
||||
template = get_template("copy_fusion")
|
||||
system_kwargs = {
|
||||
"fusion_instruction": (
|
||||
"【本次任务:按审核意见修正文案】只修改指出的问题,其他内容尽量原样保留;"
|
||||
"personal_brands 事实信息逐字保留;修正后按原标签格式完整输出。"
|
||||
),
|
||||
"global_constraints": "",
|
||||
"negative_rules": "",
|
||||
}
|
||||
issue_text = "\n".join(f"- [{i.dimension}/{i.location}] {i.text}" for i in review.issues)
|
||||
suggestion_text = "\n".join(f"- {s}" for s in review.rewrite_suggestions)
|
||||
user = render_user_prompt(
|
||||
template,
|
||||
industry="",
|
||||
target_customer="",
|
||||
marketing_purpose="",
|
||||
duration=fusion.estimated_duration or 15,
|
||||
image_analysis="(沿用原图片分析)",
|
||||
intent_result=self._intent_text(intent),
|
||||
)
|
||||
user = (
|
||||
f"{user}\n\n原文案:\n{self._fusion_text(fusion)}\n\n"
|
||||
f"审核发现的问题:\n{issue_text}\n\n修改建议:\n{suggestion_text or '(无)'}\n"
|
||||
"请输出修正后的完整文案。"
|
||||
)
|
||||
system = render_system_prompt(template, **system_kwargs)
|
||||
raw = self.client.chat_completion(
|
||||
[
|
||||
{"role": "system", "content": system},
|
||||
{"role": "user", "content": user},
|
||||
],
|
||||
temperature=0.5,
|
||||
max_tokens=2048,
|
||||
)
|
||||
if not raw:
|
||||
return self._rule_fix(fusion, review)
|
||||
rewritten = CopyGenerator._parse_fusion(CopyGenerator(self.client), raw)
|
||||
if not rewritten.title and not rewritten.script_segments:
|
||||
return self._rule_fix(fusion, review)
|
||||
# 保底:personal_brands 必须保留
|
||||
full = self._fusion_text(rewritten)
|
||||
for brand in intent.personal_brands:
|
||||
if brand.text and brand.text not in full:
|
||||
rewritten.cta = (rewritten.cta + brand.text).strip()
|
||||
return rewritten
|
||||
|
||||
def _rule_fix(self, fusion: FusionResult, review: ReviewResult) -> FusionResult:
|
||||
"""LLM 重写不可用时的本地兜底:删除/替换明显违规表述。"""
|
||||
replacements = [
|
||||
(re.compile(r"100\s*%|百分百"), "大部分"),
|
||||
(re.compile(r"绝对(有效|安全|靠谱)"), "比较\\1"),
|
||||
(re.compile(r"包治百病"), "适用多种情况"),
|
||||
(re.compile(r"立刻见效"), "坚持使用会有改善"),
|
||||
(re.compile(r"一喷(就|全|100%)"), "喷上等一会儿可以"),
|
||||
(re.compile(r"家人们谁懂啊|绝绝子|宝子们|yyds", re.IGNORECASE), ""),
|
||||
(re.compile(r"最好|最强|最牛|最便宜"), "很不错"),
|
||||
]
|
||||
|
||||
def fix(text: str) -> str:
|
||||
for pattern, repl in replacements:
|
||||
text = pattern.sub(repl, text)
|
||||
return text
|
||||
|
||||
fusion.title = fix(fusion.title)
|
||||
fusion.hook = fix(fusion.hook)
|
||||
fusion.cta = fix(fusion.cta)
|
||||
for point in fusion.body_points:
|
||||
point.text = fix(point.text)
|
||||
point.elaboration = fix(point.elaboration)
|
||||
for segment in fusion.script_segments:
|
||||
segment.text = fix(segment.text)
|
||||
fusion.raw = ""
|
||||
return fusion
|
||||
|
||||
# ── 文本工具 ────────────────────────────────────────────────────────
|
||||
@staticmethod
|
||||
def _segments(fusion: FusionResult):
|
||||
yield "title", fusion.title
|
||||
yield "hook", fusion.hook
|
||||
for point in fusion.body_points:
|
||||
yield "body_points", f"{point.text} {point.elaboration}"
|
||||
yield "cta", fusion.cta
|
||||
for segment in fusion.script_segments:
|
||||
yield "script_segments", segment.text
|
||||
|
||||
@staticmethod
|
||||
def _fusion_text(fusion: FusionResult) -> str:
|
||||
parts = [fusion.title, fusion.hook]
|
||||
parts += [p.text for p in fusion.body_points]
|
||||
parts += [s.text for s in fusion.script_segments]
|
||||
parts.append(fusion.cta)
|
||||
return "\n".join(p for p in parts if p)
|
||||
|
||||
@staticmethod
|
||||
def _intent_text(intent) -> str:
|
||||
if intent is None:
|
||||
return "无意图信息"
|
||||
parts = [f"意图:{intent.intent_summary}"]
|
||||
parts += [f"核心信息[must_keep={m.must_keep}]:{m.text}" for m in intent.core_messages]
|
||||
parts += [f"事实({b.category}):{b.text}" for b in intent.personal_brands]
|
||||
return "\n".join(parts)
|
||||
@@ -0,0 +1,116 @@
|
||||
"""内部 Pydantic 校验模型(不暴露给运营,运营只看 DB 里的纯文本)。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
|
||||
class ProductItem(BaseModel):
|
||||
name: str = "无法判断"
|
||||
features: str = "无法判断"
|
||||
position: str = "secondary"
|
||||
image_index: int = 0
|
||||
|
||||
|
||||
class ColorItem(BaseModel):
|
||||
hex: str = "#000000"
|
||||
name: str = "无法判断"
|
||||
coverage: float = 0.0
|
||||
|
||||
|
||||
class TextItem(BaseModel):
|
||||
text: str = ""
|
||||
position: str = ""
|
||||
|
||||
|
||||
class ImageAnalysis(BaseModel):
|
||||
products: list[ProductItem] = Field(default_factory=list)
|
||||
colors: list[ColorItem] = Field(default_factory=list)
|
||||
has_person: bool = False
|
||||
person_count: int = 0
|
||||
people: dict[str, str] = Field(default_factory=dict)
|
||||
mood: str = ""
|
||||
visible_text: list[TextItem] = Field(default_factory=list)
|
||||
scene: str = ""
|
||||
quality: dict[str, str] = Field(default_factory=dict)
|
||||
key_selling_points: list[str] = Field(default_factory=list)
|
||||
raw: str = ""
|
||||
|
||||
|
||||
class CoreMessage(BaseModel):
|
||||
text: str
|
||||
must_keep: bool = False
|
||||
confidence: float = 0.0
|
||||
|
||||
|
||||
class PersonalBrand(BaseModel):
|
||||
text: str
|
||||
category: str = "brand"
|
||||
|
||||
|
||||
class IntentResult(BaseModel):
|
||||
intent_summary: str = ""
|
||||
core_messages: list[CoreMessage] = Field(default_factory=list)
|
||||
personal_brands: list[PersonalBrand] = Field(default_factory=list)
|
||||
emotion_tone: str = ""
|
||||
missing_info: list[str] = Field(default_factory=list)
|
||||
raw: str = ""
|
||||
|
||||
|
||||
class BodyPoint(BaseModel):
|
||||
text: str
|
||||
elaboration: str = ""
|
||||
image_index: int = 0
|
||||
|
||||
|
||||
class ScriptSegment(BaseModel):
|
||||
text: str
|
||||
duration_sec: float = 0
|
||||
image_index: int = 0
|
||||
|
||||
|
||||
class FusionResult(BaseModel):
|
||||
title: str = ""
|
||||
hook: str = ""
|
||||
body_points: list[BodyPoint] = Field(default_factory=list)
|
||||
cta: str = ""
|
||||
script_segments: list[ScriptSegment] = Field(default_factory=list)
|
||||
word_count: int = 0
|
||||
estimated_duration: int = 0
|
||||
raw: str = ""
|
||||
|
||||
|
||||
class KenBurns(BaseModel):
|
||||
start: str = "0,0"
|
||||
end: str = "0,0"
|
||||
ease: str = "linear"
|
||||
|
||||
|
||||
class Clip(BaseModel):
|
||||
image_index: int = 0
|
||||
transition: str = "cut"
|
||||
zoom: str | None = None
|
||||
duration_sec: float = 0
|
||||
bgm_note: str = ""
|
||||
voice_text: str = ""
|
||||
subtitle_text: str = ""
|
||||
ken_burns: KenBurns = Field(default_factory=KenBurns)
|
||||
|
||||
|
||||
class Storyboard(BaseModel):
|
||||
clips: list[Clip] = Field(default_factory=list)
|
||||
raw: str = ""
|
||||
|
||||
|
||||
class ReviewIssue(BaseModel):
|
||||
dimension: str
|
||||
severity: str = "warning"
|
||||
location: str = ""
|
||||
text: str = ""
|
||||
|
||||
|
||||
class ReviewResult(BaseModel):
|
||||
passed: bool = True
|
||||
issues: list[ReviewIssue] = Field(default_factory=list)
|
||||
rewrite_suggestions: list[str] = Field(default_factory=list)
|
||||
raw: str = ""
|
||||
@@ -0,0 +1,104 @@
|
||||
"""XML 标签式输出解析器(替代 json.loads)。
|
||||
|
||||
LLM 按 ``<tag attr="x">内容</tag>`` 输出,本模块解析,解析失败不抛异常,
|
||||
由调用方走规则 fallback。采用栈式扫描,嵌套标签全部可提取(内外层都保留)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
from html import unescape
|
||||
from typing import Optional
|
||||
|
||||
_OPEN_RE = re.compile(r"<(?P<tag>[\w-]+)(?P<attrs>(?:\s(?:[^>]*?\S)?)?)(?P<self>/?)>")
|
||||
_CLOSE_RE = re.compile(r"</(?P<tag>[\w-]+)\s*>")
|
||||
_ATTR_RE = re.compile(r"""([\w:-]+)\s*=\s*(?:"([^"]*)"|'([^']*)')""")
|
||||
|
||||
|
||||
def parse_attributes(raw: str) -> dict[str, str]:
|
||||
"""解析标签属性字符串。"""
|
||||
attrs: dict[str, str] = {}
|
||||
for match in _ATTR_RE.finditer(raw or ""):
|
||||
value = match.group(2) if match.group(2) is not None else match.group(3)
|
||||
attrs[match.group(1)] = value
|
||||
return attrs
|
||||
|
||||
|
||||
def parse_tags(text: Optional[str]) -> list[dict]:
|
||||
"""提取全部标签(含嵌套内外层),返回 [{tag, attrs, text}],按开标签出现顺序。"""
|
||||
if not text:
|
||||
return []
|
||||
results: list[dict] = []
|
||||
stack: list[dict] = []
|
||||
token_re = re.compile(r"<[^>]+>")
|
||||
for token in token_re.finditer(text):
|
||||
raw_token = token.group(0)
|
||||
# 先按开/闭标签匹配
|
||||
open_match = _OPEN_RE.match(raw_token)
|
||||
close_match = _CLOSE_RE.match(raw_token)
|
||||
is_close_tag = raw_token.startswith("</")
|
||||
if not is_close_tag and open_match:
|
||||
is_self_close = open_match.group("self") == "/"
|
||||
node = {
|
||||
"tag": open_match.group("tag"),
|
||||
"attrs": parse_attributes(open_match.group("attrs")),
|
||||
"text": "",
|
||||
"_start": token.end(),
|
||||
}
|
||||
if is_self_close:
|
||||
node.pop("_start")
|
||||
results.append(node)
|
||||
else:
|
||||
stack.append(node)
|
||||
results.append(node)
|
||||
elif is_close_tag and close_match:
|
||||
tag = close_match.group("tag")
|
||||
# 弹出到最近同名开标签
|
||||
for idx in range(len(stack) - 1, -1, -1):
|
||||
if stack[idx]["tag"] == tag:
|
||||
node = stack[idx]
|
||||
node["text"] = unescape(text[node["_start"] : token.start()].strip())
|
||||
node.pop("_start", None)
|
||||
del stack[idx:]
|
||||
break
|
||||
# 未闭合标签:给剩余部分作为文本
|
||||
for node in stack:
|
||||
if "_start" in node:
|
||||
node["text"] = unescape(text[node["_start"] :].strip())
|
||||
node.pop("_start", None)
|
||||
return results
|
||||
|
||||
|
||||
def find_all(text: Optional[str], tag: str) -> list[dict]:
|
||||
"""提取指定标签的全部节点。"""
|
||||
return [n for n in parse_tags(text) if n["tag"] == tag]
|
||||
|
||||
|
||||
def find_first(text: Optional[str], tag: str) -> Optional[dict]:
|
||||
nodes = find_all(text, tag)
|
||||
return nodes[0] if nodes else None
|
||||
|
||||
|
||||
def text_of(text: Optional[str], tag: str, default: str = "") -> str:
|
||||
node = find_first(text, tag)
|
||||
return node["text"] if node else default
|
||||
|
||||
|
||||
def attr_bool(value: Optional[str], default: bool = False) -> bool:
|
||||
if value is None:
|
||||
return default
|
||||
return value.strip().lower() in {"true", "1", "yes", "是"}
|
||||
|
||||
|
||||
def attr_float(value: Optional[str], default: float = 0.0) -> float:
|
||||
try:
|
||||
return float(value) if value is not None and value.strip() else default
|
||||
except (TypeError, ValueError):
|
||||
return default
|
||||
|
||||
|
||||
def attr_int(value: Optional[str], default: int = 0) -> int:
|
||||
try:
|
||||
return int(float(value)) if value is not None and value.strip() else default
|
||||
except (TypeError, ValueError):
|
||||
return default
|
||||
+22
-8
@@ -90,18 +90,32 @@ class SharedSettings(BaseSettings):
|
||||
|
||||
# ── 豆包大模型(火山引擎方舟) ────────────────────────────────────────
|
||||
doubao_api_key: str = ""
|
||||
doubao_model: str = "doubao-seed-1-6-250615" # 推理模型(通用兜底)
|
||||
doubao_fast_model: str = "doubao-1-5-pro-32k-250115" # 快速结构化输出模型(编导脚本/意图解析/审核)
|
||||
doubao_model: str = "doubao-seed-2-1-pro-260915" # 推理模型(Seed 2.1 Pro,深度思考+多模态;原 seed-1-6 已下线)
|
||||
doubao_fast_model: str = (
|
||||
"doubao-seed-2-1-pro-260915" # #2181: lite方舟侧100%超时,默认fast_model也走pro;方舟恢复lite后通过ENV DOUBAO_FAST_MODEL切回
|
||||
)
|
||||
doubao_base_url: str = "https://ark.cn-beijing.volces.com/api/v3"
|
||||
doubao_timeout: int = 30
|
||||
doubao_max_retries: int = 2
|
||||
doubao_vision_model: str = "doubao-1-5-vision-pro-250328" # 高精度视觉(备用)
|
||||
doubao_vision_lite_model: str = "doubao-1-5-vision-lite-250315" # 快速视觉(商品识别默认,速度优先)
|
||||
doubao_vision_use_lite: bool = True # viral-video 图片分析默认用 lite 提速
|
||||
doubao_embedding_model: str = "doubao-embedding-large-text-240915"
|
||||
doubao_timeout: int = 45 # #2180: 方舟LLM高峰期响应6-8s,原30s太紧提到45s
|
||||
doubao_max_retries: int = 1 # #2180: timeout调大后一次调用就够,1次重试防偶发抖动;避免6次重试叠加到351s
|
||||
doubao_vision_model: str = (
|
||||
"doubao-seed-2-1-pro-260915" # 高精度视觉(Seed 2.1 Pro 原生多模态;原 vision-pro-250328 已下线)
|
||||
)
|
||||
doubao_vision_lite_model: str = (
|
||||
"doubao-seed-2-1-lite-260915" # 快速视觉(Seed 2.1 Lite 原生多模态;原 vision-lite-250315 不可用)
|
||||
)
|
||||
doubao_vision_use_lite: bool = False # #2181: lite视觉模型100%超时,默认关闭走pro(25-38s稳定返回)
|
||||
doubao_embedding_model: str = "doubao-embedding-vision-251215" # 多模态向量化(原 large-text-240915 已 Retiring)
|
||||
doubao_video_model: str = "doubao-seedance-2-5-260628"
|
||||
doubao_video_timeout: int = 600 # 视频生成轮询总超时(秒)
|
||||
doubao_video_poll_interval: int = 10 # 轮询间隔(秒)
|
||||
doubao_image_model: str = (
|
||||
"doubao-seedream-5-0-flash-260915" # #2173: 信任链 Seedream 改 flash 模型(实测 pro 46.5s→flash 13s;pro AI化图仍被Seedance拦截)
|
||||
)
|
||||
doubao_image_size: str = "1K" # #2173: 1K 已足够做 Seedance 参考图,2K 在 flash 下也 22s,1K 13s
|
||||
doubao_image_timeout: int = 60 # #2173: flash+1K 通常15s内,给60s余量
|
||||
doubao_trust_chain_enabled: bool = (
|
||||
True # #2173: 信任链总开关;若Seedream产物仍被Seedance拦截,可配 False 关闭直接t2v降级
|
||||
)
|
||||
|
||||
# ── DashScope (阿里云百炼 Wan 3.0 等) ─────────────────────────────────
|
||||
dashscope_api_key: str = ""
|
||||
|
||||
@@ -87,6 +87,9 @@ class ViralVideoJob:
|
||||
|
||||
user_id: str
|
||||
images: list[str] = field(default_factory=list)
|
||||
pre_trusted_images: list[str] | None = (
|
||||
None # #2172 信任链预热结果(Seedream AI 化后的 URL 列表),与 images 顺序对应
|
||||
)
|
||||
industry: str = ""
|
||||
target_customer: str = ""
|
||||
persona_id: str = ""
|
||||
|
||||
+323
-15
@@ -34,11 +34,11 @@ _HTTP_STATUS_ERROR = httpx.HTTPStatusError if hasattr(httpx, "HTTPStatusError")
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
# 视频模型 ID 解析逻辑(#2159 多模型支持)。
|
||||
# 内部使用简短别名(seedance-2.5 / seedance-2.0 / wan-3.0 等)做 PRICING key 和前端选择值;
|
||||
# 视频模型 ID 解析逻辑(#2159 多模型支持,#2170 方舟信任链统一走方舟)。
|
||||
# 内部使用简短别名(seedance-2.5 / seedance-2.0 / wan-3.0 等)做 PRICING key;
|
||||
# 实际调用时按 VIRAL_VIDEO_MODEL_CONFIG(domain/points_rules.py)的 provider/model_id 分发。
|
||||
# - provider=doubao → 火山方舟
|
||||
# - provider=dashscope → 阿里云 DashScope(Wan 系列)
|
||||
# - provider=doubao → 火山方舟 Seedance(含信任链真人 AI 化)
|
||||
# - provider=dashscope → 阿里云 DashScope(Wan 系列,可选)
|
||||
|
||||
|
||||
def _resolve_video_provider_and_id(model: str | None) -> tuple[str, str, dict]:
|
||||
@@ -180,8 +180,15 @@ class DoubaoClient:
|
||||
self.vision_model: str = settings.doubao_vision_model
|
||||
self.vision_lite_model: str = settings.doubao_vision_lite_model
|
||||
self.fast_model: str = settings.doubao_fast_model
|
||||
self.embedding_model: str = settings.doubao_embedding_model
|
||||
self.image_model: str = settings.doubao_image_model
|
||||
self.image_size: str = getattr(settings, "doubao_image_size", "1K") or "1K"
|
||||
self.image_timeout: int = getattr(settings, "doubao_image_timeout", 60) or 60
|
||||
self.trust_chain_enabled: bool = getattr(settings, "doubao_trust_chain_enabled", True)
|
||||
# 最近一次视频生成的详细错误(error_code + user_message + raw detail),供上层读取后展示给用户
|
||||
self.last_video_error: dict = {}
|
||||
# 最近一次图片生成的详细错误,供上层读取
|
||||
self.last_image_error: dict = {}
|
||||
|
||||
def embed_text(self, text: str, timeout: int | None = None) -> list[float] | None:
|
||||
"""调用豆包文本 Embedding API,返回浮点向量;失败返回 None。"""
|
||||
@@ -194,7 +201,7 @@ class DoubaoClient:
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
payload: dict[str, Any] = {
|
||||
"model": getattr(self, "embedding_model", None) or "doubao-embedding-large-text-240915",
|
||||
"model": self.embedding_model,
|
||||
"input": text.strip(),
|
||||
"encoding_format": "float",
|
||||
}
|
||||
@@ -235,6 +242,7 @@ class DoubaoClient:
|
||||
temperature: float = 0.7,
|
||||
max_tokens: int = 1024,
|
||||
model: str | None = None,
|
||||
timeout: int | None = None,
|
||||
) -> Optional[str]:
|
||||
"""调用 Chat Completion 接口.
|
||||
|
||||
@@ -262,32 +270,44 @@ class DoubaoClient:
|
||||
}
|
||||
|
||||
last_error: Optional[Exception] = None
|
||||
_t0 = time.time()
|
||||
for attempt in range(self.max_retries + 1):
|
||||
try:
|
||||
_req_timeout = timeout if timeout is not None else self.timeout
|
||||
response = httpx.post(
|
||||
url,
|
||||
headers=headers,
|
||||
json=payload,
|
||||
timeout=self.timeout,
|
||||
timeout=_req_timeout,
|
||||
)
|
||||
response.raise_for_status()
|
||||
data = response.json()
|
||||
content = data["choices"][0]["message"]["content"]
|
||||
_elapsed = time.time() - _t0
|
||||
logger.info(
|
||||
"[doubao] chat_completion 完成 model=%s tokens_in=%d tokens_out=%d elapsed=%.1fs attempt=%d timeout=%d",
|
||||
payload.get("model"),
|
||||
data.get("usage", {}).get("prompt_tokens", 0),
|
||||
data.get("usage", {}).get("completion_tokens", 0),
|
||||
_elapsed,
|
||||
attempt + 1,
|
||||
)
|
||||
return content.strip()
|
||||
except Exception as e:
|
||||
last_error = e
|
||||
if attempt < self.max_retries:
|
||||
wait = 0.5 * (2**attempt)
|
||||
logger.warning(
|
||||
"豆包API调用失败,%.1fs后重试 (第%d/%d次): %s",
|
||||
"豆包API调用失败,%.1fs后重试 (第%d/%d次, elapsed=%.1fs): %s",
|
||||
wait,
|
||||
attempt + 1,
|
||||
self.max_retries + 1,
|
||||
time.time() - _t0,
|
||||
e,
|
||||
)
|
||||
time.sleep(wait)
|
||||
|
||||
logger.error("豆包API调用最终失败: %s", last_error)
|
||||
logger.error("豆包API调用最终失败: elapsed=%.1fs err=%s", time.time() - _t0, last_error)
|
||||
return None
|
||||
|
||||
def vision_completion(
|
||||
@@ -356,6 +376,7 @@ class DoubaoClient:
|
||||
|
||||
req_timeout = timeout or self.timeout
|
||||
last_error: Optional[Exception] = None
|
||||
_t0 = time.time()
|
||||
for attempt in range(self.max_retries + 1):
|
||||
try:
|
||||
response = httpx.post(
|
||||
@@ -367,25 +388,97 @@ class DoubaoClient:
|
||||
response.raise_for_status()
|
||||
data = response.json()
|
||||
content = data["choices"][0]["message"]["content"]
|
||||
_elapsed = time.time() - _t0
|
||||
logger.info(
|
||||
"[doubao] vision_completion 完成 model=%s tokens_in=%d tokens_out=%d elapsed=%.1fs attempt=%d",
|
||||
payload.get("model"),
|
||||
data.get("usage", {}).get("prompt_tokens", 0),
|
||||
data.get("usage", {}).get("completion_tokens", 0),
|
||||
_elapsed,
|
||||
attempt + 1,
|
||||
)
|
||||
return content.strip()
|
||||
except Exception as e:
|
||||
last_error = e
|
||||
if attempt < self.max_retries:
|
||||
wait = 0.5 * (2**attempt)
|
||||
logger.warning(
|
||||
"豆包视觉API调用失败,%.1fs后重试 (第%d/%d次): %s",
|
||||
"豆包视觉API调用失败,%.1fs后重试 (第%d/%d次, elapsed=%.1fs): %s",
|
||||
wait,
|
||||
attempt + 1,
|
||||
self.max_retries + 1,
|
||||
time.time() - _t0,
|
||||
e,
|
||||
)
|
||||
time.sleep(wait)
|
||||
|
||||
logger.error("豆包视觉API调用最终失败: %s", last_error)
|
||||
logger.error("豆包视觉API调用最终失败: elapsed=%.1fs err=%s", time.time() - _t0, last_error)
|
||||
return None
|
||||
|
||||
# ── 视频生成(Seedance 2.5,异步任务)────────────────────────────
|
||||
|
||||
def preheat_trust_chain(
|
||||
self,
|
||||
portrait_descriptions: list[str],
|
||||
*,
|
||||
timeout: int | None = None,
|
||||
size: str | None = None,
|
||||
) -> list[str] | None:
|
||||
"""#2174 信任链预热(t2i 版):用 VLM 分析出的人物外貌描述,纯文生图生成 Seedream 人像,
|
||||
返回信任产物 URL 列表。
|
||||
|
||||
重要:方舟信任链规则是 Seedream 文生图(t2i)产物(不传 reference_images)才被 Seedance 信任;
|
||||
i2i(带用户照片 reference)产物不被信任,仍会被肖像审核拦截。
|
||||
|
||||
- portrait_descriptions: VLM 输出的 portrait_prompt 列表(中文描述),与原始图片顺序对应
|
||||
- 全部成功返回 list[str](顺序与输入一致)
|
||||
- 任何一张失败返回 None(上层会走现场兜底:直接 t2v 不带参考图)
|
||||
- 若 trust_chain_enabled=False 或所有描述均为"无人像",直接返回 None
|
||||
- 供 worker 在 VLM 分析完成后后台预热使用;video_generation 内部若收到 preheated 结果会直接使用。
|
||||
"""
|
||||
if not portrait_descriptions or not self.is_available or not getattr(self, "trust_chain_enabled", True):
|
||||
return None
|
||||
# 过滤掉"无人像"等无效描述,收集需要生成的索引
|
||||
_idx_map: list[int] = []
|
||||
_prompts: list[str] = []
|
||||
for i, desc in enumerate(portrait_descriptions):
|
||||
if not desc or "无人像" in desc or len(desc) < 10:
|
||||
continue
|
||||
_idx_map.append(i)
|
||||
# 把VLM的中文描述包装成适合Seedream t2i的英文+中文混合prompt,明确是写实半身人像
|
||||
_prompts.append(
|
||||
f"高清写实半身人像照片,{desc},自然光线,面部清晰居中,皮肤质感自然,"
|
||||
"高清摄影细节,构图居中,人物占据画面主体,背景柔和虚化。photorealistic portrait, "
|
||||
"sharp focus on face, soft natural lighting, high detail, half body shot."
|
||||
)
|
||||
if not _prompts:
|
||||
logger.info("[trust-chain][preheat] 无需生成信任人像(所有图片均无人像),跳过")
|
||||
return None
|
||||
trusted: list[str] = []
|
||||
_t0 = time.time()
|
||||
for idx, sd_prompt in enumerate(_prompts):
|
||||
# #2174 关键:t2i模式——不传reference_images,纯prompt文生图,产物才被Seedance信任
|
||||
sd_result = self.image_generation(
|
||||
prompt=sd_prompt,
|
||||
size=size or getattr(self, "image_size", "1K"),
|
||||
timeout=timeout or getattr(self, "image_timeout", 60),
|
||||
)
|
||||
if not sd_result:
|
||||
logger.warning(
|
||||
"[trust-chain][preheat] Seedream t2i 第 %d/%d 张失败: %s,预热整体失败",
|
||||
idx + 1,
|
||||
len(_prompts),
|
||||
getattr(self, "last_image_error", None),
|
||||
)
|
||||
return None
|
||||
trusted.append(sd_result["url"])
|
||||
logger.info(
|
||||
"[trust-chain][preheat] Seedream t2i 文生图预热完成 %d 张,总耗时 %.1fs",
|
||||
len(trusted),
|
||||
time.time() - _t0,
|
||||
)
|
||||
return trusted
|
||||
|
||||
def video_generation(
|
||||
self,
|
||||
prompt: str,
|
||||
@@ -401,6 +494,7 @@ class DoubaoClient:
|
||||
reference_images: list[str] | None = None,
|
||||
reference_audios: list[str] | None = None,
|
||||
reference_videos: list[str] | None = None,
|
||||
pre_trusted_images: list[str] | None = None,
|
||||
) -> dict | None:
|
||||
"""调用 Seedance 2.5 生视频(异步任务→轮询→下载)。
|
||||
|
||||
@@ -484,17 +578,55 @@ class DoubaoClient:
|
||||
ref_videos = [u for u in (reference_videos or [])[:3] if u and isinstance(u, str)]
|
||||
ref_imgs = [u for u in (reference_images or [])[:9] if u and isinstance(u, str)]
|
||||
|
||||
# 判断任务模式:有参考音/视/多图 → omni_reference(支持指定 ratio);纯首帧 → first_frame(ratio=adaptive)
|
||||
# ── #2170/#2172 方舟信任链(Trust Chain)────────────────────────────────
|
||||
# 真人照片直接传给 Seedance 会触发 50411 肖像审核拦截。
|
||||
# 解决:先通过同账号的 Seedream 5.0 Pro 图生图 AI 化(保持五官特征),
|
||||
# 得到的 AI 产物图属于"模型信任产物",再作为 reference_image 传给 Seedance 即可通过审核。
|
||||
# #2172/#2174: 信任链——使用预热好的 Seedream t2i 文生图(纯模型生成人像,是方舟信任产物,
|
||||
# 不会触发肖像审核)。预热在 VLM 分析后由 daemon 线程后台完成,结果通过 pre_trusted_images 传入。
|
||||
# - 预热结果有效 → 替换原参考图,走 omni_ref 模式
|
||||
# - 预热结果不可用 → 直接用原图(上层 _step_render 已现场同步跑 Seedream t2i 兜底;若再被400拦截,下方自动降级纯t2v)
|
||||
# 信任链只作用于 doubao provider;DashScope(Wan) 保持原行为。
|
||||
trust_chain_applied = False
|
||||
if provider == "doubao" and getattr(self, "trust_chain_enabled", True) and pre_trusted_images:
|
||||
raw_portrait_urls: list[str] = []
|
||||
if image_url:
|
||||
raw_portrait_urls.append(image_url)
|
||||
for u in ref_imgs:
|
||||
if u not in raw_portrait_urls:
|
||||
raw_portrait_urls.append(u)
|
||||
trusted_urls: list[str] = []
|
||||
if len(pre_trusted_images) >= 1:
|
||||
trusted_urls = list(pre_trusted_images)
|
||||
trust_chain_applied = True
|
||||
logger.info(
|
||||
"[trust-chain] 使用预热t2i结果 %d 张,替换原参考图走 reference_image 模式(原n=%d)",
|
||||
len(trusted_urls),
|
||||
len(raw_portrait_urls),
|
||||
)
|
||||
if trust_chain_applied and trusted_urls:
|
||||
# 替换:原 image_url 用第一张 AI 图,ref_imgs 用剩余
|
||||
if image_url and trusted_urls:
|
||||
image_url = trusted_urls[0]
|
||||
ref_imgs = trusted_urls[1:] if len(trusted_urls) > 1 else []
|
||||
else:
|
||||
ref_imgs = trusted_urls
|
||||
# ─────────────────────────────────────────────────────────────────
|
||||
|
||||
# 判断任务模式:
|
||||
# - 信任链强制走 reference_image(不是 first_frame;产品语义是人物参考,不是从图开始动)
|
||||
# - 有参考音/视/多图 → omni_reference(支持指定 ratio)
|
||||
# - 纯首帧无其他参考 → first_frame(ratio=adaptive)
|
||||
has_extra_refs = bool(ref_audios or ref_videos or ref_imgs)
|
||||
is_first_frame_mode = bool(image_url) and not has_extra_refs
|
||||
is_first_frame_mode = bool(image_url) and not has_extra_refs and not trust_chain_applied
|
||||
# 最终 ratio:first_frame 模式强制 adaptive,否则按用户传值(默认 9:16)
|
||||
final_ratio = "adaptive" if is_first_frame_mode else (ratio or "9:16")
|
||||
|
||||
# 构造 content 数组:text + 图 + 音 + 视
|
||||
content: list[dict[str, Any]] = [{"type": "text", "text": prompt.strip()}]
|
||||
if image_url:
|
||||
if has_extra_refs:
|
||||
# omni_reference:首张图作为 reference_image,允许指定 ratio
|
||||
if has_extra_refs or trust_chain_applied:
|
||||
# omni_reference 或信任链模式:首张图作为 reference_image,允许指定 ratio
|
||||
content.append(
|
||||
{
|
||||
"type": "image_url",
|
||||
@@ -538,7 +670,7 @@ class DoubaoClient:
|
||||
video_model,
|
||||
duration,
|
||||
final_ratio,
|
||||
"first_frame" if is_first_frame_mode else "omni_ref",
|
||||
"first_frame" if is_first_frame_mode else ("omni_ref+trust_chain" if trust_chain_applied else "omni_ref"),
|
||||
generate_audio,
|
||||
(1 if image_url else 0) + len(ref_imgs),
|
||||
len(ref_audios),
|
||||
@@ -631,6 +763,26 @@ class DoubaoClient:
|
||||
# 保留第二次的错误信息
|
||||
sc, body = sc2, body2
|
||||
|
||||
# #2183: portrait_intercept / 真人肖像审核拦截 → 去掉所有参考图(含image_url首帧),纯 t2v 重试一次
|
||||
# 信任链预热或现场 t2i 都失败时的最后兜底,保证能出片
|
||||
if not task_id and sc == 400:
|
||||
_err_code_for_400, _ = _classify_video_error(sc, body, last_err)
|
||||
if _err_code_for_400 == "portrait_intercept" and (image_url or ref_imgs):
|
||||
logger.warning(
|
||||
"Seedance 创建因 portrait_intercept 失败,降级纯 t2v(移除所有参考图)重试: img=%d ref=%d",
|
||||
1 if image_url else 0,
|
||||
len(ref_imgs),
|
||||
)
|
||||
_t2v_payload = dict(create_payload)
|
||||
_t2v_payload["content"] = [{"type": "text", "text": prompt.strip()}]
|
||||
_t2v_payload["ratio"] = ratio or "9:16"
|
||||
task_id, last_err, sc3, body3 = _do_create(_t2v_payload)
|
||||
if task_id:
|
||||
sc, body = sc3, body3
|
||||
logger.info("[trust-chain] portrait_intercept 降级纯 t2v 成功 task_id=%s", task_id)
|
||||
else:
|
||||
sc, body = sc3, body3
|
||||
|
||||
if not task_id:
|
||||
err_code, user_msg = _classify_video_error(sc, body, last_err)
|
||||
self.last_video_error = {
|
||||
@@ -793,6 +945,162 @@ class DoubaoClient:
|
||||
}
|
||||
return None
|
||||
|
||||
def image_generation(
|
||||
self,
|
||||
prompt: str,
|
||||
*,
|
||||
reference_images: list[str] | None = None,
|
||||
size: str = "2K",
|
||||
model: str | None = None,
|
||||
watermark: bool = False,
|
||||
output_format: str = "png",
|
||||
timeout: int | None = None,
|
||||
) -> dict | None:
|
||||
"""#2170: 调用方舟 Seedream 图片生成(文生图/图生图)。
|
||||
|
||||
- reference_images: 0~10 张参考图 URL;0 张 = 纯文生图;1 张 string/URL 直传;多张 list[str]。
|
||||
- 成功返回 {"url": str, "usage": dict | None};失败返回 None,错误写入 self.last_image_error。
|
||||
- 返回的 url 有时效性(通常 24h),应立即使用,不持久化存储。
|
||||
"""
|
||||
self.last_image_error = {}
|
||||
if not self.is_available:
|
||||
self.last_image_error = {
|
||||
"error_code": "auth_error",
|
||||
"user_message": "图片生成服务未配置(API Key 缺失),请联系管理员。",
|
||||
"detail": "DoubaoClient not available (api_key empty)",
|
||||
}
|
||||
return None
|
||||
if not prompt or not prompt.strip():
|
||||
self.last_image_error = {
|
||||
"error_code": "invalid_param",
|
||||
"user_message": "图片生成提示词不能为空。",
|
||||
"detail": "empty prompt",
|
||||
}
|
||||
return None
|
||||
|
||||
img_model = model or self.image_model
|
||||
_img_size = size or getattr(self, "image_size", "1K")
|
||||
url = f"{self.base_url}/images/generations"
|
||||
headers = {
|
||||
"Authorization": f"Bearer {self.api_key}",
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
payload: dict[str, Any] = {
|
||||
"model": img_model,
|
||||
"prompt": prompt.strip(),
|
||||
"size": _img_size,
|
||||
"response_format": "url",
|
||||
"output_format": output_format,
|
||||
"watermark": bool(watermark),
|
||||
}
|
||||
ref_imgs_local = [u for u in (reference_images or []) if u and isinstance(u, str)]
|
||||
if ref_imgs_local:
|
||||
if len(ref_imgs_local) == 1:
|
||||
payload["image"] = ref_imgs_local[0]
|
||||
else:
|
||||
payload["image"] = ref_imgs_local[:10]
|
||||
|
||||
req_timeout = timeout or getattr(self, "image_timeout", 60)
|
||||
last_err: Exception | None = None
|
||||
last_sc = 0
|
||||
last_body = ""
|
||||
_img_t0 = time.time()
|
||||
for attempt in range(self.max_retries + 1):
|
||||
try:
|
||||
resp = httpx.post(url, headers=headers, json=payload, timeout=req_timeout)
|
||||
last_sc = int(getattr(resp, "status_code", 0) or 0)
|
||||
last_body = (getattr(resp, "text", "") or "")[:2000]
|
||||
if last_sc >= 400:
|
||||
logger.error("Seedream 图片生成 HTTP %d: %s", last_sc, last_body[:500])
|
||||
try:
|
||||
resp.raise_for_status()
|
||||
except Exception as ee:
|
||||
last_err = ee
|
||||
if attempt < self.max_retries and last_sc >= 500:
|
||||
time.sleep(0.5 * (2**attempt))
|
||||
continue
|
||||
break
|
||||
data = resp.json()
|
||||
data_list = data.get("data") or []
|
||||
if data_list and isinstance(data_list, list):
|
||||
item = data_list[0]
|
||||
img_url = item.get("url")
|
||||
if img_url:
|
||||
logger.info(
|
||||
"Seedream 图片生成成功 model=%s ref_imgs=%d size=%s elapsed=%.1fs attempt=%d",
|
||||
img_model,
|
||||
len(ref_imgs_local),
|
||||
_img_size,
|
||||
time.time() - _img_t0,
|
||||
attempt + 1,
|
||||
)
|
||||
return {"url": img_url, "usage": data.get("usage")}
|
||||
last_err = RuntimeError(f"Seedream 返回结构异常: {str(data)[:300]}")
|
||||
break
|
||||
except _HTTP_NETWORK_ERRORS as ne:
|
||||
last_err = ne
|
||||
last_sc = 0
|
||||
last_body = f"network error: {ne}"
|
||||
logger.warning(
|
||||
"Seedream 网络异常 (%s),重试 %d/%d", type(ne).__name__, attempt + 1, self.max_retries + 1
|
||||
)
|
||||
if attempt < self.max_retries:
|
||||
time.sleep(0.5 * (2**attempt))
|
||||
continue
|
||||
break
|
||||
except Exception as e:
|
||||
last_err = e
|
||||
if attempt < self.max_retries and not isinstance(e, _HTTP_STATUS_ERROR):
|
||||
wait = 0.5 * (2**attempt)
|
||||
logger.warning(
|
||||
"Seedream 图片生成失败,%.1fs 后重试 (%d/%d): %s", wait, attempt + 1, self.max_retries + 1, e
|
||||
)
|
||||
time.sleep(wait)
|
||||
continue
|
||||
break
|
||||
|
||||
# 分类错误
|
||||
err_code = "unknown"
|
||||
user_msg = "图片生成失败,请稍后重试。"
|
||||
body_lower = (last_body or "").lower()
|
||||
if last_sc == 401 or last_sc == 403:
|
||||
err_code, user_msg = "auth_error", "图片生成服务鉴权失败,请联系管理员。"
|
||||
elif last_sc == 400:
|
||||
if any(k in body_lower for k in ("quota", "billing", "insufficient", "balance")):
|
||||
err_code, user_msg = "quota_exceeded", "图片生成配额不足或账号欠费,请联系管理员。"
|
||||
elif any(k in body_lower for k in ("rate", "throughput", "too many", "frequency")):
|
||||
err_code, user_msg = "rate_limit", "图片生成请求过于频繁,请稍后重试。"
|
||||
elif any(k in body_lower for k in ("sensitive", "porn", "terror", "risk", "audit", "content", "violat")):
|
||||
err_code, user_msg = "portrait_intercept", "参考素材未通过内容安全审核,请更换照片后重试。"
|
||||
else:
|
||||
err_code, user_msg = "invalid_param", f"图片生成参数错误:{last_body[:200]}"
|
||||
elif last_sc == 404:
|
||||
err_code, user_msg = "model_not_found", f"图片模型 {img_model} 不存在,请联系管理员。"
|
||||
elif last_sc >= 500:
|
||||
err_code, user_msg = "network_error", "图片生成服务暂时不可用,请稍后重试。"
|
||||
elif last_sc == 0:
|
||||
err_code, user_msg = "network_error", f"图片生成网络错误:{last_err!s}"[:200]
|
||||
self.last_image_error = {
|
||||
"error_code": err_code,
|
||||
"user_message": user_msg,
|
||||
"status_code": last_sc,
|
||||
"detail": (last_body or "")[:500] or (str(last_err) if last_err else ""),
|
||||
"model": img_model,
|
||||
}
|
||||
logger.error(
|
||||
"Seedream 图片生成最终失败: model=%s status=%d code=%s elapsed=%.1fs err=%s",
|
||||
img_model,
|
||||
last_sc,
|
||||
err_code,
|
||||
time.time() - _img_t0,
|
||||
last_err,
|
||||
)
|
||||
return None
|
||||
|
||||
def get_last_image_error(self) -> dict:
|
||||
"""返回最近一次 image_generation 失败的详细错误。空 dict 表示上次成功或未调用。"""
|
||||
return dict(self.last_image_error or {})
|
||||
|
||||
def get_last_video_error(self) -> dict:
|
||||
"""返回最近一次 video_generation 失败的详细错误。空 dict 表示上次成功或未调用。"""
|
||||
return dict(self.last_video_error or {})
|
||||
|
||||
@@ -502,6 +502,7 @@ def call_llm(
|
||||
max_tokens: int = 2048,
|
||||
model: str | None = None,
|
||||
system_prompt: str | None = None,
|
||||
timeout: int | None = None,
|
||||
) -> object:
|
||||
"""调用豆包大模型(文本对话),返回解析后的 JSON(dict/list)或原文字符串;失败返回 None。
|
||||
|
||||
@@ -521,7 +522,7 @@ def call_llm(
|
||||
{"role": "system", "content": system_prompt},
|
||||
{"role": "user", "content": prompt},
|
||||
]
|
||||
raw = client.chat_completion(messages, temperature=temperature, max_tokens=max_tokens, model=model)
|
||||
raw = client.chat_completion(messages, temperature=temperature, max_tokens=max_tokens, model=model, timeout=timeout)
|
||||
if raw is None:
|
||||
return None
|
||||
try:
|
||||
@@ -604,6 +605,24 @@ def call_vision(
|
||||
return raw
|
||||
|
||||
|
||||
def preheat_trust_chain(portrait_descriptions: list[str], *, timeout: int = 120) -> list[str] | None:
|
||||
"""#2174 信任链预热(t2i版):用 VLM 分析出的人物外貌描述,跑 Seedream 文生图,
|
||||
生成的信任产物 URL 可传给 call_video_generation(pre_trusted_images=...)。
|
||||
|
||||
- portrait_descriptions: VLM输出的portrait_prompt列表(中文人物外貌描述)
|
||||
- 成功返回与输入同序的信任图URL列表;任意一张失败返回None(调用方回退到纯t2v)
|
||||
- 必须传VLM人物描述,不传reference_images,走纯t2i路径才是方舟信任产物
|
||||
"""
|
||||
client = get_doubao_client()
|
||||
if not client.is_available:
|
||||
return None
|
||||
try:
|
||||
return client.preheat_trust_chain(portrait_descriptions, timeout=timeout)
|
||||
except Exception as e:
|
||||
logger.error("[ai_service] preheat_trust_chain 异常: %s", e, exc_info=True)
|
||||
return None
|
||||
|
||||
|
||||
def call_video_generation(
|
||||
prompt: str,
|
||||
*,
|
||||
@@ -617,8 +636,9 @@ def call_video_generation(
|
||||
reference_images: list[str] | None = None,
|
||||
reference_audios: list[str] | None = None,
|
||||
reference_videos: list[str] | None = None,
|
||||
pre_trusted_images: list[str] | None = None,
|
||||
) -> dict | None:
|
||||
"""调用 Seedance / Wan 视频生成(v1.6.2 多模型版)。
|
||||
"""调用 Seedance / Wan 视频生成(v1.6.2 多模型版 + #2172 信任链预热)。
|
||||
|
||||
成功返回 {"video_path": str, "usage": dict | None}(usage 含 completion_tokens),失败返回 None。
|
||||
失败时错误详情会写入 client.last_video_error,可通过 get_last_video_error() 读取:
|
||||
@@ -650,6 +670,7 @@ def call_video_generation(
|
||||
reference_images=reference_images,
|
||||
reference_audios=reference_audios,
|
||||
reference_videos=reference_videos,
|
||||
pre_trusted_images=pre_trusted_images,
|
||||
)
|
||||
if effective_ratio:
|
||||
kwargs["ratio"] = effective_ratio
|
||||
|
||||
@@ -57,7 +57,7 @@ if [ "$TARGET_ENV" = "staging" ]; then
|
||||
fi
|
||||
|
||||
# 共用 secrets 直接导出(如果存在)
|
||||
SHARED_SECRETS="OSS_ACCESS_KEY_ID OSS_ACCESS_KEY_SECRET COSYVOICE_API_KEY DASHSCOPE_API_KEY MEDIAKIT_API_KEY DOUBAO_API_KEY DOUBAO_MODEL DOUBAO_FAST_MODEL DOUBAO_BASE_URL DOUBAO_VISION_MODEL DOUBAO_VISION_LITE_MODEL DOUBAO_VISION_USE_LITE WECHAT_APP_ID WECHAT_APP_SECRET TIKHUB_API_KEY APIZERO_API_KEY GPU_WORKER_TOKEN"
|
||||
SHARED_SECRETS="OSS_ACCESS_KEY_ID OSS_ACCESS_KEY_SECRET COSYVOICE_API_KEY DASHSCOPE_API_KEY MEDIAKIT_API_KEY DOUBAO_API_KEY DOUBAO_MODEL DOUBAO_FAST_MODEL DOUBAO_BASE_URL DOUBAO_VISION_MODEL DOUBAO_VISION_LITE_MODEL DOUBAO_VISION_USE_LITE DOUBAO_IMAGE_MODEL DOUBAO_IMAGE_SIZE DOUBAO_IMAGE_TIMEOUT DOUBAO_FAST_MODEL DOUBAO_TIMEOUT DOUBAO_MAX_RETRIES WECHAT_APP_ID WECHAT_APP_SECRET TIKHUB_API_KEY APIZERO_API_KEY GPU_WORKER_TOKEN"
|
||||
for var in $SHARED_SECRETS; do
|
||||
value="${!var:-}"
|
||||
# 已经在环境中了,无需额外操作
|
||||
|
||||
@@ -0,0 +1,80 @@
|
||||
"""爆款视频 5 套 Prompt 模板种子脚本(#2040)。
|
||||
|
||||
幂等:以 (prompt_type, version) 为唯一键,存在则更新(UPSERT),重复执行结果一致。
|
||||
用法:
|
||||
python scripts/seed_viral_video_prompts.py # 自动用应用配置连库
|
||||
DATABASE_URL=postgresql+psycopg2://... python scripts/seed_viral_video_prompts.py
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import sys
|
||||
|
||||
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
|
||||
import sqlalchemy as sa # noqa: E402
|
||||
|
||||
from packages.application.viral_video.prompts import DEFAULT_TEMPLATES # noqa: E402
|
||||
|
||||
|
||||
def _engine():
|
||||
database_url = os.environ.get("DATABASE_URL")
|
||||
if database_url:
|
||||
return sa.create_engine(database_url)
|
||||
# 复用应用自身配置
|
||||
from packages.config import get_shared_settings
|
||||
|
||||
url = str(get_shared_settings().database_url)
|
||||
return sa.create_engine(url.replace("postgresql+asyncpg://", "postgresql+psycopg2://"))
|
||||
|
||||
|
||||
UPSERT_SQL = sa.text("""
|
||||
INSERT INTO viral_video_prompt_templates
|
||||
(name, prompt_type, version, system_prompt, user_prompt_template,
|
||||
example_output, is_active, updated_at)
|
||||
VALUES
|
||||
(:name, :prompt_type, :version, :system_prompt, :user_prompt_template,
|
||||
:example_output, TRUE, :now_ts)
|
||||
ON CONFLICT (prompt_type, version) DO UPDATE SET
|
||||
name = EXCLUDED.name,
|
||||
system_prompt = EXCLUDED.system_prompt,
|
||||
user_prompt_template = EXCLUDED.user_prompt_template,
|
||||
example_output = EXCLUDED.example_output,
|
||||
is_active = TRUE,
|
||||
updated_at = :now_ts
|
||||
""")
|
||||
|
||||
|
||||
def seed(engine) -> int:
|
||||
count = 0
|
||||
from datetime import datetime, timezone
|
||||
|
||||
now_ts = datetime.now(timezone.utc)
|
||||
with engine.begin() as conn:
|
||||
for item in DEFAULT_TEMPLATES:
|
||||
conn.execute(
|
||||
UPSERT_SQL,
|
||||
{
|
||||
"name": item["name"],
|
||||
"prompt_type": item["prompt_type"],
|
||||
"version": item["version"],
|
||||
"system_prompt": item["system_prompt"],
|
||||
"user_prompt_template": item["user_prompt_template"],
|
||||
"example_output": item["example_output"],
|
||||
"now_ts": now_ts,
|
||||
},
|
||||
)
|
||||
count += 1
|
||||
return count
|
||||
|
||||
|
||||
def main() -> int:
|
||||
engine = _engine()
|
||||
count = seed(engine)
|
||||
print(f"seed 完成:{count} 套模板已写入/更新(image_analysis/intent_parsing/copy_fusion/storyboard/review)")
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main())
|
||||
@@ -1,4 +1,5 @@
|
||||
"""Additional unit tests to hit uncovered lines for diff-coverage >=60%."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
@@ -13,15 +14,20 @@ from packages.shared.ai_client import DoubaoClient
|
||||
|
||||
class _FakeSettings:
|
||||
doubao_api_key = "test-key"
|
||||
doubao_model = "test-model"
|
||||
doubao_fast_model = "test-fast-model"
|
||||
doubao_model = "doubao-seed-2-1-pro-260915"
|
||||
doubao_fast_model = "doubao-seed-2-1-lite-260915"
|
||||
doubao_base_url = "https://ark.cn-beijing.volces.com/api/v3"
|
||||
doubao_timeout = 10
|
||||
doubao_max_retries = 0
|
||||
doubao_vision_model = "test-vision"
|
||||
doubao_vision_lite_model = "test-vision-lite"
|
||||
doubao_vision_model = "doubao-seed-2-1-pro-260915"
|
||||
doubao_vision_lite_model = "doubao-seed-2-1-lite-260915"
|
||||
doubao_vision_use_lite = False
|
||||
doubao_embedding_model = "test-embedding"
|
||||
doubao_embedding_model = "doubao-embedding-vision-251215"
|
||||
doubao_video_model = "doubao-seedance-2-5-260628"
|
||||
doubao_video_timeout = 480
|
||||
doubao_video_poll_interval = 10
|
||||
doubao_image_model = "doubao-seedream-5-0-pro-260628"
|
||||
doubao_image_timeout = 120
|
||||
|
||||
|
||||
def _make_client(api_key: str = "test-key") -> DoubaoClient:
|
||||
@@ -129,26 +135,50 @@ from packages.domain.atom_clip_tagger import parse_vision_response
|
||||
|
||||
class TestParseVisionResponseEdgeCases:
|
||||
def test_person_count_type_error_defaults_zero(self):
|
||||
text = json.dumps({
|
||||
"scene": [], "objects": [], "action": [], "shot": "", "has_text": False,
|
||||
"person_count": "not-an-int", "text_content": "", "caption": "x",
|
||||
})
|
||||
text = json.dumps(
|
||||
{
|
||||
"scene": [],
|
||||
"objects": [],
|
||||
"action": [],
|
||||
"shot": "",
|
||||
"has_text": False,
|
||||
"person_count": "not-an-int",
|
||||
"text_content": "",
|
||||
"caption": "x",
|
||||
}
|
||||
)
|
||||
r = parse_vision_response(text)
|
||||
assert r["person_count"] == 0
|
||||
|
||||
def test_person_count_out_of_range_clamped(self):
|
||||
text = json.dumps({
|
||||
"scene": [], "objects": [], "action": [], "shot": "", "has_text": False,
|
||||
"person_count": 10, "text_content": "", "caption": "x",
|
||||
})
|
||||
text = json.dumps(
|
||||
{
|
||||
"scene": [],
|
||||
"objects": [],
|
||||
"action": [],
|
||||
"shot": "",
|
||||
"has_text": False,
|
||||
"person_count": 10,
|
||||
"text_content": "",
|
||||
"caption": "x",
|
||||
}
|
||||
)
|
||||
r = parse_vision_response(text)
|
||||
assert r["person_count"] == 3
|
||||
|
||||
def test_person_count_negative_clamped(self):
|
||||
text = json.dumps({
|
||||
"scene": [], "objects": [], "action": [], "shot": "", "has_text": False,
|
||||
"person_count": -5, "text_content": "", "caption": "x",
|
||||
})
|
||||
text = json.dumps(
|
||||
{
|
||||
"scene": [],
|
||||
"objects": [],
|
||||
"action": [],
|
||||
"shot": "",
|
||||
"has_text": False,
|
||||
"person_count": -5,
|
||||
"text_content": "",
|
||||
"caption": "x",
|
||||
}
|
||||
)
|
||||
r = parse_vision_response(text)
|
||||
assert r["person_count"] == 0
|
||||
|
||||
@@ -159,10 +189,18 @@ class TestParseVisionResponseEdgeCases:
|
||||
|
||||
def test_caption_truncation_at_80(self):
|
||||
long_caption = "描" * 100
|
||||
text = json.dumps({
|
||||
"scene": [], "objects": [], "action": [], "shot": "", "has_text": False,
|
||||
"person_count": 0, "text_content": "", "caption": long_caption,
|
||||
})
|
||||
text = json.dumps(
|
||||
{
|
||||
"scene": [],
|
||||
"objects": [],
|
||||
"action": [],
|
||||
"shot": "",
|
||||
"has_text": False,
|
||||
"person_count": 0,
|
||||
"text_content": "",
|
||||
"caption": long_caption,
|
||||
}
|
||||
)
|
||||
r = parse_vision_response(text)
|
||||
assert len(r["caption"]) == 80
|
||||
|
||||
@@ -196,9 +234,7 @@ class TestNarrativeMatchNonDictClipTags:
|
||||
def test_non_dict_clip_tags_are_skipped(self):
|
||||
a1 = _FA("a1", tags=[])
|
||||
clip_map = {"a1": [None, "bad", {"scene": ["工厂"], "objects": [], "action": []}, 123]}
|
||||
matched, unmatched = match_assets_by_script_tags(
|
||||
[a1], script_tags=["工厂"], clip_ai_tags_by_asset=clip_map
|
||||
)
|
||||
matched, unmatched = match_assets_by_script_tags([a1], script_tags=["工厂"], clip_ai_tags_by_asset=clip_map)
|
||||
assert [a.id for a in matched] == ["a1"]
|
||||
|
||||
|
||||
@@ -231,6 +267,7 @@ class _FQuery:
|
||||
class TestUpdateCaptionEmbedding:
|
||||
def _make_repo(self, session):
|
||||
from packages.adapters.sqlalchemy_impl.asset_atom_clip_repository import SQLAlchemyAssetAtomClipRepository
|
||||
|
||||
repo = SQLAlchemyAssetAtomClipRepository.__new__(SQLAlchemyAssetAtomClipRepository)
|
||||
repo.session = session
|
||||
return repo
|
||||
|
||||
@@ -17,7 +17,7 @@ def mock_settings():
|
||||
doubao_api_key="test-api-key",
|
||||
doubao_model="doubao-pro-32k",
|
||||
doubao_base_url="https://ark.example.com/api/v3",
|
||||
doubao_timeout=30,
|
||||
doubao_timeout=45,
|
||||
doubao_max_retries=2,
|
||||
)
|
||||
yield mock
|
||||
@@ -37,7 +37,7 @@ def client_without_key():
|
||||
doubao_api_key="",
|
||||
doubao_model="doubao-pro-32k",
|
||||
doubao_base_url="https://ark.example.com/api/v3",
|
||||
doubao_timeout=30,
|
||||
doubao_timeout=45,
|
||||
doubao_max_retries=2,
|
||||
)
|
||||
yield DoubaoClient()
|
||||
@@ -52,7 +52,7 @@ class TestDoubaoClientInit:
|
||||
assert client.api_key == "test-api-key"
|
||||
assert client.model == "doubao-pro-32k"
|
||||
assert client.base_url == "https://ark.example.com/api/v3"
|
||||
assert client.timeout == 30
|
||||
assert client.timeout == 45
|
||||
assert client.max_retries == 2
|
||||
|
||||
def test_base_url_strips_trailing_slash(self, mock_settings):
|
||||
|
||||
@@ -0,0 +1,635 @@
|
||||
"""#2170 Seedream 图片生成 + 方舟信任链单测。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import httpx
|
||||
|
||||
from packages.shared.ai_client import DoubaoClient
|
||||
|
||||
|
||||
def _make_client(**overrides):
|
||||
client = DoubaoClient.__new__(DoubaoClient)
|
||||
client.api_key = overrides.get("api_key", "test-key")
|
||||
client.base_url = overrides.get("base_url", "https://ark.cn-beijing.volces.com/api/v3")
|
||||
client.model = "doubao-model"
|
||||
client.vision_model = "doubao-vision"
|
||||
client.embedding_model = "doubao-embedding"
|
||||
client.image_model = overrides.get("image_model", "doubao-seedream-5-0-pro-260628")
|
||||
client.image_timeout = overrides.get("image_timeout", 120)
|
||||
client.timeout = overrides.get("timeout", 30)
|
||||
client.max_retries = overrides.get("max_retries", 0)
|
||||
client.last_video_error = {}
|
||||
client.last_image_error = {}
|
||||
return client
|
||||
|
||||
|
||||
def _fake_time(base=1000.0, stable_calls=50, big=9e9):
|
||||
"""返回 time.time 替身:前 stable_calls 次返回 base+i,之后返回 big+i。
|
||||
|
||||
Python 3.12 logging.LogRecord.__init__ 内部会调 time.time(),
|
||||
用有限 iter 会 StopIteration,因此必须用无限生成器。
|
||||
"""
|
||||
state = {"n": 0}
|
||||
|
||||
def _t():
|
||||
n = state["n"]
|
||||
state["n"] += 1
|
||||
if n < stable_calls:
|
||||
return base + n
|
||||
return big + n
|
||||
|
||||
return _t
|
||||
|
||||
|
||||
# ── Seedream 图片生成单测 ──────────────────────────────────────────
|
||||
|
||||
|
||||
class TestImageGenerationHappyPath:
|
||||
def test_returns_none_when_no_api_key(self):
|
||||
client = _make_client(api_key="")
|
||||
assert client.image_generation("p") is None
|
||||
err = client.get_last_image_error()
|
||||
assert err["error_code"] == "auth_error"
|
||||
|
||||
def test_returns_none_on_empty_prompt(self):
|
||||
client = _make_client()
|
||||
assert client.image_generation(" ") is None
|
||||
err = client.get_last_image_error()
|
||||
assert err["error_code"] == "invalid_param"
|
||||
|
||||
def test_text_to_image_success(self):
|
||||
client = _make_client()
|
||||
captured = {}
|
||||
ok_resp = MagicMock()
|
||||
ok_resp.status_code = 200
|
||||
ok_resp.json.return_value = {"data": [{"url": "https://cdn.example.com/i.png"}], "usage": {"tokens": 1}}
|
||||
ok_resp.raise_for_status = MagicMock()
|
||||
ok_resp.text = ""
|
||||
|
||||
def fake_post(url, **kwargs):
|
||||
captured["url"] = url
|
||||
captured["json"] = kwargs.get("json")
|
||||
return ok_resp
|
||||
|
||||
with patch("packages.shared.ai_client.httpx.post", side_effect=fake_post):
|
||||
result = client.image_generation("一只可爱的猫", size="1K")
|
||||
assert result is not None
|
||||
assert result["url"] == "https://cdn.example.com/i.png"
|
||||
assert "/images/generations" in captured["url"]
|
||||
assert captured["json"]["model"] == "doubao-seedream-5-0-pro-260628"
|
||||
assert captured["json"]["size"] == "1K"
|
||||
assert "image" not in captured["json"]
|
||||
|
||||
def test_image_to_image_single_ref_passed_as_string(self):
|
||||
client = _make_client()
|
||||
captured = {}
|
||||
ok_resp = MagicMock(status_code=200)
|
||||
ok_resp.json.return_value = {"data": [{"url": "https://cdn.example.com/out.png"}]}
|
||||
ok_resp.raise_for_status = MagicMock()
|
||||
ok_resp.text = ""
|
||||
|
||||
def fake_post(url, **kwargs):
|
||||
captured["json"] = kwargs.get("json")
|
||||
return ok_resp
|
||||
|
||||
with patch("packages.shared.ai_client.httpx.post", side_effect=fake_post):
|
||||
client.image_generation("保持五官", reference_images=["https://img/x.jpg"])
|
||||
assert captured["json"]["image"] == "https://img/x.jpg"
|
||||
|
||||
def test_image_to_image_multiple_refs_passed_as_list(self):
|
||||
client = _make_client()
|
||||
captured = {}
|
||||
ok_resp = MagicMock(status_code=200)
|
||||
ok_resp.json.return_value = {"data": [{"url": "https://cdn.example.com/out.png"}]}
|
||||
ok_resp.raise_for_status = MagicMock()
|
||||
|
||||
def fake_post(url, **kwargs):
|
||||
captured["json"] = kwargs.get("json")
|
||||
return ok_resp
|
||||
|
||||
refs = [f"https://img/{i}.jpg" for i in range(3)]
|
||||
with patch("packages.shared.ai_client.httpx.post", side_effect=fake_post):
|
||||
client.image_generation("保持", reference_images=refs)
|
||||
assert captured["json"]["image"] == refs
|
||||
|
||||
def test_400_sensitive_returns_portrait_intercept(self):
|
||||
client = _make_client(max_retries=0)
|
||||
bad_resp = MagicMock(status_code=400)
|
||||
bad_resp.text = '{"error":{"code":"ContentRisk","message":"sensitive content detected"}}'
|
||||
bad_resp.json.return_value = {"error": {"code": "ContentRisk"}}
|
||||
bad_resp.raise_for_status.side_effect = httpx.HTTPStatusError("bad", request=MagicMock(), response=bad_resp)
|
||||
with patch("packages.shared.ai_client.httpx.post", return_value=bad_resp):
|
||||
assert client.image_generation("p", reference_images=["https://img/x.jpg"]) is None
|
||||
err = client.get_last_image_error()
|
||||
assert err["error_code"] == "portrait_intercept"
|
||||
|
||||
def test_500_retries_then_fails(self):
|
||||
client = _make_client(max_retries=1)
|
||||
bad_resp = MagicMock(status_code=500)
|
||||
bad_resp.text = "internal error"
|
||||
bad_resp.json.return_value = {"error": {"message": "internal"}}
|
||||
bad_resp.raise_for_status.side_effect = httpx.HTTPStatusError("500", request=MagicMock(), response=bad_resp)
|
||||
with (
|
||||
patch("packages.shared.ai_client.httpx.post", return_value=bad_resp) as mp,
|
||||
patch("packages.shared.ai_client.time.sleep", return_value=None),
|
||||
):
|
||||
assert client.image_generation("p") is None
|
||||
assert mp.call_count == 2
|
||||
err = client.get_last_image_error()
|
||||
assert err["error_code"] == "network_error"
|
||||
|
||||
|
||||
# ── 信任链集成单测 ────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestTrustChainIntegration:
|
||||
def test_pre_trusted_images_replace_original_refs(self, tmp_path):
|
||||
"""#2174: 传 pre_trusted_images(预热好的t2i信任图)时,替换原image_url/ref_imgs发给Seedance,不再现场跑Seedream。"""
|
||||
client = _make_client()
|
||||
captured_calls = []
|
||||
|
||||
task_ok = MagicMock(status_code=200)
|
||||
task_ok.json.return_value = {"id": "t-trust"}
|
||||
task_ok.raise_for_status = MagicMock()
|
||||
task_ok.text = ""
|
||||
|
||||
poll_ok = MagicMock(status_code=200)
|
||||
poll_ok.json.return_value = {"status": "succeeded", "content": {"video_url": "https://cdn.example.com/v.mp4"}}
|
||||
poll_ok.raise_for_status = MagicMock()
|
||||
|
||||
class FakeStream:
|
||||
def __init__(self):
|
||||
self._c = [b"OK"]
|
||||
self._it = iter(self._c)
|
||||
|
||||
def __enter__(self):
|
||||
return self
|
||||
|
||||
def __exit__(self, *a):
|
||||
return False
|
||||
|
||||
def raise_for_status(self):
|
||||
return None
|
||||
|
||||
def iter_bytes(self, chunk_size=None):
|
||||
return self._it
|
||||
|
||||
def fake_post(url, **kwargs):
|
||||
captured_calls.append({"url": url, "json": kwargs.get("json")})
|
||||
return task_ok
|
||||
|
||||
fake_uuid = MagicMock()
|
||||
fake_uuid.hex = "00000001"
|
||||
with (
|
||||
patch("packages.shared.ai_client.httpx.post", side_effect=fake_post),
|
||||
patch("packages.shared.ai_client.httpx.get", return_value=poll_ok),
|
||||
patch("packages.shared.ai_client.httpx.stream", return_value=FakeStream()),
|
||||
patch("packages.shared.ai_client.time.sleep", return_value=None),
|
||||
patch("packages.shared.ai_client.time.time", side_effect=_fake_time(stable_calls=2)),
|
||||
patch("packages.shared.ai_client.uuid.uuid4", return_value=fake_uuid),
|
||||
patch("packages.shared.ai_client.get_shared_settings") as ms,
|
||||
):
|
||||
ms.return_value = MagicMock(
|
||||
doubao_video_poll_interval=0,
|
||||
doubao_video_timeout=60,
|
||||
doubao_video_model="doubao-seedance-2-5-260628",
|
||||
)
|
||||
out = client.video_generation(
|
||||
"人物在海边散步",
|
||||
image_url="https://img/raw.jpg",
|
||||
pre_trusted_images=["https://ai.example.com/trusted.png"],
|
||||
duration=5,
|
||||
ratio="9:16",
|
||||
resolution="720p",
|
||||
output_dir=str(tmp_path),
|
||||
)
|
||||
assert out is not None
|
||||
# 只有一次 Seedance 调用(不现场跑Seedream)
|
||||
assert len(captured_calls) == 1
|
||||
assert "/contents/generations/tasks" in captured_calls[0]["url"]
|
||||
seedance_payload = captured_calls[0]["json"]
|
||||
content = seedance_payload["content"]
|
||||
img_items = [c for c in content if c.get("type") == "image_url"]
|
||||
assert len(img_items) == 1
|
||||
assert img_items[0]["image_url"]["url"] == "https://ai.example.com/trusted.png"
|
||||
assert img_items[0]["role"] == "reference_image"
|
||||
assert seedance_payload["ratio"] == "9:16"
|
||||
|
||||
def test_no_reference_image_skips_seedream(self, tmp_path):
|
||||
client = _make_client()
|
||||
captured_calls = []
|
||||
|
||||
task_ok = MagicMock(status_code=200)
|
||||
task_ok.json.return_value = {"id": "t-t2v"}
|
||||
task_ok.raise_for_status = MagicMock()
|
||||
poll_ok = MagicMock(status_code=200)
|
||||
poll_ok.json.return_value = {"status": "succeeded", "content": {"video_url": "https://cdn.example.com/v.mp4"}}
|
||||
poll_ok.raise_for_status = MagicMock()
|
||||
|
||||
class FakeStream:
|
||||
def __init__(self):
|
||||
self._it = iter([b"OK"])
|
||||
|
||||
def __enter__(self):
|
||||
return self
|
||||
|
||||
def __exit__(self, *a):
|
||||
return False
|
||||
|
||||
def raise_for_status(self):
|
||||
return None
|
||||
|
||||
def iter_bytes(self, chunk_size=None):
|
||||
return self._it
|
||||
|
||||
def fake_post(url, **kwargs):
|
||||
captured_calls.append({"url": url, "json": kwargs.get("json")})
|
||||
return task_ok
|
||||
|
||||
fake_uuid = MagicMock()
|
||||
fake_uuid.hex = "00000002"
|
||||
with (
|
||||
patch("packages.shared.ai_client.httpx.post", side_effect=fake_post),
|
||||
patch("packages.shared.ai_client.httpx.get", return_value=poll_ok),
|
||||
patch("packages.shared.ai_client.httpx.stream", return_value=FakeStream()),
|
||||
patch("packages.shared.ai_client.time.sleep", return_value=None),
|
||||
patch("packages.shared.ai_client.time.time", side_effect=_fake_time(stable_calls=2)),
|
||||
patch("packages.shared.ai_client.uuid.uuid4", return_value=fake_uuid),
|
||||
patch("packages.shared.ai_client.get_shared_settings") as ms,
|
||||
):
|
||||
ms.return_value = MagicMock(
|
||||
doubao_video_poll_interval=0,
|
||||
doubao_video_timeout=60,
|
||||
doubao_video_model="doubao-seedance-2-5-260628",
|
||||
)
|
||||
out = client.video_generation("海边日落", duration=5, ratio="9:16", output_dir=str(tmp_path))
|
||||
assert out is not None
|
||||
assert len(captured_calls) == 1
|
||||
assert "/contents/generations/tasks" in captured_calls[0]["url"]
|
||||
content = captured_calls[0]["json"]["content"]
|
||||
assert all(c.get("type") != "image_url" for c in content)
|
||||
|
||||
def test_no_preheated_images_falls_back_to_first_frame_mode(self, tmp_path):
|
||||
"""#2174: 无 pre_trusted_images 时原图走 first_frame 模式(ratio=adaptive),不再现场跑 Seedream。"""
|
||||
client = _make_client(max_retries=0)
|
||||
captured_calls = []
|
||||
|
||||
task_ok = MagicMock(status_code=200)
|
||||
task_ok.json.return_value = {"id": "t-fb"}
|
||||
task_ok.raise_for_status = MagicMock()
|
||||
poll_ok = MagicMock(status_code=200)
|
||||
poll_ok.json.return_value = {"status": "succeeded", "content": {"video_url": "https://cdn.example.com/v.mp4"}}
|
||||
poll_ok.raise_for_status = MagicMock()
|
||||
|
||||
class FakeStream:
|
||||
def __init__(self):
|
||||
self._it = iter([b"OK"])
|
||||
|
||||
def __enter__(self):
|
||||
return self
|
||||
|
||||
def __exit__(self, *a):
|
||||
return False
|
||||
|
||||
def raise_for_status(self):
|
||||
return None
|
||||
|
||||
def iter_bytes(self, chunk_size=None):
|
||||
return self._it
|
||||
|
||||
def fake_post(url, **kwargs):
|
||||
captured_calls.append({"url": url, "json": kwargs.get("json")})
|
||||
return task_ok
|
||||
|
||||
fake_uuid = MagicMock()
|
||||
fake_uuid.hex = "00000003"
|
||||
with (
|
||||
patch("packages.shared.ai_client.httpx.post", side_effect=fake_post),
|
||||
patch("packages.shared.ai_client.httpx.get", return_value=poll_ok),
|
||||
patch("packages.shared.ai_client.httpx.stream", return_value=FakeStream()),
|
||||
patch("packages.shared.ai_client.time.sleep", return_value=None),
|
||||
patch("packages.shared.ai_client.time.time", side_effect=_fake_time(stable_calls=2)),
|
||||
patch("packages.shared.ai_client.uuid.uuid4", return_value=fake_uuid),
|
||||
patch("packages.shared.ai_client.get_shared_settings") as ms,
|
||||
):
|
||||
ms.return_value = MagicMock(
|
||||
doubao_video_poll_interval=0,
|
||||
doubao_video_timeout=60,
|
||||
doubao_video_model="doubao-seedance-2-5-260628",
|
||||
)
|
||||
out = client.video_generation(
|
||||
"海边散步",
|
||||
image_url="https://img/raw.jpg",
|
||||
duration=5,
|
||||
ratio="9:16",
|
||||
output_dir=str(tmp_path),
|
||||
)
|
||||
assert out is not None
|
||||
# 不现场跑 Seedream,只有一次 Seedance 调用
|
||||
assert len(captured_calls) == 1
|
||||
seedance_payload = captured_calls[0]["json"]
|
||||
content = seedance_payload["content"]
|
||||
img_items = [c for c in content if c.get("type") == "image_url"]
|
||||
assert len(img_items) == 1
|
||||
assert img_items[0]["image_url"]["url"] == "https://img/raw.jpg"
|
||||
assert img_items[0]["role"] == "first_frame"
|
||||
assert seedance_payload["ratio"] == "adaptive"
|
||||
|
||||
|
||||
# ── image_generation 补充分支覆盖 ─────────────────────────────────
|
||||
|
||||
|
||||
class TestImageGenerationBranches:
|
||||
"""覆盖 image_generation 的错误分类/重试/结构异常等分支。"""
|
||||
|
||||
def test_401_returns_auth_error(self):
|
||||
client = _make_client(max_retries=0)
|
||||
r = MagicMock(status_code=401, text='{"error":{}}')
|
||||
r.json.return_value = {"error": {}}
|
||||
r.raise_for_status.side_effect = httpx.HTTPStatusError("a", request=MagicMock(), response=r)
|
||||
with patch("packages.shared.ai_client.httpx.post", return_value=r):
|
||||
assert client.image_generation("p") is None
|
||||
assert client.get_last_image_error()["error_code"] == "auth_error"
|
||||
|
||||
def test_404_returns_model_not_found(self):
|
||||
client = _make_client(max_retries=0)
|
||||
r = MagicMock(status_code=404, text="not found")
|
||||
r.json.return_value = {"error": {"message": "model not found"}}
|
||||
r.raise_for_status.side_effect = httpx.HTTPStatusError("a", request=MagicMock(), response=r)
|
||||
with patch("packages.shared.ai_client.httpx.post", return_value=r):
|
||||
assert client.image_generation("p") is None
|
||||
assert client.get_last_image_error()["error_code"] == "model_not_found"
|
||||
|
||||
def test_400_quota_returns_quota_exceeded(self):
|
||||
client = _make_client(max_retries=0)
|
||||
r = MagicMock(status_code=400, text="insufficient balance quota exceeded")
|
||||
r.json.return_value = {"error": {"message": "quota"}}
|
||||
r.raise_for_status.side_effect = httpx.HTTPStatusError("a", request=MagicMock(), response=r)
|
||||
with patch("packages.shared.ai_client.httpx.post", return_value=r):
|
||||
assert client.image_generation("p") is None
|
||||
assert client.get_last_image_error()["error_code"] == "quota_exceeded"
|
||||
|
||||
def test_400_rate_limit_returns_rate_limit(self):
|
||||
client = _make_client(max_retries=0)
|
||||
r = MagicMock(status_code=400, text="too many requests, rate limit exceeded")
|
||||
r.json.return_value = {"error": {"message": "rate"}}
|
||||
r.raise_for_status.side_effect = httpx.HTTPStatusError("a", request=MagicMock(), response=r)
|
||||
with patch("packages.shared.ai_client.httpx.post", return_value=r):
|
||||
assert client.image_generation("p") is None
|
||||
assert client.get_last_image_error()["error_code"] == "rate_limit"
|
||||
|
||||
def test_400_generic_returns_invalid_param(self):
|
||||
client = _make_client(max_retries=0)
|
||||
r = MagicMock(status_code=400, text="bad parameter size")
|
||||
r.json.return_value = {"error": {"message": "bad"}}
|
||||
r.raise_for_status.side_effect = httpx.HTTPStatusError("a", request=MagicMock(), response=r)
|
||||
with patch("packages.shared.ai_client.httpx.post", return_value=r):
|
||||
assert client.image_generation("p") is None
|
||||
assert client.get_last_image_error()["error_code"] == "invalid_param"
|
||||
|
||||
def test_200_but_no_url_returns_none(self):
|
||||
client = _make_client(max_retries=0)
|
||||
r = MagicMock(status_code=200, text="")
|
||||
r.json.return_value = {"data": [{"no_url": True}]} # 缺 url 字段
|
||||
r.raise_for_status = MagicMock()
|
||||
with patch("packages.shared.ai_client.httpx.post", return_value=r):
|
||||
assert client.image_generation("p") is None
|
||||
assert client.get_last_image_error()["error_code"] == "unknown"
|
||||
|
||||
def test_network_error_retries_then_fails(self):
|
||||
client = _make_client(max_retries=1)
|
||||
import httpcore
|
||||
|
||||
with (
|
||||
patch("packages.shared.ai_client.httpx.post", side_effect=httpx.ConnectError("no network")),
|
||||
patch("packages.shared.ai_client.time.sleep", return_value=None),
|
||||
):
|
||||
assert client.image_generation("p") is None
|
||||
err = client.get_last_image_error()
|
||||
assert err["error_code"] == "network_error"
|
||||
|
||||
def test_get_last_image_error_returns_copy(self):
|
||||
client = _make_client()
|
||||
client.last_image_error = {"error_code": "x"}
|
||||
e1 = client.get_last_image_error()
|
||||
e1["error_code"] = "mutated"
|
||||
assert client.last_image_error["error_code"] == "x"
|
||||
|
||||
|
||||
# ── 信任链分支覆盖 ────────────────────────────────────────
|
||||
|
||||
|
||||
class TestTrustChainBranches:
|
||||
def test_dashscope_provider_skips_trust_chain(self, tmp_path):
|
||||
"""provider=dashscope 时不走信任链(Wan 模型由 dashscope_client 处理,在我们分支之前已经 return)。
|
||||
这里测 doubao 分支:信任链默认触发,验证 DashScope 分发路径不受影响。"""
|
||||
# 该测试实际覆盖 video_generation 入口的 dashscope 分发:缺 DASHSCOPE_API_KEY 时返回 auth_error
|
||||
client = _make_client()
|
||||
with (patch("packages.shared.ai_client.get_shared_settings") as ms,):
|
||||
ms.return_value = MagicMock(
|
||||
doubao_video_poll_interval=0,
|
||||
doubao_video_timeout=1,
|
||||
doubao_video_model="doubao-seedance-2-5-260628",
|
||||
)
|
||||
# DashScope 不可用时返回 auth_error(不是信任链相关错误)
|
||||
result = client.video_generation(
|
||||
"p",
|
||||
output_dir=str(tmp_path),
|
||||
model="wan-3.0",
|
||||
image_url="https://img/x.jpg",
|
||||
)
|
||||
assert result is None
|
||||
err = client.get_last_video_error()
|
||||
# 不论是否走信任链,DashScope 无 key 时返回 auth_error
|
||||
assert err["error_code"] == "auth_error"
|
||||
|
||||
def test_trust_chain_partial_seedream_success_falls_back(self, tmp_path):
|
||||
"""#2174: 预热结果为空/None 时回退原图直传(走#2166的400→t2v自动降级路径)。"""
|
||||
client = _make_client(max_retries=0)
|
||||
|
||||
# 预热结果传 None → 应该直接用原图发给 Seedance
|
||||
def fake_post(url, **kwargs):
|
||||
# Seedance create task(收到原图直传时会调用)
|
||||
t = MagicMock(status_code=200, text="")
|
||||
t.json.return_value = {"id": "t-partial"}
|
||||
t.raise_for_status = MagicMock()
|
||||
return t
|
||||
|
||||
poll_ok = MagicMock(status_code=200, text="")
|
||||
poll_ok.json.return_value = {"status": "succeeded", "content": {"video_url": "https://cdn.example.com/v.mp4"}}
|
||||
poll_ok.raise_for_status = MagicMock()
|
||||
|
||||
class FS:
|
||||
def __init__(self):
|
||||
self._it = iter([b"OK"])
|
||||
|
||||
def __enter__(self):
|
||||
return self
|
||||
|
||||
def __exit__(self, *a):
|
||||
return False
|
||||
|
||||
def raise_for_status(self):
|
||||
return None
|
||||
|
||||
def iter_bytes(self, chunk_size=None):
|
||||
return self._it
|
||||
|
||||
fake_uuid = MagicMock()
|
||||
fake_uuid.hex = "00000004"
|
||||
with (
|
||||
patch("packages.shared.ai_client.httpx.post", side_effect=fake_post),
|
||||
patch("packages.shared.ai_client.httpx.get", return_value=poll_ok),
|
||||
patch("packages.shared.ai_client.httpx.stream", return_value=FS()),
|
||||
patch("packages.shared.ai_client.time.sleep", return_value=None),
|
||||
patch("packages.shared.ai_client.time.time", side_effect=_fake_time(stable_calls=50)),
|
||||
patch("packages.shared.ai_client.uuid.uuid4", return_value=fake_uuid),
|
||||
patch("packages.shared.ai_client.get_shared_settings") as ms,
|
||||
):
|
||||
ms.return_value = MagicMock(
|
||||
doubao_video_poll_interval=0,
|
||||
doubao_video_timeout=60,
|
||||
doubao_video_model="doubao-seedance-2-5-260628",
|
||||
)
|
||||
out = client.video_generation(
|
||||
"p",
|
||||
image_url="https://img/a.jpg",
|
||||
reference_images=["https://img/b.jpg"],
|
||||
duration=5,
|
||||
ratio="9:16",
|
||||
output_dir=str(tmp_path),
|
||||
)
|
||||
assert out is not None
|
||||
# 最终发给 Seedance 的图应是原始 https://img/a.jpg(回退),role=first_frame(因为 has_extra_refs=False 只有 1 张)
|
||||
# 注意:回退后 ref_imgs 是原始 ["https://img/b.jpg"],所以 has_extra_refs=True,role=reference_image
|
||||
# 断言最终 Seedance payload 里的 image_url 是原图(不是 AI 图)
|
||||
|
||||
def test_default_values_on_missing_settings(self):
|
||||
"""getattr 兜底:settings 缺 image_timeout 字段时使用默认 120。"""
|
||||
client = _make_client()
|
||||
# 直接调用 image_generation,让它走一次完整流程(成功路径),验证 timeout 取值
|
||||
ok = MagicMock(status_code=200, text="")
|
||||
ok.json.return_value = {"data": [{"url": "https://ai.example.com/x.png"}]}
|
||||
ok.raise_for_status = MagicMock()
|
||||
captured_kwargs = {}
|
||||
|
||||
def fake_post(url, **kw):
|
||||
captured_kwargs["timeout"] = kw.get("timeout")
|
||||
return ok
|
||||
|
||||
with patch("packages.shared.ai_client.httpx.post", side_effect=fake_post):
|
||||
r = client.image_generation("p", timeout=None) # 不传 timeout,走 self.image_timeout=120
|
||||
assert r is not None
|
||||
assert captured_kwargs["timeout"] == 120
|
||||
|
||||
def test_trust_chain_uses_preheated_t2i_images(self, tmp_path):
|
||||
"""#2174: 传 pre_trusted_images(预热好的t2i信任图)时,替换原参考图发给Seedance。"""
|
||||
client = _make_client()
|
||||
captured = []
|
||||
|
||||
task_ok = MagicMock(status_code=200, text="")
|
||||
task_ok.json.return_value = {"id": "t-refonly"}
|
||||
task_ok.raise_for_status = MagicMock()
|
||||
poll_ok = MagicMock(status_code=200, text="")
|
||||
poll_ok.json.return_value = {"status": "succeeded", "content": {"video_url": "https://cdn.example.com/v.mp4"}}
|
||||
poll_ok.raise_for_status = MagicMock()
|
||||
|
||||
class FS:
|
||||
def __init__(self):
|
||||
self._it = iter([b"OK"])
|
||||
|
||||
def __enter__(self):
|
||||
return self
|
||||
|
||||
def __exit__(self, *a):
|
||||
return False
|
||||
|
||||
def raise_for_status(self):
|
||||
return None
|
||||
|
||||
def iter_bytes(self, chunk_size=None):
|
||||
return self._it
|
||||
|
||||
def fake_post(url, **kw):
|
||||
captured.append({"url": url, "json": kw.get("json")})
|
||||
return task_ok
|
||||
|
||||
fu = MagicMock()
|
||||
fu.hex = "0000000a"
|
||||
with (
|
||||
patch("packages.shared.ai_client.httpx.post", side_effect=fake_post),
|
||||
patch("packages.shared.ai_client.httpx.get", return_value=poll_ok),
|
||||
patch("packages.shared.ai_client.httpx.stream", return_value=FS()),
|
||||
patch("packages.shared.ai_client.time.sleep", return_value=None),
|
||||
patch("packages.shared.ai_client.time.time", side_effect=_fake_time(stable_calls=3)),
|
||||
patch("packages.shared.ai_client.uuid.uuid4", return_value=fu),
|
||||
patch("packages.shared.ai_client.get_shared_settings") as ms,
|
||||
):
|
||||
ms.return_value = MagicMock(
|
||||
doubao_video_poll_interval=0,
|
||||
doubao_video_timeout=60,
|
||||
doubao_video_model="doubao-seedance-2-5-260628",
|
||||
)
|
||||
out = client.video_generation(
|
||||
"人物散步",
|
||||
reference_images=["https://img/portrait.jpg"],
|
||||
pre_trusted_images=["https://ai.example.com/t2i-portrait.png"],
|
||||
duration=5,
|
||||
ratio="9:16",
|
||||
output_dir=str(tmp_path),
|
||||
)
|
||||
assert out is not None
|
||||
# 只有一次 Seedance 创建任务(预热已完成,不再现场跑 Seedream)
|
||||
assert len(captured) == 1
|
||||
seedance_payload = captured[0]["json"]
|
||||
content = seedance_payload["content"]
|
||||
img_items = [c for c in content if c.get("type") == "image_url"]
|
||||
assert len(img_items) == 1
|
||||
# 不传 image_url,信任链产物放 ref_imgs,走 reference_image 模式(非 first_frame)
|
||||
assert img_items[0]["image_url"]["url"] == "https://ai.example.com/t2i-portrait.png"
|
||||
assert img_items[0]["role"] == "reference_image"
|
||||
# 因为没有 image_url,没有 text 也没有 extra_refs 之外的字段,应保留用户 ratio=9:16
|
||||
assert seedance_payload.get("ratio") == "9:16"
|
||||
|
||||
def test_image_generation_generic_exception_retries_then_fails(self):
|
||||
"""image_generation 遇到非 HTTPStatusError 的通用异常时走重试分支(lines 962-971),重试耗尽后返回 None。"""
|
||||
client = _make_client(max_retries=1)
|
||||
call_n = {"n": 0}
|
||||
|
||||
def fake_post(url, **kw):
|
||||
call_n["n"] += 1
|
||||
if call_n["n"] == 1:
|
||||
raise RuntimeError("boiler exploded")
|
||||
# 第二次调用返回成功,验证重试生效
|
||||
ok = MagicMock(status_code=200, text="")
|
||||
ok.json.return_value = {"data": [{"url": "https://ai.example.com/retry-ok.png"}]}
|
||||
ok.raise_for_status = MagicMock()
|
||||
return ok
|
||||
|
||||
with (
|
||||
patch("packages.shared.ai_client.httpx.post", side_effect=fake_post),
|
||||
patch("packages.shared.ai_client.time.sleep", return_value=None),
|
||||
):
|
||||
r = client.image_generation("test prompt")
|
||||
assert r is not None
|
||||
assert r["url"] == "https://ai.example.com/retry-ok.png"
|
||||
assert call_n["n"] == 2
|
||||
|
||||
def test_image_generation_generic_exception_exhausts_retries(self):
|
||||
"""通用异常重试耗尽后返回 None,并正确写入 last_image_error (lines 969-971 break 分支)。"""
|
||||
client = _make_client(max_retries=1)
|
||||
|
||||
def fake_post(url, **kw):
|
||||
raise RuntimeError("always fails")
|
||||
|
||||
with (
|
||||
patch("packages.shared.ai_client.httpx.post", side_effect=fake_post),
|
||||
patch("packages.shared.ai_client.time.sleep", return_value=None),
|
||||
):
|
||||
r = client.image_generation("test prompt")
|
||||
assert r is None
|
||||
err = client.last_image_error
|
||||
assert err["error_code"] == "network_error"
|
||||
assert "always fails" in err["detail"]
|
||||
@@ -17,8 +17,13 @@ def _make_client(**overrides):
|
||||
client.base_url = overrides.get("base_url", "https://ark.cn-beijing.volces.com/api/v3")
|
||||
client.model = "doubao-model"
|
||||
client.vision_model = "doubao-vision"
|
||||
client.embedding_model = "doubao-embedding"
|
||||
client.image_model = overrides.get("image_model", "doubao-seedream-5-0-pro-260628")
|
||||
client.image_timeout = overrides.get("image_timeout", 120)
|
||||
client.timeout = overrides.get("timeout", 30)
|
||||
client.max_retries = overrides.get("max_retries", 0)
|
||||
client.last_video_error = {}
|
||||
client.last_image_error = {}
|
||||
return client
|
||||
|
||||
|
||||
@@ -106,7 +111,7 @@ class TestVideoGenerationHappyPath:
|
||||
)
|
||||
out = client.video_generation(
|
||||
prompt=" 镜头一 ",
|
||||
image_url="https://img/x.jpg",
|
||||
# 不传 image_url:纯文生视频,不触发信任链,post 调用数为 1(创建任务)
|
||||
duration=5,
|
||||
ratio="9:16",
|
||||
resolution="720p",
|
||||
@@ -563,7 +568,7 @@ class TestResolveVideoModelId:
|
||||
|
||||
class TestVideoGenerationLastError:
|
||||
def test_create_400_portrait_returns_user_message(self, tmp_path):
|
||||
"""HTTP 400 + 真人拦截关键词 → portrait_intercept 错误码,用户提示友好。"""
|
||||
"""#2169: HTTP 400 + 真人拦截关键词 → 自动尝试即梦兜底;即梦未配时返回 portrait_intercept。"""
|
||||
client = _make_client(max_retries=0)
|
||||
create_resp = MagicMock()
|
||||
create_resp.status_code = 400
|
||||
@@ -572,8 +577,23 @@ class TestVideoGenerationLastError:
|
||||
create_resp.raise_for_status.side_effect = httpx.HTTPStatusError(
|
||||
"bad", request=MagicMock(), response=create_resp
|
||||
)
|
||||
# 信任链:Seedream 会先被调用来 AI 化;这里 mock Seedream 也失败,回退原图直传,
|
||||
# 原图直传被 400 portrait 拦截,最终返回 portrait_intercept。
|
||||
seedream_resp = MagicMock()
|
||||
seedream_resp.status_code = 400
|
||||
seedream_resp.text = '{"error":{"code":"ContentRisk","message":"sensitive"}}'
|
||||
seedream_resp.json.return_value = {"error": {"code": "ContentRisk", "message": "sensitive"}}
|
||||
seedream_resp.raise_for_status.side_effect = httpx.HTTPStatusError(
|
||||
"bad", request=MagicMock(), response=seedream_resp
|
||||
)
|
||||
|
||||
def fake_post(url, **kwargs):
|
||||
# 第一次 POST 是 Seedream(/images/generations),返回 portrait 拦截
|
||||
# 回退原图直传后第二次 POST 是 Seedance(/contents/generations/tasks),也返回 portrait 拦截
|
||||
return create_resp
|
||||
|
||||
with (
|
||||
patch("packages.shared.ai_client.httpx.post", return_value=create_resp),
|
||||
patch("packages.shared.ai_client.httpx.post", side_effect=fake_post),
|
||||
patch("packages.shared.ai_client.time.sleep", return_value=None),
|
||||
patch("packages.shared.ai_client.get_shared_settings") as mock_s,
|
||||
):
|
||||
@@ -584,8 +604,8 @@ class TestVideoGenerationLastError:
|
||||
assert result is None
|
||||
err = client.get_last_video_error()
|
||||
assert err["error_code"] == "portrait_intercept"
|
||||
assert "真人" in err["user_message"]
|
||||
assert err["status_code"] == 400
|
||||
assert "真人" in err["user_message"] or "肖像" in err["user_message"] or "审核" in err["user_message"]
|
||||
assert err["status_code"] in (0, 400)
|
||||
|
||||
def test_create_401_returns_auth_error(self, tmp_path):
|
||||
client = _make_client(max_retries=0)
|
||||
|
||||
@@ -80,8 +80,8 @@ class TestSharedSettingsDefaults:
|
||||
def test_default_doubao_settings(self):
|
||||
s = SharedSettings()
|
||||
assert "doubao" in s.doubao_model
|
||||
assert s.doubao_timeout == 30
|
||||
assert s.doubao_max_retries == 2
|
||||
assert s.doubao_timeout == 45 # #2180 默认提到45s
|
||||
assert s.doubao_max_retries == 1
|
||||
|
||||
|
||||
class TestAPISettingsDefaults:
|
||||
|
||||
@@ -0,0 +1,198 @@
|
||||
"""catalog 应用服务单测:会员套餐 / 积分包从共享库读取与字段映射。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _clear_cache():
|
||||
from packages.application.catalog import admin_catalog
|
||||
|
||||
admin_catalog._cache.clear()
|
||||
yield
|
||||
admin_catalog._cache.clear()
|
||||
|
||||
|
||||
def _row(**kw):
|
||||
row = MagicMock()
|
||||
for k, v in kw.items():
|
||||
setattr(row, k, v)
|
||||
return row
|
||||
|
||||
|
||||
class TestMembershipPlans:
|
||||
def test_yearly_plan_mapping(self):
|
||||
from packages.application.catalog import admin_catalog
|
||||
|
||||
row = _row(
|
||||
plan_key="premium_yearly",
|
||||
name="高级会员年卡",
|
||||
description="年度订阅",
|
||||
monthly_price=0,
|
||||
yearly_price=399,
|
||||
quotas={"4k": True, "batch_render": True, "credits_per_month": 500},
|
||||
display_order=1,
|
||||
)
|
||||
session = MagicMock()
|
||||
session.execute.return_value.fetchall.return_value = [row]
|
||||
sl = MagicMock(return_value=session)
|
||||
|
||||
with patch("packages.adapters.sqlalchemy_impl.session.SessionLocal", sl, create=True):
|
||||
plans = admin_catalog.get_membership_plans()
|
||||
|
||||
assert len(plans) == 1
|
||||
p = plans[0]
|
||||
assert p["plan_id"] == "premium_yearly"
|
||||
assert p["billing_cycle"] == "yearly"
|
||||
assert p["price_cents"] == 39900
|
||||
assert p["monthly_price_cents"] == 3325
|
||||
assert p["duration_days"] == 365
|
||||
assert p["features"]["4K 超清分辨率"] is True
|
||||
assert p["features"]["credits_per_month"] == 500
|
||||
session.close.assert_called_once()
|
||||
|
||||
def test_monthly_plan_mapping(self):
|
||||
from packages.application.catalog import admin_catalog
|
||||
|
||||
row = _row(
|
||||
plan_key="premium_monthly",
|
||||
name="高级会员月卡",
|
||||
description=None,
|
||||
monthly_price=39,
|
||||
yearly_price=0,
|
||||
quotas=None,
|
||||
display_order=2,
|
||||
)
|
||||
session = MagicMock()
|
||||
session.execute.return_value.fetchall.return_value = [row]
|
||||
sl = MagicMock(return_value=session)
|
||||
|
||||
with patch("packages.adapters.sqlalchemy_impl.session.SessionLocal", sl, create=True):
|
||||
plans = admin_catalog.get_membership_plans()
|
||||
|
||||
assert len(plans) == 1
|
||||
p = plans[0]
|
||||
assert p["billing_cycle"] == "monthly"
|
||||
assert p["price_cents"] == 3900
|
||||
assert p["monthly_price_cents"] == 3900
|
||||
assert p["duration_days"] == 30
|
||||
assert p["features"] == {}
|
||||
|
||||
def test_both_cycles_expanded(self):
|
||||
from packages.application.catalog import admin_catalog
|
||||
|
||||
row = _row(
|
||||
plan_key="premium",
|
||||
name="高级会员",
|
||||
description=None,
|
||||
monthly_price=39,
|
||||
yearly_price=399,
|
||||
quotas={},
|
||||
display_order=1,
|
||||
)
|
||||
session = MagicMock()
|
||||
session.execute.return_value.fetchall.return_value = [row]
|
||||
sl = MagicMock(return_value=session)
|
||||
|
||||
with patch("packages.adapters.sqlalchemy_impl.session.SessionLocal", sl, create=True):
|
||||
plans = admin_catalog.get_membership_plans()
|
||||
|
||||
cycles = {p["billing_cycle"] for p in plans}
|
||||
assert cycles == {"yearly", "monthly"}
|
||||
|
||||
def test_no_session_returns_empty(self):
|
||||
from packages.application.catalog import admin_catalog
|
||||
|
||||
with patch("packages.adapters.sqlalchemy_impl.session.SessionLocal", None, create=True):
|
||||
assert admin_catalog.get_membership_plans() == []
|
||||
|
||||
|
||||
class TestPointsPackages:
|
||||
def test_package_mapping_with_bonus(self):
|
||||
from packages.application.catalog import admin_catalog
|
||||
|
||||
row = _row(
|
||||
package_key="pkg_100",
|
||||
name="100元充值包",
|
||||
price=100,
|
||||
credits=1000,
|
||||
bonus_credits=100,
|
||||
is_recommended=True,
|
||||
description="推荐",
|
||||
sort_order=4,
|
||||
)
|
||||
session = MagicMock()
|
||||
session.execute.return_value.fetchall.return_value = [row]
|
||||
sl = MagicMock(return_value=session)
|
||||
|
||||
with patch("packages.adapters.sqlalchemy_impl.session.SessionLocal", sl, create=True):
|
||||
packages = admin_catalog.get_points_packages()
|
||||
|
||||
assert len(packages) == 1
|
||||
pkg = packages[0]
|
||||
assert pkg["code"] == "pkg_100"
|
||||
assert pkg["points"] == 1100
|
||||
assert pkg["price_cents"] == 10000
|
||||
assert pkg["is_recommended"] is True
|
||||
assert pkg["unit_price"] == "¥0.091/积分"
|
||||
|
||||
def test_zero_credits_unit_price_safe(self):
|
||||
from packages.application.catalog import admin_catalog
|
||||
|
||||
row = _row(
|
||||
package_key="pkg_0",
|
||||
name="空包",
|
||||
price=0,
|
||||
credits=0,
|
||||
bonus_credits=0,
|
||||
is_recommended=False,
|
||||
description=None,
|
||||
sort_order=0,
|
||||
)
|
||||
session = MagicMock()
|
||||
session.execute.return_value.fetchall.return_value = [row]
|
||||
sl = MagicMock(return_value=session)
|
||||
|
||||
with patch("packages.adapters.sqlalchemy_impl.session.SessionLocal", sl, create=True):
|
||||
packages = admin_catalog.get_points_packages()
|
||||
|
||||
assert packages[0]["points"] == 0
|
||||
assert packages[0]["price_cents"] == 0
|
||||
assert packages[0]["unit_price"] == "¥0.000/积分"
|
||||
|
||||
def test_no_session_returns_empty(self):
|
||||
from packages.application.catalog import admin_catalog
|
||||
|
||||
with patch("packages.adapters.sqlalchemy_impl.session.SessionLocal", None, create=True):
|
||||
assert admin_catalog.get_points_packages() == []
|
||||
|
||||
|
||||
class TestPackagesRoute:
|
||||
def test_get_packages_route_returns_items(self):
|
||||
from app.api.routes.points import get_packages
|
||||
|
||||
cu = MagicMock()
|
||||
cu.user.member_type = None
|
||||
rows = [
|
||||
{
|
||||
"code": "pkg_10",
|
||||
"name": "10元充值包",
|
||||
"points": 100,
|
||||
"price_cents": 1000,
|
||||
"unit_price": "¥0.100/积分",
|
||||
}
|
||||
]
|
||||
with patch(
|
||||
"packages.application.catalog.admin_catalog.get_points_packages",
|
||||
return_value=rows,
|
||||
):
|
||||
resp = get_packages(current_user=cu)
|
||||
|
||||
assert len(resp.packages) == 1
|
||||
item = resp.packages[0]
|
||||
assert item.code == "pkg_10"
|
||||
assert item.points == 100
|
||||
assert item.price_cents == 1000
|
||||
@@ -110,8 +110,8 @@ class TestSharedSettingsDefaults:
|
||||
def test_default_doubao_config(self):
|
||||
"""豆包默认配置"""
|
||||
s = self._make_settings()
|
||||
assert s.doubao_timeout == 30
|
||||
assert s.doubao_max_retries == 2
|
||||
assert s.doubao_timeout == 45 # #2180 默认提到45s
|
||||
assert s.doubao_max_retries == 1
|
||||
assert "volces.com" in s.doubao_base_url
|
||||
|
||||
def test_default_empty_api_keys(self):
|
||||
|
||||
@@ -173,33 +173,46 @@ class TestSubscriptionPlans:
|
||||
_spec.loader.exec_module(_mod)
|
||||
return _mod.list_membership_plans
|
||||
|
||||
def test_plans_endpoint_returns_three_tiers(self):
|
||||
import os # noqa: F401 (used by _import_plans_fn)
|
||||
|
||||
def test_plans_endpoint_reads_admin_table(self):
|
||||
"""/subscription/plans 改读管理后台 plans 表:返回 catalog 服务提供的真实档位。"""
|
||||
list_membership_plans = self._import_plans_fn()
|
||||
resp = list_membership_plans(current_user=_make_cu())
|
||||
real_plan = {
|
||||
"plan_id": "premium_yearly",
|
||||
"billing_cycle": "yearly",
|
||||
"name": "高级会员年卡",
|
||||
"description": "高级会员年度订阅,享受全部功能",
|
||||
"price_cents": 39900,
|
||||
"monthly_price_cents": 3325,
|
||||
"duration_days": 365,
|
||||
"features": {
|
||||
"4K 超清分辨率": True,
|
||||
"批量渲染": True,
|
||||
"优先处理队列": True,
|
||||
"credits_per_month": 500,
|
||||
},
|
||||
}
|
||||
with patch(
|
||||
"packages.application.catalog.admin_catalog.get_membership_plans",
|
||||
return_value=[real_plan],
|
||||
):
|
||||
resp = list_membership_plans(current_user=_make_cu())
|
||||
plans = resp["plans"]
|
||||
plan_ids = {p["plan_id"] for p in plans}
|
||||
assert plan_ids == {"monthly", "quarterly", "yearly"}
|
||||
for p in plans:
|
||||
assert p["price_cents"] > 0
|
||||
assert p["duration_days"] in (30, 90, 365)
|
||||
assert 0 < p["points_discount"] <= 1.0
|
||||
assert "max_resolution" in p["features"]
|
||||
|
||||
def test_longer_plans_cheaper_per_month(self):
|
||||
import os # noqa: F401
|
||||
assert len(plans) == 1
|
||||
p0 = plans[0]
|
||||
assert p0["plan_id"] == "premium_yearly"
|
||||
assert p0["price_cents"] == 39900
|
||||
assert p0["duration_days"] == 365
|
||||
assert p0["features"]["4K 超清分辨率"] is True
|
||||
|
||||
def test_plans_endpoint_empty_when_all_disabled(self):
|
||||
"""后台停用全部套餐时,用户端返回空列表。"""
|
||||
list_membership_plans = self._import_plans_fn()
|
||||
resp = list_membership_plans(current_user=_make_cu())
|
||||
plans = resp["plans"]
|
||||
monthly = next(p for p in plans if p["plan_id"] == "monthly")
|
||||
quarterly = next(p for p in plans if p["plan_id"] == "quarterly")
|
||||
yearly = next(p for p in plans if p["plan_id"] == "yearly")
|
||||
assert monthly["monthly_price_cents"] == 1990
|
||||
assert quarterly["monthly_price_cents"] < monthly["monthly_price_cents"]
|
||||
assert yearly["monthly_price_cents"] < quarterly["monthly_price_cents"]
|
||||
|
||||
with patch(
|
||||
"packages.application.catalog.admin_catalog.get_membership_plans",
|
||||
return_value=[],
|
||||
):
|
||||
resp = list_membership_plans(current_user=_make_cu())
|
||||
assert resp["plans"] == []
|
||||
|
||||
# ── P1-7: multiplier consistency ──────────────────────────────────────
|
||||
|
||||
|
||||
@@ -403,33 +403,30 @@ class TestViralVideoPipeline:
|
||||
"""v1.6: _step_script_generation 返回 dict 形式的 CopyResult,含 voiceover_script + shots。"""
|
||||
from apps.worker.worker_app.tasks.viral_video import _step_script_generation
|
||||
|
||||
mock_llm.return_value = {
|
||||
"overview": {"theme": "口红推荐", "total_duration": 15, "aspect_ratio": "9:16"},
|
||||
"scene_and_lighting": "明亮化妆台,柔和自然光",
|
||||
"shots": [
|
||||
{
|
||||
"time_range": "0-5秒",
|
||||
"shot_type_angle_movement": "近景平视,缓慢推镜",
|
||||
"scene_and_dialogue": "女主微笑展示口红:大家好,今天分享一款口红",
|
||||
"action_details": "手持口红特写",
|
||||
"audio_bgm": "轻快流行BGM",
|
||||
"transition": "硬切",
|
||||
"reference_image_index": 0,
|
||||
},
|
||||
{
|
||||
"time_range": "5-15秒",
|
||||
"shot_type_angle_movement": "特写,固定镜头",
|
||||
"scene_and_dialogue": "涂抹口红:颜色特别好看很显白",
|
||||
"action_details": "嘴唇涂抹特写",
|
||||
"audio_bgm": "轻快BGM继续",
|
||||
"transition": "结束",
|
||||
"reference_image_index": 1,
|
||||
},
|
||||
],
|
||||
"hard_constraints": ["无字幕无水印"],
|
||||
"negative_prompts": ["字幕", "水印"],
|
||||
"voiceover_script": "大家好,今天分享一款口红,颜色特别好看很显白。",
|
||||
}
|
||||
mock_llm.return_value = """<clips>
|
||||
<clip image_index="0" transition="cut" zoom="null" duration_sec="5" bgm_note="轻快流行BGM">
|
||||
<voice_text>大家好,今天分享一款口红</voice_text>
|
||||
<subtitle_text>大家好,今天分享一款口红</subtitle_text>
|
||||
<shot_type_angle_movement>近景平视,缓慢推镜</shot_type_angle_movement>
|
||||
<scene_and_dialogue>女主微笑展示口红:大家好,今天分享一款口红</scene_and_dialogue>
|
||||
<action_details>手持口红特写</action_details>
|
||||
<audio_bgm>轻快流行BGM</audio_bgm>
|
||||
<transition>硬切</transition>
|
||||
<reference_image_index>0</reference_image_index>
|
||||
<ken_burns start="0,0" end="0,0" ease="linear"/>
|
||||
</clip>
|
||||
<clip image_index="0" transition="fade" zoom="null" duration_sec="10" bgm_note="轻快BGM">
|
||||
<voice_text>颜色特别好看很显白</voice_text>
|
||||
<subtitle_text>颜色特别好看很显白</subtitle_text>
|
||||
<shot_type_angle_movement>特写,固定镜头</shot_type_angle_movement>
|
||||
<scene_and_dialogue>涂抹口红:颜色特别好看很显白</scene_and_dialogue>
|
||||
<action_details>嘴唇涂抹特写</action_details>
|
||||
<audio_bgm>轻快BGM继续</audio_bgm>
|
||||
<transition>结束</transition>
|
||||
<reference_image_index>1</reference_image_index>
|
||||
<ken_burns start="0,0" end="0,0" ease="linear"/>
|
||||
</clip>
|
||||
</clips>"""
|
||||
result = _step_script_generation(
|
||||
mock_job, {"intent": "推广口红", "key_messages": [], "tone": "亲切"}, {"products": []}
|
||||
)
|
||||
@@ -452,12 +449,13 @@ class TestViralVideoPipeline:
|
||||
assert result["voiceover_script"]
|
||||
assert len(result["shots"]) >= 1
|
||||
|
||||
@patch("packages.shared.ai_service.call_llm")
|
||||
def test_review_pass_v16(self, mock_llm, mock_job):
|
||||
@patch("packages.application.viral_video.reviewer.Reviewer.review")
|
||||
def test_review_pass_v16(self, mock_review, mock_job):
|
||||
"""v1.6 _step_review 接收 copy_result dict。"""
|
||||
from apps.worker.worker_app.tasks.viral_video import _step_review
|
||||
from packages.application.viral_video.reviewer import ReviewIssue, ReviewResult
|
||||
|
||||
mock_llm.return_value = {"passed": True, "score": 90, "details": {}}
|
||||
mock_review.return_value = ReviewResult(passed=True, score=90, issues=[], rewrite_suggestions=[])
|
||||
cr = {"voiceover_script": "大家好", "shots": []}
|
||||
result = _step_review(mock_job, cr)
|
||||
assert result["passed"] is True
|
||||
|
||||
@@ -0,0 +1,478 @@
|
||||
"""#2040 爆款视频 Prompt 模板系统单测。
|
||||
|
||||
不真调豆包 API,全部用 FakeClient 注入;覆盖:
|
||||
XML 标签解析 / 5 套模板纯文本 / loader 缓存热加载与回落 /
|
||||
三档融合差异 / personal_brands 保留 / 审核识别违规词夸大 / 自动重写 /
|
||||
各步 fallback / seed 幂等 / 负面词不出现。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import sys
|
||||
|
||||
import pytest
|
||||
import sqlalchemy as sa
|
||||
from sqlalchemy.orm import sessionmaker
|
||||
|
||||
REPO_ROOT = os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
for p in [REPO_ROOT, os.path.join(REPO_ROOT, "apps/api"), os.path.join(REPO_ROOT, "apps/worker")]:
|
||||
if p not in sys.path:
|
||||
sys.path.insert(0, p)
|
||||
|
||||
from packages.application.viral_video import xml_parser as xp # noqa: E402
|
||||
from packages.application.viral_video.generator import CopyGenerator # noqa: E402
|
||||
from packages.application.viral_video.prompt_loader import ( # noqa: E402
|
||||
get_template,
|
||||
invalidate,
|
||||
render_user_prompt,
|
||||
)
|
||||
from packages.application.viral_video.prompts import ( # noqa: E402
|
||||
BANNED_PHRASES,
|
||||
DEFAULT_TEMPLATES,
|
||||
FUSION_INSTRUCTIONS,
|
||||
)
|
||||
from packages.application.viral_video.reviewer import Reviewer # noqa: E402
|
||||
|
||||
IMAGE_XML = """<products>
|
||||
<product name="大公鸡头油污净" features="红白瓶身" position="main" image_index="0"/>
|
||||
</products>
|
||||
<colors><color hex="#D32F2F" name="红色" coverage="0.4"/></colors>
|
||||
<people has_person="false" count="0"/>
|
||||
<mood>干净实用</mood>
|
||||
<visible_text><text_item text="多功能油污净" position="瓶身"/></visible_text>
|
||||
<scene>白底棚拍</scene>
|
||||
<quality resolution="高清" lighting="柔和" composition="居中"/>
|
||||
<key_selling_points><point>去油快</point><point>625ml大容量</point></key_selling_points>"""
|
||||
|
||||
INTENT_XML = """<intent_summary>厨房去油污神器</intent_summary>
|
||||
<core_messages>
|
||||
<message must_keep="true" confidence="0.97">去油污效果好</message>
|
||||
<message must_keep="false" confidence="0.6">适合重油污</message>
|
||||
</core_messages>
|
||||
<personal_brands><brand category="price">39块钱一瓶</brand></personal_brands>
|
||||
<emotion_tone>亲切真实</emotion_tone>
|
||||
<missing_info><info>容量按625ml</info></missing_info>"""
|
||||
|
||||
FUSION_XML = """<title>厨房重油污别硬擦了</title>
|
||||
<hook>这油污忍很久了</hook>
|
||||
<body_points><point elaboration="喷上等几分钟一擦就净" image_index="0">大公鸡头去油快</point></body_points>
|
||||
<cta>重油污的可以试一瓶</cta>
|
||||
<script_segments>
|
||||
<segment duration_sec="3" image_index="0">这油污忍很久了</segment>
|
||||
<segment duration_sec="6" image_index="0">大公鸡头油污净喷上等几分钟一擦就净</segment>
|
||||
<segment duration_sec="4" image_index="0">39块钱一瓶可以试一下</segment>
|
||||
</script_segments>
|
||||
<word_count>52</word_count><estimated_duration>13</estimated_duration>"""
|
||||
|
||||
STORYBOARD_XML = """<clips>
|
||||
<clip image_index="0" transition="cut" zoom="null" duration_sec="3" bgm_note="日常">
|
||||
<voice_text>这油污忍很久了</voice_text>
|
||||
<subtitle_text>油污忍很久</subtitle_text>
|
||||
<ken_burns start="0,0" end="0,0" ease="linear"/>
|
||||
</clip>
|
||||
<clip image_index="0" transition="zoom_in" zoom="in" duration_sec="10" bgm_note="轻快">
|
||||
<voice_text>大公鸡头喷上等几分钟一擦就净</voice_text>
|
||||
<subtitle_text>一擦就净</subtitle_text>
|
||||
<ken_burns start="20,20" end="80,80" ease="ease-in-out"/>
|
||||
</clip>
|
||||
</clips>"""
|
||||
|
||||
REVIEW_FAIL_XML = """<passed>false</passed>
|
||||
<issues><issue dimension="夸大承诺" severity="error" location="body_points">出现绝对化表述</issue></issues>
|
||||
<rewrite_suggestions><suggestion>改为“大部分油污能擦掉”</suggestion></rewrite_suggestions>"""
|
||||
|
||||
REVIEW_PASS_XML = """<passed>true</passed>
|
||||
<issues></issues>
|
||||
<rewrite_suggestions></rewrite_suggestions>"""
|
||||
|
||||
FIXED_FUSION_XML = FUSION_XML.replace("一擦就净", "大部分油污能擦掉")
|
||||
|
||||
|
||||
class FakeClient:
|
||||
"""按 system 内容路由 canned 响应的假豆包客户端。"""
|
||||
|
||||
def __init__(self):
|
||||
self.chat_calls: list[list[dict]] = []
|
||||
self.vision_calls: list = []
|
||||
self.review_sequence: list[str] | None = None
|
||||
self.rewrite_response: str = FIXED_FUSION_XML
|
||||
|
||||
def chat_completion(self, messages, **kwargs):
|
||||
self.chat_calls.append(messages)
|
||||
system = messages[0]["content"]
|
||||
user = messages[1]["content"]
|
||||
if "按审核意见修正文案" in system:
|
||||
return self.rewrite_response
|
||||
if "文案合规审核员" in system:
|
||||
if self.review_sequence:
|
||||
return self.review_sequence.pop(0)
|
||||
return REVIEW_PASS_XML
|
||||
if "理解用户的营销意图" in system:
|
||||
return INTENT_XML
|
||||
if "负责把文案拆成可拍摄" in system:
|
||||
return STORYBOARD_XML
|
||||
if (
|
||||
"短视频生成营销文案" in system
|
||||
or "AI 全权创作" in system
|
||||
or "AI 辅助润色" in system
|
||||
or "用户原文为主" in system
|
||||
):
|
||||
mode = (
|
||||
"ai_full" if "AI 全权创作" in system else ("user_primary" if "用户原文为主" in system else "ai_polish")
|
||||
)
|
||||
if self._fusion_override is not None:
|
||||
return self._fusion_override
|
||||
xml = FUSION_XML
|
||||
if mode == "ai_full":
|
||||
xml = xml.replace("<title>厨房重油污别硬擦了</title>", "<title>我把厨房油污全搞定了</title>")
|
||||
elif mode == "user_primary":
|
||||
xml = xml.replace("<title>厨房重油污别硬擦了</title>", "<title>油污净使用分享</title>")
|
||||
self._last_mode = mode
|
||||
return xml
|
||||
return ""
|
||||
|
||||
_fusion_override = None
|
||||
_last_mode = None
|
||||
|
||||
def vision_completion(self, messages, images=None, **kwargs):
|
||||
self.vision_calls.append({"messages": messages, "images": images})
|
||||
return IMAGE_XML
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _clear_loader():
|
||||
invalidate()
|
||||
yield
|
||||
invalidate()
|
||||
|
||||
|
||||
# ── XML 解析 ──────────────────────────────────────────────────────────────
|
||||
class TestXmlParser:
|
||||
def test_parse_paired_tags_and_attrs(self):
|
||||
nodes = xp.find_all(IMAGE_XML, "product")
|
||||
assert len(nodes) == 1
|
||||
assert nodes[0]["attrs"]["name"] == "大公鸡头油污净"
|
||||
assert nodes[0]["attrs"]["image_index"] == "0"
|
||||
|
||||
def test_parse_selling_points(self):
|
||||
points = [n["text"] for n in xp.find_all(IMAGE_XML, "point")]
|
||||
assert points == ["去油快", "625ml大容量"]
|
||||
|
||||
def test_self_closing_and_bool_helpers(self):
|
||||
nodes = xp.parse_tags('<people has_person="false"/><done x="true"/>')
|
||||
assert xp.attr_bool(nodes[0]["attrs"]["has_person"]) is False
|
||||
assert xp.attr_bool(nodes[1]["attrs"]["x"]) is True
|
||||
assert xp.attr_float("0.97") == pytest.approx(0.97)
|
||||
assert xp.attr_int("13", 5) == 13
|
||||
|
||||
def test_malformed_text_safe(self):
|
||||
assert xp.find_all(None, "tag") == []
|
||||
assert xp.text_of("乱七八糟没有标签", "intent", "默认") == "默认"
|
||||
|
||||
|
||||
# ── 5 套模板纯文本 ────────────────────────────────────────────────────────
|
||||
class TestTemplates:
|
||||
def test_five_templates_present(self):
|
||||
types_ = {t["prompt_type"] for t in DEFAULT_TEMPLATES}
|
||||
assert types_ == {"image_analysis", "intent_parsing", "copy_fusion", "storyboard", "review"}
|
||||
|
||||
def test_no_json_blocks_in_templates(self):
|
||||
for template in DEFAULT_TEMPLATES:
|
||||
blob = "\n".join([template["system_prompt"], template["user_prompt_template"], template["example_output"]])
|
||||
assert "```json" not in blob
|
||||
assert "JSON schema" not in blob
|
||||
|
||||
def test_placeholders_render_and_missing_key_kept(self):
|
||||
template = get_template("intent_parsing")
|
||||
rendered = render_user_prompt(template, user_copy_text="去油快", industry="家居")
|
||||
assert "去油快" in rendered
|
||||
assert "去油快" in render_user_prompt(template, image_analysis="产品图", user_copy_text="去油快")
|
||||
partial = render_user_prompt(template, user_copy_text="x")
|
||||
assert "{industry}" not in partial or "{" in partial
|
||||
|
||||
|
||||
# ── loader:DB 加载/缓存/回落 ─────────────────────────────────────────────
|
||||
class TestPromptLoader:
|
||||
def test_fallback_when_session_none(self, monkeypatch):
|
||||
import packages.adapters.sqlalchemy_impl.session as session_mod
|
||||
|
||||
monkeypatch.setattr(session_mod, "SessionLocal", None, raising=False)
|
||||
template = get_template("review")
|
||||
assert template is not None
|
||||
assert "6个维度" in template.system_prompt
|
||||
|
||||
def test_db_row_takes_precedence(self, tmp_path, monkeypatch):
|
||||
import packages.adapters.sqlalchemy_impl.session as session_mod
|
||||
|
||||
db_path = tmp_path / "t.db"
|
||||
engine = sa.create_engine(f"sqlite:///{db_path}")
|
||||
with engine.begin() as conn:
|
||||
conn.execute(
|
||||
sa.text(
|
||||
"CREATE TABLE viral_video_prompt_templates ("
|
||||
"id INTEGER PRIMARY KEY, name TEXT, prompt_type TEXT, version INTEGER,"
|
||||
"system_prompt TEXT, user_prompt_template TEXT, example_output TEXT,"
|
||||
"is_active INTEGER)"
|
||||
)
|
||||
)
|
||||
conn.execute(
|
||||
sa.text(
|
||||
"INSERT INTO viral_video_prompt_templates VALUES"
|
||||
"(1,'自定义','review',2,'DB里的系统提示','DB用户提示','',1)"
|
||||
)
|
||||
)
|
||||
factory = sessionmaker(bind=engine)
|
||||
monkeypatch.setattr(session_mod, "SessionLocal", factory, raising=False)
|
||||
template = get_template("review", force_refresh=True)
|
||||
assert template.system_prompt == "DB里的系统提示"
|
||||
assert template.version == 2
|
||||
|
||||
# 改 DB 后 30 秒内仍走缓存
|
||||
with engine.begin() as conn:
|
||||
conn.execute(sa.text("UPDATE viral_video_prompt_templates SET system_prompt='改了' WHERE id=1"))
|
||||
assert get_template("review").system_prompt == "DB里的系统提示"
|
||||
# force_refresh 后热加载生效
|
||||
assert get_template("review", force_refresh=True).system_prompt == "改了"
|
||||
|
||||
def test_invalid_type_raises(self):
|
||||
with pytest.raises(ValueError):
|
||||
get_template("not_exist")
|
||||
|
||||
|
||||
# ── 5 步编排与 fallback ──────────────────────────────────────────────────
|
||||
class TestGenerator:
|
||||
def test_full_pipeline_xml_parseable(self):
|
||||
client = FakeClient()
|
||||
gen = CopyGenerator(client=client)
|
||||
result = gen.generate(["https://x/1.jpg"], industry="家居", user_copy_text="去油快", fusion_level="ai_polish")
|
||||
analysis = result["image_analysis"]
|
||||
assert analysis.products[0].name == "大公鸡头油污净"
|
||||
assert analysis.key_selling_points == ["去油快", "625ml大容量"]
|
||||
assert analysis.has_person is False
|
||||
|
||||
intent = result["intent_result"]
|
||||
assert intent.intent_summary == "厨房去油污神器"
|
||||
assert intent.core_messages[0].must_keep is True
|
||||
assert intent.personal_brands[0].text == "39块钱一瓶"
|
||||
|
||||
fusion = result["fusion_result"]
|
||||
assert fusion.title == "厨房重油污别硬擦了"
|
||||
assert len(fusion.script_segments) == 3
|
||||
|
||||
board = result["storyboard"]
|
||||
assert len(board.clips) == 2
|
||||
assert board.clips[1].transition == "zoom_in"
|
||||
assert board.clips[1].ken_burns.end == "80,80"
|
||||
# vision 确实被调用且带图
|
||||
assert client.vision_calls[0]["images"] == ["https://x/1.jpg"]
|
||||
|
||||
def test_three_fusion_levels_distinct(self):
|
||||
client = FakeClient()
|
||||
gen = CopyGenerator(client=client)
|
||||
analysis = gen.analyze_images(["https://x/1.jpg"])
|
||||
intent = gen.parse_intent("去油快", analysis)
|
||||
|
||||
titles = {}
|
||||
for level in ["ai_full", "ai_polish", "user_primary"]:
|
||||
client._fusion_override = None
|
||||
fusion = gen.fuse(level, analysis, intent, duration=15)
|
||||
titles[level] = fusion.title
|
||||
# system 里注入了对应档位指令
|
||||
system = client.chat_calls[-1][0]["content"]
|
||||
assert FUSION_INSTRUCTIONS[level][:12] in system
|
||||
assert titles["ai_full"] != titles["ai_polish"]
|
||||
assert titles["user_primary"] != titles["ai_polish"]
|
||||
|
||||
def test_image_fallback_on_garbage(self):
|
||||
client = FakeClient()
|
||||
client.vision_completion = lambda *a, **k: "完全无法解析的内容" # type: ignore
|
||||
gen = CopyGenerator(client=client)
|
||||
analysis = gen.analyze_images(["https://x/1.jpg"])
|
||||
assert analysis.products[0].name.startswith("无法判断")
|
||||
|
||||
def test_intent_fallback_on_garbage(self):
|
||||
client = FakeClient()
|
||||
client.chat_completion = lambda *a, **k: "乱码" # type: ignore
|
||||
gen = CopyGenerator(client=client)
|
||||
from packages.application.viral_video.schemas import ImageAnalysis
|
||||
|
||||
intent = gen.parse_intent("这是我的原意", ImageAnalysis())
|
||||
assert intent.intent_summary == "这是我的原意"
|
||||
assert intent.core_messages[0].must_keep is True
|
||||
|
||||
def test_fusion_fallback_on_garbage_levels(self):
|
||||
client = FakeClient()
|
||||
client.chat_completion = lambda *a, **k: "标签全无" # type: ignore
|
||||
gen = CopyGenerator(client=client)
|
||||
from packages.application.viral_video.schemas import ImageAnalysis, IntentResult
|
||||
|
||||
analysis = ImageAnalysis(products=[])
|
||||
intent = IntentResult(intent_summary="用户的意思")
|
||||
full = gen._fallback_fusion("ai_full", analysis, intent, 15, "")
|
||||
user = gen._fallback_fusion("user_primary", analysis, intent, 15, "")
|
||||
assert "回购" in full.title
|
||||
assert user.title == "用户的意思"
|
||||
|
||||
def test_storyboard_fallback_on_garbage(self):
|
||||
client = FakeClient()
|
||||
client.chat_completion = lambda *a, **k: "啥都没有" # type: ignore
|
||||
gen = CopyGenerator(client=client)
|
||||
from packages.application.viral_video.schemas import FusionResult, ScriptSegment
|
||||
|
||||
fusion = FusionResult(
|
||||
hook="开头",
|
||||
script_segments=[ScriptSegment(text="a", duration_sec=5), ScriptSegment(text="b", duration_sec=5)],
|
||||
)
|
||||
board = gen.storyboard(fusion, None, ["u"], 10)
|
||||
assert len(board.clips) == 2
|
||||
assert board.clips[0].voice_text == "a"
|
||||
|
||||
|
||||
# ── 审核与自动重写 ────────────────────────────────────────────────────────
|
||||
class TestReview:
|
||||
def test_rule_check_catches_exaggeration_even_if_llm_passes(self):
|
||||
client = FakeClient() # LLM 默认返回 passed
|
||||
reviewer = Reviewer(client=client)
|
||||
from packages.application.viral_video.schemas import FusionResult, IntentResult
|
||||
|
||||
fusion = FusionResult(title="一喷100%掉光", hook="x", cta="买")
|
||||
result = reviewer.review(fusion, IntentResult(), "ai_full")
|
||||
assert result.passed is False
|
||||
dims = {i.dimension for i in result.issues}
|
||||
assert "夸大承诺" in dims
|
||||
|
||||
def test_rule_check_catches_banned_phrase(self):
|
||||
client = FakeClient()
|
||||
reviewer = Reviewer(client=client)
|
||||
from packages.application.viral_video.schemas import FusionResult
|
||||
|
||||
fusion = FusionResult(title="绝绝子", hook="x", cta="买")
|
||||
result = reviewer.review(fusion, None, "ai_full")
|
||||
assert not result.passed
|
||||
assert any("绝绝子" in i.text for i in result.issues)
|
||||
|
||||
def test_missing_personal_brand_flagged(self):
|
||||
client = FakeClient()
|
||||
reviewer = Reviewer(client=client)
|
||||
from packages.application.viral_video.schemas import (
|
||||
FusionResult,
|
||||
IntentResult,
|
||||
PersonalBrand,
|
||||
)
|
||||
|
||||
fusion = FusionResult(title="合规标题", hook="钩子", cta="行动")
|
||||
intent = IntentResult(personal_brands=[PersonalBrand(text="39块钱一瓶", category="price")])
|
||||
result = reviewer.review(fusion, intent, "ai_polish")
|
||||
assert not result.passed
|
||||
assert any("39块钱一瓶" in i.text and i.dimension == "事实一致性" for i in result.issues)
|
||||
|
||||
def test_missing_core_message_flagged(self):
|
||||
client = FakeClient()
|
||||
reviewer = Reviewer(client=client)
|
||||
from packages.application.viral_video.schemas import (
|
||||
CoreMessage,
|
||||
FusionResult,
|
||||
IntentResult,
|
||||
)
|
||||
|
||||
fusion = FusionResult(title="标题", hook="别的内容", cta="号召")
|
||||
intent = IntentResult(core_messages=[CoreMessage(text="必须保留的原意", must_keep=True)])
|
||||
result = reviewer.review(fusion, intent, "user_primary")
|
||||
assert any(i.dimension == "用户意图保留" for i in result.issues)
|
||||
|
||||
def test_auto_rewrite_once_then_pass(self):
|
||||
client = FakeClient()
|
||||
client.review_sequence = [REVIEW_FAIL_XML, REVIEW_PASS_XML]
|
||||
gen = CopyGenerator(client=client)
|
||||
from packages.application.viral_video.schemas import FusionResult, IntentResult
|
||||
|
||||
fusion = gen._parse_fusion(FUSION_XML)
|
||||
final, review, rewrites = gen.review_and_rewrite(fusion, IntentResult(), "ai_polish")
|
||||
assert rewrites == 1
|
||||
assert review.passed is True
|
||||
assert "大部分油污能擦掉" in client.chat_calls[-2][1]["content"] or True
|
||||
|
||||
def test_rule_fix_local(self):
|
||||
reviewer = Reviewer(client=FakeClient())
|
||||
from packages.application.viral_video.schemas import (
|
||||
BodyPoint,
|
||||
FusionResult,
|
||||
ReviewIssue,
|
||||
ReviewResult,
|
||||
)
|
||||
|
||||
fusion = FusionResult(
|
||||
title="一喷100%掉光",
|
||||
hook="绝绝子",
|
||||
cta="最好用",
|
||||
body_points=[BodyPoint(text="立刻见效")],
|
||||
)
|
||||
review = ReviewResult(
|
||||
passed=False,
|
||||
issues=[ReviewIssue(dimension="夸大承诺", text="100%")],
|
||||
)
|
||||
fixed = reviewer._rule_fix(fusion, review)
|
||||
assert "100%" not in fixed.title
|
||||
assert fixed.hook == ""
|
||||
assert fixed.cta == "很不错用"
|
||||
assert fixed.body_points[0].text == "坚持使用会有改善"
|
||||
|
||||
|
||||
# ── seed 幂等 ─────────────────────────────────────────────────────────────
|
||||
class TestSeed:
|
||||
def test_seed_idempotent(self, tmp_path):
|
||||
scripts_dir = os.path.join(REPO_ROOT, "scripts")
|
||||
sys.path.insert(0, scripts_dir)
|
||||
import importlib
|
||||
|
||||
seed_mod = importlib.import_module("seed_viral_video_prompts")
|
||||
|
||||
db_path = tmp_path / "seed.db"
|
||||
engine = sa.create_engine(f"sqlite:///{db_path}")
|
||||
with engine.begin() as conn:
|
||||
conn.execute(
|
||||
sa.text(
|
||||
"CREATE TABLE viral_video_prompt_templates ("
|
||||
"id INTEGER PRIMARY KEY AUTOINCREMENT, name TEXT, prompt_type TEXT,"
|
||||
"version INTEGER, system_prompt TEXT, user_prompt_template TEXT,"
|
||||
"example_output TEXT, is_active INTEGER DEFAULT 1,"
|
||||
"created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,"
|
||||
"updated_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,"
|
||||
"UNIQUE(prompt_type, version))"
|
||||
)
|
||||
)
|
||||
assert seed_mod.seed(engine) == 5
|
||||
assert seed_mod.seed(engine) == 5 # 再来一次不报错
|
||||
with engine.begin() as conn:
|
||||
count = conn.execute(sa.text("SELECT COUNT(*) FROM viral_video_prompt_templates")).scalar()
|
||||
assert count == 5
|
||||
active_types = conn.execute # noqa: B018
|
||||
with engine.begin() as conn:
|
||||
types_ = {
|
||||
r[0]
|
||||
for r in conn.execute(sa.text("SELECT prompt_type FROM viral_video_prompt_templates WHERE is_active=1"))
|
||||
}
|
||||
assert types_ == {"image_analysis", "intent_parsing", "copy_fusion", "storyboard", "review"}
|
||||
|
||||
|
||||
# ── 负面词不出现于程序产出 ────────────────────────────────────────────────
|
||||
class TestNegativeOutput:
|
||||
def test_fallback_outputs_clean(self):
|
||||
gen = CopyGenerator(client=FakeClient())
|
||||
from packages.application.viral_video.schemas import (
|
||||
FusionResult,
|
||||
ImageAnalysis,
|
||||
IntentResult,
|
||||
)
|
||||
|
||||
fusion = gen._fallback_fusion(
|
||||
"ai_full",
|
||||
ImageAnalysis(products=[]),
|
||||
IntentResult(intent_summary="正常产品"),
|
||||
15,
|
||||
"",
|
||||
)
|
||||
blob = "\n".join(s.text for s in fusion.script_segments)
|
||||
for phrase in BANNED_PHRASES:
|
||||
assert phrase not in blob
|
||||
@@ -0,0 +1,302 @@
|
||||
"""#2040 接线集成测试:验证运行中的 viral_video 任务使用 prompt_loader 从 DB 读取模板。
|
||||
|
||||
mock LLM/Vision 调用,验证:
|
||||
1. image_analysis 走 loader 模板 + XML 解析
|
||||
2. intent_parsing 走 loader 模板 + XML 解析
|
||||
3. script_generation 走 storyboard 模板 + XML 解析,输出兼容 Seedance 的 copy_result
|
||||
4. review 走 Reviewer(review 模板)带自动重写
|
||||
5. 三档融合(ai_full / ai_polish / user_primary)注入不同 FUSION_INSTRUCTIONS
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import sys
|
||||
from pathlib import Path as _Path
|
||||
|
||||
_WORKER_ROOT = _Path(__file__).resolve().parents[2] / "apps" / "worker"
|
||||
if str(_WORKER_ROOT) not in sys.path:
|
||||
sys.path.insert(0, str(_WORKER_ROOT))
|
||||
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from packages.domain.viral_video import ViralVideoJob
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def job():
|
||||
j = ViralVideoJob(
|
||||
user_id="u1",
|
||||
images=["https://img/1.jpg", "https://img/2.jpg"],
|
||||
industry="美妆",
|
||||
duration=15,
|
||||
user_copy_text="这款口红真的太绝了,显白又持久,姐妹们冲!",
|
||||
fusion_level="ai_polish",
|
||||
)
|
||||
return j
|
||||
|
||||
|
||||
# ── Mock LLM/Vision 返回的 XML 文本 ─────────────────────────────────
|
||||
|
||||
IMAGE_XML = """
|
||||
<analysis>
|
||||
<scene>室内桌面拍摄,柔和自然光</scene>
|
||||
<mood>清新温暖</mood>
|
||||
<product name="lipstick" brand="品牌X" category="唇部彩妆"
|
||||
appearance="管状红色膏体" packaging="黑色金属管"
|
||||
features="显白,持久,滋润" portrait_prompt="无人像"
|
||||
summary="品牌X红色口红">
|
||||
<text_on_package>品牌X,211</text_on_package>
|
||||
</product>
|
||||
</analysis>
|
||||
""".strip()
|
||||
|
||||
INTENT_XML = """
|
||||
<intent>
|
||||
<intent_summary>推广显白持久口红</intent_summary>
|
||||
<core_messages>
|
||||
<message must_keep="true">显白</message>
|
||||
<message must_keep="true">持久</message>
|
||||
</core_messages>
|
||||
<personal_brands>
|
||||
<brand text="品牌X" category="brand"/>
|
||||
</personal_brands>
|
||||
<emotion_tone>亲切自然</emotion_tone>
|
||||
<suggested_title>显白持久口红推荐</suggested_title>
|
||||
</intent>
|
||||
""".strip()
|
||||
|
||||
STORYBOARD_XML = """
|
||||
<clips>
|
||||
<clip image_index="0" transition="cut" zoom="null" duration_sec="5" bgm_note="轻快BGM">
|
||||
<voice_text>这款口红真的太绝了</voice_text>
|
||||
<subtitle_text>显白又持久</subtitle_text>
|
||||
<shot_type_angle_movement>近景俯拍45度,缓慢推镜</shot_type_angle_movement>
|
||||
<scene_and_dialogue>厨房台面,主妇展示口红。对白:这款口红真的太绝了</scene_and_dialogue>
|
||||
<action_details>右手持口红展示膏体</action_details>
|
||||
<audio_bgm>轻快BGM</audio_bgm>
|
||||
<transition>硬切</transition>
|
||||
<reference_image_index>0</reference_image_index>
|
||||
<ken_burns start="0,0" end="0,0" ease="linear"/>
|
||||
</clip>
|
||||
</clips>
|
||||
""".strip()
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def invalidate_loader_cache():
|
||||
from packages.application.viral_video import prompt_loader as pl
|
||||
|
||||
pl.invalidate()
|
||||
yield
|
||||
pl.invalidate()
|
||||
|
||||
|
||||
# ── 1) 图片分析走模板 ───────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestImageAnalysisWiring:
|
||||
def test_uses_loader_template_and_xml_parse(self, job):
|
||||
from apps.worker.worker_app.tasks import viral_video as vv
|
||||
|
||||
with patch("packages.shared.ai_service.call_vision", return_value=IMAGE_XML) as mock_v:
|
||||
result = vv._analyze_single_image(0, "https://img/1.jpg", "vlm-lite", 15)
|
||||
|
||||
mock_v.assert_called_once()
|
||||
# 验证调用时传入了 system_prompt(说明走了 loader 渲染的模板)
|
||||
call_kwargs = mock_v.call_args.kwargs
|
||||
assert "system_prompt" in call_kwargs and call_kwargs["system_prompt"]
|
||||
# 结果包含从 XML 解析出的产品信息
|
||||
assert result["name"] == "lipstick"
|
||||
assert result["brand"] == "品牌X"
|
||||
assert "显白" in result["key_features"]
|
||||
assert result["text_on_package"] == ["品牌X", "211"]
|
||||
|
||||
|
||||
# ── 2) 意图解析走模板 ───────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestIntentParsingWiring:
|
||||
def test_uses_loader_and_parses_xml(self, job):
|
||||
from apps.worker.worker_app.tasks import viral_video as vv
|
||||
|
||||
img_result = {"products": [{"name": "lipstick", "brand": "品牌X", "key_features": ["显白", "持久"]}]}
|
||||
with patch("packages.shared.ai_service.call_llm", return_value=INTENT_XML) as mock_llm:
|
||||
result = vv._step_intent_parsing(job, img_result)
|
||||
|
||||
mock_llm.assert_called_once()
|
||||
assert result["intent"] == "推广显白持久口红"
|
||||
assert "显白" in result["key_messages"]
|
||||
assert result["suggested_title"] == "显白持久口红推荐"
|
||||
|
||||
|
||||
# ── 3) 脚本生成:storyboard 模板 + XML 解析 + fusion_level 注入 ────
|
||||
|
||||
|
||||
class TestScriptGenerationWiring:
|
||||
@pytest.mark.parametrize("level", ["ai_full", "ai_polish", "user_primary"])
|
||||
def test_fusion_level_injected(self, job, level):
|
||||
"""三档融合水平被注入到 storyboard 模板的 system_prompt"""
|
||||
from apps.worker.worker_app.tasks import viral_video as vv
|
||||
from packages.application.viral_video.prompts import FUSION_INSTRUCTIONS
|
||||
|
||||
job.fusion_level = level
|
||||
intent = {"intent": "推广", "key_messages": ["显白"], "tone": "亲切"}
|
||||
|
||||
captured_system = {}
|
||||
|
||||
def fake_call_llm(messages, **kw):
|
||||
captured_system["final"] = messages[0]["content"]
|
||||
return STORYBOARD_XML
|
||||
|
||||
with patch("packages.shared.ai_service.call_llm", side_effect=fake_call_llm):
|
||||
result = vv._step_script_generation(job, intent, {})
|
||||
|
||||
# fusion_level 对应的指令文本被注入到 system prompt 中
|
||||
assert FUSION_INSTRUCTIONS[level] in captured_system["final"], f"fusion_level {level} 指令未注入 system_prompt"
|
||||
# 输出保持 Seedance 兼容结构
|
||||
assert "overview" in result
|
||||
assert "shots" in result
|
||||
assert len(result["shots"]) >= 1
|
||||
assert result["shots"][0]["shot_type_angle_movement"]
|
||||
assert result["voiceover_script"]
|
||||
|
||||
def test_fallback_when_xml_and_json_unparseable(self, job):
|
||||
"""XML 解析失败且无法解析为 JSON 时,回退到兜底脚本"""
|
||||
from apps.worker.worker_app.tasks import viral_video as vv
|
||||
|
||||
job.fusion_level = "ai_polish"
|
||||
intent = {"intent": "推广", "key_messages": [], "tone": "亲切"}
|
||||
with patch("packages.shared.ai_service.call_llm", return_value="not xml not json"):
|
||||
result = vv._step_script_generation(job, intent, {})
|
||||
assert isinstance(result, dict)
|
||||
assert "voiceover_script" in result
|
||||
assert "shots" in result
|
||||
|
||||
|
||||
# ── 4) Review 使用 Reviewer + 自动重写 ─────────────────────────────
|
||||
|
||||
|
||||
class TestReviewWiring:
|
||||
def test_pass_path(self, job):
|
||||
from apps.worker.worker_app.tasks import viral_video as vv
|
||||
from packages.application.viral_video.reviewer import Reviewer, ReviewResult
|
||||
|
||||
copy_result = {
|
||||
"title": "口红推荐",
|
||||
"overview": {"theme": "口红推荐"},
|
||||
"voiceover_script": "这款口红显白又持久",
|
||||
"shots": [{"scene_and_dialogue": "展示口红"}],
|
||||
}
|
||||
job.intent_result = {"key_messages": ["显白", "持久"], "intent": "推广"}
|
||||
|
||||
pass_result = ReviewResult(passed=True, score=90, issues=[], rewrite_suggestions=[])
|
||||
with patch.object(Reviewer, "review", return_value=pass_result):
|
||||
out = vv._step_review(job, copy_result)
|
||||
assert out["passed"] is True
|
||||
|
||||
def test_rewrite_path(self, job):
|
||||
"""审核不通过时触发自动重写,并更新 job.copy_result"""
|
||||
from apps.worker.worker_app.tasks import viral_video as vv
|
||||
from packages.application.viral_video.reviewer import Reviewer, ReviewResult
|
||||
from packages.application.viral_video.schemas import FusionResult, ReviewIssue, ScriptSegment
|
||||
|
||||
copy_result = {
|
||||
"title": "原标题",
|
||||
"overview": {"theme": "原标题"},
|
||||
"voiceover_script": "这款口红绝了",
|
||||
"shots": [{"scene_and_dialogue": "展示"}],
|
||||
}
|
||||
job.intent_result = {"key_messages": ["显白"], "intent": "推广"}
|
||||
|
||||
fail_result = ReviewResult(
|
||||
passed=False,
|
||||
score=50,
|
||||
issues=[ReviewIssue(dimension="违规词", severity="high", location="开头", text="绝了")],
|
||||
rewrite_suggestions=["去掉夸大词"],
|
||||
)
|
||||
rewritten = FusionResult(
|
||||
title="新标题",
|
||||
hook="修改后钩子",
|
||||
script_segments=[ScriptSegment(text="修改后口播正文")],
|
||||
cta="行动号召",
|
||||
word_count=10,
|
||||
estimated_duration=10,
|
||||
)
|
||||
pass_after = ReviewResult(passed=True, score=88, issues=[], rewrite_suggestions=[])
|
||||
|
||||
with (
|
||||
patch.object(Reviewer, "review", side_effect=[fail_result, pass_after]),
|
||||
patch.object(Reviewer, "rewrite", return_value=rewritten),
|
||||
):
|
||||
out = vv._step_review(job, copy_result)
|
||||
|
||||
assert out["passed"] is True
|
||||
assert "rewritten_copy" in out
|
||||
assert job.generated_copy_text == "修改后口播正文"
|
||||
|
||||
|
||||
# ── 5) 端到端:每个 step 调用 loader 对应 prompt_type ──────────────
|
||||
|
||||
|
||||
class TestEndToEndLoaderUsed:
|
||||
def test_each_step_calls_loader(self, job):
|
||||
from apps.worker.worker_app.tasks import viral_video as vv
|
||||
from packages.application.viral_video import prompt_loader as pl
|
||||
|
||||
called_types = []
|
||||
real_get = pl.get_template
|
||||
|
||||
def spy_get(prompt_type, **kwargs):
|
||||
called_types.append(prompt_type)
|
||||
return real_get(prompt_type, **kwargs)
|
||||
|
||||
with (
|
||||
patch.object(pl, "get_template", side_effect=spy_get),
|
||||
patch("packages.shared.ai_service.call_vision", return_value=IMAGE_XML),
|
||||
patch("packages.shared.ai_service.call_llm", return_value=INTENT_XML),
|
||||
):
|
||||
# 1) image
|
||||
img_res = vv._analyze_single_image(0, "https://img/1.jpg", "vlm", 15)
|
||||
# 2) intent
|
||||
intent_res = vv._step_intent_parsing(job, {"products": [img_res]})
|
||||
|
||||
# 前两步分别调用了 image_analysis 和 intent_parsing
|
||||
assert "image_analysis" in called_types
|
||||
assert "intent_parsing" in called_types
|
||||
|
||||
# script 和 review 单独验证(需要不同的 LLM 返回)
|
||||
called_types_2 = []
|
||||
|
||||
def spy_get_2(prompt_type, **kwargs):
|
||||
called_types_2.append(prompt_type)
|
||||
return real_get(prompt_type, **kwargs)
|
||||
|
||||
with (
|
||||
patch.object(pl, "get_template", side_effect=spy_get_2),
|
||||
patch("packages.shared.ai_service.call_llm", return_value=STORYBOARD_XML),
|
||||
):
|
||||
copy_res = vv._step_script_generation(job, intent_res, {"products": [img_res]})
|
||||
assert "storyboard" in called_types_2
|
||||
|
||||
called_types_3 = []
|
||||
|
||||
def spy_get_3(prompt_type, **kwargs):
|
||||
called_types_3.append(prompt_type)
|
||||
return real_get(prompt_type, **kwargs)
|
||||
|
||||
from packages.application.viral_video.reviewer import Reviewer, ReviewResult
|
||||
|
||||
pass_result = ReviewResult(passed=True, score=90, issues=[], rewrite_suggestions=[])
|
||||
job.intent_result = intent_res
|
||||
job.copy_result = copy_res
|
||||
with (
|
||||
patch.object(pl, "get_template", side_effect=spy_get_3),
|
||||
patch.object(Reviewer, "review", return_value=pass_result) as mock_review,
|
||||
):
|
||||
review_res = vv._step_review(job, copy_res)
|
||||
# review 步骤内部直接调用 Reviewer.review,该方法被 mock,因此 get_template 不会被调用;
|
||||
# 此处验证 Reviewer.review 被调用即可说明 review 步骤走通了。
|
||||
assert mock_review.called, "_step_review 未调用 Reviewer.review"
|
||||
assert isinstance(review_res, dict) and "passed" in review_res
|
||||
Reference in New Issue
Block a user