diff --git a/alembic/versions/095_viral_video_prompt_templates.py b/alembic/versions/095_viral_video_prompt_templates.py new file mode 100644 index 000000000..bc57f66aa --- /dev/null +++ b/alembic/versions/095_viral_video_prompt_templates.py @@ -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(), + ), + ) diff --git a/packages/adapters/sqlalchemy_impl/models.py b/packages/adapters/sqlalchemy_impl/models.py index 1d27406be..7b3f164c3 100755 --- a/packages/adapters/sqlalchemy_impl/models.py +++ b/packages/adapters/sqlalchemy_impl/models.py @@ -991,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)) diff --git a/packages/adapters/sqlalchemy_impl/viral_video_repository.py b/packages/adapters/sqlalchemy_impl/viral_video_repository.py index f7a76b8e5..fdcf03386 100755 --- a/packages/adapters/sqlalchemy_impl/viral_video_repository.py +++ b/packages/adapters/sqlalchemy_impl/viral_video_repository.py @@ -6,7 +6,6 @@ from sqlalchemy.orm import Session from packages.adapters.sqlalchemy_impl.models import ( ViralVideoJobModel, - ViralVideoPromptTemplateModel, ViralVideoStyleTemplateModel, ) from packages.domain.viral_video import ViralVideoJob, ViralVideoStatus @@ -234,31 +233,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, - } diff --git a/packages/application/viral_video/__init__.py b/packages/application/viral_video/__init__.py new file mode 100644 index 000000000..5e6410348 --- /dev/null +++ b/packages/application/viral_video/__init__.py @@ -0,0 +1 @@ +"""应用层:爆款视频 Prompt 模板系统(#2040)。""" diff --git a/packages/application/viral_video/generator.py b/packages/application/viral_video/generator.py new file mode 100644 index 000000000..0bef9dfe4 --- /dev/null +++ b/packages/application/viral_video/generator.py @@ -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"{node['text']}", "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"{body}", "voice_text"), + subtitle_text=xp.text_of(body and f"{body}", "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) diff --git a/packages/application/viral_video/prompt_loader.py b/packages/application/viral_video/prompt_loader.py new file mode 100644 index 000000000..32af8898c --- /dev/null +++ b/packages/application/viral_video/prompt_loader.py @@ -0,0 +1,133 @@ +"""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 + + +def _load_from_db(prompt_type: str) -> Optional[PromptTemplate]: + if _session_mod.SessionLocal is None: + return None + session = None + try: + session = _session_mod.SessionLocal() + 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) diff --git a/packages/application/viral_video/prompts.py b/packages/application/viral_video/prompts.py new file mode 100644 index 000000000..2eab68e7e --- /dev/null +++ b/packages/application/viral_video/prompts.py @@ -0,0 +1,293 @@ +"""爆款视频 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} + +请严格按下面的标签格式输出,标签名一个都不能改,不要输出任何解释,不要用代码块: + 下面每个产品用一个 标签,属性 name 是产品名、features 是外观特征、position 是 main 或 secondary、image_index 是第几张图(从0开始)。 + 下面每个主要颜色用一个 标签,属性 hex 是色值、name 是颜色名、coverage 是占比小数。 + 用一个标签,属性 has_person、count、gender、age_range、pose、expression 分别描述人物情况。 + 标签写画面整体情绪氛围。 + 下面每处可见文字用一个 标签,属性 text 是文字内容、position 是位置。 + 标签写场景描述。 + 用一个标签,属性 resolution、lighting、composition、blur 描述画质。 + 下面每个卖点用一个 标签。 + +看不到或无法判断的内容,属性值填“无法判断”,布尔值填 false,不要留空标签。""" + +_IMAGE_ANALYSIS_USER = """请分析以下商品图片,共 {image_count} 张。 +所属行业:{industry} +图片地址: +{image_urls} + +按约定的标签格式输出分析结果。""" + +_IMAGE_ANALYSIS_EXAMPLE = """ + + + + + + + +干净、实用 + + + +白底棚拍产品图 + + +针对重油污设计 +大容量625ml +""" + +# ── 模板2:用户文案意图解析(LLM)────────────────────────────────────── +_INTENT_SYSTEM = f"""你负责理解用户的营销意图。用户给的文案可能只是几个关键词、碎碎念或者不完整的短句,你要读懂他真正想讲什么。 + +{GLOBAL_CONSTRAINTS} + +请严格按下面的标签格式输出,不要解释,不要用代码块: + 用用户的语言风格,一句话、30字以内概括核心意图。 + 下面每个核心信息点用一个 标签,属性 must_keep 为 true 或 false、confidence 为 0 到 1 的小数,标签内容写信息点。 + 把用户提到的具体事实——品牌名、价格、人名、地名、时间、产品名——每条用一个 标签,属性 category 取 brand、price、person、place、time、product 之一。这些事实必须原样引用,一个字都不能改。 + 写文案的情绪调性。 + 把你认为缺失、后续生成时需要合理推断的信息,每条用一个 标签;没有就输出空标签。""" + +_INTENT_USER = """用户原始文案:{user_copy_text} +所属行业:{industry} +图片分析结果(供参考): +{image_analysis} + +请理解用户意图,按标签格式输出。""" + +_INTENT_EXAMPLE = """一款厨房去油污神器,喷一喷油污就掉 + +去油污效果好,喷上等几分钟再擦 +适合厨房重油污场景 + + +大公鸡头多功能油污净 +39块钱一瓶 + +亲切、真实、带分享感 + +没有说明具体容量,按图片读出的625ml处理 +""" + +# ── 模板3:文案融合生成(LLM)────────────────────────────────────────── +_FUSION_SYSTEM = """你负责为短视频生成营销文案。请按思维链分步完成:先定人设和目标客户,再找卖点,再搭结构,再安排情绪,最后写行动号召,不要一步到位乱写。 + +{fusion_instruction} + +{global_constraints} + +{negative_rules} + +请严格按下面的标签格式输出,不要解释,不要用代码块: + 视频标题。 +<hook> 开头3秒钩子,5到15字。 +<body_points> 每个要点用一个 <point> 标签,属性 elaboration 是展开说明、image_index 是对应第几张图(从0开始),标签内容写要点。 +<cta> 口语化的行动号召。 +<script_segments> 每段配音用一个 <segment> 标签,属性 duration_sec 是秒数、image_index 是对应图片,标签内容写配音文案。 +<word_count> 配音总字数,只写数字。 +<estimated_duration> 预计时长秒数,只写数字。 + +用户在 personal_brands 中提到的品牌名、价格、人名、地名、时间、产品名等事实信息,必须原样出现在文案里,一个字都不能改。""" + +_FUSION_USER = """所属行业:{industry} +目标客户:{target_customer} +营销目的:{marketing_purpose} +视频时长:{duration}秒 +图片分析结果: +{image_analysis} +用户意图解析结果: +{intent_result} + +请按标签格式生成文案。""" + +_FUSION_EXAMPLE = """<title>厨房重油污,别再用洗洁精硬擦了 +这油污,我真的忍很久了 + +大公鸡头油污净去油快 +39块钱一瓶,性价比高 + +厨房油污重的,真的可以试一瓶 + +这油污我真的忍很久了,用洗洁精擦半天都没用 +后来换了这个大公鸡头油污净,喷上等几分钟,一擦就干净 +39块钱625ml,厨房重油污的可以试一瓶 + +58 +13""" + +# ── 模板4:编导级分镜(LLM)──────────────────────────────────────────── +_STORYBOARD_SYSTEM = f"""你是短视频编导,负责把文案拆成可拍摄的分镜。 + +工作方式: +1. 按文案的 script_segments 顺序分配镜头。 +2. 每个镜头确定画面、运镜、时长、配音和字幕。 +3. 检查所有镜头时长加起来接近目标时长,误差不超过2秒。 +4. image_index 必须在已上传图片范围内,第一张主图必须用在第一个镜头。 + +{GLOBAL_CONSTRAINTS} + +请严格按下面的标签格式输出,不要解释,不要用代码块: + 下面每个镜头用一个 标签,属性 image_index 是图片序号(从0开始)、transition 取 fade、cut、zoom_in、slide_left、dissolve、wipe 之一、zoom 取 in、out 或 null、duration_sec 是该镜头秒数、bgm_note 是该段BGM情绪。每个 里面包含: + 该镜头配音文本; + 字幕文本,可与配音一致或更精简; + 用一个空标签,属性 start、end 写“x,y”坐标、ease 写缓动方式;不需要运镜时坐标相同。""" + +_STORYBOARD_USER = """目标时长:{duration}秒 +上传图片数量:{image_count}张(第1张是主图/封面) +文案内容: +{fusion_result} +图片分析结果: +{image_analysis} + +请按标签格式输出分镜。""" + +_STORYBOARD_EXAMPLE = """ + +这油污我真的忍很久了 +这油污忍很久了 + + + +后来换了大公鸡头油污净,喷上等几分钟,一擦就干净 +喷上等几分钟,一擦就干净 + + + +39块钱625ml,厨房重油污的可以试一瓶 +39元625ml,可以试一瓶 + + +""" + +# ── 模板5:文案审核(LLM)────────────────────────────────────────────── +_REVIEW_SYSTEM = f"""你是短视频文案合规审核员,从6个维度逐条检查文案: +1. 违规词:有没有平台禁用词、敏感词。 +2. 夸大承诺:有没有“包治百病”“100%有效”“保证赚钱”等绝对化、夸大表述。 +3. 事实一致性:有没有编造价格、数据、认证,或者用户没提到的产品特性。 +4. 用户意图保留:在 ai_polish 和 user_primary 模式下,core_messages 中 must_keep=true 的点是否都保留了。 +5. 结构完整性:标题、钩子、正文、行动号召是否齐全。 +6. 语气人设:是否符合选定的人设语气,有没有“家人们谁懂啊”“绝绝子”“宝子们”等套路词。 + +{GLOBAL_CONSTRAINTS} + +请严格按下面的标签格式输出,不要解释,不要用代码块: + 整体是否通过,只写 true 或 false。 + 每个问题用一个 标签,属性 dimension 是维度名、severity 取 error 或 warning、location 是问题所在(如 hook、body_points、cta),标签内容写问题描述;没有问题就输出空标签。 + 每条具体修改建议用一个 标签;没有就输出空标签。""" + +_REVIEW_USER = """本次创作模式:{fusion_level} +待审核文案: +{fusion_result} +用户意图解析(用于核对核心信息是否保留): +{intent_result} + +请按6个维度审核,按标签格式输出。""" + +_REVIEW_EXAMPLE = """false + +出现了“一喷100%掉光”的绝对化表述,违反广告法 +用户强调的“39块钱”没有保留 + + +把“一喷100%掉光”改为“喷上等几分钟,大部分油污能擦掉” +在结尾补回“39块钱625ml” +""" + + +# 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, + }, +] diff --git a/packages/application/viral_video/reviewer.py b/packages/application/viral_video/reviewer.py new file mode 100644 index 000000000..c5d28951c --- /dev/null +++ b/packages/application/viral_video/reviewer.py @@ -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) diff --git a/packages/application/viral_video/schemas.py b/packages/application/viral_video/schemas.py new file mode 100644 index 000000000..da0126885 --- /dev/null +++ b/packages/application/viral_video/schemas.py @@ -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 = "" diff --git a/packages/application/viral_video/xml_parser.py b/packages/application/viral_video/xml_parser.py new file mode 100644 index 000000000..8bf33face --- /dev/null +++ b/packages/application/viral_video/xml_parser.py @@ -0,0 +1,104 @@ +"""XML 标签式输出解析器(替代 json.loads)。 + +LLM 按 ``内容`` 输出,本模块解析,解析失败不抛异常, +由调用方走规则 fallback。采用栈式扫描,嵌套标签全部可提取(内外层都保留)。 +""" + +from __future__ import annotations + +import re +from html import unescape +from typing import Optional + +_OPEN_RE = re.compile(r"<(?P[\w-]+)(?P(?:\s(?:[^>]*?\S)?)?)(?P/?)>") +_CLOSE_RE = re.compile(r"[\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(" 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 diff --git a/scripts/seed_viral_video_prompts.py b/scripts/seed_viral_video_prompts.py new file mode 100644 index 000000000..4cabada9f --- /dev/null +++ b/scripts/seed_viral_video_prompts.py @@ -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()) diff --git a/tests/unit/test_viral_video_prompt_system.py b/tests/unit/test_viral_video_prompt_system.py new file mode 100644 index 000000000..bef59fcef --- /dev/null +++ b/tests/unit/test_viral_video_prompt_system.py @@ -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 = """ + + + + +干净实用 + +白底棚拍 + +去油快625ml大容量""" + +INTENT_XML = """厨房去油污神器 + +去油污效果好 +适合重油污 + +39块钱一瓶 +亲切真实 +容量按625ml""" + +FUSION_XML = """厨房重油污别硬擦了 +这油污忍很久了 +大公鸡头去油快 +重油污的可以试一瓶 + +这油污忍很久了 +大公鸡头油污净喷上等几分钟一擦就净 +39块钱一瓶可以试一下 + +5213""" + +STORYBOARD_XML = """ + +这油污忍很久了 +油污忍很久 + + + +大公鸡头喷上等几分钟一擦就净 +一擦就净 + + +""" + +REVIEW_FAIL_XML = """false +出现绝对化表述 +改为“大部分油污能擦掉”""" + +REVIEW_PASS_XML = """true + +""" + +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("厨房重油污别硬擦了", "我把厨房油污全搞定了") + elif mode == "user_primary": + xml = xml.replace("厨房重油污别硬擦了", "油污净使用分享") + 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('') + 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