Compare commits
1 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| c5e0cbba6a |
+4
-13
@@ -211,24 +211,15 @@ COSYVOICE_CLONE_MODEL=voice-enrollment
|
||||
# 用于 AI 文案生成、智能剪辑等需要大模型能力的场景
|
||||
|
||||
DOUBAO_API_KEY=your-doubao-api-key
|
||||
DOUBAO_MODEL=doubao-seed-2-1-pro-260915
|
||||
DOUBAO_FAST_MODEL=doubao-seed-2-1-lite-260915
|
||||
DOUBAO_MODEL=doubao-seed-1-6-250615
|
||||
DOUBAO_FAST_MODEL=doubao-1-5-pro-32k-250115
|
||||
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-seed-2-1-pro-260915
|
||||
DOUBAO_VISION_LITE_MODEL=doubao-seed-2-1-lite-260915
|
||||
DOUBAO_VISION_MODEL=doubao-1-5-vision-pro-250328
|
||||
DOUBAO_VISION_LITE_MODEL=doubao-1-5-vision-lite-250315
|
||||
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,14 +1187,6 @@ 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 }}"
|
||||
|
||||
@@ -1,2 +0,0 @@
|
||||
Mon Oct 5 04:09:11 PM CST 2026
|
||||
2198 lite/pro并行竞速 (commit 9699a1f) — CI rebuild trigger Mon Oct 5 08:09:11 AM UTC 2026
|
||||
@@ -1,87 +0,0 @@
|
||||
"""viral_video 动态积分定价 + 积分字段从 Integer 改为 Float (#2151)
|
||||
|
||||
Revision ID: 093
|
||||
Revises: 092_viral_video_heartbeat
|
||||
Create Date: 2026-10-02
|
||||
"""
|
||||
|
||||
import sqlalchemy as sa
|
||||
|
||||
from alembic import op
|
||||
|
||||
revision = "093"
|
||||
down_revision = "092_viral_video_heartbeat"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
conn = op.get_bind()
|
||||
inspector = sa.inspect(conn)
|
||||
|
||||
# 1) points_accounts 三列 Integer -> Float
|
||||
pa_cols = {c["name"]: c for c in inspector.get_columns("points_accounts")}
|
||||
for col in ("balance", "total_earned", "total_spent"):
|
||||
if col in pa_cols:
|
||||
op.alter_column(
|
||||
"points_accounts",
|
||||
col,
|
||||
existing_type=sa.Integer(),
|
||||
type_=sa.Float(),
|
||||
existing_nullable=False,
|
||||
)
|
||||
|
||||
# 2) points_transactions amount/balance_after Integer -> Float
|
||||
pt_cols = {c["name"]: c for c in inspector.get_columns("points_transactions")}
|
||||
for col in ("amount", "balance_after"):
|
||||
if col in pt_cols:
|
||||
op.alter_column(
|
||||
"points_transactions",
|
||||
col,
|
||||
existing_type=sa.Integer(),
|
||||
type_=sa.Float(),
|
||||
existing_nullable=False,
|
||||
)
|
||||
|
||||
# 3) users.points_balance Integer -> Float
|
||||
user_cols = {c["name"]: c for c in inspector.get_columns("users")}
|
||||
if "points_balance" in user_cols:
|
||||
op.alter_column(
|
||||
"users",
|
||||
"points_balance",
|
||||
existing_type=sa.Integer(),
|
||||
type_=sa.Float(),
|
||||
existing_nullable=False,
|
||||
)
|
||||
|
||||
# 4) viral_video_jobs.credits_cost Integer -> Float
|
||||
vv_cols = {c["name"]: c for c in inspector.get_columns("viral_video_jobs")}
|
||||
if "credits_cost" in vv_cols:
|
||||
op.alter_column(
|
||||
"viral_video_jobs",
|
||||
"credits_cost",
|
||||
existing_type=sa.Integer(),
|
||||
type_=sa.Float(),
|
||||
existing_nullable=False,
|
||||
)
|
||||
|
||||
# 5) viral_video_jobs 新增列
|
||||
if "video_resolution" not in vv_cols:
|
||||
op.add_column(
|
||||
"viral_video_jobs",
|
||||
sa.Column("video_resolution", sa.String(20), nullable=False, server_default="720p"),
|
||||
)
|
||||
if "credits_prepaid" not in vv_cols:
|
||||
op.add_column(
|
||||
"viral_video_jobs",
|
||||
sa.Column("credits_prepaid", sa.Float(), nullable=False, server_default="0"),
|
||||
)
|
||||
if "credits_transaction_id" not in vv_cols:
|
||||
op.add_column(
|
||||
"viral_video_jobs",
|
||||
sa.Column("credits_transaction_id", sa.String(36), nullable=False, server_default=""),
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
pass
|
||||
@@ -1,31 +0,0 @@
|
||||
"""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")
|
||||
@@ -1,102 +0,0 @@
|
||||
"""爆款视频 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(),
|
||||
),
|
||||
)
|
||||
@@ -1,61 +0,0 @@
|
||||
"""功能计费积分字段(爆款/对口型/智能剪辑 DB 化计费)。
|
||||
|
||||
给 gpu_lipsync_tasks / generation_tasks / lipsync_jobs 三张表加积分字段:
|
||||
- credits_prepaid: 提交任务时预扣积分
|
||||
- credits_cost: 最终结算积分
|
||||
- credits_transaction_id: 预扣流水 ID
|
||||
|
||||
注意:feature_pricing_configs 配置表由 xiaoxia-admin 侧 migration 建立,
|
||||
本仓库只读,不在此创建。
|
||||
|
||||
Revision ID: 096_feature_billing_fields
|
||||
Revises: 095_viral_video_prompt_templates
|
||||
Create Date: 2026-10-05
|
||||
"""
|
||||
|
||||
import sqlalchemy as sa
|
||||
|
||||
from alembic import op
|
||||
|
||||
revision = "096_feature_billing_fields"
|
||||
down_revision = "095_viral_video_prompt_templates"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
_TABLES = ("gpu_lipsync_tasks", "generation_tasks", "lipsync_jobs")
|
||||
_COLUMNS = (
|
||||
("credits_prepaid", sa.Float(), "0"),
|
||||
("credits_cost", sa.Float(), "0"),
|
||||
("credits_transaction_id", sa.String(36), ""),
|
||||
)
|
||||
|
||||
|
||||
def _table_exists(conn, name: str) -> bool:
|
||||
return name in sa.inspect(conn).get_table_names()
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
conn = op.get_bind()
|
||||
for table in _TABLES:
|
||||
if not _table_exists(conn, table):
|
||||
continue
|
||||
existing = {c["name"] for c in sa.inspect(conn).get_columns(table)}
|
||||
for col_name, col_type, default in _COLUMNS:
|
||||
if col_name in existing:
|
||||
continue
|
||||
op.add_column(
|
||||
table,
|
||||
sa.Column(col_name, col_type, nullable=False, server_default=default),
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
conn = op.get_bind()
|
||||
for table in _TABLES:
|
||||
if not _table_exists(conn, table):
|
||||
continue
|
||||
existing = {c["name"] for c in sa.inspect(conn).get_columns(table)}
|
||||
for col_name, _col_type, _default in _COLUMNS:
|
||||
if col_name not in existing:
|
||||
continue
|
||||
op.drop_column(table, col_name)
|
||||
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
@@ -1,222 +0,0 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""099: AI 模型路由层 seed — 补齐缺失模型和能力配置.
|
||||
|
||||
幂等:所有 INSERT 先检查存在性。
|
||||
- ai_models: 补齐 qwen3.7-plus, seedream, seedance, embedding, wan3.0 等
|
||||
- ai_capability_configs: 补齐 image_generation, video_generation, embedding
|
||||
- 更新已有 capability 的 lite_model_id
|
||||
"""
|
||||
|
||||
import json
|
||||
|
||||
import sqlalchemy as sa
|
||||
|
||||
from alembic import op
|
||||
|
||||
revision = "099_ai_model_router_seed"
|
||||
down_revision = "098_viral_video_image_analysis_v5"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
conn = op.get_bind()
|
||||
|
||||
# CI 环境下 ai_models 表可能尚未创建(由 ORM 自动建表,非 migration)
|
||||
# 如果表不存在则跳过 seed,由应用启动时 ORM 建表后首次访问时生效
|
||||
table_check = conn.execute(sa.text("SELECT to_regclass('public.ai_models')")).scalar()
|
||||
if not table_check:
|
||||
# ai_models 表不存在,跳过所有 seed(CI 环境)
|
||||
return
|
||||
|
||||
# ── 1. 补齐 ai_models 缺失记录 ────────────────────────────────────────────
|
||||
existing_models = {
|
||||
row[0]
|
||||
for row in conn.execute(
|
||||
sa.text("SELECT model_key FROM ai_models WHERE deleted_at IS NULL")
|
||||
).fetchall()
|
||||
}
|
||||
|
||||
# 从已有 active 记录获取 API key(复用,不硬编码)
|
||||
dashscope_key_row = conn.execute(
|
||||
sa.text(
|
||||
"SELECT api_key FROM ai_models WHERE provider='dashscope' AND deleted_at IS NULL AND api_key IS NOT NULL AND api_key != '' LIMIT 1"
|
||||
)
|
||||
).first()
|
||||
dashscope_key = dashscope_key_row[0] if dashscope_key_row else ""
|
||||
|
||||
volcengine_key_row = conn.execute(
|
||||
sa.text(
|
||||
"SELECT api_key FROM ai_models WHERE provider='volcengine' AND deleted_at IS NULL AND api_key IS NOT NULL AND api_key != '' LIMIT 1"
|
||||
)
|
||||
).first()
|
||||
volcengine_key = volcengine_key_row[0] if volcengine_key_row else ""
|
||||
|
||||
new_models = [
|
||||
{
|
||||
"model_key": "qwen3.7-plus",
|
||||
"name": "通义千问3.7 Plus(VLM 兜底)",
|
||||
"provider": "dashscope",
|
||||
"api_key": dashscope_key,
|
||||
"api_base": "https://dashscope.aliyuncs.com/compatible-mode/v1",
|
||||
"description": "阿里云百炼 Qwen3.7 Plus 多模态模型,用于 VLM 兜底分析",
|
||||
},
|
||||
{
|
||||
"model_key": "doubao-seedream-5-0-flash-260915",
|
||||
"name": "Seedream 5.0 Flash(图片生成)",
|
||||
"provider": "volcengine",
|
||||
"api_key": volcengine_key,
|
||||
"api_base": "https://ark.cn-beijing.volces.com/api/v3",
|
||||
"description": "火山引擎 Seedream 5.0 Flash 文生图模型",
|
||||
},
|
||||
{
|
||||
"model_key": "doubao-seedance-2-5-260628",
|
||||
"name": "Seedance 2.5(视频生成)",
|
||||
"provider": "volcengine",
|
||||
"api_key": volcengine_key,
|
||||
"api_base": "https://ark.cn-beijing.volces.com/api/v3",
|
||||
"description": "火山引擎 Seedance 2.5 图/文生视频模型",
|
||||
},
|
||||
{
|
||||
"model_key": "doubao-embedding-vision-251215",
|
||||
"name": "豆包多模态向量嵌入",
|
||||
"provider": "volcengine",
|
||||
"api_key": volcengine_key,
|
||||
"api_base": "https://ark.cn-beijing.volces.com/api/v3",
|
||||
"description": "火山引擎豆包多模态向量嵌入模型",
|
||||
},
|
||||
{
|
||||
"model_key": "wan3.0-video",
|
||||
"name": "Wan 3.0 视频生成",
|
||||
"provider": "dashscope",
|
||||
"api_key": dashscope_key,
|
||||
"api_base": "https://dashscope.aliyuncs.com/api/v1",
|
||||
"description": "阿里云百炼 Wan 3.0 视频生成模型",
|
||||
},
|
||||
{
|
||||
"model_key": "doubao-seed-2-1-pro-260915",
|
||||
"name": "豆包 Seed 2.1 Pro(高精度推理)",
|
||||
"provider": "volcengine",
|
||||
"api_key": volcengine_key,
|
||||
"api_base": "https://ark.cn-beijing.volces.com/api/v3",
|
||||
"description": "火山引擎豆包 Seed 2.1 Pro 深度思考+多模态",
|
||||
},
|
||||
]
|
||||
|
||||
for m in new_models:
|
||||
if m["model_key"] not in existing_models:
|
||||
conn.execute(
|
||||
sa.text(
|
||||
"""
|
||||
INSERT INTO ai_models (id, name, provider, model_key, api_key, api_base, description, status, is_default, usage_today, created_at, updated_at)
|
||||
VALUES (gen_random_uuid()::text, :name, :provider, :model_key, :api_key, :api_base, :description, 'active', false, 0, now(), now())
|
||||
"""
|
||||
),
|
||||
m,
|
||||
)
|
||||
|
||||
# ── 2. 补齐 ai_capability_configs 缺失项 ──────────────────────────────────
|
||||
cap_table_check = conn.execute(sa.text("SELECT to_regclass('public.ai_capability_configs')")).scalar()
|
||||
if not cap_table_check:
|
||||
return
|
||||
|
||||
existing_caps = {
|
||||
row[0]
|
||||
for row in conn.execute(
|
||||
sa.text("SELECT capability_key FROM ai_capability_configs")
|
||||
).fetchall()
|
||||
}
|
||||
|
||||
def _get_model_id(model_key: str) -> str | None:
|
||||
row = conn.execute(
|
||||
sa.text(
|
||||
"SELECT id FROM ai_models WHERE model_key = :key AND deleted_at IS NULL AND status = 'active' LIMIT 1"
|
||||
),
|
||||
{"key": model_key},
|
||||
).first()
|
||||
return row[0] if row else None
|
||||
|
||||
# image_generation
|
||||
if "image_generation" not in existing_caps:
|
||||
mid = _get_model_id("doubao-seedream-5-0-flash-260915")
|
||||
if mid:
|
||||
conn.execute(
|
||||
sa.text(
|
||||
"""
|
||||
INSERT INTO ai_capability_configs (id, capability_key, capability_name, primary_model_id, timeout_seconds, max_retries, concurrency, extra_params, is_enabled, created_at, updated_at)
|
||||
VALUES (gen_random_uuid()::text, :ck, :cn, :pm, 60, 1, 2, :ep, true, now(), now())
|
||||
"""
|
||||
),
|
||||
{
|
||||
"ck": "image_generation",
|
||||
"cn": "图片生成(Seedream)",
|
||||
"pm": mid,
|
||||
"ep": json.dumps({"size": "1K"}),
|
||||
},
|
||||
)
|
||||
|
||||
# video_generation
|
||||
if "video_generation" not in existing_caps:
|
||||
mid = _get_model_id("doubao-seedance-2-5-260628")
|
||||
fb_mid = _get_model_id("wan3.0-video")
|
||||
if mid:
|
||||
conn.execute(
|
||||
sa.text(
|
||||
"""
|
||||
INSERT INTO ai_capability_configs (id, capability_key, capability_name, primary_model_id, fallback_model_id, timeout_seconds, max_retries, concurrency, extra_params, is_enabled, created_at, updated_at)
|
||||
VALUES (gen_random_uuid()::text, :ck, :cn, :pm, :fm, 600, 1, 1, :ep, true, now(), now())
|
||||
"""
|
||||
),
|
||||
{
|
||||
"ck": "video_generation",
|
||||
"cn": "视频生成(Seedance/Wan)",
|
||||
"pm": mid,
|
||||
"fm": fb_mid,
|
||||
"ep": json.dumps({}),
|
||||
},
|
||||
)
|
||||
|
||||
# embedding
|
||||
if "embedding" not in existing_caps:
|
||||
mid = _get_model_id("doubao-embedding-vision-251215")
|
||||
if mid:
|
||||
conn.execute(
|
||||
sa.text(
|
||||
"""
|
||||
INSERT INTO ai_capability_configs (id, capability_key, capability_name, primary_model_id, timeout_seconds, max_retries, concurrency, extra_params, is_enabled, created_at, updated_at)
|
||||
VALUES (gen_random_uuid()::text, :ck, :cn, :pm, 30, 2, 5, :ep, true, now(), now())
|
||||
"""
|
||||
),
|
||||
{
|
||||
"ck": "embedding",
|
||||
"cn": "向量嵌入",
|
||||
"pm": mid,
|
||||
"ep": json.dumps({}),
|
||||
},
|
||||
)
|
||||
|
||||
# ── 3. 更新 image_analysis 的 lite_model_id ─────────────────────────────
|
||||
lite_model_id = _get_model_id("qwen3.8-flash")
|
||||
if lite_model_id:
|
||||
conn.execute(
|
||||
sa.text(
|
||||
"UPDATE ai_capability_configs SET lite_model_id = :lite WHERE capability_key = 'image_analysis' AND lite_model_id IS NULL"
|
||||
),
|
||||
{"lite": lite_model_id},
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
conn = op.get_bind()
|
||||
# 安全检查表是否存在
|
||||
table_check = conn.execute(sa.text("SELECT to_regclass('public.ai_models')")).scalar()
|
||||
if not table_check:
|
||||
return
|
||||
conn.execute(
|
||||
sa.text("DELETE FROM ai_capability_configs WHERE capability_key IN ('image_generation', 'video_generation', 'embedding')")
|
||||
)
|
||||
conn.execute(
|
||||
sa.text(
|
||||
"DELETE FROM ai_models WHERE model_key IN ('qwen3.7-plus', 'doubao-seedream-5-0-flash-260915', 'doubao-seedance-2-5-260628', 'doubao-embedding-vision-251215', 'wan3.0-video', 'doubao-seed-2-1-pro-260915') AND deleted_at IS NULL"
|
||||
)
|
||||
)
|
||||
@@ -1,107 +0,0 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""100: 修正已有 capability 的模型绑定.
|
||||
|
||||
幂等:仅当 primary_model_id 当前绑定到旧模型 (doubao-seed-1-6) 时才更新,
|
||||
避免覆盖用户在后台的自定义配置。
|
||||
|
||||
- 更新 5 个 LLM capability (intent_parsing, copy_fusion, storyboard, copy_review, asset_classify)
|
||||
的 primary_model_id 从 doubao-seed-1-6 改为 doubao-seed-2-1-pro-260915
|
||||
- 更新 image_analysis 的 primary/lite/fallback 模型绑定
|
||||
"""
|
||||
|
||||
import sqlalchemy as sa
|
||||
|
||||
from alembic import op
|
||||
|
||||
revision = "100_fix_capability_model_bindings"
|
||||
down_revision = "099_ai_model_router_seed"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
conn = op.get_bind()
|
||||
|
||||
# Check tables exist
|
||||
table_check = conn.execute(sa.text("SELECT to_regclass('public.ai_models')")).scalar()
|
||||
if not table_check:
|
||||
return
|
||||
|
||||
config_table_check = conn.execute(sa.text("SELECT to_regclass('public.ai_capability_configs')")).scalar()
|
||||
if not config_table_check:
|
||||
return
|
||||
|
||||
# Look up model IDs by model_key (not hardcoded UUIDs)
|
||||
pro_model_row = conn.execute(
|
||||
sa.text(
|
||||
"SELECT id FROM ai_models WHERE model_key = 'doubao-seed-2-1-pro-260915' AND deleted_at IS NULL LIMIT 1"
|
||||
)
|
||||
).first()
|
||||
if not pro_model_row:
|
||||
return
|
||||
pro_model_id = pro_model_row[0]
|
||||
|
||||
old_model_row = conn.execute(
|
||||
sa.text("SELECT id FROM ai_models WHERE model_key = 'doubao-seed-1-6-250615' LIMIT 1")
|
||||
).first()
|
||||
old_model_id = old_model_row[0] if old_model_row else None
|
||||
|
||||
llm_capabilities = [
|
||||
"intent_parsing",
|
||||
"copy_fusion",
|
||||
"storyboard",
|
||||
"copy_review",
|
||||
"asset_classify",
|
||||
]
|
||||
|
||||
for cap_key in llm_capabilities:
|
||||
if old_model_id:
|
||||
conn.execute(
|
||||
sa.text(
|
||||
"UPDATE ai_capability_configs SET primary_model_id = :new_id, updated_at = NOW() "
|
||||
"WHERE capability_key = :cap_key AND primary_model_id = :old_id"
|
||||
),
|
||||
{"new_id": pro_model_id, "old_id": old_model_id, "cap_key": cap_key},
|
||||
)
|
||||
|
||||
# Update image_analysis
|
||||
qwen38_row = conn.execute(
|
||||
sa.text("SELECT id FROM ai_models WHERE model_key = 'qwen3.8-flash' AND deleted_at IS NULL LIMIT 1")
|
||||
).first()
|
||||
qwen37_row = conn.execute(
|
||||
sa.text("SELECT id FROM ai_models WHERE model_key = 'qwen3.7-plus' AND deleted_at IS NULL LIMIT 1")
|
||||
).first()
|
||||
|
||||
if qwen38_row and qwen37_row:
|
||||
qwen38_id = qwen38_row[0]
|
||||
qwen37_id = qwen37_row[0]
|
||||
|
||||
current_ia = conn.execute(
|
||||
sa.text(
|
||||
"SELECT primary_model_id, lite_model_id, fallback_model_id "
|
||||
"FROM ai_capability_configs WHERE capability_key = 'image_analysis'"
|
||||
)
|
||||
).first()
|
||||
|
||||
if current_ia:
|
||||
current_primary, current_lite, current_fallback = current_ia
|
||||
updates = {}
|
||||
if current_primary != qwen38_id:
|
||||
updates["primary_model_id"] = qwen38_id
|
||||
if current_lite != qwen38_id:
|
||||
updates["lite_model_id"] = qwen38_id
|
||||
if current_fallback != qwen37_id:
|
||||
updates["fallback_model_id"] = qwen37_id
|
||||
|
||||
if updates:
|
||||
set_clause = ", ".join([f"{k} = :{k}" for k in updates.keys()])
|
||||
set_clause += ", updated_at = NOW()"
|
||||
updates["cap_key"] = "image_analysis"
|
||||
conn.execute(
|
||||
sa.text(f"UPDATE ai_capability_configs SET {set_clause} WHERE capability_key = :cap_key"),
|
||||
updates,
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
pass
|
||||
@@ -1,153 +0,0 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""101: 补齐 qwen-vl-plus 视觉模型并修正 image_analysis 绑定与 max_tokens.
|
||||
|
||||
背景:
|
||||
- qwen-vl-plus 做图片识别时返回 JSON 约 500-600 tokens,旧硬编码
|
||||
max_tokens=350 导致 JSON 被截断、解析失败返回"未识别"。
|
||||
- 代码侧已移除硬编码,改由 capability 的 DB 配置决定 max_tokens。
|
||||
|
||||
幂等:
|
||||
- qwen-vl-plus 已存在则不插入;
|
||||
- 仅当 image_analysis 当前 primary_model 不是 qwen-vl-plus 时才更新绑定,
|
||||
避免覆盖后台手动配置。
|
||||
"""
|
||||
|
||||
import sqlalchemy as sa
|
||||
|
||||
from alembic import op
|
||||
|
||||
revision = "101_qwen_vl_plus_and_max_tokens"
|
||||
down_revision = "100_fix_capability_model_bindings"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
conn = op.get_bind()
|
||||
|
||||
models_table = conn.execute(sa.text("SELECT to_regclass('public.ai_models')")).scalar()
|
||||
if not models_table:
|
||||
return
|
||||
|
||||
caps_table = conn.execute(sa.text("SELECT to_regclass('public.ai_capability_configs')")).scalar()
|
||||
if not caps_table:
|
||||
return
|
||||
|
||||
# ── c. 补全其他 capability 的 max_tokens 默认值(幂等)──────────────────
|
||||
# 放在 image_analysis 特定逻辑之前,确保任何分支 return 都不会跳过本段。
|
||||
# 仅在当前值为 NULL 或过小 (<100) 时更新,不覆盖已有合理配置。
|
||||
# embedding / tts / voice_clone 不走 chat 接口,无需设置。
|
||||
default_max_tokens = {
|
||||
"intent_parsing": 500,
|
||||
"copy_fusion": 2500,
|
||||
"storyboard": 4000,
|
||||
"copy_review": 1000,
|
||||
"asset_classify": 500,
|
||||
"image_generation": 500,
|
||||
"video_generation": 500,
|
||||
}
|
||||
for cap_key, mt in default_max_tokens.items():
|
||||
conn.execute(
|
||||
sa.text(
|
||||
"UPDATE ai_capability_configs "
|
||||
"SET max_tokens = :mt, updated_at = now() "
|
||||
"WHERE capability_key = :key "
|
||||
"AND (max_tokens IS NULL OR max_tokens < 100)"
|
||||
),
|
||||
{"mt": mt, "key": cap_key},
|
||||
)
|
||||
|
||||
# ── a. 确保 qwen-vl-plus 模型存在 ────────────────────────────────────────
|
||||
conn.execute(sa.text("""
|
||||
INSERT INTO ai_models (id, name, provider, model_key, api_key, api_base,
|
||||
description, status, is_default, usage_today,
|
||||
created_at, updated_at)
|
||||
SELECT gen_random_uuid()::text,
|
||||
'通义千问VL Plus',
|
||||
'dashscope',
|
||||
'qwen-vl-plus',
|
||||
COALESCE(
|
||||
(SELECT api_key FROM ai_models
|
||||
WHERE provider = 'dashscope' AND deleted_at IS NULL
|
||||
AND api_key IS NOT NULL AND api_key != ''
|
||||
LIMIT 1),
|
||||
''
|
||||
),
|
||||
'https://dashscope.aliyuncs.com/compatible-mode/v1',
|
||||
'阿里云视觉理解模型(图片识别/分析)',
|
||||
'active', false, 0, now(), now()
|
||||
WHERE NOT EXISTS (
|
||||
SELECT 1 FROM ai_models
|
||||
WHERE model_key = 'qwen-vl-plus' AND deleted_at IS NULL
|
||||
)
|
||||
"""))
|
||||
|
||||
qwen_vl_row = conn.execute(
|
||||
sa.text(
|
||||
"SELECT id FROM ai_models WHERE model_key = 'qwen-vl-plus' "
|
||||
"AND deleted_at IS NULL AND status = 'active' LIMIT 1"
|
||||
)
|
||||
).first()
|
||||
if not qwen_vl_row:
|
||||
return
|
||||
qwen_vl_id = qwen_vl_row[0]
|
||||
|
||||
qwen37_row = conn.execute(
|
||||
sa.text(
|
||||
"SELECT id FROM ai_models WHERE model_key = 'qwen3.7-plus' "
|
||||
"AND deleted_at IS NULL AND status = 'active' LIMIT 1"
|
||||
)
|
||||
).first()
|
||||
qwen37_id = qwen37_row[0] if qwen37_row else None
|
||||
|
||||
# ── b. 仅当当前 primary 不是 qwen-vl-plus 时修正绑定与 max_tokens ───────
|
||||
current = conn.execute(
|
||||
sa.text(
|
||||
"SELECT primary_model_id, lite_model_id, fallback_model_id, max_tokens "
|
||||
"FROM ai_capability_configs WHERE capability_key = 'image_analysis'"
|
||||
)
|
||||
).first()
|
||||
|
||||
if current is None:
|
||||
# capability 不存在则创建
|
||||
conn.execute(
|
||||
sa.text("""
|
||||
INSERT INTO ai_capability_configs
|
||||
(id, capability_key, capability_name, primary_model_id,
|
||||
lite_model_id, fallback_model_id, timeout_seconds,
|
||||
max_retries, max_tokens, concurrency, extra_params,
|
||||
is_enabled, created_at, updated_at)
|
||||
VALUES (gen_random_uuid()::text, 'image_analysis', '图片分析',
|
||||
:primary, :primary, :fallback, 30, 1, 1000, 2,
|
||||
'{}'::jsonb, true, now(), now())
|
||||
"""),
|
||||
{"primary": qwen_vl_id, "fallback": qwen37_id},
|
||||
)
|
||||
return
|
||||
|
||||
current_primary = current[0]
|
||||
if current_primary == qwen_vl_id:
|
||||
# 已经绑定 qwen-vl-plus:视为后台/数据迁移已处理,不覆盖任何配置
|
||||
return
|
||||
|
||||
set_parts = [
|
||||
"primary_model_id = :vl_id",
|
||||
"lite_model_id = :vl_id",
|
||||
"max_tokens = 1000",
|
||||
"updated_at = now()",
|
||||
]
|
||||
params: dict = {"vl_id": qwen_vl_id}
|
||||
if qwen37_id is not None:
|
||||
set_parts.insert(2, "fallback_model_id = :qwen37_id")
|
||||
params["qwen37_id"] = qwen37_id
|
||||
|
||||
conn.execute(
|
||||
sa.text(
|
||||
"UPDATE ai_capability_configs SET " + ", ".join(set_parts) + " WHERE capability_key = 'image_analysis'"
|
||||
),
|
||||
params,
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
pass
|
||||
@@ -1,36 +0,0 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""102: image_analysis max_tokens 1200 -> 1500.
|
||||
|
||||
v6 prompt 更长、字段更多,旧 max_tokens 容易截断 JSON。
|
||||
仅在 image_analysis 当前 max_tokens < 1500 时更新(幂等,不覆盖后台已调到 >=1500 的配置)。
|
||||
"""
|
||||
|
||||
import sqlalchemy as sa
|
||||
|
||||
from alembic import op
|
||||
|
||||
revision = "102_image_analysis_max_tokens_1500"
|
||||
down_revision = "101_qwen_vl_plus_and_max_tokens"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
conn = op.get_bind()
|
||||
|
||||
caps_table = conn.execute(sa.text("SELECT to_regclass('public.ai_capability_configs')")).scalar()
|
||||
if not caps_table:
|
||||
return
|
||||
|
||||
conn.execute(
|
||||
sa.text(
|
||||
"UPDATE ai_capability_configs "
|
||||
"SET max_tokens = 1500, updated_at = now() "
|
||||
"WHERE capability_key = 'image_analysis' "
|
||||
"AND (max_tokens IS NULL OR max_tokens < 1500)"
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
pass
|
||||
@@ -1,194 +0,0 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""image_analysis v7 prompt + max_tokens 3000 + max_retries 3
|
||||
|
||||
Revision ID: 103_v7_prompt_and_tokens_3000
|
||||
Revises: 102_image_analysis_max_tokens_1500
|
||||
Create Date: 2026-10-07
|
||||
|
||||
变更:
|
||||
1. 插入v7精简prompt(~1KB,v6 ~4.5KB,删除few-shot/冗长规则,减少输出token占用),设为active
|
||||
2. v6停用(is_active=False),保留历史
|
||||
3. image_analysis capability: max_tokens 1500→3000,max_retries 1→3
|
||||
|
||||
ai_capability_configs 由应用 create_all 创建,全新 alembic-only 库可能不存在,
|
||||
故第3步做 to_regclass 守卫(同 102)。
|
||||
"""
|
||||
|
||||
from sqlalchemy import text
|
||||
|
||||
from alembic import op
|
||||
|
||||
revision = "103_v7_prompt_and_tokens_3000"
|
||||
down_revision = "102_image_analysis_max_tokens_1500"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
V7_SYSTEM = """# 角色
|
||||
你是一位专业的图片分析师,擅长准确识别图片中的场景、人物、物体、文字、氛围。
|
||||
|
||||
# 任务
|
||||
对用户上传的图片逐张分析,描述你看到的内容,输出JSON格式。
|
||||
|
||||
## 技能
|
||||
|
||||
### 技能1:判断图片类型
|
||||
判断图片属于哪种类型,type字段填对应的英文值:
|
||||
- 商品图(product):单个或多个商品、产品包装
|
||||
- 门店场景图(store):店铺内部、门头招牌、货架陈列
|
||||
- 人物图(person):人物形象、穿搭造型、肖像照片
|
||||
- 风景图(scene):风景、动物、美食、街景
|
||||
- 其他(other):以上都不是
|
||||
|
||||
### 技能2:描述通用信息
|
||||
不管什么图都要描述:
|
||||
- type:图片类型,填product/store/person/scene/other其中一个
|
||||
- scene:一句话描述场景,例如"理疗养生店内部,摆着多张理疗床和产品货架"
|
||||
- mood:整体氛围,2-4个词,例如"整洁专业"、"热闹温馨"
|
||||
- colors:主要颜色,最多5个,写具体颜色名(亮红色/米白色/深蓝色,不写笼统的红色蓝色)
|
||||
- visible_text:图片里看到的文字,说明什么字、在什么位置,最多5条;没看到就空数组
|
||||
- lighting:光线情况,例如"明亮柔光"、"自然光"、"室内暖黄灯"
|
||||
- composition:怎么拍的,例如"居中特写"、"中景平视"、"俯拍"
|
||||
- has_person:有没有人,true或false
|
||||
|
||||
### 技能3:描述门店场景
|
||||
如果是门店场景图(type="store"),还要描述:
|
||||
- store_type:什么类型的店,例如"养生馆"、"便利店"、"餐饮店"、"母婴店"
|
||||
- brand_signage:招牌上写了什么字、有什么品牌标识
|
||||
- visual_elements:看到哪些显眼的东西(招牌样式、灯光、货架、商品陈列、海报、收银台等),最多8个
|
||||
- product_categories:看到哪些品类的商品,例如"饮料零食"、"养生产品"
|
||||
- promotion_elements:有没有促销活动(打折海报、满减吊旗等),没有就空数组
|
||||
- atmosphere:店内什么氛围,例如"亲民生活化"、"老字号专业感"
|
||||
- cleanliness:店内干净程度,例如"干净整洁"、"货架整齐"
|
||||
- 看到顾客或店员要描述他们在做什么,has_person填true
|
||||
|
||||
### 技能4:描述商品
|
||||
如果是商品图(type="product"),逐个商品描述:
|
||||
- product_name:商品名称,尽量具体,例如"OMO奥妙除菌除螨洗衣液";看不出来填null
|
||||
- brand:什么牌子,看不出来填null
|
||||
- category:类目,从以下选一个:服饰鞋包/美妆/数码/食品/家居清洁/母婴/配饰/其他
|
||||
- package_type:什么包装,例如"瓶装"、"盒装"、"罐装"、"袋装"、"多瓶装"
|
||||
- package_color:包装主要颜色,写具体色(亮红色不写红色)
|
||||
- body_shape:瓶身或包装形状,例如"圆润胖瓶"、"竖款带把手瓶身"
|
||||
- label_design:标签设计,例如"红色标签印白色品牌logo"
|
||||
- key_text_on_package:包装上最显眼的文字(品牌名、功能词、卖点词),最多5个
|
||||
- product_features:包装特征,3-6个短语,包含颜色、瓶盖、形状、标签图案
|
||||
- key_selling_points:核心卖点,1-3个短语
|
||||
|
||||
### 技能5:描述人物
|
||||
如果是人物图(type="person"),描述:
|
||||
- person_count:几个人
|
||||
- gender:性别(男/女/无法判断)
|
||||
- age_range:年龄段(儿童/青少年/青年/中年/老年/无法判断)
|
||||
- outfit_style:穿搭风格,例如"休闲日常"、"通勤商务"、"街头潮流"
|
||||
- upper_wear:上装(颜色+款式+材质),穿裙装不填
|
||||
- lower_wear:下装(颜色+款式+版型),穿裙装不填
|
||||
- dress_wear:裙装描述,穿上下装不填
|
||||
- outerwear:外套
|
||||
- shoes:鞋子
|
||||
- bag:包袋,没有填null
|
||||
- accessories:配饰(眼镜/帽子/项链/耳环/手表/手链/围巾/腰带等),没有填空数组
|
||||
- hairstyle:发型
|
||||
- makeup:妆容,男生或看不出填null
|
||||
- expression:表情,例如"微笑看镜头"、"冷酷无表情"
|
||||
- pose:姿势动作,例如"身直立正对镜头"、"单手撩发"
|
||||
- body_type:身材,例如"纤细苗条"、"高挑身材"、"丰满匀称"
|
||||
- portrait_prompt:80-150字详细描述人物形象(后面用来AI生成肖像图),要写清年龄段、穿搭完整细节、发型发色、妆容、表情、姿势、场景、光线、风格感觉,语言要有画面感
|
||||
|
||||
### 技能6:描述风景
|
||||
如果是风景图(type="scene"),描述:
|
||||
- scene_type:什么场景,例如"自然风景"、"城市街景"、"动物"、"美食"
|
||||
- main_subject:画面主体是什么
|
||||
- key_elements:关键元素,最多8个
|
||||
- environment_objects:周围环境物体,最多8个
|
||||
- atmosphere:整体氛围,例如"秋日慵懒氛围感"、"清新自然氧气感"
|
||||
- 有人物就描述人物特征
|
||||
|
||||
## 限制
|
||||
- 只输出JSON,不要任何解释文字,不要markdown代码块包裹,不要写"好的""以下是分析结果"这种废话
|
||||
- 颜色写具体色调(亮红色/米白色/深蓝色/翠绿色),不写笼统词汇
|
||||
- 瓶身、包装、招牌上的文字尽量识别出来(品牌名、功能词、卖点词)
|
||||
- 多个商品、多个人物分开描述,不要合并
|
||||
- 看不出来、不确定的字段填null或空数组,布尔值填true/false,绝对不要瞎编
|
||||
- 确保JSON格式合法,所有大括号、中括号、引号正确闭合
|
||||
- 数组字段控制数量:colors最多5个,visible_text最多5条,visual_elements最多8个,accessories最多10个"""
|
||||
V7_USER = "请分析这张图片,按系统消息的JSON结构输出。"
|
||||
|
||||
|
||||
def _capability_table_exists(bind) -> bool:
|
||||
return bool(bind.execute(text("SELECT to_regclass('public.ai_capability_configs')")).scalar())
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
bind = op.get_bind()
|
||||
# 1. 停用旧的active image_analysis prompt(含v6)
|
||||
bind.execute(
|
||||
text(
|
||||
"UPDATE viral_video_prompt_templates SET is_active = FALSE "
|
||||
"WHERE prompt_type = 'image_analysis' AND is_active = TRUE"
|
||||
)
|
||||
)
|
||||
# 2. 幂等插入v7(存在则更新并重新激活)
|
||||
existing = bind.execute(
|
||||
text("SELECT id FROM viral_video_prompt_templates " "WHERE prompt_type = 'image_analysis' AND version = 7")
|
||||
).fetchone()
|
||||
if existing:
|
||||
bind.execute(
|
||||
text(
|
||||
"UPDATE viral_video_prompt_templates SET is_active = TRUE, "
|
||||
"system_prompt = :sys, user_prompt_template = :usr, "
|
||||
"name = 'v7 精简结构化分析', updated_at = NOW() "
|
||||
"WHERE prompt_type = 'image_analysis' AND version = 7"
|
||||
),
|
||||
{"sys": V7_SYSTEM, "usr": V7_USER},
|
||||
)
|
||||
else:
|
||||
bind.execute(
|
||||
text(
|
||||
"INSERT INTO viral_video_prompt_templates "
|
||||
"(prompt_type, version, name, system_prompt, user_prompt_template, "
|
||||
"is_active, created_at, updated_at) "
|
||||
"VALUES ('image_analysis', 7, 'v7 精简结构化分析', "
|
||||
":sys, :usr, TRUE, NOW(), NOW())"
|
||||
),
|
||||
{"sys": V7_SYSTEM, "usr": V7_USER},
|
||||
)
|
||||
# 3. capability max_tokens=3000、max_retries=3(表不存在则跳过)
|
||||
if _capability_table_exists(bind):
|
||||
bind.execute(
|
||||
text(
|
||||
"UPDATE ai_capability_configs SET max_tokens = 3000, "
|
||||
"updated_at = NOW() "
|
||||
"WHERE capability_key = 'image_analysis' AND "
|
||||
"(max_tokens IS NULL OR max_tokens < 3000)"
|
||||
)
|
||||
)
|
||||
bind.execute(
|
||||
text(
|
||||
"UPDATE ai_capability_configs SET max_retries = 3, updated_at = NOW() "
|
||||
"WHERE capability_key = 'image_analysis' AND "
|
||||
"(max_retries IS NULL OR max_retries < 3)"
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
bind = op.get_bind()
|
||||
# 删除v7
|
||||
bind.execute(
|
||||
text("DELETE FROM viral_video_prompt_templates " "WHERE prompt_type = 'image_analysis' AND version = 7")
|
||||
)
|
||||
# 恢复v6为active
|
||||
bind.execute(
|
||||
text(
|
||||
"UPDATE viral_video_prompt_templates SET is_active = TRUE "
|
||||
"WHERE prompt_type = 'image_analysis' AND version = 6"
|
||||
)
|
||||
)
|
||||
# tokens/retries回退
|
||||
if _capability_table_exists(bind):
|
||||
bind.execute(
|
||||
text(
|
||||
"UPDATE ai_capability_configs SET max_tokens = 1500, max_retries = 1, "
|
||||
"updated_at = NOW() WHERE capability_key = 'image_analysis'"
|
||||
)
|
||||
)
|
||||
@@ -1,242 +0,0 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""image_analysis v8 prompt + storyboard v3 prompt - 用户端展示格式 markdown 控制
|
||||
|
||||
Revision ID: 104_v8_display_markdown
|
||||
Revises: 103_v7_prompt_and_tokens_3000
|
||||
Create Date: 2026-10-07
|
||||
|
||||
变更:
|
||||
1. image_analysis v8: 在 v7 基础上 system_prompt 末尾追加「## 用户端展示格式」章节,
|
||||
要求 VLM 在每张图的 JSON 里输出 summary_markdown 字段(markdown 格式的图片描述),
|
||||
v8 设 is_active=true,v7 设 is_active=false。
|
||||
2. storyboard v3: 在 v2 基础上 system_prompt 追加要求 LLM 在 copy_result 中
|
||||
输出 copy_display_markdown 字段(markdown 格式的完整文案展示),
|
||||
v3 设 is_active=true,v2 设 is_active=false。
|
||||
"""
|
||||
|
||||
from sqlalchemy import text
|
||||
|
||||
from alembic import op
|
||||
|
||||
revision = "104_v8_display_markdown"
|
||||
down_revision = "103_v7_prompt_and_tokens_3000"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
# ── v8 追加的 system prompt 内容 ──────────────────────────────────────
|
||||
V8_SYSTEM_APPEND = """
|
||||
|
||||
## 用户端展示格式
|
||||
|
||||
对于每张分析的图片,在 JSON 中额外输出一个 **summary_markdown** 字段,用 markdown 格式写出给用户看的图片描述。
|
||||
|
||||
格式要求(根据图片类型自适应):
|
||||
|
||||
**商品图(type=product)**示例:
|
||||
### 商品名称
|
||||
**品牌**:品牌名 | **类目**:服饰鞋包/美妆/数码/...
|
||||
**核心特征**
|
||||
- 特征1:描述
|
||||
- 特征2:描述
|
||||
**外观**:颜色+材质+设计描述
|
||||
**包装**:包装类型描述
|
||||
**文字信息**:包装上看到的文字
|
||||
|
||||
**门店场景图(type=store)**示例:
|
||||
### 门店名称/类型
|
||||
**类型**:奶茶店/便利店/养生馆/...
|
||||
**品牌标识**:招牌文字描述
|
||||
**环境氛围**:店内整体感觉
|
||||
**陈列亮点**
|
||||
- 亮点1
|
||||
- 亮点2
|
||||
**氛围**:亲民/专业/时尚/...
|
||||
|
||||
**人物图(type=person)**示例:
|
||||
### 人物描述
|
||||
**形象**:年龄段 + 风格
|
||||
**穿搭**
|
||||
- 上装:颜色+款式
|
||||
- 下装:颜色+款式
|
||||
- 配饰:...
|
||||
**气质**:表情+姿势+整体感觉
|
||||
|
||||
**风景/场景图(type=scene)**示例:
|
||||
### 场景名称
|
||||
**类型**:自然风景/城市街景/动物/美食
|
||||
**主体**:画面主要元素
|
||||
**氛围**:整体感觉描述
|
||||
|
||||
要求:
|
||||
- 内容真实具体,从实际图片分析得出
|
||||
- 用 markdown 语法:**加粗**、列表、标题
|
||||
- 控制在 100-200 字
|
||||
- 不要编造图片中没有的信息
|
||||
"""
|
||||
|
||||
# ── storyboard v3 追加的 system prompt 内容 ──────────────────────────
|
||||
V3_STORYBOARD_APPEND = """
|
||||
|
||||
## 用户端展示格式
|
||||
|
||||
在输出分镜脚本的同时,在顶层输出一个 **copy_display_markdown** 字段(用 XML 标签 <copy_display_markdown> 包裹),用 markdown 格式写出完整文案展示。
|
||||
|
||||
格式示例:
|
||||
# 标题/主题
|
||||
|
||||
## 整体概要
|
||||
一句话描述视频内容
|
||||
|
||||
## 分镜预览
|
||||
|
||||
### 镜头1(0-3秒)
|
||||
**景别**:近景俯拍,缓慢推镜
|
||||
**画面**:场景描述
|
||||
**台词**:口播文本
|
||||
**动作**:人物动作描述
|
||||
|
||||
### 镜头2(3-9秒)
|
||||
...
|
||||
|
||||
## 完整口播
|
||||
完整口播文案文本
|
||||
|
||||
要求:
|
||||
- 把所有分镜按时间顺序整理成易读的格式
|
||||
- 用 markdown 语法组织,**加粗**标签、##二级标题、列表等
|
||||
- 控制在 300-500 字
|
||||
- 让用户一眼看懂视频会拍成什么样
|
||||
"""
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
bind = op.get_bind()
|
||||
|
||||
# ── 1. image_analysis v8 ──────────────────────────────────────────
|
||||
# 停用所有 active image_analysis prompt
|
||||
bind.execute(
|
||||
text(
|
||||
"UPDATE viral_video_prompt_templates SET is_active = FALSE "
|
||||
"WHERE prompt_type = 'image_analysis' AND is_active = TRUE"
|
||||
)
|
||||
)
|
||||
|
||||
# 读取 v7 的 prompt 内容作为基础
|
||||
v7_row = bind.execute(
|
||||
text(
|
||||
"SELECT system_prompt, user_prompt_template, COALESCE(example_output, '') "
|
||||
"FROM viral_video_prompt_templates "
|
||||
"WHERE prompt_type = 'image_analysis' "
|
||||
"ORDER BY version DESC LIMIT 1"
|
||||
)
|
||||
).fetchone()
|
||||
|
||||
if v7_row:
|
||||
v7_system = v7_row[0] or ""
|
||||
v8_system = v7_system + V8_SYSTEM_APPEND
|
||||
v8_user = v7_row[1] or "{image_url}"
|
||||
v8_example = v7_row[2] or ""
|
||||
|
||||
# 幂等:已有 v8 则更新,否则插入
|
||||
existing_v8 = bind.execute(
|
||||
text("SELECT id FROM viral_video_prompt_templates " "WHERE prompt_type = 'image_analysis' AND version = 8")
|
||||
).fetchone()
|
||||
if existing_v8:
|
||||
bind.execute(
|
||||
text(
|
||||
"UPDATE viral_video_prompt_templates SET is_active = TRUE, "
|
||||
"system_prompt = :sys, user_prompt_template = :usr, "
|
||||
"example_output = :ex, name = 'v8 用户端展示格式', "
|
||||
"updated_at = NOW() "
|
||||
"WHERE prompt_type = 'image_analysis' AND version = 8"
|
||||
),
|
||||
{"sys": v8_system, "usr": v8_user, "ex": v8_example},
|
||||
)
|
||||
else:
|
||||
bind.execute(
|
||||
text(
|
||||
"INSERT INTO viral_video_prompt_templates "
|
||||
"(prompt_type, version, name, system_prompt, user_prompt_template, "
|
||||
"example_output, is_active, created_at, updated_at) "
|
||||
"VALUES ('image_analysis', 8, 'v8 用户端展示格式', "
|
||||
":sys, :usr, :ex, TRUE, NOW(), NOW())"
|
||||
),
|
||||
{"sys": v8_system, "usr": v8_user, "ex": v8_example},
|
||||
)
|
||||
|
||||
# ── 2. storyboard v3 ─────────────────────────────────────────────
|
||||
# 停用所有 active storyboard prompt
|
||||
bind.execute(
|
||||
text(
|
||||
"UPDATE viral_video_prompt_templates SET is_active = FALSE "
|
||||
"WHERE prompt_type = 'storyboard' AND is_active = TRUE"
|
||||
)
|
||||
)
|
||||
|
||||
# 读取当前 storyboard prompt
|
||||
sb_row = bind.execute(
|
||||
text(
|
||||
"SELECT system_prompt, user_prompt_template, COALESCE(example_output, '') "
|
||||
"FROM viral_video_prompt_templates "
|
||||
"WHERE prompt_type = 'storyboard' "
|
||||
"ORDER BY version DESC LIMIT 1"
|
||||
)
|
||||
).fetchone()
|
||||
|
||||
if sb_row:
|
||||
sb_system = sb_row[0] or ""
|
||||
v3_system = sb_system + V3_STORYBOARD_APPEND
|
||||
v3_user = sb_row[1] or ""
|
||||
v3_example = sb_row[2] or ""
|
||||
|
||||
existing_v3 = bind.execute(
|
||||
text("SELECT id FROM viral_video_prompt_templates " "WHERE prompt_type = 'storyboard' AND version = 3")
|
||||
).fetchone()
|
||||
if existing_v3:
|
||||
bind.execute(
|
||||
text(
|
||||
"UPDATE viral_video_prompt_templates SET is_active = TRUE, "
|
||||
"system_prompt = :sys, user_prompt_template = :usr, "
|
||||
"example_output = :ex, name = 'v3 用户端展示格式', "
|
||||
"updated_at = NOW() "
|
||||
"WHERE prompt_type = 'storyboard' AND version = 3"
|
||||
),
|
||||
{"sys": v3_system, "usr": v3_user, "ex": v3_example},
|
||||
)
|
||||
else:
|
||||
bind.execute(
|
||||
text(
|
||||
"INSERT INTO viral_video_prompt_templates "
|
||||
"(prompt_type, version, name, system_prompt, user_prompt_template, "
|
||||
"example_output, is_active, created_at, updated_at) "
|
||||
"VALUES ('storyboard', 3, 'v3 用户端展示格式', "
|
||||
":sys, :usr, :ex, TRUE, NOW(), NOW())"
|
||||
),
|
||||
{"sys": v3_system, "usr": v3_user, "ex": v3_example},
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
bind = op.get_bind()
|
||||
|
||||
# 删除 v8
|
||||
bind.execute(
|
||||
text("DELETE FROM viral_video_prompt_templates " "WHERE prompt_type = 'image_analysis' AND version = 8")
|
||||
)
|
||||
# 恢复 v7 active
|
||||
bind.execute(
|
||||
text(
|
||||
"UPDATE viral_video_prompt_templates SET is_active = TRUE, updated_at = NOW() "
|
||||
"WHERE prompt_type = 'image_analysis' AND version = 7"
|
||||
)
|
||||
)
|
||||
|
||||
# 删除 v3
|
||||
bind.execute(text("DELETE FROM viral_video_prompt_templates " "WHERE prompt_type = 'storyboard' AND version = 3"))
|
||||
# 恢复 storyboard v2 active
|
||||
bind.execute(
|
||||
text(
|
||||
"UPDATE viral_video_prompt_templates SET is_active = TRUE, updated_at = NOW() "
|
||||
"WHERE prompt_type = 'storyboard' AND version = 2"
|
||||
)
|
||||
)
|
||||
@@ -1,116 +0,0 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""image_analysis v8 + storyboard v3 叙述优先重写版(架构大简化)
|
||||
|
||||
Revision ID: 105_narration_first
|
||||
Revises: 104_v8_display_markdown
|
||||
Create Date: 2026-10-07
|
||||
|
||||
变更:
|
||||
1. image_analysis v8:用「叙述优先」版整体替换 104 的 append 版——VLM 主交付物是
|
||||
自然叙述 summary_markdown,结构化字段仅保留 type/name/brand/has_person,
|
||||
顶层 products 改名 images;v8 active,其余 image_analysis 全部 deactivate。
|
||||
2. storyboard v3:整体替换为风格重写版(口播口语化、画面有画面感、
|
||||
copy_display_markdown 流畅叙述);v3 active,其余 storyboard deactivate。
|
||||
3. intent_parsing 类型模板全部 deactivate(意图解析步骤已删除)。
|
||||
模板内容直接取自 packages.application.viral_video.prompts.DEFAULT_TEMPLATES,
|
||||
保证代码默认值与 DB seed 完全一致。
|
||||
"""
|
||||
|
||||
from sqlalchemy import text
|
||||
|
||||
from alembic import op
|
||||
from packages.application.viral_video.prompts import DEFAULT_TEMPLATES
|
||||
|
||||
revision = "105_narration_first"
|
||||
down_revision = "104_v8_display_markdown"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def _tpl(prompt_type: str, version: int) -> dict:
|
||||
for t in DEFAULT_TEMPLATES:
|
||||
if t["prompt_type"] == prompt_type and t["version"] == version:
|
||||
return t
|
||||
raise RuntimeError("default template missing: %s v%s" % (prompt_type, version))
|
||||
|
||||
|
||||
def _upsert(bind, t: dict) -> None:
|
||||
existing = bind.execute(
|
||||
text("SELECT id FROM viral_video_prompt_templates " "WHERE prompt_type = :pt AND version = :ver"),
|
||||
{"pt": t["prompt_type"], "ver": t["version"]},
|
||||
).fetchone()
|
||||
params = {
|
||||
"pt": t["prompt_type"],
|
||||
"ver": t["version"],
|
||||
"name": t["name"],
|
||||
"sys": t["system_prompt"],
|
||||
"usr": t["user_prompt_template"],
|
||||
"ex": t.get("example_output", "") or "",
|
||||
}
|
||||
if existing:
|
||||
bind.execute(
|
||||
text(
|
||||
"UPDATE viral_video_prompt_templates SET name = :name, "
|
||||
"system_prompt = :sys, user_prompt_template = :usr, "
|
||||
"example_output = :ex, is_active = TRUE, updated_at = NOW() "
|
||||
"WHERE prompt_type = :pt AND version = :ver"
|
||||
),
|
||||
params,
|
||||
)
|
||||
else:
|
||||
bind.execute(
|
||||
text(
|
||||
"INSERT INTO viral_video_prompt_templates "
|
||||
"(prompt_type, version, name, system_prompt, user_prompt_template, "
|
||||
"example_output, is_active, created_at, updated_at) "
|
||||
"VALUES (:pt, :ver, :name, :sys, :usr, :ex, TRUE, NOW(), NOW())"
|
||||
),
|
||||
params,
|
||||
)
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
bind = op.get_bind()
|
||||
|
||||
# 1. image_analysis:停用全部后写入叙述优先 v8
|
||||
bind.execute(
|
||||
text("UPDATE viral_video_prompt_templates SET is_active = FALSE " "WHERE prompt_type = 'image_analysis'")
|
||||
)
|
||||
_upsert(bind, _tpl("image_analysis", 8))
|
||||
|
||||
# 2. storyboard:停用全部后写入重写版 v3
|
||||
bind.execute(text("UPDATE viral_video_prompt_templates SET is_active = FALSE " "WHERE prompt_type = 'storyboard'"))
|
||||
_upsert(bind, _tpl("storyboard", 3))
|
||||
|
||||
# 3. intent_parsing 已废弃:全部停用
|
||||
bind.execute(
|
||||
text("UPDATE viral_video_prompt_templates SET is_active = FALSE " "WHERE prompt_type = 'intent_parsing'")
|
||||
)
|
||||
|
||||
# 4. review 模板确保 active
|
||||
bind.execute(text("UPDATE viral_video_prompt_templates SET is_active = TRUE " "WHERE prompt_type = 'review'"))
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
bind = op.get_bind()
|
||||
# 恢复 104 的 v8/v3 无法重建(内容已替换),仅把版本 active 状态回退:
|
||||
# 停用新版,尝试恢复 v7 / v2
|
||||
bind.execute(
|
||||
text(
|
||||
"UPDATE viral_video_prompt_templates SET is_active = FALSE "
|
||||
"WHERE prompt_type IN ('image_analysis','storyboard') "
|
||||
"AND version IN (8, 3)"
|
||||
)
|
||||
)
|
||||
bind.execute(
|
||||
text(
|
||||
"UPDATE viral_video_prompt_templates SET is_active = TRUE "
|
||||
"WHERE prompt_type = 'image_analysis' AND version = 7"
|
||||
)
|
||||
)
|
||||
bind.execute(
|
||||
text(
|
||||
"UPDATE viral_video_prompt_templates SET is_active = TRUE "
|
||||
"WHERE prompt_type = 'storyboard' AND version = 2"
|
||||
)
|
||||
)
|
||||
@@ -1,43 +0,0 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""system_settings 正式建表(#2246)
|
||||
|
||||
Revision ID: 106_system_settings
|
||||
Revises: 105_narration_first
|
||||
Create Date: 2026-10-08
|
||||
|
||||
system_settings 表此前在 staging 手工创建(对应 034 占位迁移),
|
||||
此处补正式迁移保证其他环境一致。CREATE TABLE/INDEX 使用 IF NOT EXISTS,
|
||||
对已手工建表的环境幂等。
|
||||
"""
|
||||
|
||||
from alembic import op
|
||||
|
||||
revision = "106_system_settings"
|
||||
down_revision = "105_narration_first"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.execute("""
|
||||
CREATE TABLE IF NOT EXISTS system_settings (
|
||||
id VARCHAR(36) NOT NULL,
|
||||
setting_key VARCHAR(100) NOT NULL,
|
||||
setting_value TEXT,
|
||||
setting_type VARCHAR(20) NOT NULL,
|
||||
description VARCHAR(255) NOT NULL DEFAULT '',
|
||||
is_public BOOLEAN NOT NULL DEFAULT FALSE,
|
||||
updated_by VARCHAR(36),
|
||||
category VARCHAR(50) NOT NULL DEFAULT 'general',
|
||||
created_at TIMESTAMP WITH TIME ZONE NOT NULL DEFAULT NOW(),
|
||||
updated_at TIMESTAMP WITH TIME ZONE NOT NULL DEFAULT NOW(),
|
||||
CONSTRAINT pk_system_settings PRIMARY KEY (id),
|
||||
CONSTRAINT uq_system_settings_setting_key UNIQUE (setting_key)
|
||||
)
|
||||
""")
|
||||
op.execute("CREATE INDEX IF NOT EXISTS ix_system_settings_category " "ON system_settings (category)")
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.execute("DROP INDEX IF EXISTS ix_system_settings_category")
|
||||
op.execute("DROP TABLE IF EXISTS system_settings")
|
||||
@@ -1,4 +1,3 @@
|
||||
from app.api.routes.admin.ditto_emotion import router as admin_ditto_emotion_router
|
||||
from app.api.routes.ai import router as ai_router
|
||||
from app.api.routes.ai_avatar_render import router as ai_avatar_render_router
|
||||
from app.api.routes.asset_diagnosis import router as asset_diagnosis_router
|
||||
@@ -243,6 +242,3 @@ api_router.include_router(
|
||||
tags=["GPU Worker"],
|
||||
)
|
||||
api_router.include_router(viral_video_router, prefix="/viral-video", tags=["爆款视频"])
|
||||
|
||||
# #2246:后台 Ditto 表情配置(router 自带 /admin/ditto-emotion 前缀)
|
||||
api_router.include_router(admin_ditto_emotion_router)
|
||||
|
||||
@@ -1 +0,0 @@
|
||||
"""后台管理路由(#2246 起)."""
|
||||
@@ -1,163 +0,0 @@
|
||||
"""Ditto 数字人表情后台配置 — #2246.
|
||||
|
||||
路由前缀 /api/v1/admin/ditto-emotion,全部使用 _verify_internal_api_key 鉴权
|
||||
(X-API-Key header)。仅开放 5 项白名单配置:
|
||||
- ditto_emotion_enabled / ditto_emotion_model / ditto_emotion_temperature
|
||||
- ditto_emotion_prompt / ditto_blend_frames
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from fastapi import APIRouter, Depends
|
||||
from pydantic import BaseModel
|
||||
|
||||
from packages.application.system_config_service import get_system_config_service
|
||||
from packages.config import get_api_settings
|
||||
from packages.domain.system_setting import (
|
||||
SETTING_TYPE_BOOL,
|
||||
SETTING_TYPE_FLOAT,
|
||||
SETTING_TYPE_INT,
|
||||
SETTING_TYPE_STRING,
|
||||
)
|
||||
|
||||
from ..auth import _verify_internal_api_key
|
||||
|
||||
router = APIRouter(
|
||||
prefix="/admin/ditto-emotion",
|
||||
tags=["Admin"],
|
||||
dependencies=[Depends(_verify_internal_api_key)],
|
||||
)
|
||||
|
||||
MODEL_OPTIONS = [
|
||||
"doubao-seed-2-1-lite-250915",
|
||||
"doubao-seed-2-1-pro-250915",
|
||||
"deepseek-v3",
|
||||
]
|
||||
|
||||
# key → (类型, 分类)
|
||||
_WHITELIST: dict[str, str] = {
|
||||
"ditto_emotion_enabled": SETTING_TYPE_BOOL,
|
||||
"ditto_emotion_model": SETTING_TYPE_STRING,
|
||||
"ditto_emotion_temperature": SETTING_TYPE_FLOAT,
|
||||
"ditto_emotion_prompt": SETTING_TYPE_STRING,
|
||||
"ditto_blend_frames": SETTING_TYPE_INT,
|
||||
}
|
||||
|
||||
_DESCRIPTIONS: dict[str, str] = {
|
||||
"ditto_emotion_enabled": "LLM 情绪分析开关。关闭时回退到原有关键词匹配模式,不影响正常出片。",
|
||||
"ditto_emotion_model": "用于分析文案情绪的大模型。",
|
||||
"ditto_emotion_temperature": "模型温度,0-1,越低越稳定保守。",
|
||||
"ditto_emotion_prompt": "情绪分析提示词,核心调优入口,必须包含 {文案} 占位符。",
|
||||
"ditto_blend_frames": "表情切换过渡帧数(6-30),越大越柔和。",
|
||||
}
|
||||
|
||||
|
||||
def _settings():
|
||||
return get_api_settings()
|
||||
|
||||
|
||||
def _default_value(key: str) -> Any:
|
||||
return getattr(_settings(), key)
|
||||
|
||||
|
||||
def _build_config_item(key: str) -> dict[str, Any]:
|
||||
item: dict[str, Any] = {
|
||||
"key": key,
|
||||
"type": _WHITELIST[key],
|
||||
"description": _DESCRIPTIONS.get(key, ""),
|
||||
"default": _default_value(key),
|
||||
}
|
||||
service = get_system_config_service()
|
||||
item["value"] = service.get_config(key, _default_value(key))
|
||||
if key == "ditto_emotion_model":
|
||||
item["model_options"] = list(MODEL_OPTIONS)
|
||||
return item
|
||||
|
||||
|
||||
class ConfigUpdatePayload(BaseModel):
|
||||
configs: dict[str, Any]
|
||||
|
||||
|
||||
class TestPayload(BaseModel):
|
||||
test_text: str
|
||||
|
||||
|
||||
def _validate_value(key: str, value: Any) -> Any:
|
||||
st = _WHITELIST[key]
|
||||
if st == SETTING_TYPE_BOOL:
|
||||
if not isinstance(value, bool):
|
||||
raise ValueError(f"{key} 必须是布尔值")
|
||||
elif st == SETTING_TYPE_INT:
|
||||
if isinstance(value, bool) or not isinstance(value, int):
|
||||
raise ValueError(f"{key} 必须是整数")
|
||||
if not 6 <= value <= 30:
|
||||
raise ValueError(f"{key} 必须在 6-30 之间")
|
||||
elif st == SETTING_TYPE_FLOAT:
|
||||
if isinstance(value, bool):
|
||||
raise ValueError(f"{key} 必须是数字")
|
||||
try:
|
||||
value = float(value)
|
||||
except (TypeError, ValueError) as exc:
|
||||
raise ValueError(f"{key} 必须是数字") from exc
|
||||
if not 0.0 <= value <= 1.0:
|
||||
raise ValueError(f"{key} 必须在 0-1 之间")
|
||||
elif st == SETTING_TYPE_STRING:
|
||||
if not isinstance(value, str):
|
||||
raise ValueError(f"{key} 必须是字符串")
|
||||
if key == "ditto_emotion_prompt" and value.strip() and "{文案}" not in value:
|
||||
raise ValueError("提示词必须包含 {文案} 占位符")
|
||||
if key == "ditto_emotion_model" and value not in MODEL_OPTIONS:
|
||||
raise ValueError(f"模型必须是以下之一:{', '.join(MODEL_OPTIONS)}")
|
||||
return value
|
||||
|
||||
|
||||
@router.get("/config")
|
||||
def get_config() -> dict[str, Any]:
|
||||
return {"configs": [_build_config_item(k) for k in _WHITELIST]}
|
||||
|
||||
|
||||
@router.put("/config")
|
||||
def update_config(
|
||||
payload: ConfigUpdatePayload,
|
||||
x_api_key: str = Depends(_verify_internal_api_key),
|
||||
) -> dict[str, Any]:
|
||||
configs = payload.configs
|
||||
illegal = [k for k in configs if k not in _WHITELIST]
|
||||
if illegal:
|
||||
return {
|
||||
"ok": False,
|
||||
"error": f"不允许修改的配置项:{', '.join(illegal)}",
|
||||
}
|
||||
service = get_system_config_service()
|
||||
updated: dict[str, Any] = {}
|
||||
for key, raw in configs.items():
|
||||
try:
|
||||
value = _validate_value(key, raw)
|
||||
except ValueError as exc:
|
||||
return {"ok": False, "error": str(exc)}
|
||||
service.set_config(
|
||||
key,
|
||||
value,
|
||||
setting_type=_WHITELIST[key],
|
||||
updated_by=x_api_key[:8] if x_api_key else None,
|
||||
)
|
||||
updated[key] = value
|
||||
return {"ok": True, "updated": updated}
|
||||
|
||||
|
||||
@router.post("/config/test")
|
||||
def test_config(payload: TestPayload) -> dict[str, Any]:
|
||||
text = (payload.test_text or "").strip()
|
||||
if not text:
|
||||
return {"ok": False, "error": "test_text 不能为空"}
|
||||
from packages.application.ditto_emotion_service import get_ditto_emotion_service
|
||||
|
||||
service = get_ditto_emotion_service()
|
||||
segments = service.analyze(text)
|
||||
return {
|
||||
"ok": True,
|
||||
"enabled": service.enabled,
|
||||
"segments": [s.to_dict() for s in segments],
|
||||
}
|
||||
@@ -29,6 +29,8 @@ from app.services.ai_avatar_render_service import (
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from packages.middleware.points_gate import points_gate
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter()
|
||||
@@ -42,6 +44,7 @@ def _get_service(db: Session = Depends(get_db_session)) -> AiAvatarRenderService
|
||||
|
||||
|
||||
@router.post("", response_model=AiAvatarRenderJobResponse, status_code=201)
|
||||
@points_gate("ai_digital_human", per_unit=15)
|
||||
def create_render_job(
|
||||
body: CreateAiAvatarRenderRequest,
|
||||
current_user: AuthenticatedUser = Depends(get_current_user),
|
||||
|
||||
@@ -27,6 +27,7 @@ from packages.adapters.sqlalchemy_impl.generation_task_repository import (
|
||||
)
|
||||
from packages.application import ListGeneratedVideosByTaskUseCase
|
||||
from packages.domain.config_schemas import normalize_plan_config
|
||||
from packages.middleware.points_gate import points_gate
|
||||
from packages.shared.storage import get_shared_storage_service
|
||||
|
||||
from .templates_editor.dependencies import get_draft_plan_id, get_editor_services
|
||||
@@ -345,6 +346,7 @@ def _is_trusted_media_url(url: str) -> bool:
|
||||
|
||||
|
||||
@router.post("/generate-cover", response_model=GenerateCoverResponse)
|
||||
@points_gate("ai_cover")
|
||||
def generate_cover(
|
||||
body: GenerateCoverRequest,
|
||||
template_id: str = Query(..., description="模板 ID"),
|
||||
|
||||
@@ -41,6 +41,7 @@ from packages.application import (
|
||||
GetGenerationTaskUseCase,
|
||||
ListGeneratedVideosByTaskUseCase,
|
||||
)
|
||||
from packages.middleware.points_gate import points_gate
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -269,6 +270,7 @@ def _variant_value(values: list[str], index: int, fallback: str = "") -> str:
|
||||
|
||||
|
||||
@router.post("/preview", response_model=BatchPreviewGenerationTaskResponse, status_code=201)
|
||||
@points_gate("ai_video", quantity_field="preview_count")
|
||||
def create_preview_generation_task(
|
||||
request: CreatePreviewGenerationTaskRequest,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
|
||||
@@ -44,7 +44,6 @@ from packages.application import (
|
||||
GetGenerationTaskUseCase,
|
||||
ListGeneratedVideosByTaskUseCase,
|
||||
)
|
||||
from packages.domain import feature_pricing_service
|
||||
from packages.domain.smart_match import smart_select_assets
|
||||
|
||||
# #2035:文案关键词 → 素材分类 映射表(用于 smart_match category_match 维度)
|
||||
@@ -164,6 +163,8 @@ def _infer_expected_categories(script_tags: set[str] | None) -> set[str] | None:
|
||||
return matched or None
|
||||
|
||||
|
||||
from packages.middleware.points_gate import points_gate
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter()
|
||||
@@ -464,6 +465,7 @@ def _resolve_project_and_library(
|
||||
|
||||
|
||||
@router.post("/tasks", response_model=BatchGenerationTaskResponse)
|
||||
@points_gate("ai_video", quantity_field="count")
|
||||
def create_generation_task(
|
||||
request: CreateGenerationTaskRequest,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
@@ -700,17 +702,6 @@ def create_generation_task(
|
||||
logger.info("画中画已下线,strategy_id %s → one_take", effective_strategy_id)
|
||||
effective_strategy_id = "one_take"
|
||||
|
||||
# ── smart_edit 计费预扣(全局 points 开关 + 功能开关均开才扣) ──
|
||||
# 首期固定价:dynamic_cost=0,price=(0+fixed_cost)×multiplier,price_cap 封顶。
|
||||
# 预览任务不扣费;按任务条数扣费,任一任务预扣失败(余额不足)整体拒绝。
|
||||
smart_edit_charge = 0.0
|
||||
charged_task_count = 0
|
||||
if not request.is_preview and feature_pricing_service.is_feature_enabled("smart_edit"):
|
||||
unit_credits, _bd = feature_pricing_service.calculate_price("smart_edit", 0.0)
|
||||
if unit_credits > 0:
|
||||
smart_edit_charge = round(unit_credits * count, 2)
|
||||
charged_task_count = count
|
||||
|
||||
# 批量生成(count>1):每个变体必须走与单视频完全相同的独立选片流程(#1743/#1749)。
|
||||
# - 变体 0:clone 源 plan(不污染源 plan),变体 1..N-1 用 reselect_plan_for_variant
|
||||
# 完整重跑选片(素材级去重:fresh 优先 → 受控复用 overlap≤20% → 短素材禁复用);
|
||||
@@ -941,42 +932,6 @@ def create_generation_task(
|
||||
)
|
||||
# 变体序号写入 extra_meta(响应/排查时可辨识)
|
||||
task.extra_meta["variant_index"] = task_index
|
||||
|
||||
# smart_edit 逐条预扣(首期固定价,credits_cost=prepaid,不做结算)
|
||||
task_txn_id = ""
|
||||
if charged_task_count > 0:
|
||||
from packages.domain.points_service import PointsService
|
||||
|
||||
unit_credits = round(smart_edit_charge / count, 2)
|
||||
res = PointsService().deduct_points(
|
||||
user_id=user_id,
|
||||
amount=unit_credits,
|
||||
source="smart_edit",
|
||||
db=db,
|
||||
description="智能剪辑生成预扣",
|
||||
ref_id=task.id,
|
||||
)
|
||||
if not res.get("success"):
|
||||
# 余额不足:退还本次请求已扣积分后整体拒绝
|
||||
already_charged = round(unit_credits * task_index, 2)
|
||||
if already_charged > 0:
|
||||
PointsService().refund_points(
|
||||
user_id=user_id,
|
||||
amount=already_charged,
|
||||
source="smart_edit",
|
||||
db=db,
|
||||
ref_id=task.id,
|
||||
description="智能剪辑批量提交失败退回",
|
||||
)
|
||||
raise HTTPException(
|
||||
status_code=402,
|
||||
detail=(f"积分不足:智能剪辑每条需 {unit_credits:.2f} 积分,当前余额 {res.get('balance', 0)}"),
|
||||
)
|
||||
task_txn_id = str(res.get("transaction_id") or "")
|
||||
task.credits_prepaid = unit_credits
|
||||
task.credits_cost = unit_credits
|
||||
task.credits_transaction_id = task_txn_id
|
||||
generation_task_repository.update(task)
|
||||
try:
|
||||
# 兜底关联编辑计划:前端未传 source_edit_plan_id 时,
|
||||
# 通过 template_id + user_id 在 DB 层直接查找最新的 plan。
|
||||
|
||||
@@ -12,9 +12,11 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import math
|
||||
from datetime import UTC
|
||||
|
||||
from app.auth import AuthenticatedUser, get_current_user
|
||||
from app.config import settings
|
||||
from app.dependencies import (
|
||||
get_db_session,
|
||||
get_voice_clone_profile_repository,
|
||||
@@ -30,6 +32,9 @@ from app.services.mediakit_client import MediaKitError
|
||||
from fastapi import APIRouter, BackgroundTasks, Depends, HTTPException, Query
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from packages.domain.points_rules import calculate_points_cost
|
||||
from packages.domain.points_service import PointsService
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter()
|
||||
@@ -56,6 +61,37 @@ def create_lipsync_job(
|
||||
db: Session = Depends(get_db_session),
|
||||
svc: LipsyncService = Depends(_get_service),
|
||||
):
|
||||
user_id = current_user.user.id
|
||||
|
||||
# ── 积分扣点(#1895 P2) ──
|
||||
_points_deducted = 0
|
||||
_points_scene = "ai_digital_human"
|
||||
_points_svc = PointsService() if settings.points_enabled else None
|
||||
if _points_svc is not None:
|
||||
# 口型同步:TTS 模式按 script_text 估时长(240字/分钟);音频直传按 audio_duration(秒→分钟)
|
||||
if body.audio_url and body.audio_duration and body.audio_duration > 0:
|
||||
est_minutes = max(1.0, math.ceil(body.audio_duration / 60.0))
|
||||
elif body.script_text:
|
||||
est_minutes = max(1.0, math.ceil(len(body.script_text) / 240))
|
||||
else:
|
||||
est_minutes = 1.0
|
||||
_points_deducted = calculate_points_cost(
|
||||
_points_scene,
|
||||
is_member=getattr(current_user.user, "is_member", False),
|
||||
duration_minutes=est_minutes,
|
||||
member_type=getattr(current_user.user, "member_type", None),
|
||||
)
|
||||
_deduct_res = _points_svc.deduct_points(user_id, _points_deducted, _points_scene, db)
|
||||
if not _deduct_res["success"]:
|
||||
raise HTTPException(
|
||||
status_code=402,
|
||||
detail={
|
||||
"code": "INSUFFICIENT_POINTS",
|
||||
"message": f"积分不足,需要 {_points_deducted} 积分,当前余额 {_deduct_res['balance']}",
|
||||
"required": _points_deducted,
|
||||
"balance": _deduct_res["balance"],
|
||||
},
|
||||
)
|
||||
"""提交对口型任务.
|
||||
|
||||
三种模式:
|
||||
@@ -65,8 +101,6 @@ def create_lipsync_job(
|
||||
- 预合成音频(#1845 新主路径):传 {video_url, audio_url, audio_duration, sentence_timings},
|
||||
后端同步ffprobe+写入timings+直接提交MediaKit(~2-3s)。
|
||||
"""
|
||||
user_id = current_user.user.id
|
||||
|
||||
try:
|
||||
job = svc.create_job(
|
||||
user_id=user_id,
|
||||
@@ -84,8 +118,18 @@ def create_lipsync_job(
|
||||
project_id=body.project_id,
|
||||
)
|
||||
except ValueError as exc:
|
||||
if _points_deducted > 0 and _points_svc is not None:
|
||||
try:
|
||||
_points_svc.refund_points(user_id, _points_deducted, _points_scene, db)
|
||||
except Exception as refund_err:
|
||||
logger.warning(f"对口型 ValueError 退积分异常: err={refund_err}")
|
||||
raise HTTPException(status_code=400, detail=str(exc)) from exc
|
||||
except MediaKitError as exc:
|
||||
if _points_deducted > 0 and _points_svc is not None:
|
||||
try:
|
||||
_points_svc.refund_points(user_id, _points_deducted, _points_scene, db)
|
||||
except Exception as refund_err:
|
||||
logger.warning(f"对口型 MediaKitError 退积分异常: err={refund_err}")
|
||||
status_code = 502
|
||||
if exc.code in ("VoiceForbidden",):
|
||||
status_code = 403
|
||||
@@ -101,11 +145,24 @@ def create_lipsync_job(
|
||||
) from exc
|
||||
except Exception as exc:
|
||||
logger.error("创建对口型任务异常: %s", exc, exc_info=True)
|
||||
if _points_deducted > 0 and _points_svc is not None:
|
||||
try:
|
||||
_points_svc.refund_points(user_id, _points_deducted, _points_scene, db)
|
||||
except Exception as refund_err:
|
||||
logger.warning(f"对口型异常退积分异常: err={refund_err}")
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"创建对口型任务失败: {exc}",
|
||||
) from exc
|
||||
|
||||
# 创建成功但状态为 failed(同步路径失败已抛异常到上面 except;此处处理 Celery 调度失败等)
|
||||
# 若任务已创建且状态为 failed,退费
|
||||
if _points_deducted > 0 and _points_svc is not None and getattr(job, "status", None) == "failed":
|
||||
try:
|
||||
_points_svc.refund_points(user_id, _points_deducted, _points_scene, db, ref_id=job.id)
|
||||
except Exception as refund_err:
|
||||
logger.warning(f"对口型任务失败退积分异常: job_id={job.id}, err={refund_err}")
|
||||
|
||||
return job
|
||||
|
||||
|
||||
@@ -119,14 +176,37 @@ def preview_tts(
|
||||
db: Session = Depends(get_db_session),
|
||||
svc: LipsyncService = Depends(_get_service),
|
||||
):
|
||||
user_id = current_user.user.id
|
||||
|
||||
# ── 积分扣点(#1895 P2) ──
|
||||
_points_deducted = 0
|
||||
_points_scene = "ai_digital_human"
|
||||
_points_svc = PointsService() if settings.points_enabled else None
|
||||
if _points_svc is not None:
|
||||
est_minutes = max(1.0, math.ceil(len(body.script_text or "") / 240)) if body.script_text else 1.0
|
||||
_points_deducted = calculate_points_cost(
|
||||
_points_scene,
|
||||
is_member=getattr(current_user.user, "is_member", False),
|
||||
duration_minutes=est_minutes,
|
||||
member_type=getattr(current_user.user, "member_type", None),
|
||||
)
|
||||
_deduct_res = _points_svc.deduct_points(user_id, _points_deducted, _points_scene, db)
|
||||
if not _deduct_res["success"]:
|
||||
raise HTTPException(
|
||||
status_code=402,
|
||||
detail={
|
||||
"code": "INSUFFICIENT_POINTS",
|
||||
"message": f"积分不足,需要 {_points_deducted} 积分,当前余额 {_deduct_res['balance']}",
|
||||
"required": _points_deducted,
|
||||
"balance": _deduct_res["balance"],
|
||||
},
|
||||
)
|
||||
"""步骤1「生成配音」同步 TTS 预合成.
|
||||
|
||||
同步执行 TTS 合成 → 下载音频 → ffprobe 时长 → 句子时间戳计算,
|
||||
不创建 LipsyncJob、不转存 OSS,直接返回 CosyVoice 临时 URL(~24h 有效)。
|
||||
耗时约 2-3 秒。
|
||||
"""
|
||||
user_id = current_user.user.id
|
||||
|
||||
try:
|
||||
result = svc.preview_tts(
|
||||
user_id=user_id,
|
||||
@@ -138,6 +218,11 @@ def preview_tts(
|
||||
emotion=body.emotion,
|
||||
)
|
||||
except MediaKitError as exc:
|
||||
if _points_deducted > 0 and _points_svc is not None:
|
||||
try:
|
||||
_points_svc.refund_points(user_id, _points_deducted, _points_scene, db)
|
||||
except Exception as refund_err:
|
||||
logger.warning(f"TTS 预合成 MediaKitError 退积分异常: err={refund_err}")
|
||||
status_code = 400
|
||||
if exc.code in ("VoiceForbidden",):
|
||||
status_code = 403
|
||||
@@ -152,6 +237,11 @@ def preview_tts(
|
||||
) from exc
|
||||
except Exception as exc:
|
||||
logger.error("TTS 预合成异常: %s", exc, exc_info=True)
|
||||
if _points_deducted > 0 and _points_svc is not None:
|
||||
try:
|
||||
_points_svc.refund_points(user_id, _points_deducted, _points_scene, db)
|
||||
except Exception as refund_err:
|
||||
logger.warning(f"TTS 预合成异常退积分异常: err={refund_err}")
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"TTS 合成失败: {exc}",
|
||||
|
||||
@@ -145,22 +145,19 @@ def get_rules(
|
||||
def get_packages(
|
||||
current_user: AuthenticatedUser = Depends(get_current_user),
|
||||
):
|
||||
"""查询可购买的积分包列表(读管理后台 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"],
|
||||
"""查询可购买的积分包列表。"""
|
||||
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,
|
||||
)
|
||||
)
|
||||
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)
|
||||
@@ -172,7 +169,17 @@ def check_points(
|
||||
current_user: AuthenticatedUser = Depends(get_current_user),
|
||||
db: Session = Depends(get_db_session),
|
||||
):
|
||||
"""消费前检查余额是否足够。已下线/未知场景返回 cost=0(免费)。"""
|
||||
"""消费前检查余额是否足够。未知 scene_key 返回 400(而非 500)。"""
|
||||
if body.scene_key not in POINTS_SCENES:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"code": "UNKNOWN_SCENE",
|
||||
"message": f"未知场景: {body.scene_key}",
|
||||
"valid_scenes": sorted(POINTS_SCENES.keys()),
|
||||
},
|
||||
)
|
||||
|
||||
# 积分系统暂停(ENABLE_CREDIT_SYSTEM=false):所有场景直接放行,需 0 积分
|
||||
if not _credits_enabled():
|
||||
svc = _get_service()
|
||||
@@ -188,6 +195,13 @@ def check_points(
|
||||
is_mem = _is_member(current_user)
|
||||
mt = _member_type(current_user)
|
||||
|
||||
# 混剪场景先检查免费额度
|
||||
is_free_quota = False
|
||||
if body.scene_key == "ai_video" and not is_mem:
|
||||
svc = _get_service()
|
||||
if svc.check_daily_free_clip(current_user.user.id, db):
|
||||
is_free_quota = True
|
||||
|
||||
required = calculate_points_cost(
|
||||
body.scene_key,
|
||||
is_mem,
|
||||
@@ -201,11 +215,11 @@ def check_points(
|
||||
balance = account["balance"]
|
||||
|
||||
return PointsCheckResponse(
|
||||
allowed=balance >= required,
|
||||
allowed=is_free_quota or balance >= required,
|
||||
required_points=required,
|
||||
current_balance=balance,
|
||||
remaining_after=balance - required,
|
||||
is_free_quota=False,
|
||||
is_free_quota=is_free_quota,
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -44,6 +44,7 @@ from app.services.script_asr_service import (
|
||||
from fastapi import APIRouter, Depends, HTTPException, status
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from packages.middleware.points_gate import points_gate
|
||||
from packages.shared.ai_client import get_doubao_client
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -372,6 +373,7 @@ def douyin_diag():
|
||||
|
||||
|
||||
@router.post("/extract-from-douyin", response_model=ExtractFromDouyinResponse)
|
||||
@points_gate("douyin_extract")
|
||||
def extract_from_douyin(
|
||||
request: ExtractFromDouyinRequest,
|
||||
current_user: AuthenticatedUser = Depends(get_current_user),
|
||||
@@ -495,6 +497,7 @@ def extract_from_douyin(
|
||||
|
||||
|
||||
@router.post("/ai-rewrite", response_model=AiRewriteResponse)
|
||||
@points_gate("ai_rewrite")
|
||||
def ai_rewrite(
|
||||
request: AiRewriteRequest,
|
||||
current_user: AuthenticatedUser = Depends(get_current_user),
|
||||
@@ -534,6 +537,7 @@ def ai_rewrite(
|
||||
|
||||
|
||||
@router.post("/ai-generate-titles", response_model=AiGenerateTitlesResponse)
|
||||
@points_gate("ai_title")
|
||||
def ai_generate_titles(
|
||||
request: AiGenerateTitlesRequest,
|
||||
current_user: AuthenticatedUser = Depends(get_current_user),
|
||||
|
||||
@@ -86,13 +86,33 @@ async def get_current_subscription(
|
||||
def list_membership_plans(
|
||||
current_user: AuthenticatedUser = Depends(get_current_user),
|
||||
) -> dict[str, list[dict[str, Any]]]:
|
||||
"""查询可购买的会员套餐(读管理后台 plans 表真实数据)。
|
||||
"""查询所有会员档位(供前端会员购买页展示)。
|
||||
|
||||
仅返回 is_enabled=true 的套餐;后台启停/改价后最多 30 秒生效。
|
||||
返回 points 积分体系下的会员档位(月卡/季卡/年卡),含价格、时长、积分折扣等信息。
|
||||
"""
|
||||
from packages.application.catalog.admin_catalog import get_membership_plans
|
||||
from packages.domain.points_rules import MEMBER_DISCOUNT, MEMBERSHIP_PRICES
|
||||
|
||||
return {"plans": 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}
|
||||
|
||||
|
||||
@router.get("/billing-records", response_model=list[BillingRecord])
|
||||
|
||||
@@ -4,12 +4,14 @@ from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import math
|
||||
import subprocess
|
||||
import tempfile
|
||||
from pathlib import Path
|
||||
from typing import Any, Optional
|
||||
|
||||
from app.auth import AuthenticatedUser, get_current_user
|
||||
from app.config import settings
|
||||
from app.core.celery_app import celery_app
|
||||
from app.core.storage import get_storage_service
|
||||
from app.dependencies import (
|
||||
@@ -51,6 +53,8 @@ from packages.application.tts_job.use_cases import (
|
||||
)
|
||||
from packages.application.tts_job.workflow import TTSWorkflowService
|
||||
from packages.domain import Asset, AssetLibrary, AssetLibraryKind, AssetStatus, ClassificationStatus
|
||||
from packages.domain.points_rules import calculate_points_cost
|
||||
from packages.domain.points_service import PointsService
|
||||
from packages.domain.voice_presets import list_voices
|
||||
from packages.ports.asset_library_repository import AssetLibraryRepository
|
||||
from packages.ports.asset_repository import AssetRepository
|
||||
@@ -140,6 +144,31 @@ def synthesize(
|
||||
"""
|
||||
user_id = authenticated_user.user.id
|
||||
|
||||
# ── 积分扣点(#1895 P2) ──
|
||||
_points_deducted = 0
|
||||
_points_scene = "ai_voice"
|
||||
_points_svc = PointsService() if settings.points_enabled else None
|
||||
if _points_svc is not None:
|
||||
# 中文按 ~240 字/分钟粗估时长,至少按 1 分钟扣 1 分
|
||||
est_minutes = max(1.0, math.ceil(len(request.text) / 240))
|
||||
_points_deducted = calculate_points_cost(
|
||||
_points_scene,
|
||||
is_member=getattr(authenticated_user.user, "is_member", False),
|
||||
duration_minutes=est_minutes,
|
||||
member_type=getattr(authenticated_user.user, "member_type", None),
|
||||
)
|
||||
_deduct_res = _points_svc.deduct_points(user_id, _points_deducted, _points_scene, db)
|
||||
if not _deduct_res["success"]:
|
||||
raise HTTPException(
|
||||
status_code=402,
|
||||
detail={
|
||||
"code": "INSUFFICIENT_POINTS",
|
||||
"message": f"积分不足,需要 {_points_deducted} 积分,当前余额 {_deduct_res['balance']}",
|
||||
"required": _points_deducted,
|
||||
"balance": _deduct_res["balance"],
|
||||
},
|
||||
)
|
||||
|
||||
# 解析 voice_id:前端可能传克隆音色 profile UUID(而非 CosyVoice voice_id),
|
||||
# 与 /tts/preview 保持一致:命中 profile → 校验归属 → 取 CosyVoice voice_id
|
||||
actual_voice_id = request.voice_id
|
||||
@@ -202,6 +231,7 @@ def synthesize(
|
||||
cosyvoice_service=cosyvoice_service,
|
||||
)
|
||||
|
||||
synthesis_error: Exception | None = None
|
||||
try:
|
||||
job = workflow.start_synthesis(job.id)
|
||||
except Exception as e:
|
||||
@@ -209,10 +239,18 @@ def synthesize(
|
||||
# 但 DB 异常、网络异常等意外错误可能逃逸。
|
||||
# 与音色克隆接口保持一致:标记 failed,返回 201,不抛 500。
|
||||
logger.error(f"TTS 合成异常: job_id={job.id}, error={e}", exc_info=True)
|
||||
synthesis_error = e
|
||||
try:
|
||||
job = workflow.process_synthesis_failure(job.id, str(e))
|
||||
except Exception as inner_e:
|
||||
logger.error(f"标记 TTS job 失败时出错: job_id={job.id}, error={inner_e}")
|
||||
# 合成失败且已扣积分 → 退费
|
||||
if synthesis_error is not None and _points_deducted > 0 and _points_svc is not None:
|
||||
try:
|
||||
_points_svc.refund_points(user_id, _points_deducted, _points_scene, db, ref_id=job.id)
|
||||
except Exception as refund_err:
|
||||
logger.warning(f"TTS 合失败退积分异常: job_id={job.id}, err={refund_err}")
|
||||
|
||||
# 若任务处于 processing 状态(异步模式),触发 Celery 后台轮询
|
||||
if job.status.value == "processing":
|
||||
# 分段合成任务 vs 普通单段任务
|
||||
@@ -231,6 +269,13 @@ def synthesize(
|
||||
workflow.process_synthesis_failure(job.id, f"Celery 任务调度失败: {e}")
|
||||
except Exception as inner_e:
|
||||
logger.error(f"Celery 调度后标记失败时出错: job_id={job.id}, error={inner_e}")
|
||||
# 调度失败退费
|
||||
if _points_deducted > 0 and _points_svc is not None:
|
||||
try:
|
||||
_points_svc.refund_points(user_id, _points_deducted, _points_scene, db, ref_id=job.id)
|
||||
except Exception as refund_err:
|
||||
logger.warning(f"Celery 调度失败退积分异常: job_id={job.id}, err={refund_err}")
|
||||
|
||||
return TTSSynthesizeResponse(
|
||||
job_id=job.id,
|
||||
status=job.status,
|
||||
@@ -565,6 +610,31 @@ def preview_tts(
|
||||
用于前端预览配音效果,限制文本长度 200 字以内。
|
||||
支持预设音色和克隆音色:克隆音色传的是 profile UUID,需解析为 CosyVoice voice_id。
|
||||
"""
|
||||
user_id = authenticated_user.user.id
|
||||
# ── 积分扣点(#1895 P2) ──
|
||||
_points_deducted = 0
|
||||
_points_scene = "ai_voice"
|
||||
_points_svc = PointsService() if settings.points_enabled else None
|
||||
if _points_svc is not None:
|
||||
est_minutes = max(1.0, math.ceil(len(request.text) / 240))
|
||||
_points_deducted = calculate_points_cost(
|
||||
_points_scene,
|
||||
is_member=getattr(authenticated_user.user, "is_member", False),
|
||||
duration_minutes=est_minutes,
|
||||
member_type=getattr(authenticated_user.user, "member_type", None),
|
||||
)
|
||||
_deduct_res = _points_svc.deduct_points(user_id, _points_deducted, _points_scene, db)
|
||||
if not _deduct_res["success"]:
|
||||
raise HTTPException(
|
||||
status_code=402,
|
||||
detail={
|
||||
"code": "INSUFFICIENT_POINTS",
|
||||
"message": f"积分不足,需要 {_points_deducted} 积分,当前余额 {_deduct_res['balance']}",
|
||||
"required": _points_deducted,
|
||||
"balance": _deduct_res["balance"],
|
||||
},
|
||||
)
|
||||
|
||||
# 解析 voice_id:前端可能传 VoiceCloneProfile UUID 或预设音色 ID
|
||||
actual_voice_id = request.voice_id
|
||||
profile = voice_clone_repo.get(request.voice_id)
|
||||
@@ -594,6 +664,12 @@ def preview_tts(
|
||||
language=getattr(request, "language", "zh-CN"),
|
||||
)
|
||||
except (CosyVoiceError, ValueError) as e:
|
||||
# 合成失败退费
|
||||
if _points_deducted > 0 and _points_svc is not None:
|
||||
try:
|
||||
_points_svc.refund_points(user_id, _points_deducted, _points_scene, db)
|
||||
except Exception as refund_err:
|
||||
logger.warning(f"TTS 预览失败退积分异常: {refund_err}")
|
||||
if isinstance(e, CosyVoiceError):
|
||||
raise HTTPException(status_code=status.HTTP_502_BAD_GATEWAY, detail=f"TTS 合成失败: {e}") from e
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=str(e)) from e
|
||||
|
||||
@@ -32,11 +32,7 @@ from app.schemas.viral_video import (
|
||||
ConfirmCopyRequest,
|
||||
ConfirmIntentRequest,
|
||||
CreateViralVideoRequest,
|
||||
CreditsFormulaBreakdown,
|
||||
EstimateCreditsRequest,
|
||||
EstimateCreditsResponse,
|
||||
GenerateCopyRequest,
|
||||
RetryViralVideoRequest,
|
||||
StyleTemplateListResponse,
|
||||
StyleTemplateResponse,
|
||||
ViralVideoHistoryResponse,
|
||||
@@ -49,9 +45,7 @@ from packages.adapters.sqlalchemy_impl.viral_video_repository import (
|
||||
SQLAlchemyViralVideoJobRepository,
|
||||
SQLAlchemyViralVideoStyleTemplateRepository,
|
||||
)
|
||||
from packages.domain.points_rules import list_viral_video_models
|
||||
from packages.domain.viral_video import ViralVideoStatus
|
||||
from packages.shared.dashscope_client import get_dashscope_client
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -146,10 +140,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),
|
||||
credits_cost=job.credits_cost,
|
||||
error_msg=job.error_msg,
|
||||
retry_count=job.retry_count,
|
||||
started_at=job.started_at,
|
||||
@@ -202,7 +193,6 @@ def create_viral_video(
|
||||
voice_source=getattr(request, "voice_source", "") or "",
|
||||
video_ratio=getattr(request, "video_ratio", "9:16") or "9:16",
|
||||
video_model=getattr(request, "video_model", "") or "",
|
||||
video_resolution=getattr(request, "video_resolution", "720p") or "720p",
|
||||
copy_result=None,
|
||||
)
|
||||
|
||||
@@ -245,7 +235,6 @@ def analyze_images(
|
||||
voice_source=request.voice_source or "",
|
||||
video_ratio=request.video_ratio or "9:16",
|
||||
video_model=request.video_model or "",
|
||||
video_resolution=getattr(request, "video_resolution", "720p") or "720p",
|
||||
duration=request.duration or 15,
|
||||
)
|
||||
repo.save(job)
|
||||
@@ -280,18 +269,11 @@ def generate_copy(
|
||||
raise HTTPException(status_code=404, detail="任务不存在")
|
||||
if job.user_id != authenticated_user.user.id:
|
||||
raise HTTPException(status_code=403, detail="无权操作此任务")
|
||||
# 允许首次进入(IMAGE_ANALYZED/PENDING)、失败重试(FAILED)、文案重新生成(COPY_GENERATED/COMPLETED)
|
||||
if job.status not in (
|
||||
ViralVideoStatus.IMAGE_ANALYZED,
|
||||
ViralVideoStatus.PENDING,
|
||||
ViralVideoStatus.FAILED,
|
||||
ViralVideoStatus.COPY_GENERATED,
|
||||
ViralVideoStatus.COMPLETED,
|
||||
):
|
||||
if job.status not in (ViralVideoStatus.IMAGE_ANALYZED, ViralVideoStatus.PENDING, ViralVideoStatus.FAILED):
|
||||
raise HTTPException(status_code=409, detail=f"任务当前状态 {job.status} 不能生成文案")
|
||||
|
||||
# 失败重试 / 重新生成:retry_count 自增
|
||||
if job.status in (ViralVideoStatus.FAILED, ViralVideoStatus.COPY_GENERATED, ViralVideoStatus.COMPLETED):
|
||||
# 允许失败任务重试:重置
|
||||
if job.status == ViralVideoStatus.FAILED:
|
||||
job.retry_count += 1
|
||||
job.error_msg = ""
|
||||
|
||||
@@ -316,7 +298,6 @@ def generate_copy(
|
||||
job.voice_source = request.voice_source or job.voice_source
|
||||
job.video_ratio = request.video_ratio or job.video_ratio or "9:16"
|
||||
job.video_model = request.video_model or job.video_model or ""
|
||||
job.video_resolution = getattr(request, "video_resolution", "") or job.video_resolution or "720p"
|
||||
|
||||
job.resume_from_image_analyzed()
|
||||
repo.update(job)
|
||||
@@ -348,108 +329,6 @@ def confirm_copy(
|
||||
raise HTTPException(status_code=403, detail="无权操作此任务")
|
||||
if job.status != ViralVideoStatus.COPY_GENERATED:
|
||||
raise HTTPException(status_code=409, detail=f"任务当前状态 {job.status} 不能确认文案(需 copy_generated)")
|
||||
# #2218: 额外校验 copy_result 完整性,防止孤儿/脏数据进入渲染
|
||||
if not isinstance(job.copy_result, dict) or not job.copy_result:
|
||||
raise HTTPException(status_code=409, detail="文案数据缺失,请先点击「生成文案」")
|
||||
|
||||
# Bug1 fix: 用户 confirm 时允许修改 video_model/video_resolution/video_ratio/duration
|
||||
old_duration = int(getattr(job, "duration", 15) or 15)
|
||||
old_resolution = getattr(job, "video_resolution", "720p") or "720p"
|
||||
old_ratio = getattr(job, "video_ratio", "9:16") or "9:16"
|
||||
old_model = getattr(job, "video_model", None) or "seedance-2.5"
|
||||
|
||||
if request.duration is not None:
|
||||
job.duration = max(5, min(30, int(request.duration)))
|
||||
if request.video_resolution is not None:
|
||||
job.video_resolution = request.video_resolution
|
||||
if request.video_ratio is not None:
|
||||
job.video_ratio = request.video_ratio
|
||||
if request.video_model is not None:
|
||||
job.video_model = request.video_model
|
||||
|
||||
param_changed = (
|
||||
(request.duration is not None and int(request.duration) != old_duration)
|
||||
or (request.video_resolution is not None and request.video_resolution != old_resolution)
|
||||
or (request.video_ratio is not None and request.video_ratio != old_ratio)
|
||||
or (request.video_model is not None and request.video_model != old_model)
|
||||
)
|
||||
|
||||
# 积分预扣(已扣过/重试任务跳过)
|
||||
from app.config import settings as _settings
|
||||
|
||||
if _settings.points_enabled:
|
||||
already_paid = (float(getattr(job, "credits_prepaid", 0) or 0) > 0) or (
|
||||
float(getattr(job, "credits_cost", 0) or 0) > 0
|
||||
)
|
||||
if param_changed and already_paid:
|
||||
# 参数变更:回退旧预扣,按新参数重新预扣
|
||||
from packages.domain.points_rules import calculate_viral_video_credits, resolve_video_dimensions
|
||||
from packages.domain.points_service import PointsService
|
||||
|
||||
old_w, old_h = resolve_video_dimensions(old_resolution, old_ratio)
|
||||
old_est = calculate_viral_video_credits(old_duration, old_w, old_h, old_model)
|
||||
new_w, new_h = resolve_video_dimensions(
|
||||
getattr(job, "video_resolution", "720p") or "720p",
|
||||
job.video_ratio or "9:16",
|
||||
)
|
||||
new_est = calculate_viral_video_credits(
|
||||
int(job.duration or 15), new_w, new_h, job.video_model or "seedance-2.5"
|
||||
)
|
||||
svc = PointsService()
|
||||
# 退回旧预扣
|
||||
if getattr(job, "credits_transaction_id", None):
|
||||
svc.refund_points(
|
||||
user_id=authenticated_user.user.id,
|
||||
amount=float(job.credits_prepaid),
|
||||
source="viral_video",
|
||||
db=session,
|
||||
ref_id=job.credits_transaction_id,
|
||||
description="confirm-copy 参数变更退还旧预扣",
|
||||
)
|
||||
# 预扣新金额
|
||||
if new_est > 0:
|
||||
res = svc.deduct_viral_video(authenticated_user.user.id, new_est, job.id, session)
|
||||
if not res.get("success"):
|
||||
balance = res.get("balance", 0)
|
||||
raise HTTPException(
|
||||
status_code=402,
|
||||
detail={
|
||||
"code": "INSUFFICIENT_POINTS",
|
||||
"message": f"积分不足,需要 {new_est} 积分,当前余额 {balance}",
|
||||
"required": new_est,
|
||||
"balance": balance,
|
||||
},
|
||||
)
|
||||
job.credits_prepaid = new_est
|
||||
job.credits_transaction_id = res.get("transaction_id", "") or ""
|
||||
logger.info("[爆款视频][confirm-copy] 参数变更,积分重算: old=%d new=%d job_id=%s", old_est, new_est, job.id)
|
||||
elif not already_paid:
|
||||
from packages.domain.points_rules import calculate_viral_video_credits, resolve_video_dimensions
|
||||
from packages.domain.points_service import PointsService
|
||||
|
||||
w, h = resolve_video_dimensions(
|
||||
getattr(job, "video_resolution", "720p") or "720p",
|
||||
job.video_ratio or "9:16",
|
||||
)
|
||||
est_credits = calculate_viral_video_credits(
|
||||
int(job.duration or 15), w, h, job.video_model or "seedance-2.5"
|
||||
)
|
||||
svc = PointsService()
|
||||
res = svc.deduct_viral_video(authenticated_user.user.id, est_credits, job.id, session)
|
||||
if not res.get("success"):
|
||||
balance = res.get("balance", 0)
|
||||
raise HTTPException(
|
||||
status_code=402,
|
||||
detail={
|
||||
"code": "INSUFFICIENT_POINTS",
|
||||
"message": f"积分不足,需要 {est_credits} 积分,当前余额 {balance}",
|
||||
"required": est_credits,
|
||||
"balance": balance,
|
||||
},
|
||||
)
|
||||
job.credits_prepaid = est_credits
|
||||
job.credits_transaction_id = res.get("transaction_id", "") or ""
|
||||
repo.update(job)
|
||||
|
||||
job.resume_from_copy_generated(edited_copy=request.edited_copy or None)
|
||||
repo.update(job)
|
||||
@@ -465,38 +344,6 @@ def confirm_copy(
|
||||
return _to_response(job)
|
||||
|
||||
|
||||
@router.post("/estimate-credits", response_model=EstimateCreditsResponse)
|
||||
def estimate_credits(
|
||||
request: EstimateCreditsRequest,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
) -> EstimateCreditsResponse:
|
||||
"""爆款视频积分预估(纯计算,不扣费、不创建任务)。
|
||||
|
||||
返回 estimated_credits 与 formula_breakdown(tokens / video_cost / fixed_cost /
|
||||
profit_multiplier / model_price / width / height / fps),便于前端展示计费明细。
|
||||
同时兼容前端传 model 或 video_model、resolution 或 video_resolution、ratio 或 video_ratio。
|
||||
"""
|
||||
from packages.domain.points_rules import (
|
||||
calculate_viral_video_credits_with_breakdown,
|
||||
resolve_video_dimensions,
|
||||
)
|
||||
|
||||
model = (request.model or "").strip() or "seedance-2.5"
|
||||
resolution = (request.resolution or "").strip() or "720p"
|
||||
ratio = (request.ratio or "").strip() or "9:16"
|
||||
duration = int(request.duration or 15)
|
||||
|
||||
w, h = resolve_video_dimensions(resolution, ratio)
|
||||
credits, bd = calculate_viral_video_credits_with_breakdown(
|
||||
duration,
|
||||
w,
|
||||
h,
|
||||
model,
|
||||
)
|
||||
breakdown = CreditsFormulaBreakdown(**bd)
|
||||
return EstimateCreditsResponse(estimated_credits=credits, formula_breakdown=breakdown)
|
||||
|
||||
|
||||
@router.get("/history", response_model=ViralVideoHistoryResponse)
|
||||
def list_viral_video_history(
|
||||
limit: int = 50,
|
||||
@@ -531,17 +378,6 @@ def list_style_templates(
|
||||
return StyleTemplateListResponse(items=items)
|
||||
|
||||
|
||||
@router.get("/models")
|
||||
def list_available_models() -> dict:
|
||||
"""返回爆款视频可用模型列表(供前端模型选择器使用)。"""
|
||||
dashscope_available = get_dashscope_client() is not None
|
||||
models = list_viral_video_models(
|
||||
include_placeholder=False,
|
||||
dashscope_available=dashscope_available,
|
||||
)
|
||||
return {"models": models}
|
||||
|
||||
|
||||
@router.get("/{job_id}", response_model=ViralVideoJobResponse)
|
||||
def get_viral_video_job(
|
||||
job_id: str,
|
||||
@@ -561,17 +397,10 @@ def get_viral_video_job(
|
||||
@router.post("/{job_id}/retry", response_model=ViralVideoJobResponse)
|
||||
def retry_viral_video_job(
|
||||
job_id: str,
|
||||
request: RetryViralVideoRequest | None = None,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
session: Session = Depends(get_db_session),
|
||||
) -> ViralVideoJobResponse:
|
||||
"""重试失败的爆款视频任务(也支持对僵尸/超时 running 任务强制重置后重试)。
|
||||
|
||||
可选 body (RetryViralVideoRequest):若传入新的 duration/video_resolution/video_ratio/
|
||||
video_model,会重新预估积分并与原 credits_prepaid 做差额多退少补(不足抛 402 阻止重试);
|
||||
不传 body 或参数无变化时,保持原参数、原预扣金额不变,仅重置状态并入队。
|
||||
credits_prepaid 为 0 的老任务首次重试会走预扣流程(与 confirm-copy 一致)。
|
||||
"""
|
||||
"""重试失败的爆款视频任务(也支持对僵尸/超时 running 任务强制重置后重试)。"""
|
||||
from datetime import datetime, timezone
|
||||
|
||||
repo = _get_job_repo(session)
|
||||
@@ -592,104 +421,6 @@ def retry_viral_video_job(
|
||||
if job.status != ViralVideoStatus.FAILED and not is_stale_running:
|
||||
raise HTTPException(status_code=409, detail="只有失败或超时的任务可以重试")
|
||||
|
||||
# ── 参数变更检测 + 积分多退少补 ──────────────────────────────────────
|
||||
req = request or RetryViralVideoRequest()
|
||||
new_duration = req.duration
|
||||
new_resolution = (req.video_resolution or "").strip() or None
|
||||
new_ratio = (req.video_ratio or "").strip() or None
|
||||
new_model = (req.video_model or "").strip() or None
|
||||
|
||||
old_duration = int(getattr(job, "duration", 15) or 15)
|
||||
old_resolution = (getattr(job, "video_resolution", "720p") or "720p").strip() or "720p"
|
||||
old_ratio = (getattr(job, "video_ratio", "9:16") or "9:16").strip() or "9:16"
|
||||
old_model = (getattr(job, "video_model", "") or "").strip()
|
||||
|
||||
# 仅当有任意字段传入且值不同才算"参数变更"
|
||||
param_changed = bool(
|
||||
(new_duration is not None and int(new_duration) != old_duration)
|
||||
or (new_resolution is not None and new_resolution != old_resolution)
|
||||
or (new_ratio is not None and new_ratio != old_ratio)
|
||||
or (new_model is not None and new_model != old_model)
|
||||
)
|
||||
|
||||
from app.config import settings as _settings
|
||||
|
||||
need_points_settle = False
|
||||
new_est = 0.0
|
||||
if _settings.points_enabled and param_changed:
|
||||
from packages.domain.points_rules import (
|
||||
calculate_viral_video_credits_with_breakdown,
|
||||
resolve_video_dimensions,
|
||||
)
|
||||
|
||||
eff_dur = int(new_duration if new_duration is not None else old_duration)
|
||||
eff_res = new_resolution if new_resolution is not None else old_resolution
|
||||
eff_ratio = new_ratio if new_ratio is not None else old_ratio
|
||||
eff_model = new_model if new_model is not None else (old_model or "seedance-2.5")
|
||||
w, h = resolve_video_dimensions(eff_res, eff_ratio)
|
||||
new_est, _ = calculate_viral_video_credits_with_breakdown(eff_dur, w, h, eff_model or "seedance-2.5")
|
||||
need_points_settle = True
|
||||
|
||||
# 写入新参数(即使不开 points 也要允许用户重试时改参数)
|
||||
if new_duration is not None:
|
||||
job.duration = max(5, min(30, int(new_duration)))
|
||||
if new_resolution is not None:
|
||||
job.video_resolution = new_resolution
|
||||
if new_ratio is not None:
|
||||
job.video_ratio = new_ratio
|
||||
if new_model is not None:
|
||||
job.video_model = new_model
|
||||
|
||||
if need_points_settle:
|
||||
from packages.domain.points_service import PointsService
|
||||
|
||||
old_prepaid = float(getattr(job, "credits_prepaid", 0) or 0)
|
||||
svc = PointsService()
|
||||
diff = round(new_est - old_prepaid, 2)
|
||||
if abs(diff) >= 0.01:
|
||||
if diff > 0:
|
||||
# 新预扣更多:补扣差额
|
||||
res = svc.deduct_viral_video(authenticated_user.user.id, diff, job.id, session)
|
||||
if not res.get("success"):
|
||||
balance = res.get("balance", 0)
|
||||
raise HTTPException(
|
||||
status_code=402,
|
||||
detail={
|
||||
"code": "INSUFFICIENT_POINTS",
|
||||
"message": f"重试参数变更后需补扣 {diff} 积分,余额不足(当前 {balance},需 {new_est})",
|
||||
"required": new_est,
|
||||
"balance": balance,
|
||||
"delta": diff,
|
||||
},
|
||||
)
|
||||
job.credits_prepaid = round(old_prepaid + diff, 2)
|
||||
logger.info(
|
||||
"[爆款视频][retry] 补扣差额 job_id=%s diff=%.2f new_prepaid=%.2f",
|
||||
job.id,
|
||||
diff,
|
||||
job.credits_prepaid,
|
||||
)
|
||||
else:
|
||||
# 新预扣更少:退还差额
|
||||
refund = round(-diff, 2)
|
||||
txn_id = getattr(job, "credits_transaction_id", "") or ""
|
||||
svc.refund_points(
|
||||
user_id=authenticated_user.user.id,
|
||||
amount=refund,
|
||||
source="viral_video",
|
||||
db=session,
|
||||
ref_id=txn_id or job.id,
|
||||
description="爆款视频重试参数变更退费",
|
||||
)
|
||||
job.credits_prepaid = round(old_prepaid - refund, 2)
|
||||
logger.info(
|
||||
"[爆款视频][retry] 退还差额 job_id=%s refund=%.2f new_prepaid=%.2f",
|
||||
job.id,
|
||||
refund,
|
||||
job.credits_prepaid,
|
||||
)
|
||||
# 差额为 0 则不调整
|
||||
|
||||
# 重置状态
|
||||
job.retry_count += 1
|
||||
job.status = ViralVideoStatus.PENDING
|
||||
@@ -704,13 +435,7 @@ def retry_viral_video_job(
|
||||
# 重新入队
|
||||
try:
|
||||
celery_app.send_task("worker.run_viral_video_pipeline", args=[job.id])
|
||||
logger.info(
|
||||
"[爆款视频] 重试入队: job_id=%s retry_count=%d stale=%s params_changed=%s",
|
||||
job.id,
|
||||
job.retry_count,
|
||||
is_stale_running,
|
||||
param_changed,
|
||||
)
|
||||
logger.info("[爆款视频] 重试入队: job_id=%s retry_count=%d stale=%s", job.id, job.retry_count, is_stale_running)
|
||||
except Exception as e:
|
||||
logger.error("[爆款视频] 重试入队失败: %s", e, exc_info=True)
|
||||
job.mark_failed(f"重试入队失败: {e}")
|
||||
@@ -954,28 +679,13 @@ async def viral_video_websocket(websocket: WebSocket, job_id: str) -> None:
|
||||
job = job_repo.get(job_id)
|
||||
if job is not None:
|
||||
status_val = job.status.value if hasattr(job.status, "value") else str(job.status)
|
||||
# #P0: 初始快照必须包含前端重连/刷新所需的业务字段,
|
||||
# 结构对齐 worker 推送的 image_analyzed / copy_generated 事件。
|
||||
data: dict = {"status": status_val}
|
||||
ia = getattr(job, "image_analysis", None)
|
||||
if isinstance(ia, dict) and ia:
|
||||
data["image_analysis"] = ia
|
||||
cr = _build_copy_result(job)
|
||||
if isinstance(cr, dict) and cr:
|
||||
data["copy_result"] = cr
|
||||
gct = getattr(job, "generated_copy_text", "") or ""
|
||||
if gct:
|
||||
data["generated_copy_text"] = gct
|
||||
sb = getattr(job, "storyboard", None) or []
|
||||
if sb:
|
||||
data["storyboard"] = sb
|
||||
initial = {
|
||||
"type": "viral_video:progress",
|
||||
"job_id": job_id,
|
||||
"stage": _stage_from_status(job),
|
||||
"progress": _estimate_progress(job),
|
||||
"message": _initial_message(job),
|
||||
"data": data,
|
||||
"data": {"status": status_val},
|
||||
}
|
||||
await websocket.send_json(initial)
|
||||
# 已经终态 → 再发一条终态事件后立即关闭,避免占连接
|
||||
|
||||
@@ -13,9 +13,9 @@ from pydantic import BaseModel, Field
|
||||
class PointsBalanceResponse(BaseModel):
|
||||
"""积分余额 + 会员状态"""
|
||||
|
||||
balance: float = Field(..., description="当前积分余额")
|
||||
total_earned: float = Field(..., description="累计获得积分")
|
||||
total_spent: float = Field(..., description="累计消耗积分")
|
||||
balance: int = Field(..., description="当前积分余额")
|
||||
total_earned: int = Field(..., description="累计获得积分")
|
||||
total_spent: int = Field(..., description="累计消耗积分")
|
||||
is_member: bool = Field(default=False, description="是否付费会员")
|
||||
member_type: Optional[str] = Field(None, description="会员类型: monthly/quarterly/yearly")
|
||||
member_expires_at: Optional[datetime] = Field(None, description="会员到期时间")
|
||||
@@ -30,8 +30,8 @@ class PointsTransactionItem(BaseModel):
|
||||
id: str
|
||||
type: str = Field(..., description="类型: add/deduct")
|
||||
source: str = Field(..., description="来源场景")
|
||||
amount: float
|
||||
balance_after: float
|
||||
amount: int
|
||||
balance_after: int
|
||||
description: str = ""
|
||||
ref_id: str = ""
|
||||
created_at: Optional[str] = None
|
||||
@@ -99,9 +99,9 @@ class PointsCheckResponse(BaseModel):
|
||||
"""消费前余额检查响应"""
|
||||
|
||||
allowed: bool
|
||||
required_points: float
|
||||
current_balance: float
|
||||
remaining_after: float
|
||||
required_points: int
|
||||
current_balance: int
|
||||
remaining_after: int
|
||||
is_free_quota: bool = False
|
||||
|
||||
|
||||
@@ -112,7 +112,7 @@ class PointsDeductRequest(BaseModel):
|
||||
"""积分扣减请求"""
|
||||
|
||||
scene_key: str
|
||||
amount: float
|
||||
amount: int
|
||||
description: Optional[str] = ""
|
||||
ref_id: Optional[str] = ""
|
||||
|
||||
@@ -170,7 +170,7 @@ class MembershipStatusResponse(BaseModel):
|
||||
is_member: bool
|
||||
member_type: Optional[str] = None
|
||||
member_expires_at: Optional[datetime] = None
|
||||
points_balance: float
|
||||
points_balance: int
|
||||
max_resolution: str = Field(
|
||||
default="1080p",
|
||||
description="可用最高分辨率: 720p(free) / 1080p(paid)",
|
||||
|
||||
@@ -22,7 +22,6 @@ VALID_STAGES = (
|
||||
)
|
||||
VALID_VIDEO_RATIOS = ("9:16", "16:9", "1:1", "4:3", "3:4", "21:9")
|
||||
VALID_DURATIONS = (5, 10, 15, 20, 25, 30)
|
||||
VALID_VIDEO_RESOLUTIONS = ("480p", "720p", "1080p", "普清", "高清", "超清")
|
||||
|
||||
|
||||
# -- 编导脚本结构(v1.6) --
|
||||
@@ -89,7 +88,6 @@ class CreateViralVideoRequest(BaseModel):
|
||||
voice_source: str = ""
|
||||
video_ratio: str = "9:16"
|
||||
video_model: str = ""
|
||||
video_resolution: str = "720p"
|
||||
|
||||
@field_validator("fusion_level")
|
||||
@classmethod
|
||||
@@ -119,7 +117,6 @@ class AnalyzeImagesRequest(BaseModel):
|
||||
voice_source: str = ""
|
||||
video_ratio: str = "9:16"
|
||||
video_model: str = ""
|
||||
video_resolution: str = "720p"
|
||||
duration: int = Field(default=15, ge=5, le=30)
|
||||
|
||||
|
||||
@@ -144,7 +141,6 @@ class GenerateCopyRequest(BaseModel):
|
||||
voice_source: str = ""
|
||||
video_ratio: str = "9:16"
|
||||
video_model: str = ""
|
||||
video_resolution: str = "720p"
|
||||
|
||||
@field_validator("fusion_level")
|
||||
@classmethod
|
||||
@@ -167,10 +163,6 @@ class ConfirmCopyRequest(BaseModel):
|
||||
"""v1.5+ 阶段3:用户确认/编辑口播后开始渲染(TTS+单次Seedance)。"""
|
||||
|
||||
edited_copy: str = Field(default="", description="用户编辑后的口播文案;为空则用 AI 生成的 voiceover_script")
|
||||
video_model: str | None = Field(default=None, description="用户选定的视频生成模型(confirm时可选)")
|
||||
video_resolution: str | None = Field(default=None, description="用户选定的分辨率(confirm时可选)")
|
||||
video_ratio: str | None = Field(default=None, description="用户选定的比例(confirm时可选)")
|
||||
duration: int | None = Field(default=None, ge=5, le=30, description="用户选定的时长秒数(confirm时可选,5~30)")
|
||||
|
||||
|
||||
class ConfirmIntentRequest(BaseModel):
|
||||
@@ -226,10 +218,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
|
||||
credits_cost: int = 0
|
||||
error_msg: str = ""
|
||||
retry_count: int = 0
|
||||
started_at: datetime | None = None
|
||||
@@ -261,57 +250,6 @@ class AnalyzeStyleResponse(BaseModel):
|
||||
style_guide: dict | None = None
|
||||
|
||||
|
||||
# -- 积分预估 --
|
||||
|
||||
|
||||
class EstimateCreditsRequest(BaseModel):
|
||||
"""爆款视频积分预估请求。
|
||||
|
||||
前端可传 model 或 video_model(兼容老字段);resolution/ratio/duration 为预估所需参数。
|
||||
"""
|
||||
|
||||
model: str = Field(default="", alias="video_model")
|
||||
resolution: str = Field(default="720p", alias="video_resolution")
|
||||
ratio: str = Field(default="9:16", alias="video_ratio")
|
||||
duration: int = Field(default=15, ge=5, le=30)
|
||||
|
||||
model_config = {"populate_by_name": True}
|
||||
|
||||
|
||||
class CreditsFormulaBreakdown(BaseModel):
|
||||
"""爆款视频积分计费公式明细(前端展示用)。"""
|
||||
|
||||
tokens: float = Field(..., description="估算视频 tokens 数 (duration*width*height*fps/1024)")
|
||||
video_cost: float = Field(..., description="视频生成成本(元)= tokens/1e6 * model_price")
|
||||
fixed_cost: float = Field(..., description="固定成本(元),含 VLM/LLM/TTS/OSS/服务器")
|
||||
profit_multiplier: float = Field(..., description="利润系数(默认 1.3)")
|
||||
model_price: float = Field(..., description="模型单价(元/百万 tokens)")
|
||||
width: int = Field(..., description="视频宽度像素")
|
||||
height: int = Field(..., description="视频高度像素")
|
||||
fps: int = Field(..., description="视频帧率")
|
||||
|
||||
|
||||
class EstimateCreditsResponse(BaseModel):
|
||||
"""爆款视频积分预估响应。"""
|
||||
|
||||
estimated_credits: float
|
||||
formula_breakdown: CreditsFormulaBreakdown = Field(..., description="计费公式明细")
|
||||
|
||||
|
||||
class RetryViralVideoRequest(BaseModel):
|
||||
"""重试爆款视频任务的请求体(可选,允许改参数重新预估积分多退少补)。
|
||||
|
||||
不传 body 或字段全缺省:保持原参数、不重新扣点,走默认重置+入队逻辑。
|
||||
传入新的 duration/video_resolution/video_ratio/video_model:重新预估积分,
|
||||
与原 credits_prepaid 比较后多退少补(差额补扣不足抛 402)。
|
||||
"""
|
||||
|
||||
duration: int | None = Field(default=None, ge=5, le=30, description="重试时新的视频时长(秒)")
|
||||
video_resolution: str | None = Field(default=None, description="重试时新的分辨率,如 720p/1080p")
|
||||
video_ratio: str | None = Field(default=None, description="重试时新的画幅比,如 9:16/16:9")
|
||||
video_model: str | None = Field(default=None, description="重试时新的视频模型,如 seedance-2.5")
|
||||
|
||||
|
||||
# -- WebSocket 事件 Schema --
|
||||
|
||||
|
||||
|
||||
@@ -228,8 +228,6 @@ class GpuLipsyncService:
|
||||
lipsync_job_id: str = "",
|
||||
user_id: str = "",
|
||||
project_id: str = "",
|
||||
credits_prepaid: float = 0.0,
|
||||
credits_transaction_id: str = "",
|
||||
) -> GpuLipsyncTaskModel:
|
||||
task_id = str(uuid.uuid4())
|
||||
now = datetime.now(UTC)
|
||||
@@ -242,8 +240,6 @@ class GpuLipsyncService:
|
||||
audio_url=audio_url,
|
||||
status="pending",
|
||||
attempt=0,
|
||||
credits_prepaid=float(credits_prepaid or 0.0),
|
||||
credits_transaction_id=str(credits_transaction_id or ""),
|
||||
created_at=now,
|
||||
updated_at=now,
|
||||
)
|
||||
|
||||
@@ -38,7 +38,6 @@ from sqlalchemy.orm import Session
|
||||
from packages.adapters.sqlalchemy_impl.models import LipsyncJobModel
|
||||
from packages.application.cosyvoice_service import CosyVoiceError
|
||||
from packages.config import get_api_settings
|
||||
from packages.domain import feature_pricing_service
|
||||
from packages.domain.sentence_timings import (
|
||||
compute_sentence_timings,
|
||||
probe_audio_duration,
|
||||
@@ -223,47 +222,7 @@ class LipsyncService:
|
||||
if timings:
|
||||
job.sentence_timings = timings
|
||||
|
||||
# 4. 检查是否走 Ditto(蚂蚁数字人,#2076):开关 + 配置完整
|
||||
use_ditto = False
|
||||
if self.settings.use_ditto_lipsync:
|
||||
try:
|
||||
from packages.application.ditto_service import get_ditto_client
|
||||
|
||||
ditto = get_ditto_client()
|
||||
if ditto.is_configured:
|
||||
use_ditto = True
|
||||
logger.info("[lipsync] 优先走 Ditto 蚂蚁数字人: job_id=%s", job.id)
|
||||
else:
|
||||
logger.info(
|
||||
"[lipsync] Ditto 开关已开但配置不完整(base_url=%s, template=%s),继续判断 GPU: job_id=%s",
|
||||
bool(ditto.base_url),
|
||||
bool(ditto.default_video_url),
|
||||
job.id,
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.warning("[lipsync] Ditto 初始化失败,继续判断 GPU: job_id=%s err=%s", job.id, exc)
|
||||
|
||||
if use_ditto:
|
||||
try:
|
||||
# Ditto 使用预置人物模板视频,不用用户上传的 video_url;
|
||||
# 但保留用户 video_url 以便失败回退到 GPU/MediaKit。
|
||||
job.status = "processing"
|
||||
job.mediakit_task_id = "ditto:submitted"
|
||||
job.updated_at = datetime.now(UTC)
|
||||
self.db.commit()
|
||||
from app.tasks.lipsync_ditto import lipsync_ditto_process_async
|
||||
|
||||
lipsync_ditto_process_async.apply_async(args=(job.id, job.user_id))
|
||||
logger.info("[lipsync] Ditto 任务已异步派发: job_id=%s", job.id)
|
||||
return
|
||||
except Exception as exc:
|
||||
logger.warning("[lipsync] Ditto 派发失败,回退 GPU/MediaKit: job_id=%s err=%s", job.id, exc)
|
||||
try:
|
||||
self.db.rollback()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# 5. 检查是否走 GPU 路径:开关打开 + 有可用 Worker
|
||||
# 4. 检查是否走 GPU 路径:开关打开 + 有可用 Worker
|
||||
use_gpu = False
|
||||
if self.settings.use_gpu_lipsync:
|
||||
try:
|
||||
@@ -409,8 +368,6 @@ class LipsyncService:
|
||||
lipsync_job_id=job.id,
|
||||
user_id=job.user_id,
|
||||
project_id=job.project_id,
|
||||
credits_prepaid=float(getattr(job, "credits_prepaid", 0) or 0),
|
||||
credits_transaction_id=str(getattr(job, "credits_transaction_id", "") or ""),
|
||||
)
|
||||
logger.info(
|
||||
"[lipsync] 已创建 GPU 任务(异步): job_id=%s gpu_task=%s",
|
||||
@@ -458,121 +415,6 @@ class LipsyncService:
|
||||
job.output_duration,
|
||||
)
|
||||
|
||||
# ── lip_sync 计费辅助 ────────────────────────────────────────────────
|
||||
|
||||
@staticmethod
|
||||
def _estimate_duration(
|
||||
*,
|
||||
audio_duration: Optional[float] = None,
|
||||
sentence_timings: Optional[list] = None,
|
||||
script_text: str = "",
|
||||
) -> float:
|
||||
"""预估音频/成片秒数。
|
||||
|
||||
优先级:audio_duration(预合成前端已 ffprobe)> timings 末句 end_time >
|
||||
脚本字数 / 5 字每秒 > 默认 10 秒。
|
||||
"""
|
||||
if audio_duration and float(audio_duration) > 0:
|
||||
return float(audio_duration)
|
||||
if sentence_timings:
|
||||
max_end = 0.0
|
||||
for item in sentence_timings:
|
||||
if isinstance(item, dict):
|
||||
end = item.get("end_time") or item.get("end") or 0.0
|
||||
else:
|
||||
end = 0.0
|
||||
try:
|
||||
max_end = max(max_end, float(end))
|
||||
except (TypeError, ValueError):
|
||||
continue
|
||||
if max_end > 0:
|
||||
return max_end
|
||||
text = (script_text or "").strip()
|
||||
if text:
|
||||
return max(1.0, len(text) / 5.0)
|
||||
return 10.0
|
||||
|
||||
def _settle_lip_sync(self, job: LipsyncJobModel, actual_duration: float) -> None:
|
||||
"""按实际时长结算(首期只退不补:final < prepaid 退差额,> 不补)。
|
||||
|
||||
幂等:credits_cost 已 > 0 说明结算过,直接跳过。
|
||||
结算失败不阻塞业务(结果已产出),仅记录日志。
|
||||
"""
|
||||
try:
|
||||
prepaid = float(getattr(job, "credits_prepaid", 0) or 0)
|
||||
if prepaid <= 0:
|
||||
return
|
||||
if float(getattr(job, "credits_cost", 0) or 0) > 0:
|
||||
return
|
||||
feature_cfg = feature_pricing_service.get_feature_config("lip_sync")
|
||||
unit_cost = float(feature_cfg.dynamic_unit_cost) if feature_cfg is not None else 0.0
|
||||
duration = float(actual_duration or 0.0)
|
||||
if duration <= 0:
|
||||
duration = self._estimate_duration(
|
||||
sentence_timings=job.sentence_timings,
|
||||
script_text=job.script_text,
|
||||
)
|
||||
final_price, _bd = feature_pricing_service.calculate_price("lip_sync", duration * unit_cost)
|
||||
final_price = round(float(final_price), 2)
|
||||
job.credits_cost = final_price
|
||||
if final_price < prepaid - 0.009:
|
||||
refund = round(prepaid - final_price, 2)
|
||||
from packages.domain.points_service import PointsService
|
||||
|
||||
res = PointsService().refund_points(
|
||||
user_id=job.user_id,
|
||||
amount=refund,
|
||||
source="lip_sync",
|
||||
db=self.db,
|
||||
ref_id=str(job.credits_transaction_id or job.id),
|
||||
description="对口型结算退费",
|
||||
)
|
||||
if not res.get("success"):
|
||||
logger.warning(
|
||||
"[lip_sync] 结算退费失败 job_id=%s refund=%.2f(不阻塞)",
|
||||
job.id,
|
||||
refund,
|
||||
)
|
||||
# final > prepaid:首期只退不补,不补扣
|
||||
self.db.commit()
|
||||
except Exception: # noqa: BLE001
|
||||
logger.exception("[lip_sync] 结算异常 job_id=%s(不阻塞结果)", job.id)
|
||||
try:
|
||||
self.db.rollback()
|
||||
except Exception: # noqa: BLE001
|
||||
pass
|
||||
|
||||
def _refund_lip_sync(self, job: LipsyncJobModel) -> None:
|
||||
"""任务失败/取消时全额退还预扣积分(credits_cost 已结算则退实际未消耗部分)。"""
|
||||
try:
|
||||
prepaid = float(getattr(job, "credits_prepaid", 0) or 0)
|
||||
if prepaid <= 0:
|
||||
return
|
||||
txn_id = str(getattr(job, "credits_transaction_id", "") or "")
|
||||
cost = float(getattr(job, "credits_cost", 0) or 0)
|
||||
refund = round(prepaid - cost, 2) if cost > 0 else round(prepaid, 2)
|
||||
if refund <= 0:
|
||||
return
|
||||
from packages.domain.points_service import PointsService
|
||||
|
||||
res = PointsService().refund_points(
|
||||
user_id=job.user_id,
|
||||
amount=refund,
|
||||
source="lip_sync",
|
||||
db=self.db,
|
||||
ref_id=txn_id or job.id,
|
||||
description="对口型失败/取消退款",
|
||||
)
|
||||
if res.get("success"):
|
||||
job.credits_cost = prepaid # 标记已全额退回,防重复退
|
||||
self.db.commit()
|
||||
except Exception: # noqa: BLE001
|
||||
logger.exception("[lip_sync] 退款异常 job_id=%s", job.id)
|
||||
try:
|
||||
self.db.rollback()
|
||||
except Exception: # noqa: BLE001
|
||||
pass
|
||||
|
||||
# ── 创建任务 ──────────────────────────────────────────────────────────
|
||||
|
||||
def create_job(
|
||||
@@ -624,35 +466,6 @@ class LipsyncService:
|
||||
if not isinstance(sentence_timings, list) or len(sentence_timings) == 0:
|
||||
raise MediaKitError("预合成模式 sentence_timings 不能为空", code="InvalidInput")
|
||||
|
||||
# 0.5 lip_sync 计费预扣(全局 points 开关 + 功能开关均开才扣)
|
||||
prepaid_credits = 0.0
|
||||
prepaid_txn_id = ""
|
||||
if feature_pricing_service.is_feature_enabled("lip_sync"):
|
||||
est_duration = self._estimate_duration(
|
||||
audio_duration=audio_duration,
|
||||
sentence_timings=sentence_timings,
|
||||
script_text=script_text,
|
||||
)
|
||||
feature_cfg = feature_pricing_service.get_feature_config("lip_sync")
|
||||
unit_cost = float(feature_cfg.dynamic_unit_cost) if feature_cfg is not None else 0.0
|
||||
dynamic_cost = est_duration * unit_cost
|
||||
prepaid_credits, _bd = feature_pricing_service.calculate_price("lip_sync", dynamic_cost)
|
||||
if prepaid_credits > 0:
|
||||
from packages.domain.points_service import PointsService
|
||||
|
||||
res = PointsService().deduct_points(
|
||||
user_id=user_id,
|
||||
amount=prepaid_credits,
|
||||
source="lip_sync",
|
||||
db=self.db,
|
||||
description="对口型生成预扣",
|
||||
)
|
||||
if not res.get("success"):
|
||||
raise ValueError(
|
||||
f"积分不足:本次对口型需 {prepaid_credits:.2f} 积分,当前余额 {res.get('balance', 0)}"
|
||||
)
|
||||
prepaid_txn_id = str(res.get("transaction_id") or "")
|
||||
|
||||
# 1. 创建数据库记录
|
||||
job_id = str(uuid.uuid4())
|
||||
job = LipsyncJobModel(
|
||||
@@ -669,8 +482,6 @@ class LipsyncService:
|
||||
emotion=emotion or "",
|
||||
# 音频直传(含预合成)直接进入 pending(后续同步改为 submitted);TTS 模式进入 tts_processing
|
||||
status="tts_processing" if is_tts_mode else "pending",
|
||||
credits_prepaid=prepaid_credits,
|
||||
credits_transaction_id=prepaid_txn_id,
|
||||
)
|
||||
self.db.add(job)
|
||||
self.db.flush()
|
||||
@@ -845,30 +656,6 @@ class LipsyncService:
|
||||
if job.status in (STATUS_COMPLETED, "failed"):
|
||||
return job
|
||||
|
||||
# Ditto 异步路径:mediakit_task_id 以 "ditto:" 开头,由 Celery 任务异步更新
|
||||
# 不做 MediaKit 轮询,只检查是否卡住太久(>10 分钟)则标失败
|
||||
if job.mediakit_task_id and job.mediakit_task_id.startswith("ditto:"):
|
||||
if job.status in ("processing", "submitted"):
|
||||
_now = datetime.now(UTC)
|
||||
_upd = job.updated_at
|
||||
if _upd is not None and _upd.tzinfo is None:
|
||||
_upd = _upd.replace(tzinfo=UTC)
|
||||
stale_minutes = 10
|
||||
if _upd and (_now - _upd).total_seconds() > stale_minutes * 60:
|
||||
logger.warning(
|
||||
"Ditto 异步任务超时(>%d 分钟),标记失败: job_id=%s",
|
||||
stale_minutes,
|
||||
job_id,
|
||||
)
|
||||
job.status = "failed"
|
||||
job.error_message = f"Ditto 处理超时(>{stale_minutes} 分钟)"
|
||||
job.error_code = "DittoTimeout"
|
||||
job.completed_at = _now
|
||||
job.updated_at = _now
|
||||
self.db.commit()
|
||||
self._refund_lip_sync(job)
|
||||
return job
|
||||
|
||||
# GPU 异步路径:mediakit_task_id 以 "gpu:" 开头,由 Celery 任务异步更新
|
||||
# 不做 MediaKit 轮询,只检查是否卡住太久(>30 分钟)则标失败
|
||||
if job.mediakit_task_id and job.mediakit_task_id.startswith("gpu:"):
|
||||
@@ -890,8 +677,6 @@ class LipsyncService:
|
||||
job.completed_at = _now
|
||||
job.updated_at = _now
|
||||
self.db.commit()
|
||||
# lip_sync 超时全额退款
|
||||
self._refund_lip_sync(job)
|
||||
return job
|
||||
|
||||
# 未提交的任务不轮询
|
||||
@@ -917,8 +702,6 @@ class LipsyncService:
|
||||
job.completed_at = datetime.now(UTC)
|
||||
job.updated_at = datetime.now(UTC)
|
||||
self.db.commit()
|
||||
# lip_sync 结算(只退不补)
|
||||
self._settle_lip_sync(job, float(job.output_duration or 0.0))
|
||||
# 异步转存自家 OSS
|
||||
try:
|
||||
from app.tasks.lipsync_tts import persist_output_video_task
|
||||
@@ -936,8 +719,6 @@ class LipsyncService:
|
||||
job.error_message = error.get("message", "任务执行失败")
|
||||
job.error_code = error.get("code", "TaskFailed")
|
||||
job.completed_at = datetime.now(UTC)
|
||||
# lip_sync 失败全额退款(先退款再统一 commit)
|
||||
self._refund_lip_sync(job)
|
||||
else:
|
||||
# 中间状态(running/processing/queued 等)同步到 DB,避免前端永远卡在 submitted
|
||||
if isinstance(mk_status, str) and mk_status:
|
||||
@@ -1031,8 +812,6 @@ class LipsyncService:
|
||||
job.status = "cancelled"
|
||||
job.updated_at = datetime.now(UTC)
|
||||
self.db.commit()
|
||||
# lip_sync 取消全额退款
|
||||
self._refund_lip_sync(job)
|
||||
self.db.refresh(job)
|
||||
|
||||
return job
|
||||
|
||||
@@ -11,12 +11,14 @@
|
||||
存储路径与元信息约定),返回 asset_id —— 下游仍以 voice_library_id(实为
|
||||
audio asset id)消费,渲染链路零改动。
|
||||
|
||||
积分扣点与 /tts 合成端点保持一致(ai_voice 场景),失败退费。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import math
|
||||
import subprocess
|
||||
import tempfile
|
||||
from dataclasses import dataclass
|
||||
@@ -30,10 +32,13 @@ from packages.application.cosyvoice_service import CosyVoiceService
|
||||
from packages.application.tts_job.use_cases import CreateTTSJobUseCase
|
||||
from packages.application.tts_job.workflow import TTSWorkflowService
|
||||
from packages.domain import Asset, AssetLibrary, AssetLibraryKind, AssetStatus, ClassificationStatus
|
||||
from packages.domain.points_rules import calculate_points_cost
|
||||
from packages.domain.points_service import PointsService
|
||||
from packages.shared.storage import SharedStorageService
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_POINTS_SCENE = "ai_voice"
|
||||
_SYNTH_TIMEOUT = 180.0 # 叙事配音在 HTTP 请求内同步等待,长文案分段合成时留出余量
|
||||
_CONTENT_TYPE_MAP = {"mp3": "audio/mpeg", "wav": "audio/wav", "pcm": "audio/pcm", "opus": "audio/opus"}
|
||||
|
||||
@@ -268,6 +273,24 @@ def prepare_narrative_voice(
|
||||
voice_clone_repository=voice_clone_repository,
|
||||
)
|
||||
|
||||
# 积分扣点(与 /tts 合成端点同口径),失败时在合成失败分支退费
|
||||
points_svc = PointsService() if points_enabled else None
|
||||
points_deducted = 0
|
||||
if points_svc is not None:
|
||||
est_minutes = max(1.0, math.ceil(len(content) / 240))
|
||||
points_deducted = calculate_points_cost(
|
||||
_POINTS_SCENE,
|
||||
is_member=is_member,
|
||||
duration_minutes=est_minutes,
|
||||
member_type=member_type,
|
||||
)
|
||||
deduct_res = points_svc.deduct_points(user_id, points_deducted, _POINTS_SCENE, db)
|
||||
if not deduct_res["success"]:
|
||||
raise NarrativeError(
|
||||
f"积分不足,需要 {points_deducted} 积分,当前余额 {deduct_res['balance']}",
|
||||
status_code=402,
|
||||
)
|
||||
|
||||
use_case = CreateTTSJobUseCase(tts_repository)
|
||||
job = use_case.execute(
|
||||
user_id=user_id,
|
||||
@@ -288,9 +311,19 @@ def prepare_narrative_voice(
|
||||
workflow.process_synthesis_failure(job.id, str(e))
|
||||
except Exception: # noqa: BLE001
|
||||
logger.warning("标记叙事 TTS job 失败出错: job_id=%s", job.id, exc_info=True)
|
||||
if points_deducted and points_svc is not None:
|
||||
try:
|
||||
points_svc.refund_points(user_id, points_deducted, _POINTS_SCENE, db, ref_id=job.id)
|
||||
except Exception: # noqa: BLE001
|
||||
logger.warning("叙事 TTS 失败退积分异常: job_id=%s", job.id, exc_info=True)
|
||||
raise NarrativeError(f"配音合成失败:{e}", status_code=502) from e
|
||||
|
||||
if not job.is_completed:
|
||||
if points_deducted and points_svc is not None:
|
||||
try:
|
||||
points_svc.refund_points(user_id, points_deducted, _POINTS_SCENE, db, ref_id=job.id)
|
||||
except Exception: # noqa: BLE001
|
||||
logger.warning("叙事 TTS 未完成退积分异常: job_id=%s", job.id, exc_info=True)
|
||||
raise NarrativeError("配音合成未完成,请稍后重试", status_code=504)
|
||||
|
||||
asset = _save_tts_job_as_voice_asset(
|
||||
|
||||
@@ -1,351 +0,0 @@
|
||||
"""Ditto 蚂蚁数字人口型异步任务 — #2076.
|
||||
|
||||
把 Ditto 同步 HTTP 调用(30-120s)从 API 请求移到 Celery 后台执行:
|
||||
1. 加载 LipsyncJob
|
||||
2. 调 DittoClient.generate_and_persist(video_url=默认模板, audio_url=job.audio_url, script=job.script_text)
|
||||
3. 成功:标记 completed,写入 output_video_url(Ditto 输出自带音频,无需二次混流/超分)
|
||||
4. 失败:回退 GPU MuseTalk → 再失败回退 MediaKit
|
||||
|
||||
注意:
|
||||
- 保留 MuseTalk 代码不动;Ditto 优先,失败按原链路兜底
|
||||
- Ditto 使用预置的人物模板视频(settings.ditto_default_video_url),不用用户上传的 video_url
|
||||
- 不传 GFPGAN 超分,不需要 ffmpeg 音视频混流
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from datetime import UTC, datetime
|
||||
from typing import TYPE_CHECKING, Optional
|
||||
|
||||
from celery import shared_task
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from packages.adapters.sqlalchemy_impl.models import LipsyncJobModel
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_DITTO_URL_TTL_SECONDS = 7 * 24 * 3600 # Ditto 结果 OSS URL 7 天有效
|
||||
|
||||
|
||||
def _get_db_session() -> Session:
|
||||
try:
|
||||
from worker_app.db import SessionLocal # type: ignore
|
||||
except ImportError:
|
||||
from app.db import SessionLocal # type: ignore
|
||||
return SessionLocal()
|
||||
|
||||
|
||||
def _sign_media_url(url: str) -> str:
|
||||
"""对自家 OSS URL 签 7 天预签名。"""
|
||||
if not url:
|
||||
return url
|
||||
try:
|
||||
from urllib.parse import urlparse
|
||||
|
||||
from packages.shared.storage import get_shared_storage_service
|
||||
|
||||
storage = get_shared_storage_service()
|
||||
public_base = getattr(storage, "public_url", "")
|
||||
if not isinstance(public_base, str) or not public_base:
|
||||
return url
|
||||
own_host = urlparse(public_base).netloc.lower()
|
||||
host = urlparse(url).netloc.lower()
|
||||
if not own_host or host != own_host:
|
||||
return url
|
||||
return storage.get_download_url(url, expires_seconds=_DITTO_URL_TTL_SECONDS)
|
||||
except Exception:
|
||||
return url
|
||||
|
||||
|
||||
def _probe_video_duration(video_bytes: bytes) -> float:
|
||||
"""用 ffprobe 探测视频时长(秒);失败返回 0。"""
|
||||
try:
|
||||
import os
|
||||
import subprocess
|
||||
import tempfile
|
||||
|
||||
with tempfile.NamedTemporaryFile(suffix=".mp4", delete=False) as tmp:
|
||||
tmp.write(video_bytes)
|
||||
tmp_path = tmp.name
|
||||
try:
|
||||
out = subprocess.check_output(
|
||||
[
|
||||
"ffprobe",
|
||||
"-v",
|
||||
"error",
|
||||
"-show_entries",
|
||||
"format=duration",
|
||||
"-of",
|
||||
"default=noprint_wrappers=1:nokey=1",
|
||||
tmp_path,
|
||||
],
|
||||
stderr=subprocess.STDOUT,
|
||||
timeout=10,
|
||||
)
|
||||
return float(out.decode().strip() or 0)
|
||||
finally:
|
||||
os.unlink(tmp_path)
|
||||
except Exception as exc:
|
||||
logger.warning("[ditto_task] ffprobe 失败: %s", exc)
|
||||
return 0.0
|
||||
|
||||
|
||||
def _refund_lip_sync(db: Session, job: "LipsyncJobModel") -> None:
|
||||
"""Ditto 失败/取消时全额退款(复用 lipsync_service 的退款逻辑)。"""
|
||||
try:
|
||||
from app.services.lipsync_service import LipsyncService
|
||||
|
||||
LipsyncService(db)._refund_lip_sync(job)
|
||||
except Exception:
|
||||
logger.exception("[ditto_task] lip_sync 退款异常 job_id=%s", job.id)
|
||||
|
||||
|
||||
def _settle_lip_sync(db: Session, job: "LipsyncJobModel", duration: float) -> None:
|
||||
"""Ditto 成功后按实际时长结算。"""
|
||||
try:
|
||||
from app.services.lipsync_service import LipsyncService
|
||||
|
||||
LipsyncService(db)._settle_lip_sync(job, duration)
|
||||
except Exception:
|
||||
logger.exception("[ditto_task] lip_sync 结算异常 job_id=%s(不阻塞)", job.id)
|
||||
|
||||
|
||||
def _fallback_to_gpu_then_mediakit(db: Session, job: "LipsyncJobModel") -> None:
|
||||
"""Ditto 失败后:优先回退 GPU MuseTalk,再回退 MediaKit 云端。
|
||||
|
||||
复用 lipsync_service 现有路径逻辑以保证兜底一致性。
|
||||
"""
|
||||
# 先尝试走 GPU MuseTalk(若可用)
|
||||
try:
|
||||
from app.services.gpu_lipsync_service import GpuLipsyncService
|
||||
from app.tasks.lipsync_gpu import lipsync_gpu_process_async
|
||||
|
||||
gpu_svc = GpuLipsyncService(db)
|
||||
if gpu_svc.has_available_worker():
|
||||
logger.info("[ditto_task] 回退 GPU MuseTalk: job_id=%s", job.id)
|
||||
# 复用 lipsync_service._submit_to_gpu_create 逻辑
|
||||
from app.services.lipsync_service import LipsyncService
|
||||
|
||||
svc = LipsyncService(db)
|
||||
storage = _shared_storage()
|
||||
persisted_audio = None
|
||||
try:
|
||||
persisted_audio = svc._persist_external_audio_for_gpu(job=job, storage=storage)
|
||||
except Exception as exc:
|
||||
logger.warning("[ditto_task] GPU 外部音频转存失败: %s", exc)
|
||||
audio_url_for_task = persisted_audio or job.audio_url
|
||||
gpu_task = gpu_svc.create_task(
|
||||
video_url=job.video_url,
|
||||
audio_url=audio_url_for_task,
|
||||
lipsync_job_id=job.id,
|
||||
user_id=job.user_id,
|
||||
)
|
||||
if gpu_task is not None:
|
||||
job.mediakit_task_id = f"gpu:{gpu_task.id}"
|
||||
job.status = "processing"
|
||||
job.updated_at = datetime.now(UTC)
|
||||
db.commit()
|
||||
lipsync_gpu_process_async.apply_async(args=(job.id, job.user_id, gpu_task.id))
|
||||
return
|
||||
db.rollback()
|
||||
except Exception as exc:
|
||||
logger.warning("[ditto_task] GPU MuseTalk 回退失败,转 MediaKit: %s", exc)
|
||||
try:
|
||||
db.rollback()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# 最后兜底:MediaKit 云端
|
||||
try:
|
||||
from app.services.mediakit_client import get_mediakit_client
|
||||
|
||||
client = get_mediakit_client()
|
||||
video_url = _sign_media_url(job.video_url)
|
||||
signed_audio_url = _sign_media_url(job.audio_url)
|
||||
result = client.submit_lipsync(
|
||||
video_url=video_url,
|
||||
audio_url=signed_audio_url,
|
||||
enable_video_loop=job.enable_video_loop,
|
||||
client_token=job.id,
|
||||
)
|
||||
job.mediakit_task_id = result["task_id"]
|
||||
job.status = "submitted"
|
||||
job.submitted_at = datetime.now(UTC)
|
||||
job.updated_at = datetime.now(UTC)
|
||||
db.commit()
|
||||
logger.info("[ditto_task] 已回退 MediaKit: job_id=%s task_id=%s", job.id, result["task_id"])
|
||||
except Exception as exc:
|
||||
job.status = "failed"
|
||||
job.error_message = f"Ditto/GPU/MediaKit 均失败: {exc}"
|
||||
job.error_code = "AllBackendsFailed"
|
||||
job.updated_at = datetime.now(UTC)
|
||||
db.commit()
|
||||
logger.error("[ditto_task] 所有兜底均失败: job_id=%s err=%s", job.id, exc)
|
||||
|
||||
|
||||
def _shared_storage():
|
||||
from packages.shared.storage import get_shared_storage_service
|
||||
|
||||
return get_shared_storage_service()
|
||||
|
||||
|
||||
@shared_task(
|
||||
name="lipsync_ditto_process_async",
|
||||
bind=True,
|
||||
max_retries=0,
|
||||
acks_late=True,
|
||||
time_limit=600,
|
||||
soft_time_limit=540,
|
||||
)
|
||||
def lipsync_ditto_process_async(self, job_id: str, user_id: str) -> None:
|
||||
"""异步调用 Ditto 生成口型视频。
|
||||
|
||||
Args:
|
||||
job_id: LipsyncJob ID
|
||||
user_id: 用户 ID
|
||||
"""
|
||||
from packages.application.ditto_emotion_service import get_ditto_emotion_service
|
||||
from packages.application.ditto_service import DittoError, get_ditto_client
|
||||
|
||||
db: Session = _get_db_session()
|
||||
job: Optional[LipsyncJobModel] = None
|
||||
try:
|
||||
from packages.adapters.sqlalchemy_impl.models import LipsyncJobModel
|
||||
|
||||
job = db.query(LipsyncJobModel).filter_by(id=job_id, user_id=user_id).first()
|
||||
if job is None:
|
||||
logger.error("[ditto_task] job 不存在: job_id=%s", job_id)
|
||||
return
|
||||
|
||||
if job.status != "processing":
|
||||
logger.warning(
|
||||
"[ditto_task] job 状态异常(非 processing),跳过: job_id=%s status=%s",
|
||||
job_id,
|
||||
job.status,
|
||||
)
|
||||
return
|
||||
|
||||
audio_url = job.audio_url or ""
|
||||
script = job.script_text or ""
|
||||
if not audio_url:
|
||||
raise DittoError("job.audio_url 为空,无法调用 Ditto", code="InvalidParam")
|
||||
|
||||
logger.info(
|
||||
"[ditto_task] 开始 Ditto 生成: job_id=%s audio=%s script_len=%d",
|
||||
job_id,
|
||||
audio_url[:100],
|
||||
len(script),
|
||||
)
|
||||
# ── LLM 情绪分析(#2076 后续):生成 emo_timeline ──
|
||||
emo_timeline = ""
|
||||
try:
|
||||
emo_svc = get_ditto_emotion_service()
|
||||
if emo_svc.enabled and script:
|
||||
# 探测音频时长用于时间对齐
|
||||
try:
|
||||
from packages.domain.sentence_timings import probe_audio_duration
|
||||
from packages.shared.url_security import safe_download_bytes
|
||||
|
||||
audio_bytes = safe_download_bytes(
|
||||
audio_url,
|
||||
allowed_mime_types=("audio/mpeg", "audio/wav", "audio/x-wav", "audio/mp3"),
|
||||
timeout=30,
|
||||
)
|
||||
audio_duration = probe_audio_duration(audio_bytes)
|
||||
except Exception as audio_exc:
|
||||
logger.warning("[ditto_task] 音频时长探测失败,emo_timeline 降级空: %s", audio_exc)
|
||||
audio_duration = 0.0
|
||||
if audio_duration > 0:
|
||||
sentence_timings = getattr(job, "sentence_timings", None)
|
||||
emo_timeline = emo_svc.build_timeline(
|
||||
text=script,
|
||||
audio_duration=audio_duration,
|
||||
sentence_timings=sentence_timings,
|
||||
)
|
||||
if emo_timeline:
|
||||
logger.info("[ditto_task] 情绪时间线已生成: segments=%d", len(emo_timeline) // 50)
|
||||
except Exception as emo_exc:
|
||||
logger.warning("[ditto_task] 情绪分析异常(降级中性): %s", emo_exc)
|
||||
emo_timeline = ""
|
||||
client = get_ditto_client()
|
||||
result = client.generate_and_persist(
|
||||
job_id=job_id,
|
||||
user_id=user_id,
|
||||
audio_url=audio_url,
|
||||
script=script,
|
||||
emo_timeline=emo_timeline,
|
||||
# video_url 不传则用默认模板
|
||||
)
|
||||
|
||||
# Ditto 返回的 MP4 自带音频,签名 OSS URL(7天有效)后标记完成
|
||||
job.output_video_url = _sign_media_url(result.video_url)
|
||||
# 探测时长(用于计费)
|
||||
duration = _probe_video_duration(result.video_bytes)
|
||||
if duration <= 0:
|
||||
# 兜底:按音频时长估算(1秒≈1秒)
|
||||
try:
|
||||
from packages.domain.sentence_timings import probe_audio_duration
|
||||
from packages.shared.url_security import safe_download_bytes
|
||||
|
||||
audio_data = safe_download_bytes(
|
||||
audio_url, allowed_mime_types=("audio/mpeg", "audio/wav", "audio/x-wav"), timeout=30
|
||||
)
|
||||
duration = probe_audio_duration(audio_data)
|
||||
except Exception:
|
||||
duration = 0.0
|
||||
job.output_duration = duration
|
||||
job.status = "completed"
|
||||
job.completed_at = datetime.now(UTC)
|
||||
job.updated_at = datetime.now(UTC)
|
||||
db.commit()
|
||||
logger.info(
|
||||
"[ditto_task] Ditto 完成: job_id=%s url=%s duration=%.2fs rtf=%.2f frames=%d",
|
||||
job_id,
|
||||
result.video_url[:100],
|
||||
duration,
|
||||
result.rtf,
|
||||
result.frames,
|
||||
)
|
||||
_settle_lip_sync(db, job, duration)
|
||||
|
||||
except DittoError as exc:
|
||||
logger.error("[ditto_task] Ditto 失败,回退: job_id=%s code=%s err=%s", job_id, exc.code, exc)
|
||||
if job is not None:
|
||||
try:
|
||||
db.rollback()
|
||||
job = db.query(type(job)).filter_by(id=job_id).first() if hasattr(job, "id") else job
|
||||
# 回退 GPU/MediaKit
|
||||
_fallback_to_gpu_then_mediakit(db, job)
|
||||
except Exception as fallback_exc:
|
||||
logger.exception("[ditto_task] 回退也失败 job_id=%s err=%s", job_id, fallback_exc)
|
||||
try:
|
||||
if job:
|
||||
job.status = "failed"
|
||||
job.error_message = f"Ditto 失败且回退异常: {exc}; fallback: {fallback_exc}"
|
||||
job.error_code = "FallbackError"
|
||||
job.updated_at = datetime.now(UTC)
|
||||
db.commit()
|
||||
except Exception:
|
||||
pass
|
||||
except Exception as exc:
|
||||
logger.exception("[ditto_task] 未预期异常: job_id=%s err=%s", job_id, exc)
|
||||
if job is not None:
|
||||
try:
|
||||
db.rollback()
|
||||
job = db.query(type(job)).filter_by(id=job_id).first()
|
||||
_fallback_to_gpu_then_mediakit(db, job)
|
||||
except Exception as fallback_exc:
|
||||
logger.exception("[ditto_task] 回退也失败 job_id=%s err=%s", job_id, fallback_exc)
|
||||
try:
|
||||
if job:
|
||||
job.status = "failed"
|
||||
job.error_message = f"Ditto 异常: {exc}"
|
||||
job.error_code = "DittoAsyncError"
|
||||
job.updated_at = datetime.now(UTC)
|
||||
db.commit()
|
||||
except Exception:
|
||||
pass
|
||||
finally:
|
||||
db.close()
|
||||
@@ -104,7 +104,6 @@ def lipsync_gpu_process_async(self, job_id: str, user_id: str, gpu_task_id: str)
|
||||
job.updated_at = datetime.now(UTC)
|
||||
db.commit()
|
||||
logger.info("[lipsync_gpu_async] GPU 任务已被用户取消: job_id=%s", job_id)
|
||||
_refund_lip_sync(db, job)
|
||||
return
|
||||
|
||||
if final_task.status != "done":
|
||||
@@ -142,7 +141,6 @@ def lipsync_gpu_process_async(self, job_id: str, user_id: str, gpu_task_id: str)
|
||||
job_id,
|
||||
job.output_duration,
|
||||
)
|
||||
_settle_lip_sync(db, job, final_task)
|
||||
except Exception as exc:
|
||||
logger.exception("[lipsync_gpu_async] 异常: job_id=%s err=%s", job_id, exc)
|
||||
try:
|
||||
@@ -159,33 +157,6 @@ def lipsync_gpu_process_async(self, job_id: str, user_id: str, gpu_task_id: str)
|
||||
db.close()
|
||||
|
||||
|
||||
def _settle_lip_sync(db: Session, job: LipsyncJobModel, gpu_task) -> None:
|
||||
"""GPU 成功后结算:同步 credits_cost 到 gpu 任务并按实际时长多退少不补。"""
|
||||
try:
|
||||
from app.services.lipsync_service import LipsyncService
|
||||
|
||||
# GPU 任务表先同步结算结果(标记用)
|
||||
LipsyncService._settle_lip_sync(job, float(getattr(gpu_task, "result_duration", 0) or 0.0))
|
||||
gpu_task.credits_cost = float(job.credits_cost or 0.0)
|
||||
db.commit()
|
||||
except Exception: # noqa: BLE001
|
||||
logger.exception("[lipsync_gpu_async] lip_sync 结算异常 job_id=%s(不阻塞)", job.id)
|
||||
try:
|
||||
db.rollback()
|
||||
except Exception: # noqa: BLE001
|
||||
pass
|
||||
|
||||
|
||||
def _refund_lip_sync(db: Session, job: LipsyncJobModel) -> None:
|
||||
"""GPU 取消/失败路径全额退款。"""
|
||||
try:
|
||||
from app.services.lipsync_service import LipsyncService
|
||||
|
||||
LipsyncService(db)._refund_lip_sync(job)
|
||||
except Exception: # noqa: BLE001
|
||||
logger.exception("[lipsync_gpu_async] lip_sync 退款异常 job_id=%s", job.id)
|
||||
|
||||
|
||||
def _fallback_to_mediakit(db: Session, job: LipsyncJobModel) -> None:
|
||||
"""GPU 失败时回退到 MediaKit 云端渲染。"""
|
||||
try:
|
||||
|
||||
@@ -261,33 +261,7 @@ def tts_synthesize_and_submit(
|
||||
"[lipsync_tts] 句子时间戳计算失败(不影响主流程): job_id=%s err=%s", job_id, _st_err, exc_info=True
|
||||
)
|
||||
|
||||
# 3. 优先走 Ditto(#2076):开关打开且配置完整时,派发 Ditto 异步任务,不再走 MediaKit
|
||||
ditto_dispatched = False
|
||||
try:
|
||||
from packages.config import get_api_settings as _get_settings
|
||||
|
||||
_settings = _get_settings()
|
||||
if _settings.use_ditto_lipsync and _settings.ditto_api_base_url and _settings.ditto_default_video_url:
|
||||
from app.tasks.lipsync_ditto import lipsync_ditto_process_async
|
||||
|
||||
job.status = "processing"
|
||||
job.mediakit_task_id = "ditto:tts-submitted"
|
||||
job.updated_at = datetime.now(UTC)
|
||||
db.commit()
|
||||
lipsync_ditto_process_async.apply_async(args=(job_id, user_id))
|
||||
logger.info("[lipsync_tts] TTS 完成,已派发 Ditto 任务: job_id=%s", job_id)
|
||||
ditto_dispatched = True
|
||||
except Exception as _ditto_err:
|
||||
logger.warning("[lipsync_tts] Ditto 派发失败,回退 MediaKit: job_id=%s err=%s", job_id, _ditto_err)
|
||||
try:
|
||||
db.rollback()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
if ditto_dispatched:
|
||||
return
|
||||
|
||||
# 4. 签名 URL 并提交到 MediaKit(复用模块内 _sign_media_url,避免对 LipsyncService 的耦合)
|
||||
# 3. 签名 URL 并提交到 MediaKit(复用模块内 _sign_media_url,避免对 LipsyncService 的耦合)
|
||||
audio_url = _sign_media_url(job.audio_url)
|
||||
video_url = _sign_media_url(job.video_url)
|
||||
|
||||
|
||||
Generated
-12
@@ -14,7 +14,6 @@
|
||||
"axios": "^1.7.2",
|
||||
"classnames": "^2.5.1",
|
||||
"dayjs": "^1.11.23",
|
||||
"marked": "^12.0.2",
|
||||
"mp4box": "^2.4.1",
|
||||
"react": "^18.3.1",
|
||||
"react-dom": "^18.3.1",
|
||||
@@ -4503,17 +4502,6 @@
|
||||
"url": "https://github.com/sponsors/sindresorhus"
|
||||
}
|
||||
},
|
||||
"node_modules/marked": {
|
||||
"version": "12.0.2",
|
||||
"resolved": "https://registry.npmmirror.com/marked/-/marked-12.0.2.tgz",
|
||||
"integrity": "sha512-qXUm7e/YKFoqFPYPa3Ukg9xlI5cyAtGmyEIzMfW//m6kXwCy2Ps9DYf5ioijFKQ8qyuscrHoY04iJGctu2Kg0Q==",
|
||||
"bin": {
|
||||
"marked": "bin/marked.js"
|
||||
},
|
||||
"engines": {
|
||||
"node": ">= 18"
|
||||
}
|
||||
},
|
||||
"node_modules/math-intrinsics": {
|
||||
"version": "1.1.0",
|
||||
"resolved": "https://registry.npmjs.org/math-intrinsics/-/math-intrinsics-1.1.0.tgz",
|
||||
|
||||
@@ -25,7 +25,6 @@
|
||||
"axios": "^1.7.2",
|
||||
"classnames": "^2.5.1",
|
||||
"dayjs": "^1.11.23",
|
||||
"marked": "^12.0.2",
|
||||
"mp4box": "^2.4.1",
|
||||
"react": "^18.3.1",
|
||||
"react-dom": "^18.3.1",
|
||||
|
||||
@@ -1,76 +0,0 @@
|
||||
/**
|
||||
* 后台管理 API client(#2246)
|
||||
*
|
||||
* 独立 axios 实例:不经过主 apiClient 的 Bearer token / 401 刷新逻辑,
|
||||
* 后台鉴权使用 X-API-Key(存 localStorage,不硬编码)。
|
||||
*/
|
||||
import axios from "axios"
|
||||
|
||||
export const ADMIN_API_KEY_STORAGE = "ditto_admin_api_key"
|
||||
|
||||
export type ConfigType = "bool" | "int" | "float" | "string" | "json"
|
||||
|
||||
export interface ConfigItem {
|
||||
key: string
|
||||
value: unknown
|
||||
default: unknown
|
||||
type: ConfigType
|
||||
description: string
|
||||
model_options?: string[]
|
||||
}
|
||||
|
||||
export interface TestResult {
|
||||
ok: boolean
|
||||
enabled?: boolean
|
||||
segments?: Array<{ text: string; emo: number; intensity: number }>
|
||||
error?: string
|
||||
}
|
||||
|
||||
export function getAdminApiKey(): string {
|
||||
return localStorage.getItem(ADMIN_API_KEY_STORAGE) || ""
|
||||
}
|
||||
|
||||
export function setAdminApiKey(key: string): void {
|
||||
localStorage.setItem(ADMIN_API_KEY_STORAGE, key)
|
||||
}
|
||||
|
||||
export function clearAdminApiKey(): void {
|
||||
localStorage.removeItem(ADMIN_API_KEY_STORAGE)
|
||||
}
|
||||
|
||||
function createAdminClient() {
|
||||
const client = axios.create({
|
||||
baseURL: "/api/v1",
|
||||
timeout: 60000,
|
||||
headers: { "Content-Type": "application/json" },
|
||||
})
|
||||
client.interceptors.request.use((config) => {
|
||||
const key = getAdminApiKey()
|
||||
if (key && config.headers) {
|
||||
config.headers["X-API-Key"] = key
|
||||
}
|
||||
return config
|
||||
})
|
||||
return client
|
||||
}
|
||||
|
||||
const adminClient = createAdminClient()
|
||||
|
||||
export async function fetchConfig(): Promise<ConfigItem[]> {
|
||||
const { data } = await adminClient.get("/admin/ditto-emotion/config")
|
||||
return data.configs as ConfigItem[]
|
||||
}
|
||||
|
||||
export async function updateConfig(
|
||||
configs: Record<string, unknown>,
|
||||
): Promise<{ ok: boolean; updated?: Record<string, unknown>; error?: string }> {
|
||||
const { data } = await adminClient.put("/admin/ditto-emotion/config", { configs })
|
||||
return data
|
||||
}
|
||||
|
||||
export async function testConfig(testText: string): Promise<TestResult> {
|
||||
const { data } = await adminClient.post("/admin/ditto-emotion/config/test", {
|
||||
test_text: testText,
|
||||
})
|
||||
return data as TestResult
|
||||
}
|
||||
@@ -9,8 +9,6 @@ import type {
|
||||
AnalyzeImagesRequest,
|
||||
GenerateCopyRequest,
|
||||
ConfirmCopyRequest,
|
||||
ViralVideoModel,
|
||||
ViralVideoModelsResponse,
|
||||
} from "./types"
|
||||
|
||||
/** 创建爆款视频任务 */
|
||||
@@ -53,29 +51,6 @@ export function analyzeViralStyle(id: string) {
|
||||
return apiClient.post<ViralVideoJob>(`/viral-video/${id}/analyze-style`).then((r) => r.data)
|
||||
}
|
||||
|
||||
/** 动态预估积分消耗(STEP3 参数变化时调用) */
|
||||
export function estimateViralVideoCredits(params: {
|
||||
video_model: string
|
||||
resolution: string
|
||||
video_ratio: string
|
||||
duration: number
|
||||
}) {
|
||||
return apiClient
|
||||
.post<{ estimated_credits: number }>("/viral-video/estimate-credits", params)
|
||||
.then((r) => r.data)
|
||||
}
|
||||
|
||||
/** 获取支持的视频模型列表(GET /viral-video/models)。后端返回 {models: [...]} 包装 */
|
||||
export function getViralVideoModels() {
|
||||
return apiClient.get<ViralVideoModelsResponse>("/viral-video/models").then((r) => {
|
||||
const data = r.data as ViralVideoModelsResponse | ViralVideoModel[] | null | undefined
|
||||
if (Array.isArray(data)) return data
|
||||
if (data && Array.isArray((data as ViralVideoModelsResponse).models)) {
|
||||
return (data as ViralVideoModelsResponse).models
|
||||
}
|
||||
return []
|
||||
})
|
||||
}
|
||||
/** ── 三步拆分:前端 mock 辅助函数(后端新接口上线后可替换) ── */
|
||||
|
||||
/**
|
||||
|
||||
@@ -65,23 +65,27 @@ export function isAnalysisStage(stage: ViralVideoStage | undefined): boolean {
|
||||
return isImageAnalysisStage(stage) || isCopyStage(stage)
|
||||
}
|
||||
|
||||
/** 单张图片 VLM 识别结果(v8 叙述优先,仅保留最少结构化字段) */
|
||||
/** 单张图片 VLM 识别出的商品信息 */
|
||||
export interface ImageProductAnalysis {
|
||||
/** store / product / person / scene */
|
||||
type?: string
|
||||
name?: string
|
||||
brand?: string
|
||||
has_person?: boolean
|
||||
/** v8: 用户端展示用的叙述 markdown(由提示词控制排版) */
|
||||
summary_markdown?: string
|
||||
/** 标题行兼容字段 */
|
||||
category?: string
|
||||
brand?: string
|
||||
colors?: string[]
|
||||
material_or_texture?: string
|
||||
key_features?: string[]
|
||||
visual_style?: string
|
||||
scene?: string
|
||||
target_audience_hint?: string
|
||||
text_on_image?: string
|
||||
/** 旧字段兼容 */
|
||||
spec?: string
|
||||
features?: string[] | string
|
||||
label_text?: string
|
||||
selling_points?: string
|
||||
image_index?: number
|
||||
}
|
||||
|
||||
export interface ImageAnalysisResult {
|
||||
/** v8 字段 */
|
||||
images?: ImageProductAnalysis[]
|
||||
/** 老数据兼容 */
|
||||
products?: ImageProductAnalysis[]
|
||||
}
|
||||
|
||||
@@ -128,8 +132,6 @@ export interface CopyResult {
|
||||
/** 向后兼容:= voiceover_script */
|
||||
suggested_copy?: string
|
||||
title?: string
|
||||
/** v3 storyboard: 用户端展示用的 markdown 文案(由提示词控制排版) */
|
||||
copy_display_markdown?: string
|
||||
/** v1.5 旧字段兼容(老数据降级时可能出现) */
|
||||
scenes?: Array<{ shot: string; narration: string; duration?: number }>
|
||||
}
|
||||
@@ -174,7 +176,7 @@ export interface ViralVideoJob {
|
||||
fusion_level?: FusionLevel
|
||||
voice_id?: string
|
||||
voice_mode?: "global" | "per_video"
|
||||
voice_source?: "preset" | "library" | "clone" | "upload" | "my_voice"
|
||||
voice_source?: "preset" | "library" | "clone" | "upload"
|
||||
bgm_preference?: string
|
||||
intent_result?: IntentResult
|
||||
intent_text?: string
|
||||
@@ -210,7 +212,7 @@ export interface GenerateViralVideoRequest {
|
||||
user_copy_text?: string
|
||||
fusion_level?: FusionLevel
|
||||
voice_id?: string
|
||||
voice_source?: "preset" | "library" | "clone" | "upload" | "my_voice"
|
||||
voice_source?: "preset" | "library" | "clone" | "upload"
|
||||
bgm_preference?: string
|
||||
industry?: string
|
||||
target_customer?: string
|
||||
@@ -242,7 +244,7 @@ export interface AnalyzeImagesRequest {
|
||||
/** TTS 音色 ID(STEP1 已选音色时传) */
|
||||
voice_id?: string
|
||||
/** 音色来源:preset | library | clone | upload */
|
||||
voice_source?: "preset" | "library" | "clone" | "upload" | "my_voice"
|
||||
voice_source?: "preset" | "library" | "clone" | "upload"
|
||||
/** Seedance 视频比例:9:16 | 16:9 | 1:1 */
|
||||
video_ratio?: string
|
||||
/** Seedance 模型 ID(空则使用服务端默认) */
|
||||
@@ -271,36 +273,17 @@ export interface GenerateCopyRequest {
|
||||
/** TTS 音色 ID(优先级高于 persona_id) */
|
||||
voice_id?: string
|
||||
/** 音色来源:preset | library | clone | upload */
|
||||
voice_source?: "preset" | "library" | "clone" | "upload" | "my_voice"
|
||||
voice_source?: "preset" | "library" | "clone" | "upload"
|
||||
/** Seedance 视频比例(9:16/16:9/1:1 等) */
|
||||
video_ratio?: string
|
||||
/** Seedance 模型 ID(空则使用服务端默认) */
|
||||
video_model?: string
|
||||
}
|
||||
|
||||
/** 视频模型描述(GET /viral-video/models) */
|
||||
export interface ViralVideoModel {
|
||||
key: string
|
||||
display_name: string
|
||||
supports_audio: boolean
|
||||
supported_resolutions: string[]
|
||||
max_duration: number
|
||||
/** 计费模式(可选):per_second / per_video / token 等 */
|
||||
billing_mode?: string
|
||||
is_default?: boolean
|
||||
}
|
||||
|
||||
/** GET /viral-video/models 响应包装 */
|
||||
export interface ViralVideoModelsResponse {
|
||||
models: ViralVideoModel[]
|
||||
}
|
||||
|
||||
/** v1.6 阶段3请求:用户确认/编辑口播文案后开始单次 Seedance 出片(POST /viral-video/{id}/confirm-copy) */
|
||||
export interface ConfirmCopyRequest {
|
||||
/** 用户编辑后的口播文案;为空则使用 AI 生成的 voiceover_script */
|
||||
edited_copy?: string
|
||||
/** 视频模型 key,覆盖默认 */
|
||||
video_model?: string
|
||||
}
|
||||
|
||||
/** 旧分镜片段结构(保留兼容;新代码请使用 ShotScript) */
|
||||
|
||||
@@ -18,8 +18,6 @@ export interface VoiceClone {
|
||||
language: string
|
||||
gender: string
|
||||
error_message: string | null
|
||||
/** CosyVoice 实际使用的音色 ID(status=ready 时由后端填充,用于 TTS 调用) */
|
||||
voice_id?: string | null
|
||||
created_at: string
|
||||
updated_at: string
|
||||
}
|
||||
|
||||
@@ -18,7 +18,6 @@ export const toVoiceClone = (profile: VoiceCloneProfile): VoiceClone => ({
|
||||
language: profile.language || "",
|
||||
gender: profile.gender || "",
|
||||
error_message: profile.error_message || null,
|
||||
voice_id: profile.voice_id,
|
||||
created_at: profile.created_at,
|
||||
updated_at: profile.updated_at,
|
||||
})
|
||||
|
||||
@@ -1,182 +0,0 @@
|
||||
/* DurationWheelPicker —— 弹层式滚轮选择器(样式与表单一致) */
|
||||
|
||||
/* 触发按钮:外观复用 .vv-select 风格 */
|
||||
.dw-trigger {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
justify-content: space-between;
|
||||
width: 100%;
|
||||
height: 36px;
|
||||
padding: 0 12px;
|
||||
background: #fff;
|
||||
border: 1px solid #e0e0e8;
|
||||
border-radius: 8px;
|
||||
font-size: 13px;
|
||||
color: #1f2937;
|
||||
cursor: pointer;
|
||||
box-sizing: border-box;
|
||||
transition: all 0.15s;
|
||||
user-select: none;
|
||||
}
|
||||
.dw-trigger:hover {
|
||||
border-color: #c0c0d0;
|
||||
}
|
||||
.dw-trigger-open,
|
||||
.dw-trigger:focus-within {
|
||||
border-color: #7c3aed !important;
|
||||
box-shadow: 0 0 0 2px rgba(124, 58, 237, 0.12);
|
||||
}
|
||||
.dw-trigger-disabled {
|
||||
opacity: 0.5;
|
||||
pointer-events: none;
|
||||
cursor: not-allowed;
|
||||
}
|
||||
.dw-trigger-val {
|
||||
flex: 1;
|
||||
overflow: hidden;
|
||||
text-overflow: ellipsis;
|
||||
white-space: nowrap;
|
||||
}
|
||||
.dw-trigger-placeholder {
|
||||
color: #9ca3af;
|
||||
}
|
||||
.dw-trigger-arrow {
|
||||
font-size: 10px;
|
||||
color: #9ca3af;
|
||||
margin-left: 8px;
|
||||
transition: transform 0.2s;
|
||||
}
|
||||
.dw-trigger-arrow-up {
|
||||
transform: rotate(180deg);
|
||||
}
|
||||
|
||||
/* 弹层容器 */
|
||||
.dw-popup {
|
||||
padding: 8px;
|
||||
min-width: 140px;
|
||||
}
|
||||
|
||||
/* 滚轮 */
|
||||
.dw-picker {
|
||||
position: relative;
|
||||
width: 100%;
|
||||
overflow: hidden;
|
||||
border-radius: 8px;
|
||||
background: #fafafe;
|
||||
border: 1px solid #e5e7eb;
|
||||
}
|
||||
.dw-picker-list {
|
||||
margin: 0;
|
||||
padding: 0;
|
||||
list-style: none;
|
||||
height: 100%;
|
||||
overflow-y: scroll;
|
||||
scroll-snap-type: y mandatory;
|
||||
-webkit-overflow-scrolling: touch;
|
||||
scrollbar-width: none;
|
||||
}
|
||||
.dw-picker-list::-webkit-scrollbar {
|
||||
display: none;
|
||||
}
|
||||
.dw-picker-item {
|
||||
display: flex;
|
||||
align-items: baseline;
|
||||
justify-content: center;
|
||||
gap: 3px;
|
||||
scroll-snap-align: center;
|
||||
cursor: pointer;
|
||||
font-size: 15px;
|
||||
color: #9ca3af;
|
||||
font-weight: 400;
|
||||
transition:
|
||||
color 0.15s,
|
||||
transform 0.15s,
|
||||
font-weight 0.15s;
|
||||
}
|
||||
.dw-picker-item-val {
|
||||
font-variant-numeric: tabular-nums;
|
||||
}
|
||||
.dw-picker-item-unit {
|
||||
font-size: 13px;
|
||||
color: inherit;
|
||||
}
|
||||
.dw-picker-item-active {
|
||||
color: #7c3aed;
|
||||
font-weight: 600;
|
||||
}
|
||||
.dw-picker-item-active .dw-picker-item-val {
|
||||
font-size: 18px;
|
||||
}
|
||||
.dw-picker-item-active .dw-picker-item-unit {
|
||||
font-size: 14px;
|
||||
}
|
||||
|
||||
/* 中心选中条 */
|
||||
.dw-picker-mask {
|
||||
position: absolute;
|
||||
left: 6px;
|
||||
right: 6px;
|
||||
pointer-events: none;
|
||||
background: #f5f0ff;
|
||||
border-radius: 6px;
|
||||
z-index: 1;
|
||||
}
|
||||
.dw-picker-mask::before,
|
||||
.dw-picker-mask::after {
|
||||
content: "";
|
||||
position: absolute;
|
||||
left: 0;
|
||||
right: 0;
|
||||
height: 1px;
|
||||
background: #d8c4ff;
|
||||
}
|
||||
.dw-picker-mask::before {
|
||||
top: 0;
|
||||
}
|
||||
.dw-picker-mask::after {
|
||||
bottom: 0;
|
||||
}
|
||||
|
||||
/* 上下渐变 */
|
||||
.dw-picker-fade {
|
||||
position: absolute;
|
||||
left: 0;
|
||||
right: 0;
|
||||
height: 40%;
|
||||
pointer-events: none;
|
||||
z-index: 2;
|
||||
}
|
||||
.dw-picker-fade-top {
|
||||
top: 0;
|
||||
background: linear-gradient(to bottom, #fafafe 25%, rgba(250, 250, 254, 0));
|
||||
}
|
||||
.dw-picker-fade-bottom {
|
||||
bottom: 0;
|
||||
background: linear-gradient(to top, #fafafe 25%, rgba(250, 250, 254, 0));
|
||||
}
|
||||
|
||||
/* 弹层按钮区 */
|
||||
.dw-popup-actions {
|
||||
display: flex;
|
||||
gap: 8px;
|
||||
justify-content: flex-end;
|
||||
margin-top: 8px;
|
||||
}
|
||||
.dw-popup-actions .ant-btn {
|
||||
border-radius: 6px;
|
||||
}
|
||||
.dw-popup-actions .ant-btn-primary {
|
||||
background: #7c3aed;
|
||||
}
|
||||
.dw-popup-actions .ant-btn-primary:hover {
|
||||
background: #6d28d9 !important;
|
||||
}
|
||||
|
||||
/* 覆盖 antd Popover 默认内边距 */
|
||||
.dw-popover .ant-popover-inner {
|
||||
padding: 0 !important;
|
||||
overflow: hidden;
|
||||
}
|
||||
.dw-popover .ant-popover-arrow {
|
||||
display: none;
|
||||
}
|
||||
@@ -1,180 +0,0 @@
|
||||
/**
|
||||
* DurationWheelPicker —— 竖屏滚轮式时长选择器(弹层版)
|
||||
*
|
||||
* 设计:
|
||||
* - 外观是和其他表单 Select 一致的输入框(白色底+1px灰边+紫色focus ring)
|
||||
* - 点击输入框弹出 Popover,内部是滚轮 picker(原生 scroll-snap,零依赖)
|
||||
* - 滚轮样式:白底容器,选中行 #7c3aed 紫字加粗+浅紫背景条
|
||||
* - 支持触摸/鼠标滚轮/点击;松手吸附;底部"确认/取消"按钮
|
||||
* - 默认范围 15–30 秒,步长 1 秒
|
||||
*/
|
||||
import React, { useEffect, useMemo, useRef, useState, useCallback } from "react"
|
||||
import { Popover, Button } from "antd"
|
||||
import { DownOutlined } from "@ant-design/icons"
|
||||
import "./DurationWheelPicker.css"
|
||||
|
||||
export interface DurationWheelPickerProps {
|
||||
value?: number
|
||||
min?: number
|
||||
max?: number
|
||||
step?: number
|
||||
unit?: string
|
||||
onChange?: (value: number) => void
|
||||
placeholder?: string
|
||||
disabled?: boolean
|
||||
/** 弹层宽度,默认 160px */
|
||||
popupWidth?: number
|
||||
/** 弹层内滚轮高度,默认 180px */
|
||||
wheelHeight?: number
|
||||
}
|
||||
|
||||
const ITEM_HEIGHT = 36
|
||||
|
||||
const DurationWheelPicker: React.FC<DurationWheelPickerProps> = ({
|
||||
value = 20,
|
||||
min = 15,
|
||||
max = 30,
|
||||
step = 1,
|
||||
unit = "秒",
|
||||
onChange,
|
||||
placeholder = "请选择时长",
|
||||
disabled = false,
|
||||
popupWidth = 160,
|
||||
wheelHeight = 180,
|
||||
}) => {
|
||||
const options = useMemo(() => {
|
||||
const arr: number[] = []
|
||||
for (let v = min; v <= max; v += step) arr.push(v)
|
||||
return arr
|
||||
}, [min, max, step])
|
||||
|
||||
const [open, setOpen] = useState(false)
|
||||
// 弹层内暂存值,点确认才提交
|
||||
const [draft, setDraft] = useState<number>(value)
|
||||
const listRef = useRef<HTMLUListElement>(null)
|
||||
const scrollTimerRef = useRef<ReturnType<typeof setTimeout> | null>(null)
|
||||
|
||||
useEffect(() => {
|
||||
if (open) {
|
||||
setDraft(value)
|
||||
// 下一帧滚到当前值
|
||||
requestAnimationFrame(() => scrollToValue(value, false))
|
||||
}
|
||||
// eslint-disable-next-line react-hooks/exhaustive-deps
|
||||
}, [open])
|
||||
|
||||
const scrollToValue = useCallback(
|
||||
(v: number, smooth = true) => {
|
||||
const list = listRef.current
|
||||
if (!list) return
|
||||
const idx = options.indexOf(v)
|
||||
if (idx < 0) return
|
||||
list.scrollTo({ top: idx * ITEM_HEIGHT, behavior: smooth ? "smooth" : "auto" })
|
||||
},
|
||||
[options],
|
||||
)
|
||||
|
||||
const handleScroll = () => {
|
||||
if (scrollTimerRef.current) clearTimeout(scrollTimerRef.current)
|
||||
scrollTimerRef.current = setTimeout(() => {
|
||||
const list = listRef.current
|
||||
if (!list) return
|
||||
const idx = Math.round(list.scrollTop / ITEM_HEIGHT)
|
||||
const clamped = Math.max(0, Math.min(options.length - 1, idx))
|
||||
const targetTop = clamped * ITEM_HEIGHT
|
||||
if (Math.abs(list.scrollTop - targetTop) > 1) {
|
||||
list.scrollTo({ top: targetTop, behavior: "smooth" })
|
||||
}
|
||||
setDraft(options[clamped])
|
||||
}, 100)
|
||||
}
|
||||
|
||||
const handleConfirm = () => {
|
||||
onChange?.(draft)
|
||||
setOpen(false)
|
||||
}
|
||||
|
||||
const handleCancel = () => {
|
||||
setOpen(false)
|
||||
}
|
||||
|
||||
const handleItemClick = (v: number) => {
|
||||
setDraft(v)
|
||||
scrollToValue(v, true)
|
||||
}
|
||||
|
||||
const maskTop = wheelHeight / 2 - ITEM_HEIGHT / 2
|
||||
|
||||
const wheel = (
|
||||
<div className="dw-popup">
|
||||
<div
|
||||
className="dw-picker"
|
||||
style={{ height: wheelHeight, width: popupWidth - 24 /* padding */ }}
|
||||
>
|
||||
<div className="dw-picker-mask" style={{ top: maskTop, height: ITEM_HEIGHT }} aria-hidden />
|
||||
<div className="dw-picker-fade dw-picker-fade-top" aria-hidden />
|
||||
<div className="dw-picker-fade dw-picker-fade-bottom" aria-hidden />
|
||||
<ul
|
||||
ref={listRef}
|
||||
className="dw-picker-list"
|
||||
onScroll={handleScroll}
|
||||
style={{
|
||||
paddingTop: wheelHeight / 2 - ITEM_HEIGHT / 2,
|
||||
paddingBottom: wheelHeight / 2 - ITEM_HEIGHT / 2,
|
||||
}}
|
||||
>
|
||||
{options.map((v) => {
|
||||
const isActive = v === draft
|
||||
return (
|
||||
<li
|
||||
key={v}
|
||||
className={`dw-picker-item${isActive ? " dw-picker-item-active" : ""}`}
|
||||
style={{ height: ITEM_HEIGHT, lineHeight: `${ITEM_HEIGHT}px` }}
|
||||
onClick={() => handleItemClick(v)}
|
||||
aria-selected={isActive}
|
||||
role="option"
|
||||
>
|
||||
<span className="dw-picker-item-val">{v}</span>
|
||||
<span className="dw-picker-item-unit">{unit}</span>
|
||||
</li>
|
||||
)
|
||||
})}
|
||||
</ul>
|
||||
</div>
|
||||
<div className="dw-popup-actions">
|
||||
<Button size="small" onClick={handleCancel}>
|
||||
取消
|
||||
</Button>
|
||||
<Button size="small" type="primary" onClick={handleConfirm}>
|
||||
确认
|
||||
</Button>
|
||||
</div>
|
||||
</div>
|
||||
)
|
||||
|
||||
return (
|
||||
<Popover
|
||||
open={!disabled && open}
|
||||
onOpenChange={(v) => setOpen(v)}
|
||||
content={wheel}
|
||||
trigger="click"
|
||||
placement="bottomLeft"
|
||||
overlayClassName="dw-popover"
|
||||
overlayStyle={{ padding: 0 }}
|
||||
overlayInnerStyle={{ padding: 0, borderRadius: 10 }}
|
||||
destroyTooltipOnHide
|
||||
>
|
||||
<div
|
||||
className={`dw-trigger${disabled ? " dw-trigger-disabled" : ""}${open ? " dw-trigger-open" : ""}`}
|
||||
style={{ height: 36 }}
|
||||
>
|
||||
<span className={`dw-trigger-val${value != null ? "" : " dw-trigger-placeholder"}`}>
|
||||
{value != null ? `${value}${unit}` : placeholder}
|
||||
</span>
|
||||
<DownOutlined className={`dw-trigger-arrow${open ? " dw-trigger-arrow-up" : ""}`} />
|
||||
</div>
|
||||
</Popover>
|
||||
)
|
||||
}
|
||||
|
||||
export default DurationWheelPicker
|
||||
@@ -10,4 +10,4 @@
|
||||
* 功能流程不做积分预校验,直接走生成。
|
||||
* - true:展示完整积分系统 UI。
|
||||
*/
|
||||
export const ENABLE_CREDIT_SYSTEM = true
|
||||
export const ENABLE_CREDIT_SYSTEM = false
|
||||
|
||||
@@ -1,311 +0,0 @@
|
||||
import React, { useEffect, useMemo, useState } from "react"
|
||||
import {
|
||||
Alert,
|
||||
Button,
|
||||
Card,
|
||||
Form,
|
||||
Input,
|
||||
InputNumber,
|
||||
Modal,
|
||||
Select,
|
||||
Slider,
|
||||
Space,
|
||||
Spin,
|
||||
Switch,
|
||||
message,
|
||||
} from "antd"
|
||||
import {
|
||||
clearAdminApiKey,
|
||||
fetchConfig,
|
||||
getAdminApiKey,
|
||||
setAdminApiKey,
|
||||
testConfig,
|
||||
updateConfig,
|
||||
type ConfigItem,
|
||||
} from "@/api/admin/dittoEmotion"
|
||||
import "./Admin.css"
|
||||
|
||||
const EMO_LABELS: Record<number, string> = {
|
||||
3: "开心",
|
||||
4: "中性",
|
||||
5: "伤心",
|
||||
6: "惊讶",
|
||||
}
|
||||
|
||||
const DittoEmotionConfig: React.FC = () => {
|
||||
const [hasKey, setHasKey] = useState<boolean>(!!getAdminApiKey())
|
||||
const [keyInput, setKeyInput] = useState<string>("")
|
||||
const [loading, setLoading] = useState<boolean>(false)
|
||||
const [saving, setSaving] = useState<boolean>(false)
|
||||
const [testing, setTesting] = useState<boolean>(false)
|
||||
const [items, setItems] = useState<ConfigItem[]>([])
|
||||
const [modelOptions, setModelOptions] = useState<string[]>([])
|
||||
const [testResult, setTestResult] = useState<string>("")
|
||||
const [form] = Form.useForm()
|
||||
|
||||
const [testInput, setTestInput] = useState<string>("")
|
||||
|
||||
const load = React.useCallback(async () => {
|
||||
setLoading(true)
|
||||
try {
|
||||
const configs = await fetchConfig()
|
||||
setItems(configs)
|
||||
const values: Record<string, unknown> = {}
|
||||
configs.forEach((c) => {
|
||||
values[c.key] = c.value
|
||||
if (c.model_options) setModelOptions(c.model_options)
|
||||
})
|
||||
form.setFieldsValue(values)
|
||||
} catch {
|
||||
// 401/403 等 → 提示 key 可能无效
|
||||
message.error("加载配置失败,请检查 X-API-Key 是否正确")
|
||||
} finally {
|
||||
setLoading(false)
|
||||
}
|
||||
}, [form])
|
||||
|
||||
useEffect(() => {
|
||||
if (hasKey) {
|
||||
void load()
|
||||
}
|
||||
}, [hasKey, load])
|
||||
|
||||
const defaults = useMemo(() => {
|
||||
const m: Record<string, unknown> = {}
|
||||
items.forEach((c) => {
|
||||
m[c.key] = c.default
|
||||
})
|
||||
return m
|
||||
}, [items])
|
||||
|
||||
const submitKey = () => {
|
||||
if (!keyInput.trim()) {
|
||||
message.warning("请输入 X-API-Key")
|
||||
return
|
||||
}
|
||||
setAdminApiKey(keyInput.trim())
|
||||
setHasKey(true)
|
||||
}
|
||||
|
||||
const changeKey = () => {
|
||||
clearAdminApiKey()
|
||||
setHasKey(false)
|
||||
setKeyInput("")
|
||||
}
|
||||
|
||||
const resetDefaults = () => {
|
||||
form.setFieldsValue(defaults)
|
||||
message.info("已填入默认值,点击「保存配置」后生效")
|
||||
}
|
||||
|
||||
const validateBeforeSave = async (): Promise<Record<string, unknown> | null> => {
|
||||
try {
|
||||
const values = await form.validateFields()
|
||||
const prompt = (values.ditto_emotion_prompt || "") as string
|
||||
if (prompt.trim() && !prompt.includes("{文案}")) {
|
||||
message.error("提示词必须包含 {文案} 占位符")
|
||||
return null
|
||||
}
|
||||
return values as Record<string, unknown>
|
||||
} catch {
|
||||
return null
|
||||
}
|
||||
}
|
||||
|
||||
const onSave = async () => {
|
||||
const values = await validateBeforeSave()
|
||||
if (!values) return
|
||||
setSaving(true)
|
||||
try {
|
||||
const res = await updateConfig(values)
|
||||
if (res.ok) {
|
||||
message.success("配置已保存并立即生效")
|
||||
await load()
|
||||
} else {
|
||||
message.error(res.error || "保存失败")
|
||||
}
|
||||
} catch {
|
||||
message.error("保存失败,请检查网络或 X-API-Key")
|
||||
} finally {
|
||||
setSaving(false)
|
||||
}
|
||||
}
|
||||
|
||||
const onTest = async () => {
|
||||
const testText = (testInput || "").trim()
|
||||
if (!testText) {
|
||||
message.warning("请先在下方输入测试文案")
|
||||
return
|
||||
}
|
||||
setTesting(true)
|
||||
setTestResult("")
|
||||
try {
|
||||
const res = await testConfig(testText)
|
||||
if (!res.ok) {
|
||||
message.error(res.error || "测试失败")
|
||||
} else if (!res.enabled) {
|
||||
message.info("当前表情开关为关闭状态,无情绪结果,可先开启后再测")
|
||||
} else {
|
||||
const lines = (res.segments || []).map(
|
||||
(s) => `【${EMO_LABELS[s.emo] ?? s.emo} ${s.intensity}】${s.text}`,
|
||||
)
|
||||
setTestResult(lines.join("\n") || "未解析到情绪结果")
|
||||
}
|
||||
} catch {
|
||||
message.error("测试失败,请检查网络或 X-API-Key")
|
||||
} finally {
|
||||
setTesting(false)
|
||||
}
|
||||
}
|
||||
|
||||
if (!hasKey) {
|
||||
return (
|
||||
<div className="admin-coming-soon-page">
|
||||
<Modal
|
||||
title="请输入后台 X-API-Key"
|
||||
open
|
||||
closable={false}
|
||||
footer={[
|
||||
<Button type="primary" key="ok" onClick={submitKey}>
|
||||
确认
|
||||
</Button>,
|
||||
]}
|
||||
>
|
||||
<Alert
|
||||
type="info"
|
||||
showIcon
|
||||
style={{ marginBottom: 12 }}
|
||||
message="Key 仅保存在本机浏览器 localStorage,用于后台接口鉴权(X-API-Key)。"
|
||||
/>
|
||||
<Input.Password
|
||||
autoFocus
|
||||
placeholder="X-API-Key"
|
||||
value={keyInput}
|
||||
onChange={(e) => setKeyInput(e.target.value)}
|
||||
onPressEnter={submitKey}
|
||||
/>
|
||||
</Modal>
|
||||
</div>
|
||||
)
|
||||
}
|
||||
|
||||
return (
|
||||
<div className="admin-coming-soon-page">
|
||||
<div style={{ maxWidth: 860, width: "100%" }}>
|
||||
<Card
|
||||
title="Ditto 数字人表情设置"
|
||||
extra={
|
||||
<Button size="small" onClick={changeKey}>
|
||||
更换 X-API-Key
|
||||
</Button>
|
||||
}
|
||||
className="xx-card"
|
||||
>
|
||||
<Alert
|
||||
type="success"
|
||||
showIcon
|
||||
style={{ marginBottom: 16 }}
|
||||
message="修改保存后立即生效,无需重启或发版。"
|
||||
action={
|
||||
<Button size="small" onClick={resetDefaults}>
|
||||
重置默认
|
||||
</Button>
|
||||
}
|
||||
/>
|
||||
|
||||
<Spin spinning={loading}>
|
||||
<Form form={form} layout="vertical">
|
||||
<Form.Item
|
||||
name="ditto_emotion_enabled"
|
||||
label="启用 LLM 情绪分析"
|
||||
valuePropName="checked"
|
||||
extra="关闭后立即回退到原有关键词匹配模式,不影响正常出片。"
|
||||
>
|
||||
<Switch />
|
||||
</Form.Item>
|
||||
|
||||
<Form.Item
|
||||
name="ditto_emotion_model"
|
||||
label="情绪分析模型"
|
||||
rules={[{ required: true, message: "请选择模型" }]}
|
||||
>
|
||||
<Select
|
||||
options={(modelOptions.length ? modelOptions : []).map((m) => ({
|
||||
label: m,
|
||||
value: m,
|
||||
}))}
|
||||
/>
|
||||
</Form.Item>
|
||||
|
||||
<Form.Item name="ditto_emotion_temperature" label="温度(0-1,越低越稳定)">
|
||||
<Space style={{ width: "100%" }} align="center">
|
||||
<Slider min={0} max={1} step={0.1} style={{ width: 320 }} />
|
||||
<InputNumber min={0} max={1} step={0.1} />
|
||||
</Space>
|
||||
</Form.Item>
|
||||
|
||||
<Form.Item
|
||||
name="ditto_emotion_prompt"
|
||||
label="情绪分析提示词(必须包含 {文案} 占位符)"
|
||||
rules={[
|
||||
{
|
||||
validator: (_, value) =>
|
||||
!value || !String(value).trim() || String(value).includes("{文案}")
|
||||
? Promise.resolve()
|
||||
: Promise.reject(new Error("必须包含 {文案} 占位符")),
|
||||
},
|
||||
]}
|
||||
>
|
||||
<Input.TextArea
|
||||
rows={15}
|
||||
placeholder="留空则使用系统默认提示词"
|
||||
style={{ fontFamily: "monospace" }}
|
||||
/>
|
||||
</Form.Item>
|
||||
|
||||
<Form.Item
|
||||
name="ditto_blend_frames"
|
||||
label="表情过渡帧数(6-30,越大越柔和)"
|
||||
rules={[{ required: true, message: "请输入过渡帧数" }]}
|
||||
>
|
||||
<InputNumber min={6} max={30} step={1} precision={0} />
|
||||
</Form.Item>
|
||||
</Form>
|
||||
|
||||
<Space style={{ marginTop: 8 }}>
|
||||
<Button type="primary" loading={saving} onClick={onSave}>
|
||||
保存配置
|
||||
</Button>
|
||||
</Space>
|
||||
</Spin>
|
||||
</Card>
|
||||
|
||||
<Card title="配置测试" className="xx-card" style={{ marginTop: 16 }}>
|
||||
<Input.TextArea
|
||||
rows={3}
|
||||
placeholder="输入测试文案,例如:这款面膜超级好用!今天补水效果太棒了。"
|
||||
value={testInput}
|
||||
onChange={(e) => setTestInput(e.target.value)}
|
||||
/>
|
||||
<Space style={{ marginTop: 12 }}>
|
||||
<Button loading={testing} onClick={onTest}>
|
||||
用当前配置测试
|
||||
</Button>
|
||||
</Space>
|
||||
{testResult && (
|
||||
<Input.TextArea
|
||||
readOnly
|
||||
rows={6}
|
||||
value={testResult}
|
||||
style={{ marginTop: 12, fontFamily: "monospace", whiteSpace: "pre-wrap" }}
|
||||
/>
|
||||
)}
|
||||
</Card>
|
||||
</div>
|
||||
</div>
|
||||
)
|
||||
}
|
||||
|
||||
export default DittoEmotionConfig
|
||||
|
||||
export const Component = DittoEmotionConfig
|
||||
@@ -80,8 +80,6 @@ const AiAvatarPage: React.FC = () => {
|
||||
const [finalizeLoading, setFinalizeLoading] = useState(false)
|
||||
|
||||
/* ── 对口型轮询 ── */
|
||||
/** 对口型轮询总时长上限(10分钟):超过后停止轮询并提示去历史记录查看 */
|
||||
const LIPSYNC_POLL_MAX_MS = 10 * 60 * 1000
|
||||
const lipsyncTimerRef = useRef<ReturnType<typeof setInterval> | null>(null)
|
||||
/* ── 渲染进度轮询 ── */
|
||||
const renderTimerRef = useRef<ReturnType<typeof setInterval> | null>(null)
|
||||
@@ -272,24 +270,7 @@ const AiAvatarPage: React.FC = () => {
|
||||
// 如果是预合成模式,后端会同步把状态置为 submitted(甚至可能已返回 running),
|
||||
// 但仍需轮询等 completed
|
||||
if (lipsyncTimerRef.current) clearInterval(lipsyncTimerRef.current)
|
||||
// 轮询间隔 5 秒;单请求超时 5 分钟(见 api/aiAvatar.ts);总轮询上限 10 分钟
|
||||
// 单次请求失败/超时不中断轮询,继续下一轮;超过总上限后停止并提示用户去历史记录查看
|
||||
lipsyncTimerRef.current = setInterval(async () => {
|
||||
// 总时长保护:超过 10 分钟停止轮询
|
||||
if (Date.now() - lipsyncStartAtRef.current > LIPSYNC_POLL_MAX_MS) {
|
||||
if (lipsyncTimerRef.current) {
|
||||
clearInterval(lipsyncTimerRef.current)
|
||||
lipsyncTimerRef.current = null
|
||||
}
|
||||
if (lipsyncTickRef.current) {
|
||||
clearInterval(lipsyncTickRef.current)
|
||||
lipsyncTickRef.current = null
|
||||
}
|
||||
setLipsyncStatus("failed")
|
||||
setLipsyncErrorMessage("渲染时间较长,请稍后在历史记录中查看")
|
||||
message.warning("对口型渲染时间较长,已停止自动刷新,请稍后在历史记录中查看")
|
||||
return
|
||||
}
|
||||
try {
|
||||
const updated = await getLipsyncJob(job.id)
|
||||
state.setLipsyncJob(updated)
|
||||
@@ -315,10 +296,9 @@ const AiAvatarPage: React.FC = () => {
|
||||
setLipsyncErrorMessage(updated.error_message || "对口型生成失败")
|
||||
}
|
||||
} catch (err) {
|
||||
// 单次轮询失败(含 timeout):不中断轮询,打印日志后等下一轮
|
||||
console.warn("[对口型] 轮询请求失败,将继续下一轮:", err)
|
||||
console.error("[对口型] 轮询错误:", err)
|
||||
}
|
||||
}, 5000)
|
||||
}, 3000)
|
||||
} catch (err) {
|
||||
console.error("[对口型] 创建失败:", {
|
||||
status: (err as { response?: { status?: number } })?.response?.status,
|
||||
|
||||
@@ -72,8 +72,7 @@ export const previewTts = async (data: {
|
||||
}
|
||||
|
||||
export const getLipsyncJob = async (id: string): Promise<LipsyncJob> => {
|
||||
// MuseTalk 渲染 8s 视频约 54s + 排队时间,给足 5 分钟超时避免单次轮询 AxiosError 中断
|
||||
const response = await apiClient.get<LipsyncJob>(`/lipsync/jobs/${id}`, { timeout: 300_000 })
|
||||
const response = await apiClient.get<LipsyncJob>(`/lipsync/jobs/${id}`, { timeout: 60000 })
|
||||
return response.data
|
||||
}
|
||||
|
||||
@@ -92,10 +91,7 @@ export const submitRender = async (data: {
|
||||
}
|
||||
|
||||
export const getRenderJob = async (jobId: string): Promise<RenderJob> => {
|
||||
// 渲染链路(对口型+B-roll+标题+合成+上传)耗时较长,给足 5 分钟超时
|
||||
const response = await apiClient.get<RenderJob>(`/ai-avatar/render/${jobId}`, {
|
||||
timeout: 300_000,
|
||||
})
|
||||
const response = await apiClient.get<RenderJob>(`/ai-avatar/render/${jobId}`, { timeout: 60000 })
|
||||
return response.data
|
||||
}
|
||||
|
||||
|
||||
@@ -1002,11 +1002,6 @@
|
||||
.vv-form-row {
|
||||
margin-bottom: 10px;
|
||||
}
|
||||
.vv-form-hint {
|
||||
font-size: 12px;
|
||||
color: #9ca3af;
|
||||
line-height: 1.4;
|
||||
}
|
||||
@media (max-width: 500px) {
|
||||
.vv-form-grid {
|
||||
grid-template-columns: 1fr;
|
||||
@@ -1050,13 +1045,10 @@
|
||||
flex-direction: column;
|
||||
align-items: center;
|
||||
justify-content: center;
|
||||
height: 360px;
|
||||
padding: 28px 16px;
|
||||
gap: 10px;
|
||||
background: #fff;
|
||||
border: 1px solid #e5e7eb;
|
||||
border-radius: 10px;
|
||||
margin-top: 8px;
|
||||
}
|
||||
.vv-copy-loading .vv-spinner {
|
||||
width: 28px;
|
||||
@@ -1085,41 +1077,21 @@
|
||||
|
||||
/* ── Storyboard (linear doc style) ── */
|
||||
.vv-storyboard {
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
height: 360px;
|
||||
padding: 10px 12px;
|
||||
background: #fff;
|
||||
border: 1px solid #e5e7eb;
|
||||
border-radius: 10px;
|
||||
margin-top: 8px;
|
||||
overflow: hidden;
|
||||
padding: 6px 2px;
|
||||
background: transparent;
|
||||
border: none;
|
||||
}
|
||||
|
||||
.vv-sb-doc {
|
||||
flex: 1 1 auto;
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
gap: 3px;
|
||||
color: #1f2937;
|
||||
font-size: 13px;
|
||||
line-height: 1.55;
|
||||
overflow-y: auto;
|
||||
padding-right: 4px;
|
||||
margin-right: -4px;
|
||||
}
|
||||
.vv-sb-doc::-webkit-scrollbar {
|
||||
width: 6px;
|
||||
}
|
||||
.vv-sb-doc::-webkit-scrollbar-thumb {
|
||||
background: #d8c4ff;
|
||||
border-radius: 3px;
|
||||
}
|
||||
.vv-sb-doc::-webkit-scrollbar-track {
|
||||
background: transparent;
|
||||
}
|
||||
.vv-sb-h {
|
||||
margin: 6px 0 2px;
|
||||
margin: 8px 0 3px;
|
||||
padding: 0;
|
||||
font-size: 14px;
|
||||
font-weight: 700;
|
||||
@@ -1130,20 +1102,48 @@
|
||||
.vv-sb-h:first-child {
|
||||
margin-top: 0;
|
||||
}
|
||||
|
||||
/* 总览 —— 每行一段 */
|
||||
.vv-sb-inline-row {
|
||||
.vv-sb-kv {
|
||||
display: flex;
|
||||
align-items: baseline;
|
||||
flex-wrap: wrap;
|
||||
align-items: flex-start;
|
||||
gap: 4px;
|
||||
font-size: 13px;
|
||||
line-height: 1.7;
|
||||
line-height: 1.55;
|
||||
}
|
||||
.vv-sb-k {
|
||||
flex-shrink: 0;
|
||||
color: #6b7280;
|
||||
font-weight: 500;
|
||||
min-width: 120px;
|
||||
}
|
||||
.vv-sb-kv-ref {
|
||||
align-items: center;
|
||||
}
|
||||
.vv-sb-inline {
|
||||
flex: 1;
|
||||
border: none;
|
||||
background: transparent;
|
||||
color: #1f2937;
|
||||
margin: 2px 0;
|
||||
font-size: 13px;
|
||||
font-family: inherit;
|
||||
padding: 1px 4px;
|
||||
outline: none;
|
||||
border-radius: 4px;
|
||||
}
|
||||
.vv-sb-inline-input {
|
||||
border-bottom: 1px dashed transparent;
|
||||
transition: border-color 0.15s;
|
||||
}
|
||||
.vv-sb-inline-input:hover,
|
||||
.vv-sb-inline-input:focus {
|
||||
border-bottom-color: #7c3aed;
|
||||
background: #f5f0ff;
|
||||
}
|
||||
.vv-sb-inline:disabled {
|
||||
color: #9ca3af;
|
||||
cursor: default;
|
||||
}
|
||||
|
||||
.vv-sb-inline-select {
|
||||
min-width: 80px;
|
||||
min-width: 120px;
|
||||
}
|
||||
.vv-sb-inline-select .ant-select-selector {
|
||||
background: transparent !important;
|
||||
@@ -1158,15 +1158,17 @@
|
||||
line-height: 24px !important;
|
||||
padding-left: 0 !important;
|
||||
}
|
||||
|
||||
/* 段落式 textarea 基础样式(仅编辑态使用) */
|
||||
.vv-sb-doc-ta {
|
||||
flex: 1;
|
||||
min-height: 32px;
|
||||
background: transparent !important;
|
||||
border: 1px dashed transparent !important;
|
||||
padding: 2px 4px !important;
|
||||
font-size: 13px !important;
|
||||
line-height: 1.55 !important;
|
||||
color: #1f2937 !important;
|
||||
border-radius: 6px;
|
||||
border-radius: 4px;
|
||||
resize: vertical;
|
||||
outline: none;
|
||||
transition:
|
||||
border-color 0.15s,
|
||||
background 0.15s;
|
||||
@@ -1176,39 +1178,9 @@
|
||||
border-color: #7c3aed !important;
|
||||
background: #f5f0ff !important;
|
||||
}
|
||||
|
||||
/* 场景与光线 —— 段落样式 */
|
||||
.vv-sb-para {
|
||||
margin: 4px 0;
|
||||
.vv-sb-doc-ta-sm {
|
||||
min-height: 24px;
|
||||
}
|
||||
|
||||
/* 内联编辑 textarea(点击后弹出) */
|
||||
.vv-sb-inline-edit-ta {
|
||||
display: block;
|
||||
width: 100%;
|
||||
margin-top: 4px;
|
||||
min-height: 28px;
|
||||
background: #fafafe !important;
|
||||
border: 1px solid #d8cafc !important;
|
||||
border-radius: 6px !important;
|
||||
padding: 6px 8px !important;
|
||||
font-size: 13px !important;
|
||||
line-height: 1.6 !important;
|
||||
font-family: inherit;
|
||||
color: #1f2937 !important;
|
||||
outline: none;
|
||||
resize: vertical;
|
||||
}
|
||||
.vv-sb-inline-edit-ta-sm {
|
||||
max-width: 120px;
|
||||
}
|
||||
.vv-sb-time-ta {
|
||||
max-width: 160px;
|
||||
font-weight: 600;
|
||||
color: #7c3aed !important;
|
||||
}
|
||||
|
||||
/* 逐镜头 */
|
||||
.vv-sb-doc-shots {
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
@@ -1216,60 +1188,25 @@
|
||||
margin-top: 2px;
|
||||
}
|
||||
.vv-sb-doc-shot {
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
gap: 2px;
|
||||
margin-bottom: 6px;
|
||||
padding-left: 8px;
|
||||
border-left: 2px solid rgba(124, 58, 237, 0.3);
|
||||
}
|
||||
.vv-sb-doc-shot-head {
|
||||
margin-bottom: 1px;
|
||||
}
|
||||
.vv-sb-time-doc {
|
||||
display: block;
|
||||
font-size: 14px;
|
||||
font-weight: 600;
|
||||
color: #7c3aed;
|
||||
margin: 6px 0 2px;
|
||||
cursor: text;
|
||||
}
|
||||
.vv-sb-time-doc:hover {
|
||||
background: rgba(124, 58, 237, 0.06);
|
||||
border-radius: 3px;
|
||||
}
|
||||
|
||||
/* 字段段落 */
|
||||
.vv-sb-field {
|
||||
display: flex;
|
||||
flex-wrap: wrap;
|
||||
align-items: baseline;
|
||||
margin: 2px 0;
|
||||
font-size: 13px;
|
||||
line-height: 1.6;
|
||||
font-weight: 700;
|
||||
color: #7c3aed;
|
||||
background: transparent;
|
||||
border: none;
|
||||
outline: none;
|
||||
padding: 0 4px 1px;
|
||||
font-family: inherit;
|
||||
border-radius: 4px;
|
||||
}
|
||||
.vv-sb-field-k {
|
||||
color: #1f2937;
|
||||
font-weight: 600;
|
||||
margin-right: 0;
|
||||
white-space: nowrap;
|
||||
}
|
||||
.vv-sb-field-val {
|
||||
color: #374151;
|
||||
cursor: text;
|
||||
border-radius: 3px;
|
||||
padding: 0 2px;
|
||||
transition: background 0.15s;
|
||||
word-break: break-word;
|
||||
flex: 1;
|
||||
min-width: 0;
|
||||
}
|
||||
.vv-sb-field:hover .vv-sb-field-val {
|
||||
background: rgba(124, 58, 237, 0.06);
|
||||
}
|
||||
|
||||
/* 参考图片行 */
|
||||
.vv-sb-ref-row {
|
||||
margin-top: 4px;
|
||||
display: flex;
|
||||
align-items: center;
|
||||
flex-wrap: wrap;
|
||||
gap: 4px;
|
||||
.vv-sb-time-doc:focus {
|
||||
background: #f5f0ff;
|
||||
}
|
||||
|
||||
.vv-sb-ref-badge {
|
||||
@@ -1337,10 +1274,6 @@
|
||||
font-size: 12px;
|
||||
}
|
||||
|
||||
/* 标签组段落间距 */
|
||||
.vv-sb-tags-block {
|
||||
margin-top: 4px;
|
||||
}
|
||||
.vv-sb-taglist {
|
||||
display: flex;
|
||||
flex-wrap: wrap;
|
||||
@@ -1403,17 +1336,46 @@
|
||||
background: #f5f0ff;
|
||||
}
|
||||
|
||||
/* 口播稿 —— 复用 vv-sb-field 样式,无额外需求 */
|
||||
.vv-sb-doc-collapse {
|
||||
background: transparent;
|
||||
border: none;
|
||||
color: inherit;
|
||||
font: inherit;
|
||||
cursor: pointer;
|
||||
padding: 0;
|
||||
display: inline-flex;
|
||||
align-items: center;
|
||||
gap: 6px;
|
||||
}
|
||||
.vv-sb-caret {
|
||||
display: inline-block;
|
||||
transition: transform 0.2s;
|
||||
font-size: 10px;
|
||||
color: #9ca3af;
|
||||
}
|
||||
.vv-sb-caret.open {
|
||||
transform: rotate(180deg);
|
||||
}
|
||||
.vv-sb-vo {
|
||||
margin-top: 4px !important;
|
||||
background: #fafafe !important;
|
||||
border-left: 3px solid #7c3aed !important;
|
||||
border-radius: 0 6px 6px 0 !important;
|
||||
padding: 6px 10px !important;
|
||||
min-height: 44px !important;
|
||||
font-size: 13px !important;
|
||||
line-height: 1.7 !important;
|
||||
color: #1f2937 !important;
|
||||
border: none !important;
|
||||
}
|
||||
|
||||
.vv-sb-actions {
|
||||
flex-shrink: 0;
|
||||
display: flex;
|
||||
align-items: center;
|
||||
gap: 8px;
|
||||
margin-top: 8px;
|
||||
padding-top: 8px;
|
||||
border-top: 1px solid #e5e7eb;
|
||||
background: #fff;
|
||||
}
|
||||
.vv-sb-actions .vv-btn-ghost {
|
||||
padding: 6px 14px;
|
||||
@@ -1494,7 +1456,7 @@
|
||||
}
|
||||
|
||||
/* ── Asset/voice picker modal styles (in page) ───────────── */
|
||||
.vv-modal-mask {
|
||||
.vv-modal {
|
||||
position: fixed;
|
||||
inset: 0;
|
||||
background: rgba(0, 0, 0, 0.45);
|
||||
@@ -1504,59 +1466,21 @@
|
||||
justify-content: center;
|
||||
padding: 20px;
|
||||
}
|
||||
.vv-modal {
|
||||
.vv-modal-body {
|
||||
background: #fff;
|
||||
border: 1px solid #e5e7eb;
|
||||
border-radius: 12px;
|
||||
max-width: 720px;
|
||||
width: 100%;
|
||||
max-height: 80vh;
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
overflow: hidden;
|
||||
position: relative;
|
||||
box-shadow: 0 12px 40px rgba(15, 23, 42, 0.18);
|
||||
}
|
||||
.vv-modal.vv-modal-lg {
|
||||
max-width: 860px;
|
||||
}
|
||||
.vv-modal-head {
|
||||
flex-shrink: 0;
|
||||
display: flex;
|
||||
align-items: center;
|
||||
justify-content: space-between;
|
||||
padding: 14px 20px;
|
||||
border-bottom: 1px solid #e5e7eb;
|
||||
}
|
||||
.vv-modal-title {
|
||||
font-size: 15px;
|
||||
font-weight: 600;
|
||||
color: #1f2937;
|
||||
}
|
||||
.vv-modal-body {
|
||||
flex: 1 1 auto;
|
||||
overflow-y: auto;
|
||||
padding: 16px 20px;
|
||||
min-height: 0;
|
||||
}
|
||||
.vv-modal-foot {
|
||||
flex-shrink: 0;
|
||||
display: flex;
|
||||
align-items: center;
|
||||
justify-content: flex-end;
|
||||
gap: 10px;
|
||||
padding: 12px 20px;
|
||||
border-top: 1px solid #e5e7eb;
|
||||
background: #fff;
|
||||
}
|
||||
.vv-modal-foot .vv-btn-primary {
|
||||
width: auto;
|
||||
padding: 8px 18px;
|
||||
}
|
||||
.vv-modal-foot .vv-btn-ghost {
|
||||
padding: 8px 18px;
|
||||
padding: 20px;
|
||||
position: relative;
|
||||
}
|
||||
.vv-modal-close {
|
||||
position: absolute;
|
||||
top: 14px;
|
||||
right: 14px;
|
||||
background: transparent;
|
||||
border: none;
|
||||
color: #6b7280;
|
||||
@@ -1565,11 +1489,6 @@
|
||||
width: 28px;
|
||||
height: 28px;
|
||||
border-radius: 6px;
|
||||
display: flex;
|
||||
align-items: center;
|
||||
justify-content: center;
|
||||
padding: 0;
|
||||
line-height: 1;
|
||||
}
|
||||
.vv-modal-close:hover {
|
||||
color: #ef4444;
|
||||
@@ -1859,213 +1778,3 @@
|
||||
color: #7c3aed;
|
||||
font-weight: 500;
|
||||
}
|
||||
|
||||
/* ── v1.6 我的音色(默认主路径) ── */
|
||||
.vv-voice-section {
|
||||
margin-top: 4px;
|
||||
border: 1px solid #e5e7eb;
|
||||
border-radius: 8px;
|
||||
background: #fff;
|
||||
padding: 10px;
|
||||
}
|
||||
.vv-voice-section-head {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
justify-content: space-between;
|
||||
margin-bottom: 8px;
|
||||
}
|
||||
.vv-voice-section-title {
|
||||
font-size: 13px;
|
||||
font-weight: 600;
|
||||
color: #1f2937;
|
||||
display: inline-flex;
|
||||
align-items: center;
|
||||
gap: 6px;
|
||||
}
|
||||
.vv-voice-section-title .anticon {
|
||||
color: #7c3aed;
|
||||
}
|
||||
.vv-link-btn-sm {
|
||||
font-size: 12px;
|
||||
padding: 2px 6px;
|
||||
display: inline-flex;
|
||||
align-items: center;
|
||||
gap: 4px;
|
||||
}
|
||||
.vv-my-voice-list {
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
gap: 4px;
|
||||
max-height: 220px;
|
||||
overflow-y: auto;
|
||||
}
|
||||
.vv-my-voice-item {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
gap: 10px;
|
||||
padding: 8px 10px;
|
||||
border-radius: 6px;
|
||||
cursor: pointer;
|
||||
transition: background 0.15s;
|
||||
}
|
||||
.vv-my-voice-item:hover {
|
||||
background: #f5f0ff;
|
||||
}
|
||||
.vv-my-voice-item.selected {
|
||||
background: #f5f0ff;
|
||||
}
|
||||
.vv-my-voice-item.selected .vv-voice-radio {
|
||||
border-color: #7c3aed;
|
||||
background: #7c3aed;
|
||||
box-shadow: inset 0 0 0 2px #fff;
|
||||
}
|
||||
.vv-voice-empty {
|
||||
padding: 16px 12px;
|
||||
text-align: center;
|
||||
color: #9ca3af;
|
||||
font-size: 12px;
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
align-items: center;
|
||||
gap: 8px;
|
||||
}
|
||||
.vv-voice-empty-ic {
|
||||
font-size: 28px;
|
||||
color: #d1d5db;
|
||||
}
|
||||
.vv-voice-empty-text {
|
||||
font-size: 12px;
|
||||
}
|
||||
.vv-voice-alt-row {
|
||||
display: flex;
|
||||
gap: 8px;
|
||||
margin-top: 8px;
|
||||
}
|
||||
.vv-voice-alt-btn {
|
||||
flex: 1;
|
||||
display: inline-flex;
|
||||
align-items: center;
|
||||
justify-content: center;
|
||||
gap: 6px;
|
||||
padding: 8px 10px;
|
||||
background: #fafafe;
|
||||
border: 1px solid #e5e7eb;
|
||||
border-radius: 6px;
|
||||
color: #6b7280;
|
||||
font-size: 12px;
|
||||
cursor: pointer;
|
||||
transition: all 0.15s;
|
||||
}
|
||||
.vv-voice-alt-btn:hover {
|
||||
border-color: #7c3aed;
|
||||
color: #7c3aed;
|
||||
background: #f5f0ff;
|
||||
}
|
||||
.vv-voice-alt-btn.selected {
|
||||
background: #f5f0ff;
|
||||
border-color: #7c3aed;
|
||||
color: #7c3aed;
|
||||
}
|
||||
.vv-voice-panel-actions {
|
||||
display: flex;
|
||||
gap: 12px;
|
||||
margin-bottom: 6px;
|
||||
padding-bottom: 6px;
|
||||
border-bottom: 1px dashed #e5e7eb;
|
||||
}
|
||||
|
||||
/* ─────────── markdown 渲染(提示词控制展示格式) ─────────── */
|
||||
.vv-recog-md {
|
||||
padding: 4px 0;
|
||||
}
|
||||
.vv-copy-preview {
|
||||
margin-bottom: 14px;
|
||||
padding: 12px 14px;
|
||||
background: linear-gradient(180deg, #faf7ff 0%, #f6f2ff 100%);
|
||||
border: 1px solid #ece4fb;
|
||||
border-radius: 10px;
|
||||
}
|
||||
.vv-copy-preview-h {
|
||||
margin: 0 0 8px;
|
||||
border-bottom: none;
|
||||
padding-bottom: 0;
|
||||
}
|
||||
.vv-md-body {
|
||||
font-size: 13px;
|
||||
line-height: 1.7;
|
||||
color: #374151;
|
||||
word-break: break-word;
|
||||
}
|
||||
.vv-md-body h1,
|
||||
.vv-md-body h2,
|
||||
.vv-md-body h3,
|
||||
.vv-md-body h4 {
|
||||
margin: 10px 0 6px;
|
||||
font-weight: 600;
|
||||
color: #1f2937;
|
||||
line-height: 1.4;
|
||||
}
|
||||
.vv-md-body h1 {
|
||||
font-size: 18px;
|
||||
}
|
||||
.vv-md-body h2 {
|
||||
font-size: 16px;
|
||||
}
|
||||
.vv-md-body h3 {
|
||||
font-size: 15px;
|
||||
}
|
||||
.vv-md-body h4 {
|
||||
font-size: 14px;
|
||||
}
|
||||
.vv-md-body p {
|
||||
margin: 6px 0;
|
||||
}
|
||||
.vv-md-body ul,
|
||||
.vv-md-body ol {
|
||||
margin: 6px 0;
|
||||
padding-left: 20px;
|
||||
}
|
||||
.vv-md-body li {
|
||||
margin: 3px 0;
|
||||
}
|
||||
.vv-md-body strong {
|
||||
color: #111827;
|
||||
font-weight: 600;
|
||||
}
|
||||
.vv-md-body blockquote {
|
||||
margin: 8px 0;
|
||||
padding: 4px 12px;
|
||||
border-left: 3px solid #7c3aed;
|
||||
background: rgba(124, 58, 237, 0.05);
|
||||
color: #4b5563;
|
||||
}
|
||||
.vv-md-body code {
|
||||
padding: 1px 5px;
|
||||
background: #f3f4f6;
|
||||
border-radius: 4px;
|
||||
font-size: 12px;
|
||||
color: #be185d;
|
||||
}
|
||||
.vv-md-body a {
|
||||
color: #7c3aed;
|
||||
text-decoration: none;
|
||||
}
|
||||
.vv-md-body a:hover {
|
||||
text-decoration: underline;
|
||||
}
|
||||
.vv-md-body table {
|
||||
border-collapse: collapse;
|
||||
margin: 8px 0;
|
||||
width: 100%;
|
||||
}
|
||||
.vv-md-body th,
|
||||
.vv-md-body td {
|
||||
border: 1px solid #e5e7eb;
|
||||
padding: 6px 10px;
|
||||
text-align: left;
|
||||
}
|
||||
.vv-md-body hr {
|
||||
border: none;
|
||||
border-top: 1px solid #e5e7eb;
|
||||
margin: 12px 0;
|
||||
}
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -245,11 +245,12 @@ export default function AssetPickerModal({
|
||||
</div>
|
||||
{multiple && (
|
||||
<div className="vv-modal-foot">
|
||||
<button className="vv-btn vv-btn-ghost" onClick={onClose}>
|
||||
<button className="vv-btn vv-btn-ghost vv-btn-sm" onClick={onClose}>
|
||||
取消
|
||||
</button>
|
||||
<button
|
||||
className="vv-btn vv-btn-primary"
|
||||
style={{ width: "auto", marginTop: 0, padding: "8px 18px" }}
|
||||
onClick={handleConfirm}
|
||||
disabled={picked.size === 0}
|
||||
>
|
||||
|
||||
@@ -121,11 +121,7 @@ const appChildren: RouteObject[] = [
|
||||
children: [
|
||||
{
|
||||
index: true,
|
||||
element: <Navigate to="/app/admin/ditto-emotion" replace />,
|
||||
},
|
||||
{
|
||||
path: "ditto-emotion",
|
||||
lazy: lazyRoute(() => import("@/pages/admin/DittoEmotionConfig")),
|
||||
lazy: lazyRoute(() => import("@/pages/admin/AdminComingSoon")),
|
||||
},
|
||||
{
|
||||
path: "users",
|
||||
|
||||
@@ -53,9 +53,6 @@ celery_app.conf.imports = (
|
||||
# #1998 GPU MuseTalk 异步推理:wait_for_result→签名 URL→回写 lipsync_jobs
|
||||
# 必须在 Worker 侧注册,否则 apply_async 消息无人消费,job 永远卡在 processing
|
||||
"app.tasks.lipsync_gpu",
|
||||
# #2076 Ditto 蚂蚁数字人异步推理:同步 HTTP 调用 Ditto → MP4 流转存 OSS → 回写 lipsync_jobs
|
||||
# 必须在 Worker 侧注册;失败回退 GPU MuseTalk → MediaKit
|
||||
"app.tasks.lipsync_ditto",
|
||||
)
|
||||
|
||||
# Celery Beat 定时任务调度
|
||||
|
||||
@@ -387,41 +387,6 @@ BATCH_RENDER_SIMILARITY_LIMIT = 0.20
|
||||
"""批次内成片查重相似度阈值:超过则重选独立 plan 重渲一次(20%)。"""
|
||||
|
||||
|
||||
def _refund_smart_edit_prepaid(task_id: str) -> None:
|
||||
"""智能剪辑任务最终失败时退还预扣积分(幂等)。"""
|
||||
session = SessionLocal()
|
||||
try:
|
||||
from packages.adapters.sqlalchemy_impl.generation_task_repository import (
|
||||
SQLAlchemyGenerationTaskRepository,
|
||||
)
|
||||
from packages.domain.points_service import PointsService
|
||||
|
||||
repo = SQLAlchemyGenerationTaskRepository(session)
|
||||
task = repo.get(task_id)
|
||||
if not task:
|
||||
return
|
||||
prepaid = float(getattr(task, "credits_prepaid", 0) or 0)
|
||||
if prepaid <= 0:
|
||||
return
|
||||
txn_id = getattr(task, "credits_transaction_id", "") or ""
|
||||
res = PointsService().refund_points(
|
||||
user_id=task.user_id,
|
||||
amount=prepaid,
|
||||
source="smart_edit",
|
||||
db=session,
|
||||
ref_id=task.id,
|
||||
related_transaction_id=txn_id or None,
|
||||
description="智能剪辑任务失败退回",
|
||||
)
|
||||
task.credits_cost = 0.0
|
||||
task.credits_prepaid = 0.0
|
||||
repo.update(task)
|
||||
if not res.get("success"):
|
||||
logger.warning("[task_id=%s] 失败退积分未成功: %s", task_id, res)
|
||||
finally:
|
||||
session.close()
|
||||
|
||||
|
||||
def should_rerender_for_batch_dedup(*, batch_id: str, render_attempt: int, batch_similarity) -> bool:
|
||||
"""批次内查重后判定是否需要重选 plan 重渲。
|
||||
|
||||
@@ -1202,10 +1167,6 @@ def generate_video(self, task_id: str) -> dict:
|
||||
"mark_failed",
|
||||
error_message="source_edit_plan_id is required. Please create a preview task first.",
|
||||
)
|
||||
try:
|
||||
_refund_smart_edit_prepaid(task_id)
|
||||
except Exception:
|
||||
logger.warning("[task_id=%s] 失败退积分异常", task_id, exc_info=True)
|
||||
return {
|
||||
"status": "failed",
|
||||
"task_id": task_id,
|
||||
@@ -1244,7 +1205,6 @@ def generate_video(self, task_id: str) -> dict:
|
||||
)
|
||||
|
||||
# ── 自动重试逻辑 ──────────────────────────────────────────────────
|
||||
will_retry = False
|
||||
try:
|
||||
from packages.adapters.sqlalchemy_impl.generation_task_repository import (
|
||||
SQLAlchemyGenerationTaskRepository,
|
||||
@@ -1257,7 +1217,6 @@ def generate_video(self, task_id: str) -> dict:
|
||||
if _task and _task.auto_retry_enabled and _task.auto_retry_max > 0:
|
||||
current_retry = _task.retry_count or 0
|
||||
if current_retry < _task.auto_retry_max:
|
||||
will_retry = True
|
||||
logger.info(
|
||||
"[task_id=%s] 触发自动重试: 当前重试次数=%d, 最大重试次数=%d",
|
||||
task_id,
|
||||
@@ -1291,13 +1250,6 @@ def generate_video(self, task_id: str) -> dict:
|
||||
exc_info=True,
|
||||
)
|
||||
|
||||
# 最终失败(不再重试):退还 smart_edit 预扣积分
|
||||
if not will_retry:
|
||||
try:
|
||||
_refund_smart_edit_prepaid(task_id)
|
||||
except Exception:
|
||||
logger.warning("[task_id=%s] 失败退积分异常", task_id, exc_info=True)
|
||||
|
||||
return {
|
||||
"status": "failed",
|
||||
"task_id": task_id,
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,4 +0,0 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""V2 图片分析:火山OCR专用API + doubao-lite强约束JSON并行,单次pro VLM兜底。"""
|
||||
|
||||
from .fast_path import analyze_image_v2, analyze_images_v2 # noqa: F401
|
||||
@@ -1,95 +0,0 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""V2 prompt 解析:优先读后台 viral_video_prompt_templates(prompt_type='image_analysis'
|
||||
且 is_active=true),30s TTL 热加载;DB 无有效记录/异常时,fallback 到 prompts.py 的
|
||||
image_analysis v8 默认 system/user。
|
||||
|
||||
规则:
|
||||
- DB 有 is_active=true 的 image_analysis 记录:system 原样用 DB.system_prompt
|
||||
(自带完整输出格式,不追加任何硬编码 schema),user 用 DB.user_prompt_template
|
||||
渲染(填入 image_url / ocr_text);
|
||||
- DB 无记录/异常:system/user 用 prompts.py 里的 v8 默认模板。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import threading
|
||||
import time
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _default_template() -> dict:
|
||||
# 延迟导入:避免模块加载时拉起整个 packages 依赖链(也便于旧 Python 收集测试)
|
||||
from packages.application.viral_video.prompts import DEFAULT_TEMPLATES
|
||||
|
||||
for item in DEFAULT_TEMPLATES:
|
||||
if item["prompt_type"] == "image_analysis":
|
||||
return item
|
||||
raise RuntimeError("image_analysis 默认模板缺失")
|
||||
|
||||
|
||||
_cache_lock = threading.Lock()
|
||||
_cache: dict[str, tuple[float, tuple[str, str]]] = {}
|
||||
_CACHE_TTL = 30.0
|
||||
|
||||
|
||||
def _load_db_template():
|
||||
"""查 DB is_active=true 的 image_analysis 记录;不可达/无记录返回 None。"""
|
||||
try:
|
||||
from packages.application.viral_video.prompt_loader import _load_from_db
|
||||
|
||||
return _load_from_db("image_analysis")
|
||||
except Exception as e: # noqa: BLE001
|
||||
logger.warning("[vision.v2] 查询DB image_analysis prompt失败: %s", e)
|
||||
return None
|
||||
|
||||
|
||||
def _render_user(user_tpl: str, image_url: str, ocr_text: str) -> str:
|
||||
try:
|
||||
return user_tpl.format(image_url=image_url, ocr_text=ocr_text or "无")
|
||||
except Exception: # noqa: BLE001
|
||||
return user_tpl
|
||||
|
||||
|
||||
def _resolve(kind: str, image_url: str = "", ocr_text: str = "") -> tuple[str, str]:
|
||||
now = time.time()
|
||||
cache_key = f"prompt_{kind}"
|
||||
with _cache_lock:
|
||||
hit = _cache.get(cache_key)
|
||||
if hit and now - hit[0] < _CACHE_TTL:
|
||||
sys_prompt, usr_prompt = hit[1]
|
||||
return sys_prompt, _render_user(usr_prompt, image_url, ocr_text)
|
||||
|
||||
default = _default_template()
|
||||
sys_prompt = default["system_prompt"]
|
||||
usr_prompt = default["user_prompt_template"]
|
||||
|
||||
tpl = _load_db_template()
|
||||
if tpl is not None:
|
||||
db_sys = (getattr(tpl, "system_prompt", "") or "").strip()
|
||||
if db_sys:
|
||||
sys_prompt = db_sys
|
||||
db_usr = getattr(tpl, "user_prompt_template", "") or usr_prompt
|
||||
usr_prompt = db_usr or usr_prompt
|
||||
logger.info(
|
||||
"[vision.v2] 使用DB image_analysis prompt version=%s",
|
||||
getattr(tpl, "version", "?"),
|
||||
)
|
||||
|
||||
with _cache_lock:
|
||||
_cache[cache_key] = (now, (sys_prompt, usr_prompt))
|
||||
return sys_prompt, _render_user(usr_prompt, image_url, ocr_text)
|
||||
|
||||
|
||||
def resolve_fast_prompt(image_url: str = "", ocr_text: str = "") -> tuple[str, str]:
|
||||
return _resolve("fast", image_url, ocr_text)
|
||||
|
||||
|
||||
def resolve_pro_prompt(image_url: str = "", ocr_text: str = "") -> tuple[str, str]:
|
||||
return _resolve("pro", image_url, ocr_text)
|
||||
|
||||
|
||||
def invalidate_cache() -> None:
|
||||
with _cache_lock:
|
||||
_cache.clear()
|
||||
@@ -1,106 +0,0 @@
|
||||
"""V2 结果组装(v8 叙述优先,大幅精简)。
|
||||
|
||||
设计原则:VLM 直接输出最终给用户看的 summary_markdown,assembler 只负责
|
||||
- 解析 fast JSON(兼容顶层 {"images":[...]} 与 {"products":[...]} 两种键);
|
||||
- 补齐 5 个必备字段(type/name/brand/has_person/summary_markdown);
|
||||
- summary_markdown 缺失(异常)时才拼一句最基础的兜底文字。
|
||||
|
||||
正常情况下不改写、不“润色”VLM 输出,不做 brand 多级兜底,不处理任何
|
||||
colors/material/key_features 等细分字段。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from typing import Any
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_VALID_TYPES = ("store", "product", "person", "scene")
|
||||
|
||||
|
||||
def _coerce_bool(value: Any) -> bool:
|
||||
if isinstance(value, bool):
|
||||
return value
|
||||
if isinstance(value, (int, float)):
|
||||
return value != 0
|
||||
if isinstance(value, str):
|
||||
return value.strip().lower() in ("true", "1", "yes", "是")
|
||||
return False
|
||||
|
||||
|
||||
def _basic_markdown(image: dict[str, Any]) -> str:
|
||||
"""异常兜底:VLM 没给 summary_markdown 时只拼一句基础文字。"""
|
||||
name = (image.get("name") or "").strip() or "未识别"
|
||||
brand = (image.get("brand") or "").strip()
|
||||
typ = image.get("type") or "scene"
|
||||
label = f"{brand}{name}" if brand and brand not in name else (brand or name)
|
||||
if typ == "store":
|
||||
return f"这是{label}的门店场景,画面细节识别不完整。"
|
||||
if typ == "person":
|
||||
return f"画面中的人物与{label}相关,细节识别不完整。"
|
||||
if typ == "product":
|
||||
return f"这是{label}的商品图片,具体外观细节识别不完整。"
|
||||
return f"画面内容为{label},细节识别不完整。"
|
||||
|
||||
|
||||
def _normalize_image(raw: Any, idx: int) -> dict[str, Any]:
|
||||
"""把一条 VLM 输出归一化为 5 字段 dict。"""
|
||||
if not isinstance(raw, dict):
|
||||
raw = {}
|
||||
|
||||
typ = str(raw.get("type") or "").strip().lower()
|
||||
if typ not in _VALID_TYPES:
|
||||
typ = "scene"
|
||||
|
||||
name = str(raw.get("name") or "").strip() or "未识别"
|
||||
brand = str(raw.get("brand") or "").strip()
|
||||
has_person = _coerce_bool(raw.get("has_person"))
|
||||
if typ == "person" and not has_person:
|
||||
# type=person 通常意味着主体是人,保持一致(仅异常补全)
|
||||
has_person = True
|
||||
|
||||
summary = raw.get("summary_markdown")
|
||||
summary = summary.strip() if isinstance(summary, str) else ""
|
||||
|
||||
image: dict[str, Any] = {
|
||||
"type": typ,
|
||||
"name": name,
|
||||
"brand": brand,
|
||||
"has_person": has_person,
|
||||
"summary_markdown": summary,
|
||||
}
|
||||
if not summary:
|
||||
image["summary_markdown"] = _basic_markdown(image)
|
||||
image["_source"] = "summary_missing"
|
||||
logger.info("[assembler] 图片 #%s 缺少 summary_markdown,使用基础兜底", idx)
|
||||
return image
|
||||
|
||||
|
||||
def _extract_items(fast_json: Any) -> list[Any]:
|
||||
"""从 fast JSON 中取出图片条目:优先 images,兼容 products。"""
|
||||
if not isinstance(fast_json, dict):
|
||||
return []
|
||||
items = fast_json.get("images")
|
||||
if not isinstance(items, list):
|
||||
items = fast_json.get("products")
|
||||
return items if isinstance(items, list) else []
|
||||
|
||||
|
||||
def assemble_result(idx: int, fast_json: Any, ocr_texts: list[str] | None = None) -> dict[str, Any]:
|
||||
"""组装单张图片分析结果。
|
||||
|
||||
每次调用对应一张图片;fast_json 形如 {"images": [{...}]}(v8)。
|
||||
返回单条 image dict(5 字段,必要时带 _source)。
|
||||
"""
|
||||
items = _extract_items(fast_json)
|
||||
if items:
|
||||
image = _normalize_image(items[0], idx)
|
||||
else:
|
||||
# 极端异常:fast 无任何可用条目,OCR 文字可作为名称线索
|
||||
ocr_hint = ""
|
||||
if ocr_texts:
|
||||
ocr_hint = "、".join(t for t in ocr_texts if t)[:40]
|
||||
image = _normalize_image({"name": ocr_hint or "未识别", "summary_markdown": ""}, idx)
|
||||
image["_source"] = "empty_fast_json"
|
||||
return image
|
||||
@@ -1,124 +0,0 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""V2 图片分析主路径(v8 叙述优先):每图并行 OCR(火山 MediaKit,未配置自动跳过)
|
||||
+ fast VLM 强约束 JSON;失败时单次 pro VLM 兜底。
|
||||
|
||||
架构:
|
||||
- 单图 2 路并行(OCR + fast VLM),外层 N 图全并发(workers=8);
|
||||
- 兜底单次 pro VLM,无竞速/复杂重试;
|
||||
- 输出统一为 5 字段 image dict(type/name/brand/has_person/summary_markdown)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import os
|
||||
import time
|
||||
from concurrent.futures import ThreadPoolExecutor, as_completed
|
||||
from typing import Any
|
||||
|
||||
from . import assembler, ocr_volc, vlm_fallback, vlm_fast_json
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_IMG_WORKERS = int(os.environ.get("VISION_V2_IMG_WORKERS", "8"))
|
||||
_FAST_TIMEOUT = float(os.environ.get("VISION_V2_FAST_TIMEOUT", "20"))
|
||||
_FAST_JSON_TIMEOUT = float(os.environ.get("VISION_V2_FAST_JSON_TIMEOUT", "20"))
|
||||
_OCR_TIMEOUT = float(os.environ.get("VISION_V2_OCR_TIMEOUT", "6"))
|
||||
_PRO_TIMEOUT = float(os.environ.get("VISION_V2_PRO_TIMEOUT", "45"))
|
||||
|
||||
|
||||
def _is_usable(r: dict[str, Any] | None) -> bool:
|
||||
if not isinstance(r, dict):
|
||||
return False
|
||||
return bool((r.get("summary_markdown") or "").strip())
|
||||
|
||||
|
||||
def _basic_failure(ocr_result: list[str], fast_elapsed: float, source: str) -> dict[str, Any]:
|
||||
image = assembler.assemble_result(-1, {}, ocr_result)
|
||||
image["_source"] = source
|
||||
image["_fast_elapsed"] = round(fast_elapsed, 2)
|
||||
return image
|
||||
|
||||
|
||||
def analyze_image_v2(idx: int, img_url: str) -> dict[str, Any]:
|
||||
t0 = time.time()
|
||||
|
||||
fj_result: dict[str, Any] | None = None
|
||||
ocr_result: list[str] = []
|
||||
pool = ThreadPoolExecutor(max_workers=2)
|
||||
f_fj = pool.submit(vlm_fast_json.call_fast_json, img_url, timeout=_FAST_JSON_TIMEOUT)
|
||||
f_ocr = pool.submit(ocr_volc.call_ocr, img_url, timeout=_OCR_TIMEOUT)
|
||||
try:
|
||||
for fut in as_completed([f_fj, f_ocr], timeout=_FAST_TIMEOUT):
|
||||
try:
|
||||
res = fut.result(timeout=1)
|
||||
except Exception as e: # noqa: BLE001
|
||||
logger.warning("[vision.v2] 图片 #%d 子任务异常: %s", idx, e)
|
||||
continue
|
||||
if fut is f_fj and isinstance(res, dict):
|
||||
fj_result = res
|
||||
elif fut is f_ocr and isinstance(res, list):
|
||||
ocr_result = res
|
||||
except TimeoutError:
|
||||
for f in (f_fj, f_ocr):
|
||||
if not f.done():
|
||||
f.cancel()
|
||||
logger.warning("[vision.v2] 图片 #%d fast路径超时(%.0fs),走pro兜底", idx, _FAST_TIMEOUT)
|
||||
finally:
|
||||
fast_elapsed = time.time() - t0
|
||||
pool.shutdown(wait=False)
|
||||
|
||||
if fj_result:
|
||||
assembled = assembler.assemble_result(idx, fj_result, ocr_result)
|
||||
if _is_usable(assembled):
|
||||
assembled["_fast_elapsed"] = round(fast_elapsed, 2)
|
||||
logger.info("[vision.v2] 图片 #%d fast命中 elapsed=%.2fs", idx, fast_elapsed)
|
||||
return assembled
|
||||
|
||||
pro_result = vlm_fallback.call_pro_vlm(img_url, idx, ocr_hint=ocr_result, timeout=_PRO_TIMEOUT)
|
||||
if _is_usable(pro_result):
|
||||
pro_result["_fallback_used"] = True
|
||||
pro_result["_fast_elapsed"] = round(fast_elapsed, 2)
|
||||
logger.info("[vision.v2] 图片 #%d pro兜底命中 total=%.2fs", idx, time.time() - t0)
|
||||
return pro_result
|
||||
|
||||
logger.warning("[vision.v2] 图片 #%d 全路径失败 elapsed=%.2fs", idx, time.time() - t0)
|
||||
return _basic_failure(ocr_result, fast_elapsed, "v2_all_failed")
|
||||
|
||||
|
||||
def analyze_images_v2(img_urls: list[str]) -> list[dict[str, Any]]:
|
||||
if not img_urls:
|
||||
return []
|
||||
workers = min(_IMG_WORKERS, len(img_urls), 16)
|
||||
results: list[dict[str, Any] | None] = [None] * len(img_urls)
|
||||
|
||||
logger.info(
|
||||
"[vision.v2] 开始图片分析 n=%d workers=%d fast_timeout=%.0fs pro_timeout=%.0fs",
|
||||
len(img_urls),
|
||||
workers,
|
||||
_FAST_TIMEOUT,
|
||||
_PRO_TIMEOUT,
|
||||
)
|
||||
t0 = time.time()
|
||||
with ThreadPoolExecutor(max_workers=workers) as pool:
|
||||
future_to_idx = {pool.submit(analyze_image_v2, idx, url): idx for idx, url in enumerate(img_urls)}
|
||||
for fut in as_completed(future_to_idx):
|
||||
idx = future_to_idx[fut]
|
||||
try:
|
||||
results[idx] = fut.result()
|
||||
except Exception as e: # noqa: BLE001
|
||||
logger.warning("[vision.v2] 图片 #%d future异常: %s", idx, e, exc_info=True)
|
||||
results[idx] = assembler.assemble_result(idx, {}, [])
|
||||
results[idx]["_source"] = "v2_future_exception" # type: ignore[index]
|
||||
|
||||
elapsed = time.time() - t0
|
||||
succ = sum(1 for r in results if _is_usable(r))
|
||||
fb = sum(1 for r in results if r and r.get("_fallback_used"))
|
||||
logger.info(
|
||||
"[vision.v2] 完成 n=%d usable=%d pro_fallback=%d elapsed=%.2fs",
|
||||
len(img_urls),
|
||||
succ,
|
||||
fb,
|
||||
elapsed,
|
||||
)
|
||||
return [r for r in results if r is not None] # type: ignore[misc]
|
||||
@@ -1,138 +0,0 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""VLM 返回文本的稳健 JSON 提取工具。
|
||||
|
||||
背景:复杂门店图 VLM 输出经常被 max_tokens 截断(finish_reason=length),
|
||||
json.loads 失败后整个结果被丢弃,导致"未识别"。本工具提供:
|
||||
1. markdown 代码块剥离(含只开不闭的截断场景)
|
||||
2. 最外层 { } 切片
|
||||
3. 非法控制字符清理
|
||||
4. 直接 json.loads
|
||||
5. 截断 JSON 括号/引号栈补全修复
|
||||
6. 尾部逐字符截断重试(去除最后一个不完整 token 后修复)
|
||||
|
||||
成功返回 dict;截断修复产物带 _partial=True 标记;彻底失败返回 None。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import re
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_CODE_FENCE_RE = re.compile(r"^```(?:json)?\s*\n?(.*?)\n?```\s*$", re.DOTALL)
|
||||
|
||||
|
||||
def _strip_code_fence(s: str) -> str:
|
||||
s = s.strip()
|
||||
m = _CODE_FENCE_RE.match(s)
|
||||
if m:
|
||||
return m.group(1).strip()
|
||||
# 兼容开头 ```json 但结尾无 ```(截断场景)
|
||||
if s.startswith("```"):
|
||||
lines = s.split("\n")
|
||||
if lines and lines[0].startswith("```"):
|
||||
lines = lines[1:]
|
||||
s = "\n".join(lines).strip()
|
||||
return s
|
||||
|
||||
|
||||
def _repair_truncated_json(text: str) -> str:
|
||||
"""尝试补全被截断的JSON:维护 bracket/quote 栈,在末尾补闭合符。"""
|
||||
stack: list[str] = []
|
||||
in_string = False
|
||||
escape = False
|
||||
for ch in text:
|
||||
if escape:
|
||||
escape = False
|
||||
continue
|
||||
if ch == "\\" and in_string:
|
||||
escape = True
|
||||
continue
|
||||
if ch == '"':
|
||||
in_string = not in_string
|
||||
continue
|
||||
if in_string:
|
||||
continue
|
||||
if ch in "{[":
|
||||
stack.append(ch)
|
||||
elif ch == "}":
|
||||
if stack and stack[-1] == "{":
|
||||
stack.pop()
|
||||
elif ch == "]":
|
||||
if stack and stack[-1] == "[":
|
||||
stack.pop()
|
||||
repair = ""
|
||||
if in_string:
|
||||
repair += '"'
|
||||
for opener in reversed(stack):
|
||||
repair += "}" if opener == "{" else "]"
|
||||
if repair:
|
||||
logger.info(
|
||||
"[json_utils] 截断JSON修复: 补全%d个闭合符 in_string=%s",
|
||||
len(repair),
|
||||
in_string,
|
||||
)
|
||||
return text + repair
|
||||
|
||||
|
||||
def _clean_invalid_chars(text: str) -> str:
|
||||
"""清理JSON中非法的控制字符(tab/newline 之外的 0x00-0x1f 段)。"""
|
||||
return re.sub(r"[\x00-\x08\x0b\x0c\x0e-\x1f]", "", text)
|
||||
|
||||
|
||||
def extract_json_object(text: str) -> dict | None:
|
||||
"""从VLM返回文本中稳健提取JSON对象。
|
||||
|
||||
返回 dict 或 None。成功的 dict 可能带 _partial=True 标记,
|
||||
表示原始文本被截断、经括号补全后得到的产物。
|
||||
"""
|
||||
if not text or not isinstance(text, str):
|
||||
return None
|
||||
# 1. 剥离 markdown
|
||||
text = _strip_code_fence(text)
|
||||
# 2. 找最外层 { }
|
||||
lpos = text.find("{")
|
||||
if lpos < 0:
|
||||
return None
|
||||
rpos = text.rfind("}")
|
||||
if rpos > lpos:
|
||||
text = text[lpos : rpos + 1]
|
||||
else:
|
||||
# 截断场景:无任何闭合 },取到末尾交给修复器
|
||||
text = text[lpos:]
|
||||
# 3. 清理非法控制字符
|
||||
text = _clean_invalid_chars(text)
|
||||
# 4. 直接 loads
|
||||
try:
|
||||
obj = json.loads(text)
|
||||
return obj if isinstance(obj, dict) else None
|
||||
except json.JSONDecodeError:
|
||||
pass
|
||||
# 5. 尝试截断修复
|
||||
repaired = _repair_truncated_json(text)
|
||||
try:
|
||||
obj = json.loads(repaired)
|
||||
if isinstance(obj, dict):
|
||||
obj["_partial"] = True
|
||||
return obj
|
||||
except json.JSONDecodeError:
|
||||
pass
|
||||
# 6. 尾部逐字符截断重试(去除最后一个不完整 token)
|
||||
for _ in range(50):
|
||||
last_comma = repaired.rfind(",")
|
||||
last_brace = max(repaired.rfind("}"), repaired.rfind("]"))
|
||||
cut = max(last_comma, last_brace)
|
||||
if cut < 10:
|
||||
break
|
||||
repaired = repaired[: cut + 1]
|
||||
repaired = _repair_truncated_json(repaired)
|
||||
try:
|
||||
obj = json.loads(repaired)
|
||||
if isinstance(obj, dict):
|
||||
obj["_partial"] = True
|
||||
return obj
|
||||
except json.JSONDecodeError:
|
||||
continue
|
||||
return None
|
||||
@@ -1,109 +0,0 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""火山引擎 AI MediaKit OCR(同步)调用封装。
|
||||
|
||||
接口:POST {mediakit_base_url}/tools-sync/ocr
|
||||
鉴权:Bearer {mediakit_api_key}
|
||||
请求体:{"image_url": "<公网可访问URL>"} (部分版本也支持 image_base64)
|
||||
响应:{"code":0,"data":{"texts":[{"text":"...","bbox":[x,y,w,h],...},...],...}}
|
||||
|
||||
目标:识别商品包装/Logo/水印上的文字,作为 fast_json VLM 的补充。
|
||||
返回值:识别到的文本字符串列表(失败返回 [])。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import time
|
||||
from typing import Any
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
DEFAULT_TIMEOUT = 8 # OCR 秒级返回,8s 绰绰有余
|
||||
|
||||
|
||||
def call_ocr(img_url: str, *, timeout: int = DEFAULT_TIMEOUT) -> list[str]:
|
||||
"""调用 MediaKit 同步 OCR,返回去重后的纯文本列表。
|
||||
|
||||
不做重试(外层降级逻辑负责)。失败/未配置返回空列表,不抛异常。
|
||||
"""
|
||||
t0 = time.time()
|
||||
try:
|
||||
import httpx
|
||||
|
||||
from packages.shared.mediakit_client import get_mediakit_client
|
||||
|
||||
client = get_mediakit_client()
|
||||
if not client.is_available:
|
||||
logger.info("[vision.v2] mediakit 未配置,跳过 OCR")
|
||||
return []
|
||||
|
||||
url = f"{client.base_url}/tools-sync/ocr"
|
||||
headers = {
|
||||
"Authorization": f"Bearer {client.api_key}",
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
payload: dict[str, Any] = {"image_url": img_url}
|
||||
# 部分文档版本用 image_base64,但公网 URL 场景下 image_url 最简
|
||||
resp = httpx.post(url, headers=headers, json=payload, timeout=timeout)
|
||||
elapsed = time.time() - t0
|
||||
if resp.status_code != 200:
|
||||
logger.warning(
|
||||
"[vision.v2] OCR HTTP %d elapsed=%.1fs body=%s",
|
||||
resp.status_code,
|
||||
elapsed,
|
||||
resp.text[:200],
|
||||
)
|
||||
return []
|
||||
data = resp.json()
|
||||
# 兼容几种可能的响应结构
|
||||
code = data.get("code", data.get("status", 0))
|
||||
if code not in (0, "OK", "success", 200):
|
||||
logger.warning("[vision.v2] OCR 业务错误 code=%s elapsed=%.1fs resp=%s", code, elapsed, str(data)[:200])
|
||||
return []
|
||||
texts = _extract_texts(data)
|
||||
# 去重 + 过滤空
|
||||
seen: set[str] = set()
|
||||
out: list[str] = []
|
||||
for t in texts:
|
||||
t = (t or "").strip()
|
||||
if t and t not in seen and len(t) <= 100: # 过滤过长的误识别
|
||||
seen.add(t)
|
||||
out.append(t)
|
||||
logger.info("[vision.v2] OCR 完成 elapsed=%.1fs n=%d texts=%s", elapsed, len(out), out[:5])
|
||||
return out
|
||||
except Exception as e:
|
||||
elapsed = time.time() - t0
|
||||
logger.warning("[vision.v2] OCR 异常 elapsed=%.1fs err=%s", elapsed, e, exc_info=True)
|
||||
return []
|
||||
|
||||
|
||||
def _extract_texts(data: dict) -> list[str]:
|
||||
"""从 OCR 响应中抽取文本,兼容多种结构。"""
|
||||
out: list[str] = []
|
||||
# 常见结构1: data.texts = [{"text": "..."}, ...]
|
||||
d = data.get("data") or data
|
||||
if isinstance(d, dict):
|
||||
for key in ("texts", "lines", "words", "items", "result"):
|
||||
items = d.get(key)
|
||||
if isinstance(items, list):
|
||||
for it in items:
|
||||
if isinstance(it, dict):
|
||||
txt = it.get("text") or it.get("content") or it.get("word")
|
||||
if txt:
|
||||
out.append(str(txt))
|
||||
elif isinstance(it, str):
|
||||
out.append(it)
|
||||
break
|
||||
# 结构2: data.text = "..."
|
||||
if not out:
|
||||
t = d.get("text")
|
||||
if isinstance(t, str):
|
||||
out.append(t)
|
||||
# 结构3: data.ocr_text / data.content
|
||||
if not out:
|
||||
for key in ("ocr_text", "content", "raw_text"):
|
||||
v = d.get(key)
|
||||
if isinstance(v, str) and v.strip():
|
||||
out.append(v)
|
||||
break
|
||||
return out
|
||||
@@ -1,109 +0,0 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""V2 兜底路径:vision client(fallback 变体)单图调用,走 v8 叙述优先 prompt。
|
||||
|
||||
fast 超时/非 JSON/为空时单次调用;输出统一走 assembler.assemble_result 组装,
|
||||
与 fast 路径同为 5 字段 image dict。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import time
|
||||
from typing import Any
|
||||
|
||||
from . import _prompt, assembler
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_DEFAULT_TIMEOUT = 45
|
||||
|
||||
|
||||
def call_pro_vlm(
|
||||
img_url: str,
|
||||
idx: int,
|
||||
*,
|
||||
ocr_hint: list[str] | None = None,
|
||||
timeout: int = _DEFAULT_TIMEOUT,
|
||||
max_tokens: int | None = None,
|
||||
) -> dict[str, Any] | None:
|
||||
t0 = time.time()
|
||||
ocr_text = "、".join(t for t in (ocr_hint or []) if t)[:200]
|
||||
|
||||
try:
|
||||
from packages.shared.ai_router import ai_router
|
||||
|
||||
client = ai_router.get_vision_client("image_analysis", variant="fallback")
|
||||
if not client or not client.is_available:
|
||||
logger.warning("[vision.v2] pro vision client 不可用,跳过")
|
||||
return None
|
||||
except Exception as e: # noqa: BLE001
|
||||
logger.warning("[vision.v2] ai_router 获取失败: %s", e)
|
||||
return None
|
||||
|
||||
system_prompt, user_prompt = _prompt.resolve_pro_prompt(img_url, ocr_text)
|
||||
|
||||
messages: list[dict[str, Any]] = [
|
||||
{"role": "system", "content": system_prompt},
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "image_url", "image_url": {"url": img_url}},
|
||||
{"type": "text", "text": user_prompt},
|
||||
],
|
||||
},
|
||||
]
|
||||
|
||||
try:
|
||||
call_kwargs: dict[str, Any] = {
|
||||
"messages": messages,
|
||||
"images": None,
|
||||
"temperature": 0.3,
|
||||
"timeout": timeout,
|
||||
"enable_thinking": False,
|
||||
"response_format": {"type": "json_object"},
|
||||
"max_tokens": max_tokens if max_tokens is not None else 4000,
|
||||
}
|
||||
|
||||
from .json_utils import extract_json_object
|
||||
|
||||
obj = None
|
||||
for _outer in range(2):
|
||||
kw = dict(call_kwargs)
|
||||
if _outer == 1:
|
||||
kw.pop("response_format", None)
|
||||
msgs2 = [dict(messages[0]), dict(messages[1])]
|
||||
cont = [dict(c) for c in msgs2[1]["content"]]
|
||||
cont[-1] = {
|
||||
"type": "text",
|
||||
"text": user_prompt + "\n严格只输出JSON对象,不要解释或markdown。",
|
||||
}
|
||||
msgs2[1] = {"role": "user", "content": cont}
|
||||
kw["messages"] = msgs2
|
||||
raw = client.vision_completion(**kw)
|
||||
if not raw:
|
||||
continue
|
||||
obj = extract_json_object(raw)
|
||||
if obj is not None:
|
||||
break
|
||||
logger.warning("[vision.v2] pro 非JSON(100字) outer=%s: %s", _outer, raw[:100])
|
||||
|
||||
elapsed = time.time() - t0
|
||||
if obj is None:
|
||||
logger.warning("[vision.v2] pro 两次均未得到JSON elapsed=%.1fs", elapsed)
|
||||
return None
|
||||
if obj.get("_partial"):
|
||||
logger.warning("[vision.v2] pro 截断JSON(partial) elapsed=%.1fs", elapsed)
|
||||
|
||||
result = assembler.assemble_result(idx, obj, ocr_hint or [])
|
||||
result["_source"] = "vlm_pro"
|
||||
result["_fallback_used"] = True
|
||||
logger.info("[vision.v2] pro 完成 model=%s elapsed=%.1fs", client.model, elapsed)
|
||||
return result
|
||||
except Exception as e: # noqa: BLE001
|
||||
logger.warning(
|
||||
"[vision.v2] pro 异常 elapsed=%.1fs err=%s",
|
||||
time.time() - t0,
|
||||
e,
|
||||
exc_info=True,
|
||||
)
|
||||
return None
|
||||
@@ -1,111 +0,0 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""V2 快速路径:vision client(默认 image_analysis 能力)强约束 JSON-only 调用。
|
||||
|
||||
要点:
|
||||
- 通过 ai_router.get_vision_client() 获取 client;
|
||||
- enable_thinking=False 关闭推理链,response_format=json_object 强约束 JSON;
|
||||
- system/user prompt 优先读后台模板(v8 叙述优先),DB 不可用时用 prompts.py 默认;
|
||||
- temperature=0.1(稳定输出 JSON);两次尝试(第二次去 json_object 约束)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import time
|
||||
from typing import Any
|
||||
|
||||
from . import _prompt
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_DEFAULT_TIMEOUT = 20
|
||||
|
||||
|
||||
def call_fast_json(
|
||||
img_url: str,
|
||||
*,
|
||||
timeout: int = _DEFAULT_TIMEOUT,
|
||||
max_tokens: int | None = None,
|
||||
) -> dict[str, Any] | None:
|
||||
t0 = time.time()
|
||||
|
||||
try:
|
||||
from packages.shared.ai_router import ai_router
|
||||
|
||||
client = ai_router.get_vision_client("image_analysis", variant="primary")
|
||||
if not client or not client.is_available:
|
||||
logger.warning("[vision.v2] vision client 不可用,跳过 fast_json")
|
||||
return None
|
||||
except Exception as e: # noqa: BLE001
|
||||
logger.warning("[vision.v2] ai_router 获取失败: %s", e)
|
||||
return None
|
||||
|
||||
system_prompt, user_prompt = _prompt.resolve_fast_prompt(img_url, "")
|
||||
|
||||
messages: list[dict[str, Any]] = [
|
||||
{"role": "system", "content": system_prompt},
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "image_url", "image_url": {"url": img_url}},
|
||||
{"type": "text", "text": user_prompt},
|
||||
],
|
||||
},
|
||||
]
|
||||
|
||||
try:
|
||||
call_kwargs: dict[str, Any] = {
|
||||
"messages": messages,
|
||||
"images": None,
|
||||
"temperature": 0.1,
|
||||
"timeout": timeout,
|
||||
"enable_thinking": False,
|
||||
"response_format": {"type": "json_object"},
|
||||
}
|
||||
if max_tokens is not None:
|
||||
call_kwargs["max_tokens"] = max_tokens
|
||||
|
||||
from .json_utils import extract_json_object
|
||||
|
||||
obj = None
|
||||
for _outer in range(2):
|
||||
kw = dict(call_kwargs)
|
||||
if _outer == 1:
|
||||
kw.pop("response_format", None)
|
||||
msgs2 = [dict(messages[0]), dict(messages[1])]
|
||||
cont = [dict(c) for c in msgs2[1]["content"]]
|
||||
cont[-1] = {
|
||||
"type": "text",
|
||||
"text": user_prompt + "\n严格只输出JSON对象,不要解释或markdown。",
|
||||
}
|
||||
msgs2[1] = {"role": "user", "content": cont}
|
||||
kw["messages"] = msgs2
|
||||
raw = client.vision_completion(**kw)
|
||||
if not raw:
|
||||
continue
|
||||
obj = extract_json_object(raw)
|
||||
if obj is not None:
|
||||
break
|
||||
logger.warning("[vision.v2] fast_json 非JSON(100字) outer=%s: %s", _outer, raw[:100])
|
||||
|
||||
elapsed = time.time() - t0
|
||||
if obj is None:
|
||||
logger.warning("[vision.v2] fast_json 两次均未得到JSON elapsed=%.1fs", elapsed)
|
||||
return None
|
||||
if obj.get("_partial"):
|
||||
logger.warning("[vision.v2] fast_json 截断JSON(partial) elapsed=%.1fs", elapsed)
|
||||
logger.info(
|
||||
"[vision.v2] fast_json 完成 model=%s elapsed=%.1fs type=%s",
|
||||
client.model,
|
||||
elapsed,
|
||||
obj.get("type"),
|
||||
)
|
||||
return obj
|
||||
except Exception as e: # noqa: BLE001
|
||||
logger.warning(
|
||||
"[vision.v2] fast_json 异常 elapsed=%.1fs err=%s",
|
||||
time.time() - t0,
|
||||
e,
|
||||
exc_info=True,
|
||||
)
|
||||
return None
|
||||
@@ -241,35 +241,19 @@ 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=${DOUBAO_TIMEOUT}
|
||||
DOUBAO_TIMEOUT=60
|
||||
|
||||
# 最大重试次数
|
||||
DOUBAO_MAX_RETRIES=${DOUBAO_MAX_RETRIES}
|
||||
DOUBAO_MAX_RETRIES=2
|
||||
|
||||
# 视觉模型(支持图片/视频理解的模型,model name 格式)
|
||||
# 视觉模型 Endpoint ID(支持图片/视频理解的模型)
|
||||
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(微信开放平台已配置)
|
||||
@@ -311,10 +295,3 @@ GPU_ENCODE_CRF=23
|
||||
GPU_ENCODE_FALLBACK_CPU=true
|
||||
GPU_ENCODE_MEZZANINE_TRANSPORT=oss
|
||||
GPU_ENCODE_OSS_TMP_PREFIX=tmp/gpu-mezzanine/
|
||||
|
||||
# ==================== Ditto 蚂蚁数字人口型 ====================
|
||||
# 注意:这些值必须写死在模板里(不是 CI Secret),否则每次 CI 重新渲染 .env 都会被丢弃,
|
||||
# 导致 staging 发版后 Ditto 口型服务静默降级到 GPU/MediaKit(P0 防复发)。
|
||||
USE_DITTO_LIPSYNC=true
|
||||
DITTO_API_BASE_URL=http://100.76.80.23:8000
|
||||
DITTO_DEFAULT_VIDEO_URL=https://xiaoxia-autocut.oss-cn-hangzhou.aliyuncs.com/uploads/default_avatar.mp4
|
||||
|
||||
@@ -56,7 +56,7 @@ class UserModel(Base):
|
||||
is_member = Column(Boolean, nullable=False, default=False)
|
||||
member_type = Column(String(20), nullable=True)
|
||||
member_expires_at = Column(DateTime, nullable=True)
|
||||
points_balance = Column(Float, nullable=False, default=0)
|
||||
points_balance = Column(Integer, nullable=False, default=0)
|
||||
created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(UTC))
|
||||
|
||||
|
||||
@@ -335,10 +335,6 @@ class GenerationTaskModel(Base):
|
||||
bgm_config = Column(JSON, nullable=False, default=dict)
|
||||
extra_meta = Column("metadata", JSON, nullable=False, default=dict)
|
||||
logs = Column(Text, nullable=False, default="[]", server_default="[]")
|
||||
# 功能计费(smart_edit):预扣积分 / 最终积分 / 预扣流水 ID
|
||||
credits_prepaid = Column(Float, nullable=False, default=0.0, server_default="0")
|
||||
credits_cost = Column(Float, nullable=False, default=0.0, server_default="0")
|
||||
credits_transaction_id = Column(String(36), nullable=False, default="", server_default="")
|
||||
created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(UTC))
|
||||
updated_at = Column(
|
||||
DateTime,
|
||||
@@ -731,11 +727,6 @@ class LipsyncJobModel(Base):
|
||||
# 精确句子时间戳(TTS 合成后由 silencedetect 计算,用于 B-roll 精确定位)
|
||||
sentence_timings = Column(JSON, nullable=True) # list[{index,text,start_time,end_time}]
|
||||
|
||||
# 功能计费(lip_sync):预扣积分 / 最终积分 / 预扣流水 ID
|
||||
credits_prepaid = Column(Float, nullable=False, default=0.0, server_default="0")
|
||||
credits_cost = Column(Float, nullable=False, default=0.0, server_default="0")
|
||||
credits_transaction_id = Column(String(36), nullable=False, default="", server_default="")
|
||||
|
||||
# 时间戳
|
||||
submitted_at = Column(DateTime, nullable=True)
|
||||
completed_at = Column(DateTime, nullable=True)
|
||||
@@ -784,9 +775,9 @@ class PointsAccountModel(Base):
|
||||
|
||||
id = Column(String(36), primary_key=True)
|
||||
user_id = Column(String(36), nullable=False, unique=True, index=True)
|
||||
balance = Column(Float, nullable=False, default=0)
|
||||
total_earned = Column(Float, nullable=False, default=0)
|
||||
total_spent = Column(Float, nullable=False, default=0)
|
||||
balance = Column(Integer, nullable=False, default=0)
|
||||
total_earned = Column(Integer, nullable=False, default=0)
|
||||
total_spent = Column(Integer, nullable=False, default=0)
|
||||
created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(UTC))
|
||||
updated_at = Column(DateTime, nullable=False, default=lambda: datetime.now(UTC))
|
||||
|
||||
@@ -801,8 +792,8 @@ class PointsTransactionModel(Base):
|
||||
account_id = Column(String(36), nullable=False, index=True)
|
||||
type = Column(String(20), nullable=False, index=True) # earn / spend / refund
|
||||
source = Column(String(50), nullable=False, index=True)
|
||||
amount = Column(Float, nullable=False)
|
||||
balance_after = Column(Float, nullable=False)
|
||||
amount = Column(Integer, nullable=False)
|
||||
balance_after = Column(Integer, nullable=False)
|
||||
description = Column(String(255), nullable=False, default="")
|
||||
ref_id = Column(String(100), nullable=False, default="")
|
||||
created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(UTC))
|
||||
@@ -914,11 +905,6 @@ class GpuLipsyncTaskModel(Base):
|
||||
# 心跳:worker 最近一次 poll/result 的时间,用于判定 worker 失联
|
||||
last_heartbeat_at = Column(DateTime, nullable=True)
|
||||
|
||||
# 功能计费(lip_sync):预扣积分 / 最终积分 / 预扣流水 ID
|
||||
credits_prepaid = Column(Float, nullable=False, default=0.0, server_default="0")
|
||||
credits_cost = Column(Float, nullable=False, default=0.0, server_default="0")
|
||||
credits_transaction_id = Column(String(36), nullable=False, default="", server_default="")
|
||||
|
||||
|
||||
class GpuWorkerModel(Base):
|
||||
"""GPU Worker 注册表 — 反向轮询模式下用于心跳与监控."""
|
||||
@@ -942,7 +928,6 @@ 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="")
|
||||
@@ -976,10 +961,7 @@ class ViralVideoJobModel(Base):
|
||||
JSON, nullable=True
|
||||
) # v1.6: 编导脚本结构{overview,scene_and_lighting,shots,hard_constraints,negative_prompts,voiceover_script}
|
||||
result_video_url = Column(String(1000), nullable=False, default="")
|
||||
credits_cost = Column(Float, nullable=False, default=0)
|
||||
video_resolution = Column(String(20), nullable=False, default="720p")
|
||||
credits_prepaid = Column(Float, nullable=False, default=0.0)
|
||||
credits_transaction_id = Column(String(36), nullable=False, default="")
|
||||
credits_cost = Column(Integer, nullable=False, default=0)
|
||||
error_msg = Column(Text, nullable=False, default="")
|
||||
retry_count = Column(Integer, nullable=False, default=0)
|
||||
started_at = Column(DateTime(timezone=True), nullable=True)
|
||||
@@ -1005,34 +987,16 @@ class ViralVideoStyleTemplateModel(Base):
|
||||
|
||||
|
||||
class ViralVideoPromptTemplateModel(Base):
|
||||
"""爆款视频 Prompt 模板表(#2040:纯文本 XML 标签模板,运营可直接编辑)"""
|
||||
"""爆款视频 Prompt 模板表(由 #2040 seed)"""
|
||||
|
||||
__tablename__ = "viral_video_prompt_templates"
|
||||
|
||||
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)
|
||||
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))
|
||||
|
||||
|
||||
class SystemSettingModel(Base):
|
||||
"""系统配置表 ORM 模型(#2246:后台可配置项,表已手工存在于 staging)."""
|
||||
|
||||
__tablename__ = "system_settings"
|
||||
|
||||
id = Column(String(36), primary_key=True)
|
||||
setting_key = Column(String(100), nullable=False, unique=True)
|
||||
setting_value = Column(Text, nullable=True)
|
||||
setting_type = Column(String(20), nullable=False)
|
||||
description = Column(String(255), nullable=False, default="", server_default="")
|
||||
is_public = Column(Boolean, nullable=False, default=False, server_default="false")
|
||||
updated_by = Column(String(36), nullable=True)
|
||||
category = Column(String(50), nullable=False, default="general", server_default="general")
|
||||
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)
|
||||
version = Column(Integer, nullable=False, default=1)
|
||||
is_active = Column(Boolean, nullable=False, default=True, index=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))
|
||||
|
||||
@@ -1,65 +0,0 @@
|
||||
"""system_settings 表 SQLAlchemy Repository — #2246."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.models import SystemSettingModel
|
||||
from packages.domain.system_setting import SystemSetting
|
||||
|
||||
|
||||
class SQLAlchemySystemSettingRepository:
|
||||
def __init__(self, session: Session):
|
||||
self.session = session
|
||||
|
||||
def get_by_key(self, setting_key: str) -> SystemSetting | None:
|
||||
model = self.session.query(SystemSettingModel).filter(SystemSettingModel.setting_key == setting_key).first()
|
||||
if model is None:
|
||||
return None
|
||||
return self._to_domain(model)
|
||||
|
||||
def list_all(self, category: str | None = None) -> list[SystemSetting]:
|
||||
query = self.session.query(SystemSettingModel)
|
||||
if category is not None:
|
||||
query = query.filter(SystemSettingModel.category == category)
|
||||
return [self._to_domain(m) for m in query.all()]
|
||||
|
||||
def upsert(self, setting: SystemSetting) -> SystemSetting:
|
||||
model = (
|
||||
self.session.query(SystemSettingModel).filter(SystemSettingModel.setting_key == setting.setting_key).first()
|
||||
)
|
||||
if model is None:
|
||||
model = SystemSettingModel(id=setting.id)
|
||||
self.session.add(model)
|
||||
model.setting_key = setting.setting_key
|
||||
model.setting_value = setting.setting_value
|
||||
model.setting_type = setting.setting_type
|
||||
model.description = setting.description
|
||||
model.is_public = setting.is_public
|
||||
model.category = setting.category
|
||||
model.updated_by = setting.updated_by
|
||||
self.session.commit()
|
||||
return setting
|
||||
|
||||
def delete_by_key(self, setting_key: str) -> bool:
|
||||
model = self.session.query(SystemSettingModel).filter(SystemSettingModel.setting_key == setting_key).first()
|
||||
if model is None:
|
||||
return False
|
||||
self.session.delete(model)
|
||||
self.session.commit()
|
||||
return True
|
||||
|
||||
@staticmethod
|
||||
def _to_domain(model: SystemSettingModel) -> SystemSetting:
|
||||
return SystemSetting(
|
||||
id=model.id,
|
||||
setting_key=model.setting_key,
|
||||
setting_value=model.setting_value,
|
||||
setting_type=model.setting_type,
|
||||
description=model.description or "",
|
||||
is_public=bool(model.is_public),
|
||||
category=model.category or "general",
|
||||
updated_by=model.updated_by,
|
||||
created_at=model.created_at,
|
||||
updated_at=model.updated_at,
|
||||
)
|
||||
@@ -39,9 +39,6 @@ class SQLAlchemyUserRepository(UserRepository):
|
||||
model.phone_verified = user.phone_verified
|
||||
model.binding_completed_at = user.binding_completed_at
|
||||
model.profile_completed = user.profile_completed
|
||||
model.is_member = user.is_member
|
||||
model.member_type = user.member_type
|
||||
model.member_expires_at = user.member_expires_at
|
||||
model.created_at = user.created_at
|
||||
|
||||
self.session.commit()
|
||||
@@ -118,8 +115,5 @@ class SQLAlchemyUserRepository(UserRepository):
|
||||
phone_verified=model.phone_verified or False,
|
||||
binding_completed_at=model.binding_completed_at,
|
||||
profile_completed=model.profile_completed if model.profile_completed is not None else True,
|
||||
is_member=model.is_member if model.is_member is not None else False,
|
||||
member_type=model.member_type,
|
||||
member_expires_at=model.member_expires_at,
|
||||
created_at=model.created_at,
|
||||
)
|
||||
|
||||
@@ -6,33 +6,18 @@ 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 → 领域实体。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
|
||||
"""ORM → 领域实体。"""
|
||||
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 "",
|
||||
@@ -61,10 +46,7 @@ def _to_domain(model: ViralVideoJobModel) -> ViralVideoJob:
|
||||
generated_copy_text=getattr(model, "generated_copy_text", "") or "",
|
||||
copy_result=dict(model.copy_result) if getattr(model, "copy_result", None) else None,
|
||||
result_video_url=model.result_video_url or "",
|
||||
video_resolution=getattr(model, "video_resolution", "720p") or "720p",
|
||||
credits_prepaid=float(getattr(model, "credits_prepaid", 0) or 0),
|
||||
credits_transaction_id=getattr(model, "credits_transaction_id", "") or "",
|
||||
credits_cost=float(model.credits_cost or 0),
|
||||
credits_cost=model.credits_cost or 0,
|
||||
error_msg=model.error_msg or "",
|
||||
retry_count=model.retry_count or 0,
|
||||
started_at=model.started_at,
|
||||
@@ -85,7 +67,6 @@ 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,
|
||||
@@ -114,10 +95,7 @@ class SQLAlchemyViralVideoJobRepository:
|
||||
generated_copy_text=job.generated_copy_text,
|
||||
copy_result=job.copy_result,
|
||||
result_video_url=job.result_video_url,
|
||||
video_resolution=getattr(job, "video_resolution", "720p") or "720p",
|
||||
credits_prepaid=float(getattr(job, "credits_prepaid", 0) or 0),
|
||||
credits_transaction_id=getattr(job, "credits_transaction_id", "") or "",
|
||||
credits_cost=float(getattr(job, "credits_cost", 0) or 0),
|
||||
credits_cost=job.credits_cost,
|
||||
error_msg=job.error_msg,
|
||||
retry_count=job.retry_count,
|
||||
started_at=job.started_at,
|
||||
@@ -142,12 +120,8 @@ 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)
|
||||
model.credits_transaction_id = getattr(job, "credits_transaction_id", "") or ""
|
||||
model.credits_cost = float(getattr(job, "credits_cost", 0) or 0)
|
||||
model.credits_cost = job.credits_cost
|
||||
model.error_msg = job.error_msg
|
||||
model.retry_count = job.retry_count
|
||||
model.started_at = job.started_at
|
||||
@@ -244,3 +218,31 @@ 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,
|
||||
}
|
||||
|
||||
@@ -3,24 +3,10 @@
|
||||
使用 bcrypt 安全存储密码
|
||||
"""
|
||||
|
||||
import hashlib
|
||||
from typing import Optional
|
||||
|
||||
import bcrypt
|
||||
|
||||
# bcrypt 只对前 72 字节有效,且 bcrypt>=4.1 会对超长输入直接抛 ValueError。
|
||||
# 超长密码先做一次 SHA-256(定长 hex),再交给 bcrypt,
|
||||
# 既绕过长度限制又保持对超长不同密码的区分度。
|
||||
_BCRYPT_MAX_BYTES = 72
|
||||
|
||||
|
||||
def _prepare_password_bytes(password: str) -> bytes:
|
||||
raw = password.encode("utf-8")
|
||||
if len(raw) > _BCRYPT_MAX_BYTES:
|
||||
return hashlib.sha256(raw).hexdigest().encode("utf-8")
|
||||
return raw
|
||||
|
||||
|
||||
from packages.domain.auth.password_hasher import PasswordHasherPort, PasswordValidatorPort
|
||||
|
||||
|
||||
@@ -56,8 +42,8 @@ class PasswordHasher(PasswordHasherPort):
|
||||
if not password:
|
||||
raise ValueError("Password cannot be empty")
|
||||
|
||||
# bcrypt 需要 bytes(超长密码先 SHA-256 以兼容 72 字节限制)
|
||||
password_bytes = _prepare_password_bytes(password)
|
||||
# bcrypt 需要 bytes
|
||||
password_bytes = password.encode("utf-8")
|
||||
|
||||
# 生成 salt 并哈希
|
||||
salt = bcrypt.gensalt(rounds=self.rounds)
|
||||
@@ -81,7 +67,7 @@ class PasswordHasher(PasswordHasherPort):
|
||||
return False
|
||||
|
||||
try:
|
||||
password_bytes = _prepare_password_bytes(password)
|
||||
password_bytes = password.encode("utf-8")
|
||||
hashed_bytes = hashed_password.encode("utf-8")
|
||||
|
||||
return bcrypt.checkpw(password_bytes, hashed_bytes)
|
||||
|
||||
@@ -1 +0,0 @@
|
||||
"""应用层:对外展示目录(套餐/积分包)。"""
|
||||
@@ -1,152 +0,0 @@
|
||||
"""读取管理后台配置的会员套餐 / 积分充值包(共享库真实数据)。
|
||||
|
||||
替代旧的硬编码 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)
|
||||
@@ -351,25 +351,12 @@ class CosyVoiceService:
|
||||
用于私有 bucket 下,将裸 URL 转为预签名 URL,
|
||||
确保 CosyVoice 服务器能下载参考音频.
|
||||
"""
|
||||
# 优先从 ai_router 获取 DB 配置
|
||||
_router_key, _router_url, _router_model = "", "", ""
|
||||
try:
|
||||
from packages.shared.ai_router import ai_router
|
||||
|
||||
tts_client = ai_router.get_tts_client("tts")
|
||||
if tts_client and tts_client.is_available:
|
||||
_router_key = tts_client.api_key
|
||||
_router_url = tts_client.base_url
|
||||
_router_model = tts_client.model
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
settings = get_shared_settings()
|
||||
|
||||
self._api_key = api_key or _router_key or settings.cosyvoice_api_key
|
||||
self._base_url = base_url or _router_url or settings.cosyvoice_base_url
|
||||
self._model = model or _router_model or settings.cosyvoice_model
|
||||
self._clone_model = clone_model or getattr(settings, "cosyvoice_clone_model", "")
|
||||
self._api_key = api_key or settings.cosyvoice_api_key
|
||||
self._base_url = base_url or settings.cosyvoice_base_url
|
||||
self._model = model or settings.cosyvoice_model
|
||||
self._clone_model = clone_model or getattr(settings, "cosyvoice_clone_model", "voice-enrollment")
|
||||
self._audio_url_signer = audio_url_signer
|
||||
|
||||
# base_url 规范化:去掉末尾的路径残留(兼容旧版配置)
|
||||
|
||||
@@ -1,332 +0,0 @@
|
||||
"""Ditto LLM 情绪分析服务 — #2076 后续:根据文案生成 emo_timeline.
|
||||
|
||||
职责:
|
||||
1. 正则按 。!?; 初步分句
|
||||
2. 调 DoubaoClient.chat_completion 分析每句表情(emo: 0-7, intensity: 0-1)
|
||||
3. 结果 LRU 缓存(文案 hash → 情绪列表)
|
||||
4. LLM 失败/超时/格式错 → 返回空列表(降级中性表情,不阻塞生成)
|
||||
5. TTS 完成后按字数比例或 sentence_timings 对齐成秒级 timeline
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import json
|
||||
import logging
|
||||
import re
|
||||
from functools import lru_cache
|
||||
from pathlib import Path
|
||||
from typing import Any, Optional
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# ── 表情常量 ─────────────────────────────────────────────────────
|
||||
EMO_ANGER = 0
|
||||
EMO_DISGUST = 1
|
||||
EMO_FEAR = 2
|
||||
EMO_HAPPY = 3
|
||||
EMO_NEUTRAL = 4
|
||||
EMO_SAD = 5
|
||||
EMO_SURPRISE = 6
|
||||
EMO_CONTEMPT = 7
|
||||
ALLOWED_EMOS = {EMO_HAPPY, EMO_NEUTRAL, EMO_SAD, EMO_SURPRISE} # 营销场景白名单
|
||||
|
||||
# ── 分句正则 ─────────────────────────────────────────────────────
|
||||
_SENT_SPLIT_RE = re.compile(r"(?<=[。!?;!?;])\s*")
|
||||
|
||||
# ── 默认 prompt 模板文件路径 ──────────────────────────────────────
|
||||
_DEFAULT_PROMPT_PATH = Path(__file__).parent / "prompts" / "ditto_emotion.txt"
|
||||
|
||||
|
||||
def _load_default_prompt() -> str:
|
||||
try:
|
||||
return _DEFAULT_PROMPT_PATH.read_text(encoding="utf-8").strip()
|
||||
except Exception:
|
||||
# 文件不存在时用极简兜底
|
||||
return (
|
||||
"分析文案每句话表情,输出JSON数组:"
|
||||
'[{"text":"句子","emo":4,"intensity":0.2}],emo:3开心4中性5伤心6惊讶,'
|
||||
"禁止0/1/2/7。\n【文案】\n{文案}"
|
||||
)
|
||||
|
||||
|
||||
# ── 数据结构 ─────────────────────────────────────────────────────
|
||||
class EmotionSegment:
|
||||
"""单句情绪结果(LLM 输出的原始结构)."""
|
||||
|
||||
__slots__ = ("text", "emo", "intensity")
|
||||
|
||||
def __init__(self, text: str, emo: int, intensity: float):
|
||||
self.text = text
|
||||
self.emo = emo
|
||||
self.intensity = intensity
|
||||
|
||||
def to_dict(self) -> dict[str, Any]:
|
||||
return {"text": self.text, "emo": self.emo, "intensity": self.intensity}
|
||||
|
||||
|
||||
class EmotionTimelineEntry:
|
||||
"""对齐到音频时间轴后的情绪片段(传给 Ditto)."""
|
||||
|
||||
__slots__ = ("start", "end", "emo", "intensity")
|
||||
|
||||
def __init__(self, start: float, end: float, emo: int, intensity: float):
|
||||
self.start = round(start, 2)
|
||||
self.end = round(end, 2)
|
||||
self.emo = emo
|
||||
self.intensity = round(intensity, 2)
|
||||
|
||||
def to_dict(self) -> dict[str, Any]:
|
||||
return {
|
||||
"start": self.start,
|
||||
"end": self.end,
|
||||
"emo": self.emo,
|
||||
"intensity": self.intensity,
|
||||
}
|
||||
|
||||
|
||||
# ── 分句 ─────────────────────────────────────────────────────────
|
||||
def split_sentences(text: str) -> list[str]:
|
||||
"""按中文句末标点切分,过滤空串."""
|
||||
if not text:
|
||||
return []
|
||||
parts = _SENT_SPLIT_RE.split(text.strip())
|
||||
return [p.strip() for p in parts if p and p.strip()]
|
||||
|
||||
|
||||
# ── 解析 LLM 返回的 JSON ─────────────────────────────────────────
|
||||
def _parse_emotion_json(raw: str) -> list[EmotionSegment]:
|
||||
"""解析 LLM 返回,容错处理:
|
||||
- 去掉 markdown 代码块包裹
|
||||
- 只取第一个 JSON 数组
|
||||
- 逐行校验 emo/intensity 合法性,过滤无效项
|
||||
"""
|
||||
if not raw:
|
||||
return []
|
||||
text = raw.strip()
|
||||
# 去掉 ```json ... ``` 包裹
|
||||
if text.startswith("```"):
|
||||
text = re.sub(r"^```(?:json)?\s*", "", text)
|
||||
text = re.sub(r"\s*```$", "", text)
|
||||
# 找第一个 [ 到最后一个 ]
|
||||
lb = text.find("[")
|
||||
rb = text.rfind("]")
|
||||
if lb == -1 or rb == -1 or rb <= lb:
|
||||
return []
|
||||
try:
|
||||
data = json.loads(text[lb : rb + 1])
|
||||
except (json.JSONDecodeError, ValueError):
|
||||
return []
|
||||
if not isinstance(data, list):
|
||||
return []
|
||||
|
||||
results: list[EmotionSegment] = []
|
||||
for item in data:
|
||||
if not isinstance(item, dict):
|
||||
continue
|
||||
try:
|
||||
emo = int(item.get("emo", EMO_NEUTRAL))
|
||||
intensity = float(item.get("intensity", 0.2))
|
||||
except (TypeError, ValueError):
|
||||
continue
|
||||
if emo not in ALLOWED_EMOS:
|
||||
emo = EMO_NEUTRAL
|
||||
intensity = max(0.05, min(1.0, intensity))
|
||||
sent_text = str(item.get("text", "")).strip()
|
||||
if not sent_text:
|
||||
continue
|
||||
results.append(EmotionSegment(text=sent_text, emo=emo, intensity=intensity))
|
||||
return results
|
||||
|
||||
|
||||
# ── 时间对齐(按字数比例)────────────────────────────────────────
|
||||
def align_timeline_by_length(
|
||||
segments: list[EmotionSegment],
|
||||
audio_duration: float,
|
||||
) -> list[EmotionTimelineEntry]:
|
||||
"""按各句字数占总字数比例分配 audio_duration 时长."""
|
||||
if not segments or audio_duration <= 0:
|
||||
return []
|
||||
total_chars = sum(len(s.text) for s in segments)
|
||||
if total_chars <= 0:
|
||||
return []
|
||||
entries: list[EmotionTimelineEntry] = []
|
||||
pos = 0.0
|
||||
for i, seg in enumerate(segments):
|
||||
if i == len(segments) - 1:
|
||||
end = audio_duration # 最后一段到结尾,避免浮点误差
|
||||
else:
|
||||
end = pos + (len(seg.text) / total_chars) * audio_duration
|
||||
if end > pos:
|
||||
entries.append(
|
||||
EmotionTimelineEntry(
|
||||
start=pos,
|
||||
end=end,
|
||||
emo=seg.emo,
|
||||
intensity=seg.intensity,
|
||||
)
|
||||
)
|
||||
pos = end
|
||||
return entries
|
||||
|
||||
|
||||
def align_timeline_by_timings(
|
||||
segments: list[EmotionSegment],
|
||||
sentence_timings: list[dict[str, Any]],
|
||||
audio_duration: float,
|
||||
) -> list[EmotionTimelineEntry]:
|
||||
"""使用 TTS sentence_timings 精确对齐(优先方案).
|
||||
|
||||
sentence_timings 格式:[{"start":0.0,"end":1.2,"text":"句子"}, ...]
|
||||
按句序匹配 segments 和 timings,长度不一致时回退到按字数比例。
|
||||
"""
|
||||
if not sentence_timings or len(sentence_timings) != len(segments):
|
||||
return align_timeline_by_length(segments, audio_duration)
|
||||
entries: list[EmotionTimelineEntry] = []
|
||||
for seg, timing in zip(segments, sentence_timings, strict=False):
|
||||
try:
|
||||
start = float(timing.get("start", 0))
|
||||
end = float(timing.get("end", 0))
|
||||
except (TypeError, ValueError):
|
||||
return align_timeline_by_length(segments, audio_duration)
|
||||
if end <= start:
|
||||
continue
|
||||
entries.append(
|
||||
EmotionTimelineEntry(
|
||||
start=start,
|
||||
end=end,
|
||||
emo=seg.emo,
|
||||
intensity=seg.intensity,
|
||||
)
|
||||
)
|
||||
return entries
|
||||
|
||||
|
||||
# ── LLM 情绪分析服务 ─────────────────────────────────────────────
|
||||
class DittoEmotionService:
|
||||
"""Ditto 情绪分析服务(带 LRU 缓存)."""
|
||||
|
||||
def __init__(self, settings=None):
|
||||
from packages.config import get_api_settings
|
||||
|
||||
self.settings = settings or get_api_settings()
|
||||
self._client = None
|
||||
|
||||
def _cfg(self, key: str) -> Any:
|
||||
"""优先读后台 system_config,未配置则回退到 settings(env 默认)."""
|
||||
try:
|
||||
from packages.application.system_config_service import get_config
|
||||
|
||||
return get_config(key, getattr(self.settings, key, None))
|
||||
except Exception:
|
||||
return getattr(self.settings, key, None)
|
||||
|
||||
@property
|
||||
def enabled(self) -> bool:
|
||||
return bool(self._cfg("ditto_emotion_enabled"))
|
||||
|
||||
def _get_prompt_template(self) -> str:
|
||||
"""优先用配置(环境变量),否则读文件."""
|
||||
cfg_prompt = self._cfg("ditto_emotion_prompt") or ""
|
||||
if cfg_prompt.strip():
|
||||
return cfg_prompt.strip()
|
||||
return _load_default_prompt()
|
||||
|
||||
def _cache_key(self, text: str) -> str:
|
||||
return hashlib.md5(text.strip().encode("utf-8")).hexdigest()
|
||||
|
||||
def _get_llm_client(self):
|
||||
if self._client is None:
|
||||
from packages.shared.ai_client import get_doubao_client
|
||||
|
||||
self._client = get_doubao_client()
|
||||
return self._client
|
||||
|
||||
def _call_llm(self, text: str) -> list[EmotionSegment]:
|
||||
"""调 LLM 分析情绪,失败返回空列表."""
|
||||
template = self._get_prompt_template()
|
||||
prompt = template.replace("{文案}", text)
|
||||
messages = [{"role": "user", "content": prompt}]
|
||||
model = self._cfg("ditto_emotion_model") or None
|
||||
temperature = self._cfg("ditto_emotion_temperature")
|
||||
timeout = getattr(self.settings, "ditto_emotion_timeout", 10)
|
||||
max_tokens = getattr(self.settings, "ditto_emotion_max_tokens", 1024)
|
||||
try:
|
||||
client = self._get_llm_client()
|
||||
result = client.chat_completion(
|
||||
messages=messages,
|
||||
model=model,
|
||||
temperature=temperature,
|
||||
max_tokens=max_tokens,
|
||||
timeout=timeout,
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.warning("[ditto_emotion] LLM 调用异常: %s", exc)
|
||||
return []
|
||||
if not result:
|
||||
return []
|
||||
segments = _parse_emotion_json(result)
|
||||
if not segments:
|
||||
logger.warning("[ditto_emotion] LLM 返回解析失败: %s", result[:200])
|
||||
return segments
|
||||
|
||||
def analyze(self, text: str) -> list[EmotionSegment]:
|
||||
"""分析文案情绪(带缓存),失败返回空列表."""
|
||||
if not self.enabled or not text or not text.strip():
|
||||
return []
|
||||
key = self._cache_key(text)
|
||||
return _cached_analyze(self, key, text)
|
||||
|
||||
def build_timeline(
|
||||
self,
|
||||
text: str,
|
||||
audio_duration: float,
|
||||
sentence_timings: Optional[list[dict[str, Any]]] = None,
|
||||
) -> str:
|
||||
"""完整流程:分句→LLM分析→时间对齐→序列化为JSON字符串.
|
||||
|
||||
返回: JSON 字符串(可直接传 Ditto emo_timeline 参数);空字符串表示降级中性。
|
||||
"""
|
||||
segments = self.analyze(text)
|
||||
if not segments:
|
||||
return ""
|
||||
if sentence_timings:
|
||||
entries = align_timeline_by_timings(segments, sentence_timings, audio_duration)
|
||||
else:
|
||||
entries = align_timeline_by_length(segments, audio_duration)
|
||||
if not entries:
|
||||
return ""
|
||||
return json.dumps([e.to_dict() for e in entries], ensure_ascii=False)
|
||||
|
||||
|
||||
# ── 模块级 LRU 缓存实例 ─────────────────────────────────────────
|
||||
# 每个 service 实例共享缓存(按 cache_key 区分)
|
||||
@lru_cache(maxsize=512)
|
||||
def _cached_analyze(service: DittoEmotionService, cache_key: str, text: str) -> list[EmotionSegment]:
|
||||
"""LRU 缓存包装:cache_key 由文案 hash 生成,maxsize 从配置读."""
|
||||
# 注意:service 参数仅用于传递调用,缓存由 cache_key 驱动
|
||||
segments = service._call_llm(text)
|
||||
# 如果 LLM 返回空(比如分句数量不匹配),尝试直接对预分句结果分析
|
||||
if not segments:
|
||||
pre_splits = split_sentences(text)
|
||||
if len(pre_splits) > 1:
|
||||
# 用预分句结果兜底:全中性低强度
|
||||
segments = [EmotionSegment(text=s, emo=EMO_NEUTRAL, intensity=0.1) for s in pre_splits]
|
||||
return segments
|
||||
|
||||
|
||||
_singleton: Optional[DittoEmotionService] = None
|
||||
|
||||
|
||||
def get_ditto_emotion_service() -> DittoEmotionService:
|
||||
global _singleton
|
||||
if _singleton is None:
|
||||
_singleton = DittoEmotionService()
|
||||
return _singleton
|
||||
|
||||
|
||||
def reset_ditto_emotion_service() -> None:
|
||||
"""#2246:后台配置变更后重置单例并清空 LLM 结果 LRU 缓存."""
|
||||
global _singleton
|
||||
_singleton = None
|
||||
_cached_analyze.cache_clear()
|
||||
@@ -1,302 +0,0 @@
|
||||
"""蚂蚁 Ditto 数字人口型 API 客户端 — #2076.
|
||||
|
||||
封装 Ditto FastAPI(部署在 5060Ti GPU 节点,Tailscale 内网可达):
|
||||
- GET /health 健康检查
|
||||
- POST /generate 生成口型视频(同步返回 MP4 流)
|
||||
|
||||
关键特性:
|
||||
- 入参:video_url(人物模板视频 URL) + audio_url(TTS 音频 URL) + script(文案原文)
|
||||
- 出参:直接返回 video/mp4 字节流(自带音频,无需二次混流)
|
||||
- 429 时指数退避重试(最多 ditto_max_retries 次)
|
||||
- 500/超时视为失败
|
||||
- 输出 MP4 字节流转存到自家 OSS,返回公网 URL
|
||||
|
||||
注意:
|
||||
- 保留 MuseTalk/GPU 路径不变;本服务作为更高优先级的第三条口型路径
|
||||
- 不传 emotion/表情精细控制,使用默认 emo_global=4(中性)+ use_script_emo=true(关键词驱动表情)
|
||||
- Ditto 输出自带音视频,不需要 GFPGAN 超分,不需要 ffmpeg 音视频混流
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import io
|
||||
import logging
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
from typing import Optional
|
||||
|
||||
import httpx
|
||||
|
||||
from packages.config import get_api_settings
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class DittoError(Exception):
|
||||
"""Ditto API 调用失败."""
|
||||
|
||||
def __init__(self, message: str, code: str = "DittoError", status_code: int = 0):
|
||||
self.code = code
|
||||
self.status_code = status_code
|
||||
super().__init__(message)
|
||||
|
||||
|
||||
@dataclass
|
||||
class DittoResult:
|
||||
"""Ditto 生成结果."""
|
||||
|
||||
video_bytes: bytes
|
||||
video_url: str = "" # 转存 OSS 后填充
|
||||
elapsed_seconds: float = 0.0
|
||||
rtf: float = 0.0 # 实时率(响应头 X-RTF)
|
||||
frames: int = 0 # 帧数(响应头 X-Frames)
|
||||
|
||||
|
||||
class DittoClient:
|
||||
"""蚂蚁 Ditto 数字人口型 API 客户端."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
base_url: Optional[str] = None,
|
||||
default_video_url: Optional[str] = None,
|
||||
max_retries: Optional[int] = None,
|
||||
timeout: Optional[int] = None,
|
||||
):
|
||||
s = get_api_settings()
|
||||
self.base_url = (base_url or s.ditto_api_base_url or "").rstrip("/")
|
||||
self.default_video_url = default_video_url or s.ditto_default_video_url or ""
|
||||
self.max_retries = int(max_retries if max_retries is not None else s.ditto_max_retries)
|
||||
self.timeout = int(timeout if timeout is not None else s.ditto_request_timeout)
|
||||
try:
|
||||
from packages.application.system_config_service import get_config
|
||||
|
||||
self.blend_frames = int(get_config("ditto_blend_frames", s.ditto_blend_frames))
|
||||
except Exception:
|
||||
self.blend_frames = int(s.ditto_blend_frames)
|
||||
|
||||
@property
|
||||
def is_configured(self) -> bool:
|
||||
"""配置是否完整(base_url + 默认模板视频都有值)."""
|
||||
return bool(self.base_url) and bool(self.default_video_url)
|
||||
|
||||
def health(self) -> bool:
|
||||
"""健康检查;成功返回 True,失败返回 False(不抛异常)."""
|
||||
if not self.base_url:
|
||||
return False
|
||||
url = f"{self.base_url}/health"
|
||||
try:
|
||||
with httpx.Client(timeout=5.0) as client:
|
||||
resp = client.get(url)
|
||||
ok = resp.status_code == 200
|
||||
if ok:
|
||||
logger.info("[ditto] health check OK: %s", url)
|
||||
else:
|
||||
logger.warning("[ditto] health check status=%d: %s", resp.status_code, url)
|
||||
return ok
|
||||
except Exception as exc:
|
||||
logger.warning("[ditto] health check failed: %s", exc)
|
||||
return False
|
||||
|
||||
def generate(
|
||||
self,
|
||||
*,
|
||||
audio_url: str,
|
||||
script: str,
|
||||
video_url: Optional[str] = None,
|
||||
emo_global: int = 4,
|
||||
use_script_emo: bool = True,
|
||||
blend_frames: Optional[int] = None,
|
||||
emo_timeline: str = "",
|
||||
) -> DittoResult:
|
||||
"""调用 Ditto /generate 接口,返回 MP4 字节流结果.
|
||||
|
||||
Raises DittoError on failure.
|
||||
"""
|
||||
if not self.base_url:
|
||||
raise DittoError("DITTO_API_BASE_URL 未配置", code="ConfigMissing")
|
||||
driver_url = video_url or self.default_video_url
|
||||
if not driver_url:
|
||||
raise DittoError("Ditto 人物模板视频 URL 未配置", code="ConfigMissing")
|
||||
if not audio_url:
|
||||
raise DittoError("audio_url 不能为空", code="InvalidParam")
|
||||
if not script:
|
||||
script = " "
|
||||
|
||||
_blend = blend_frames if blend_frames is not None else self.blend_frames
|
||||
payload = {
|
||||
"video_url": driver_url,
|
||||
"audio_url": audio_url,
|
||||
"script": script,
|
||||
"emo_global": emo_global,
|
||||
"use_script_emo": use_script_emo,
|
||||
"blend_frames": _blend,
|
||||
}
|
||||
if emo_timeline:
|
||||
payload["emo_timeline"] = emo_timeline
|
||||
url = f"{self.base_url}/generate"
|
||||
|
||||
last_exc: Optional[Exception] = None
|
||||
for attempt in range(self.max_retries + 1):
|
||||
try:
|
||||
start = time.monotonic()
|
||||
# 精细化超时:connect=10s(网络不通快速失败),read=120s(最长音频~45s按RTF=2.8推算)
|
||||
_timeout = httpx.Timeout(connect=10.0, read=self.timeout, write=30.0, pool=10.0)
|
||||
with httpx.Client(timeout=_timeout, follow_redirects=True) as client:
|
||||
resp = client.post(url, json=payload)
|
||||
elapsed = time.monotonic() - start
|
||||
|
||||
if resp.status_code == 429:
|
||||
wait = min(2**attempt, 30)
|
||||
logger.warning(
|
||||
"[ditto] GPU 繁忙 (429),%ds 后重试 (%d/%d)",
|
||||
wait,
|
||||
attempt + 1,
|
||||
self.max_retries,
|
||||
)
|
||||
if attempt >= self.max_retries:
|
||||
raise DittoError(
|
||||
f"Ditto GPU 繁忙,重试 {self.max_retries} 次仍失败",
|
||||
code="BusyRetriesExhausted",
|
||||
status_code=429,
|
||||
)
|
||||
time.sleep(wait)
|
||||
continue
|
||||
|
||||
if resp.status_code != 200:
|
||||
_text = (resp.text or "")[:300]
|
||||
logger.error(
|
||||
"[ditto] generate 失败 status=%d attempt=%d body=%s",
|
||||
resp.status_code,
|
||||
attempt + 1,
|
||||
_text,
|
||||
)
|
||||
if resp.status_code >= 500 and attempt < self.max_retries:
|
||||
time.sleep(min(2**attempt, 15))
|
||||
continue
|
||||
raise DittoError(
|
||||
f"Ditto 返回 {resp.status_code}: {_text}",
|
||||
code="DittoAPIError",
|
||||
status_code=resp.status_code,
|
||||
)
|
||||
|
||||
video_bytes = resp.content
|
||||
if not video_bytes or len(video_bytes) < 1024:
|
||||
raise DittoError(
|
||||
f"Ditto 返回内容异常(size={len(video_bytes) if video_bytes else 0})",
|
||||
code="EmptyResponse",
|
||||
)
|
||||
try:
|
||||
rtf = float(resp.headers.get("X-RTF", "0") or 0)
|
||||
except ValueError:
|
||||
rtf = 0.0
|
||||
try:
|
||||
frames = int(resp.headers.get("X-Frames", "0") or 0)
|
||||
except ValueError:
|
||||
frames = 0
|
||||
try:
|
||||
x_time = float(resp.headers.get("X-Time", "0") or 0)
|
||||
if x_time > 0:
|
||||
elapsed = x_time
|
||||
except ValueError:
|
||||
pass
|
||||
|
||||
logger.info(
|
||||
"[ditto] generate 成功 size=%d rtf=%.2f frames=%d elapsed=%.1fs attempt=%d",
|
||||
len(video_bytes),
|
||||
rtf,
|
||||
frames,
|
||||
elapsed,
|
||||
attempt + 1,
|
||||
)
|
||||
return DittoResult(
|
||||
video_bytes=video_bytes,
|
||||
elapsed_seconds=elapsed,
|
||||
rtf=rtf,
|
||||
frames=frames,
|
||||
)
|
||||
|
||||
except DittoError:
|
||||
raise
|
||||
except (httpx.ConnectError, httpx.NetworkError, ConnectionError, OSError) as exc:
|
||||
# 网络不通/连接被拒(如 GPU 断网/Tailscale 掉线),不重试,直接快速回退
|
||||
logger.warning("[ditto] 网络不可达 attempt=%d err=%s", attempt + 1, exc)
|
||||
raise DittoError(
|
||||
f"Ditto 网络不可达: {exc}",
|
||||
code="NetworkUnreachable",
|
||||
) from exc
|
||||
except httpx.TimeoutException as exc:
|
||||
last_exc = exc
|
||||
logger.warning("[ditto] 请求超时 attempt=%d err=%s", attempt + 1, exc)
|
||||
if attempt < self.max_retries:
|
||||
time.sleep(min(2**attempt, 15))
|
||||
continue
|
||||
raise DittoError(
|
||||
f"Ditto 请求超时(read={self.timeout}s),重试耗尽",
|
||||
code="Timeout",
|
||||
) from exc
|
||||
except Exception as exc:
|
||||
last_exc = exc
|
||||
logger.warning("[ditto] 请求异常 attempt=%d err=%s", attempt + 1, exc)
|
||||
if attempt < self.max_retries:
|
||||
time.sleep(min(2**attempt, 10))
|
||||
continue
|
||||
raise DittoError(f"Ditto 调用异常: {exc}", code="NetworkError") from exc
|
||||
|
||||
raise DittoError("Ditto 未知错误", code="Unknown") from last_exc
|
||||
|
||||
def generate_and_persist(
|
||||
self,
|
||||
*,
|
||||
job_id: str,
|
||||
user_id: str,
|
||||
audio_url: str,
|
||||
script: str,
|
||||
video_url: Optional[str] = None,
|
||||
emo_timeline: str = "",
|
||||
blend_frames: Optional[int] = None,
|
||||
) -> DittoResult:
|
||||
"""调用 generate 并把 MP4 转存到自家 OSS,返回带 video_url 的结果."""
|
||||
result = self.generate(
|
||||
audio_url=audio_url,
|
||||
script=script,
|
||||
video_url=video_url,
|
||||
emo_timeline=emo_timeline,
|
||||
blend_frames=blend_frames,
|
||||
)
|
||||
try:
|
||||
from packages.shared.storage import get_shared_storage_service
|
||||
|
||||
storage = get_shared_storage_service()
|
||||
storage_key = f"ditto-output/{user_id}/{job_id}.mp4"
|
||||
public_url = storage.upload_file(
|
||||
io.BytesIO(result.video_bytes),
|
||||
storage_key,
|
||||
content_type="video/mp4",
|
||||
)
|
||||
result.video_url = public_url
|
||||
logger.info(
|
||||
"[ditto] 转存 OSS 完成 job=%s key=%s",
|
||||
job_id,
|
||||
storage_key,
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.error("[ditto] 转存 OSS 失败 job=%s err=%s", job_id, exc, exc_info=True)
|
||||
raise DittoError(f"Ditto 结果转存 OSS 失败: {exc}", code="StorageError") from exc
|
||||
return result
|
||||
|
||||
|
||||
_ditto_client_singleton: Optional[DittoClient] = None
|
||||
|
||||
|
||||
def get_ditto_client() -> DittoClient:
|
||||
"""获取 DittoClient 单例(简易工厂,便于单测 mock)."""
|
||||
global _ditto_client_singleton
|
||||
if _ditto_client_singleton is None:
|
||||
_ditto_client_singleton = DittoClient()
|
||||
return _ditto_client_singleton
|
||||
|
||||
|
||||
def reset_ditto_client() -> None:
|
||||
"""#2246:后台配置变更后重置 DittoClient 单例."""
|
||||
global _ditto_client_singleton
|
||||
_ditto_client_singleton = None
|
||||
@@ -1,27 +0,0 @@
|
||||
你是一个数字人视频表情导演。给定一段口播文案,分析每句话应该用什么表情和强度,让数字人说话时表情自然有变化,不僵硬。
|
||||
【表情编号】
|
||||
0=愤怒(营销场景禁用)
|
||||
1=厌恶(禁用)
|
||||
2=害怕(禁用)
|
||||
3=开心:介绍优点、优惠、好消息、号召行动时用
|
||||
4=中性:默认表情,陈述事实、平铺直叙时用
|
||||
5=伤心:仅在共情用户痛点时低强度使用(如"是不是经常遇到…")
|
||||
6=惊讶:惊喜、意外、强调价值时用(如"居然""只要""竟然")
|
||||
7=轻蔑(禁用)
|
||||
【强度说明】
|
||||
0.1-0.2:几乎看不出变化,比中性多一点情绪色彩
|
||||
0.3-0.4:有明显但自然的情绪,正常说话的波动
|
||||
0.5-0.6:较强情绪,感叹句/重点强调
|
||||
0.7+:极强情绪,极少使用
|
||||
【规则】
|
||||
1. 按自然语义分句,以。!?;为主要分界,逗号不分
|
||||
2. 60-70%的句子应该用中性(4),不要每句都标情绪
|
||||
3. 情绪和内容匹配:卖点→开心(3),痛点共情→伤心(5)低强度,惊喜/划算→惊讶(6),陈述→中性(4)
|
||||
4. 相邻句子情绪不要剧烈跳变
|
||||
5. 感叹号结尾强度0.4-0.6,句号结尾一般0.1-0.3
|
||||
6. 开头结尾句用中性(4)或低强度开心(3)
|
||||
7. 禁止使用0/1/2/7
|
||||
【输出格式】严格JSON数组,不要输出其他内容
|
||||
[{"text":"句子原文","emo":3,"intensity":0.4}]
|
||||
【文案】
|
||||
{文案}
|
||||
@@ -1,189 +0,0 @@
|
||||
"""系统配置应用服务 — #2246.
|
||||
|
||||
- get_config(key, default) / set_config(...) / list_configs(category)
|
||||
- 首次访问时从 DB 加载并进程内缓存;set_config 后失效缓存
|
||||
- set_config 成功后重置 Ditto 情绪服务与 Ditto 客户端单例(含 LRU 缓存),
|
||||
保证后台改动立即生效;worker 不直接 HTTP 读配置,统一通过本模块。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import threading
|
||||
from typing import Any
|
||||
|
||||
from packages.adapters.sqlalchemy_impl import session as db_session
|
||||
from packages.adapters.sqlalchemy_impl.system_setting_repository import (
|
||||
SQLAlchemySystemSettingRepository,
|
||||
)
|
||||
from packages.domain.system_setting import (
|
||||
SystemSetting,
|
||||
infer_setting_type,
|
||||
serialize_setting_value,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class SystemConfigService:
|
||||
"""进程内缓存的系统配置服务."""
|
||||
|
||||
def __init__(self, session_factory=None):
|
||||
self._session_factory = session_factory
|
||||
self._cache: dict[str, Any] | None = None
|
||||
self._lock = threading.Lock()
|
||||
|
||||
# ── 会话 ─────────────────────────────────────────────────────
|
||||
def _get_session_factory(self):
|
||||
# 必须运行时读取模块属性:模块导入时 SessionLocal 还是 None,
|
||||
# initialize_database() 之后才被赋值,import 时绑定会拿到旧值。
|
||||
factory = self._session_factory or db_session.SessionLocal
|
||||
if factory is None:
|
||||
raise RuntimeError("数据库会话工厂未初始化")
|
||||
return factory
|
||||
|
||||
# ── 缓存 ─────────────────────────────────────────────────────
|
||||
def _load_cache(self) -> dict[str, Any]:
|
||||
with self._lock:
|
||||
if self._cache is not None:
|
||||
return self._cache
|
||||
cache: dict[str, Any] = {}
|
||||
factory = self._get_session_factory()
|
||||
session = factory()
|
||||
try:
|
||||
repo = SQLAlchemySystemSettingRepository(session)
|
||||
for setting in repo.list_all():
|
||||
try:
|
||||
cache[setting.setting_key] = setting.get_typed_value()
|
||||
except Exception as exc: # 损坏配置不阻塞启动
|
||||
logger.warning(
|
||||
"[system_config] 跳过损坏配置 %s: %s",
|
||||
setting.setting_key,
|
||||
exc,
|
||||
)
|
||||
finally:
|
||||
session.close()
|
||||
self._cache = cache
|
||||
return cache
|
||||
|
||||
def reload(self) -> dict[str, Any]:
|
||||
"""强制重新从 DB 加载配置,返回新缓存."""
|
||||
with self._lock:
|
||||
self._cache = None
|
||||
return self._load_cache()
|
||||
|
||||
def invalidate(self) -> None:
|
||||
with self._lock:
|
||||
self._cache = None
|
||||
|
||||
# ── 读 ───────────────────────────────────────────────────────
|
||||
def get_config(self, key: str, default: Any = None) -> Any:
|
||||
cache = self._load_cache()
|
||||
return cache.get(key, default)
|
||||
|
||||
def list_configs(self, category: str | None = None) -> list[SystemSetting]:
|
||||
factory = self._get_session_factory()
|
||||
session = factory()
|
||||
try:
|
||||
return SQLAlchemySystemSettingRepository(session).list_all(category)
|
||||
finally:
|
||||
session.close()
|
||||
|
||||
# ── 写 ───────────────────────────────────────────────────────
|
||||
def set_config(
|
||||
self,
|
||||
key: str,
|
||||
value: Any,
|
||||
setting_type: str | None = None,
|
||||
updated_by: str | None = None,
|
||||
*,
|
||||
description: str | None = None,
|
||||
category: str = "general",
|
||||
is_public: bool = False,
|
||||
) -> Any:
|
||||
st = setting_type or infer_setting_type(value)
|
||||
serialized = serialize_setting_value(value, st)
|
||||
|
||||
factory = self._get_session_factory()
|
||||
session = factory()
|
||||
try:
|
||||
repo = SQLAlchemySystemSettingRepository(session)
|
||||
existing = repo.get_by_key(key)
|
||||
if existing is not None:
|
||||
existing.setting_value = serialized
|
||||
existing.setting_type = st
|
||||
if updated_by is not None:
|
||||
existing.updated_by = updated_by
|
||||
if description is not None:
|
||||
existing.description = description
|
||||
setting = existing
|
||||
else:
|
||||
setting = SystemSetting(
|
||||
setting_key=key,
|
||||
setting_value=serialized,
|
||||
setting_type=st,
|
||||
description=description or "",
|
||||
is_public=is_public,
|
||||
category=category,
|
||||
updated_by=updated_by,
|
||||
)
|
||||
repo.upsert(setting)
|
||||
finally:
|
||||
session.close()
|
||||
|
||||
self.invalidate()
|
||||
self._reset_ditto_singletons(key)
|
||||
return value
|
||||
|
||||
def delete_config(self, key: str) -> bool:
|
||||
factory = self._get_session_factory()
|
||||
session = factory()
|
||||
try:
|
||||
deleted = SQLAlchemySystemSettingRepository(session).delete_by_key(key)
|
||||
finally:
|
||||
session.close()
|
||||
if deleted:
|
||||
self.invalidate()
|
||||
self._reset_ditto_singletons(key)
|
||||
return deleted
|
||||
|
||||
# ── Ditto 联动 ───────────────────────────────────────────────
|
||||
@staticmethod
|
||||
def _reset_ditto_singletons(key: str) -> None:
|
||||
try:
|
||||
from packages.application import ditto_emotion_service as emo_mod
|
||||
from packages.application import ditto_service as ditto_mod
|
||||
|
||||
emo_mod.reset_ditto_emotion_service()
|
||||
ditto_mod.reset_ditto_client()
|
||||
logger.info("[system_config] %s 更新,已重置 Ditto 单例", key)
|
||||
except Exception as exc: # 联动失败不影响配置落库
|
||||
logger.warning("[system_config] 重置 Ditto 单例失败: %s", exc)
|
||||
|
||||
|
||||
_service_singleton: SystemConfigService | None = None
|
||||
|
||||
|
||||
def get_system_config_service() -> SystemConfigService:
|
||||
global _service_singleton
|
||||
if _service_singleton is None:
|
||||
_service_singleton = SystemConfigService()
|
||||
return _service_singleton
|
||||
|
||||
|
||||
def get_config(key: str, default: Any = None) -> Any:
|
||||
"""便捷读取:优先 DB 配置,未配置时返回 default(调用方传 env 值兜底)."""
|
||||
try:
|
||||
return get_system_config_service().get_config(key, default)
|
||||
except Exception as exc:
|
||||
logger.warning("[system_config] 读取 %s 失败,使用默认值: %s", key, exc)
|
||||
return default
|
||||
|
||||
|
||||
def set_config(
|
||||
key: str,
|
||||
value: Any,
|
||||
setting_type: str | None = None,
|
||||
updated_by: str | None = None,
|
||||
) -> Any:
|
||||
return get_system_config_service().set_config(key, value, setting_type=setting_type, updated_by=updated_by)
|
||||
@@ -1 +0,0 @@
|
||||
"""应用层:爆款视频 Prompt 模板系统(#2040)。"""
|
||||
@@ -1,427 +0,0 @@
|
||||
"""爆款视频 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)
|
||||
@@ -1,160 +0,0 @@
|
||||
"""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)
|
||||
@@ -1,252 +0,0 @@
|
||||
"""爆款视频 Prompt 模板默认值(v8 / v3 叙述优先重构)。
|
||||
|
||||
设计原则(灵应 2026-10-07):LLM 直接输出最终给用户看的文案,代码尽量薄。
|
||||
- image_analysis v8:VLM 主交付物是自然叙述风格的 summary_markdown,结构化
|
||||
字段仅保留 type/name/brand/has_person,顶层 products 改名 images;
|
||||
- storyboard v3:口播台词口语化、画面描述有画面感,copy_display_markdown 是
|
||||
LLM 直接写给用户看的流畅叙述文案,代码只做解析不改写;
|
||||
- intent_parsing 步骤整体删除,意图理解并入 storyboard 一次调用。
|
||||
|
||||
模板字段与 DB 表 viral_video_prompt_templates、prompt_loader 完全对应:
|
||||
name / prompt_type / version(int) / system_prompt / user_prompt_template /
|
||||
example_output / is_active。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
# 所有文案类 Prompt 自动注入的硬约束
|
||||
GLOBAL_CONSTRAINTS = """【必须遵守的硬约束】
|
||||
1. 不编造时间:不写“今年最新”“2024 爆款”等会过时的时间表述。
|
||||
2. 不承诺效果:不写“保证”“一定”“100%有效”“包治百病”等绝对化用语。
|
||||
3. 不编造价格、销量、认证、奖项:除非用户在文案中明确给出,否则一律不写。
|
||||
4. 符合广告法及平台社区规范。
|
||||
5. 只描述图片中真实可见的内容,看不到的不瞎猜。"""
|
||||
|
||||
# 负向提示(注入 storyboard / 视频生成负面词)
|
||||
NEGATIVE_RULES = """【反套路化要求】
|
||||
禁止使用“家人们谁懂啊”“绝绝子”“宝子们”“家人们”“太绝了”“yyds”等烂大街网络词;
|
||||
禁止固定模板化开头;语言要像真人朋友之间的分享,自然、具体、有信息量。"""
|
||||
|
||||
# 输出禁用套路词(测试会检查)
|
||||
BANNED_PHRASES = ["家人们谁懂啊", "绝绝子", "宝子们", "yyds", "太绝了"]
|
||||
|
||||
# 文案融合三档独立指令段(storyboard 一次生成,按档位注入风格指令)
|
||||
FUSION_INSTRUCTIONS = {
|
||||
"ai_full": """【本次创作模式:AI 全权创作】
|
||||
你是资深短视频编导。用户只提供了产品/门店图片,没有给出具体文案方向。请根据图片的真实观察和营销参数,自由发挥创作完整成片级方案,口播自然、画面可拍。""",
|
||||
"ai_polish": """【本次创作模式:AI 辅助润色】
|
||||
用户已给出方向或碎碎念。以用户的意思为主,保留其所有核心信息,在此基础上润色、补衔接、优化表达,让口播更自然、画面更具体;绝不改变用户核心意思,不添加用户没提到的卖点,品牌名、价格、人名等事实原样保留。""",
|
||||
"user_primary": """【本次创作模式:以用户原文为主】
|
||||
最小化修改:只做必要的通顺、合规修正与衔接补全,用户的核心句子与事实一律不改;用户文案已经很好就直接用,不为改而改。""",
|
||||
}
|
||||
|
||||
# ── 模板1:图片多模态分析 v8(叙述优先)───────────────────────────────
|
||||
_IMAGE_ANALYSIS_SYSTEM = """你是一名擅长观察和写作的品牌内容编导。面对一张真实图片,先用眼睛仔细看,再用自然、流畅、具体的中文把画面写成一段可以直接读给人听的描述。
|
||||
|
||||
## 输出格式(严格 JSON,不要输出 JSON 以外的任何内容)
|
||||
{
|
||||
"images": [
|
||||
{
|
||||
"type": "store 或 product 或 person 或 scene,四选一",
|
||||
"name": "主体名称,看不出就写“未识别”",
|
||||
"brand": "品牌名,看不出就留空字符串",
|
||||
"has_person": false,
|
||||
"summary_markdown": "用 Markdown 写成的自然叙述,这是最主要的交付物"
|
||||
}
|
||||
]
|
||||
}
|
||||
|
||||
## summary_markdown 写作要求(最重要)
|
||||
1. 写成完整、通顺的句子,像在跟朋友认真描述你看到的画面;不要用分号堆砌关键词,不要罗列“核心特征:xxx”“主色调:xxx”这类填表式标签。
|
||||
2. 开头先给一句整体定性,让读者立刻明白这是什么场景、什么主体。
|
||||
3. 颜色、材质、形状、部件要具体可感,写到位置和搭配;画面里出现的文字原样读出并自然融进句子,数字、规格、价格精确引用,看不清的不要编造。
|
||||
4. 只写真实看到的内容,不脑补功能、疗效、销量或画面之外的信息。
|
||||
5. 长度控制在 200-500 字。
|
||||
|
||||
## 按类型组织内容
|
||||
- type=store(门店/店内环境):用以下小标题分段,小标题下写连贯的句子而不是清单:
|
||||
###店铺主体
|
||||
###周边物品
|
||||
1.家具陈设
|
||||
2.商品与标识
|
||||
- type=product(商品):按自然段从整体到局部描写——先说是什么、什么品牌,再写包装/外形、颜色与材质、标签文字、可见部件与规格。
|
||||
- type=person(人物):描述人物身份感、姿态、穿着(上下装/颜色/款式)、动作与所处环境;用于品牌宣传时突出其精神状态。
|
||||
- type=scene(纯场景/风景):描述空间或风景的构成、色彩、光线、氛围与关键物件。
|
||||
|
||||
## 判断规则
|
||||
- has_person:画面中出现可辨识的真实人物(脸或完整上半身)才为 true,海报/模特立牌/照片里的人不算。
|
||||
- 一张图只描述其本身;多张图属于同一场景时可呼应,但不编造对应关系。
|
||||
- 输出必须是严格 JSON,summary_markdown 是字符串,内部换行用 \\n 表示。"""
|
||||
|
||||
_IMAGE_ANALYSIS_USER = """请分析这张图片。
|
||||
图片地址:{image_url}
|
||||
OCR 辅助文字(可能为空,仅供参考,不要照抄错误识别):{ocr_text}
|
||||
|
||||
严格按系统要求只输出 JSON。"""
|
||||
|
||||
_IMAGE_ANALYSIS_EXAMPLE = """{
|
||||
"images": [
|
||||
{
|
||||
"type": "store",
|
||||
"name": "御众堂门店",
|
||||
"brand": "御众堂",
|
||||
"has_person": false,
|
||||
"summary_markdown": "###店铺主体\\n这是一家名为“御众堂”的线下门店内部,整体暖木色调……"
|
||||
}
|
||||
]
|
||||
}"""
|
||||
|
||||
# ── 模板2:编导级分镜 v3(意图理解 + 分镜一次完成)────────────────────
|
||||
_STORYBOARD_SYSTEM = (
|
||||
"""你是一名懂短视频的编导和口播文案高手。你会拿到图片的真实观察、营销目的和用户参数,请一次性完成对营销意图的理解,并产出可直接拍摄/生成的分镜脚本。不要单独输出“意图解析”,意图要直接体现在台词和分镜里。
|
||||
|
||||
## 输出格式(XML,严格按结构输出,不要输出额外解释)
|
||||
<script>
|
||||
<copy_display_markdown><![CDATA[直接展示给用户看的成片文案,用 Markdown 写成流畅叙述]]></copy_display_markdown>
|
||||
<clips>
|
||||
<clip index="1">
|
||||
<time_range>0-3秒</time_range>
|
||||
<voiceover>这一镜的口播台词</voiceover>
|
||||
<visual>具体、有画面感的镜头描述(主体/动作/镜头运动/景别/光线)</visual>
|
||||
<reference_image_index>0</reference_image_index>
|
||||
</clip>
|
||||
</clips>
|
||||
<voiceover_script>把所有 clip 的 voiceover 连成完整口播稿</voiceover_script>
|
||||
<theme>一句话主题</theme>
|
||||
<negative>"""
|
||||
+ NEGATIVE_RULES
|
||||
+ """</negative>
|
||||
</script>
|
||||
|
||||
## 写作要求
|
||||
1. 口播台词:像真人面对镜头说话,短句、口语化、有停顿有情绪,开头 3 秒给出钩子;不要书面腔,不要机械报参数。
|
||||
2. 画面描述:写清“观众会看到什么”,有动作、有镜头运动、有景别和光线,具体可拍;不堆砌形容词,不写无法实现的画面。
|
||||
3. copy_display_markdown:直接展示给最终用户的文案,用 Markdown 写成自然、流畅、有感染力的成片成片文案,可用小标题与短句组织;不要做字段列表,不要出现“镜头一/台词:”这类制作说明。
|
||||
4. 内容必须来自图片观察与用户给出的信息,不编造卖点、不夸大、不使用绝对化用语和虚假承诺。
|
||||
5. reference_image_index 填本镜参考图片序号(从 0 开始),没有合适参考图填 -1。
|
||||
6. 分镜数量与时长匹配总时长,节奏紧凑。
|
||||
7. 口播字数硬约束(必须严格遵守):按每秒约 2.5~3 个中文字(正常口播语速)计算:
|
||||
- 5秒视频:voiceover_script 总字数 12~15 字
|
||||
- 10秒视频:voiceover_script 总字数 25~30 字
|
||||
- 15秒视频:voiceover_script 总字数 35~45 字
|
||||
- 20秒视频:voiceover_script 总字数 50~60 字
|
||||
- 30秒视频:voiceover_script 总字数 75~90 字
|
||||
- 每个 clip 的 voiceover 字数按该镜头时长比例分配
|
||||
- 所有 clip 的 voiceover 字数之和必须等于总 voiceover_script 字数
|
||||
- 宁可少写也不要多写,超长会导致 TTS 音频超出视频时长限制
|
||||
8. 镜头数量硬约束(必须严格遵守):
|
||||
- 5秒视频:1~2 个镜头
|
||||
- 10秒视频:3 个镜头
|
||||
- 15秒视频:3~4 个镜头
|
||||
- 20秒视频:4~5 个镜头
|
||||
- 30秒视频:6~8 个镜头
|
||||
9. 时间轴硬约束(必须严格遵守):
|
||||
- 每个 clip 的 time_range 必须写成 "X-Y秒" 格式,X 和 Y 是具体数字
|
||||
- 第一个 clip 必须从 0 秒开始
|
||||
- 最后一个 clip 必须结束于 total_duration 秒
|
||||
- 相邻 clip 首尾相接,不能有间隙也不能重叠
|
||||
- 每个 clip 的时长 = Y - X,必须 >= 2 秒
|
||||
10. 每个 clip 必须分配一个 reference_image_index(从 0 开始的图片序号),没有合适图片填 -1"""
|
||||
)
|
||||
|
||||
_STORYBOARD_USER = """<marketing_purpose>{marketing_purpose}</marketing_purpose>
|
||||
<image_analysis>
|
||||
{image_summary}
|
||||
</image_analysis>
|
||||
<user_parameters>
|
||||
<theme_hint>{theme_hint}</theme_hint>
|
||||
<duration>{duration}秒</duration>
|
||||
<aspect_ratio>{aspect_ratio}</aspect_ratio>
|
||||
<tone>{tone}</tone>
|
||||
<target_audience>{target_audience}</target_audience>
|
||||
<extra_requirements>{extra_requirements}</extra_requirements>
|
||||
</user_parameters>
|
||||
{video_style_section}
|
||||
请严格按 XML 结构输出分镜脚本。"""
|
||||
|
||||
_STORYBOARD_EXAMPLE = """<script>
|
||||
<copy_display_markdown><![CDATA[# 在御众堂,把松弛的自己一点点找回来
|
||||
产后妈妈最懂那种力不从心,推开门,暖光和一杯热茶先接住了你……]]></copy_display_markdown>
|
||||
<clips>
|
||||
<clip index="1">
|
||||
<time_range>0-3秒</time_range>
|
||||
<voiceover>生完娃,是不是连照镜子的勇气都没了?</voiceover>
|
||||
<visual>中近景,暖光下一位妈妈略显疲惫地看向镜中,镜头缓缓推近</visual>
|
||||
<reference_image_index>0</reference_image_index>
|
||||
</clip>
|
||||
</clips>
|
||||
<voiceover_script>生完娃,是不是连照镜子的勇气都没了?</voiceover_script>
|
||||
<theme>产后妈妈走进御众堂重拾状态</theme>
|
||||
<negative>模糊、畸变、夸大疗效、绝对化用语</negative>
|
||||
</script>"""
|
||||
|
||||
# ── 模板3:文案审核(合规/质量门禁)───────────────────────────────────
|
||||
_REVIEW_SYSTEM = """你是一名短视频广告合规审核与文案优化专家。审核待审文案:
|
||||
1) 广告法与平台合规(绝对化用语、虚假承诺、医疗功效宣称、导流违规);
|
||||
2) 卖点是否聚焦、逻辑是否通顺、口播是否自然;
|
||||
3) 是否有机械堆砌、书面腔、标签化表述。
|
||||
|
||||
只输出 XML,结构:
|
||||
<review>
|
||||
<passed>true 或 false</passed>
|
||||
<issues>
|
||||
<issue>
|
||||
<severity>high 或 medium 或 low</severity>
|
||||
<field>问题所在位置/字段</field>
|
||||
<problem>具体问题</problem>
|
||||
<suggestion>可直接替换的修改</suggestion>
|
||||
</issue>
|
||||
</issues>
|
||||
<rewrite>整体重写后的合规流畅版本(无问题时留空)</rewrite>
|
||||
</review>
|
||||
没有问题时 issues 留空、passed 为 true、rewrite 留空。"""
|
||||
|
||||
_REVIEW_USER = """<fusion_text>
|
||||
{fusion_text}
|
||||
</fusion_text>
|
||||
|
||||
请审核以上文案。"""
|
||||
|
||||
_REVIEW_EXAMPLE = """<review>
|
||||
<passed>false</passed>
|
||||
<issues>
|
||||
<issue>
|
||||
<severity>high</severity>
|
||||
<field>opening</field>
|
||||
<problem>使用绝对化用语“全网第一”</problem>
|
||||
<suggestion>改为“很多老客户回购的一款”</suggestion>
|
||||
</issue>
|
||||
</issues>
|
||||
<rewrite>……</rewrite>
|
||||
</review>"""
|
||||
|
||||
|
||||
DEFAULT_TEMPLATES: list[dict] = [
|
||||
{
|
||||
"name": "图片多模态分析 v8",
|
||||
"prompt_type": "image_analysis",
|
||||
"version": 8,
|
||||
"system_prompt": _IMAGE_ANALYSIS_SYSTEM,
|
||||
"user_prompt_template": _IMAGE_ANALYSIS_USER,
|
||||
"example_output": _IMAGE_ANALYSIS_EXAMPLE,
|
||||
"is_active": True,
|
||||
},
|
||||
{
|
||||
"name": "编导级分镜 v3",
|
||||
"prompt_type": "storyboard",
|
||||
"version": 3,
|
||||
"system_prompt": _STORYBOARD_SYSTEM,
|
||||
"user_prompt_template": _STORYBOARD_USER,
|
||||
"example_output": _STORYBOARD_EXAMPLE,
|
||||
"is_active": True,
|
||||
},
|
||||
{
|
||||
"name": "文案审核",
|
||||
"prompt_type": "review",
|
||||
"version": 1,
|
||||
"system_prompt": _REVIEW_SYSTEM,
|
||||
"user_prompt_template": _REVIEW_USER,
|
||||
"example_output": _REVIEW_EXAMPLE,
|
||||
"is_active": True,
|
||||
},
|
||||
]
|
||||
@@ -1,337 +0,0 @@
|
||||
"""文案审核 + 自动重写(#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:
|
||||
# markdown展示字段不参与合规审核(避免格式字符误判)
|
||||
_MARKDOWN_FIELDS = {"summary_markdown", "copy_display_markdown"}
|
||||
|
||||
def __init__(self, client=None):
|
||||
if client is None:
|
||||
try:
|
||||
from packages.shared.ai_router import ai_router
|
||||
|
||||
client = ai_router.get_llm_client("copy_review")
|
||||
except Exception:
|
||||
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:
|
||||
# LLM审核失败(超时/网络错误等),降级放行,不阻断渲染
|
||||
return ReviewResult(
|
||||
passed=True,
|
||||
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]:
|
||||
try:
|
||||
return self._llm_review_inner(fusion, intent, fusion_level)
|
||||
except Exception as e:
|
||||
import logging
|
||||
|
||||
logging.getLogger(__name__).warning("[Reviewer] LLM审核调用异常,降级放行: %s", e)
|
||||
return None
|
||||
|
||||
def _llm_review_inner(
|
||||
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,
|
||||
timeout=25,
|
||||
)
|
||||
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,
|
||||
timeout=25,
|
||||
)
|
||||
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:
|
||||
_MARKDOWN_FIELDS = {"summary_markdown", "copy_display_markdown"}
|
||||
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)
|
||||
# 过滤掉markdown展示字段,避免格式字符被误判
|
||||
parts = [p for p in parts if not any(mk in p for mk in _MARKDOWN_FIELDS)]
|
||||
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)
|
||||
@@ -1,116 +0,0 @@
|
||||
"""内部 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 = ""
|
||||
@@ -1,113 +0,0 @@
|
||||
"""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*(?:"([^"]*)"|'([^']*)')""")
|
||||
_CDATA_RE = re.compile(r"^<!\[CDATA\[(.*)\]\]>$", re.DOTALL)
|
||||
|
||||
|
||||
def _strip_cdata(s: str) -> str:
|
||||
"""剥离 LLM 可能照抄示例输出的 ``<![CDATA[...]]>`` 包裹层。"""
|
||||
m = _CDATA_RE.match(s.strip())
|
||||
return m.group(1) if m else 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["text"] = _strip_cdata(node["text"])
|
||||
node.pop("_start", None)
|
||||
del stack[idx:]
|
||||
break
|
||||
# 未闭合标签:给剩余部分作为文本
|
||||
for node in stack:
|
||||
if "_start" in node:
|
||||
node["text"] = unescape(text[node["_start"] :].strip())
|
||||
node["text"] = _strip_cdata(node["text"])
|
||||
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
|
||||
+19
-98
@@ -80,46 +80,34 @@ class SharedSettings(BaseSettings):
|
||||
|
||||
# ── CosyVoice (阿里云百炼语音合成) ───────────────────────────────────
|
||||
cosyvoice_api_key: str = ""
|
||||
cosyvoice_base_url: str = ""
|
||||
cosyvoice_model: str = ""
|
||||
cosyvoice_voice: str = "longxiaochun_v3"
|
||||
cosyvoice_base_url: str = "https://dashscope.aliyuncs.com/api/v1"
|
||||
cosyvoice_model: str = "cosyvoice-v3-flash"
|
||||
cosyvoice_voice: str = "longxiaochun_v3" # 默认音色(v3 系列系统音色带 _v3 后缀)
|
||||
cosyvoice_sample_rate: int = 22050
|
||||
cosyvoice_format: str = "mp3"
|
||||
cosyvoice_format: str = "mp3" # 输出格式:mp3/wav/pcm
|
||||
# 音色克隆模型名(固定为 voice-enrollment)
|
||||
cosyvoice_clone_model: str = ""
|
||||
cosyvoice_clone_model: str = "voice-enrollment"
|
||||
|
||||
# ── 豆包大模型(火山引擎方舟) ────────────────────────────────────────
|
||||
# AI模型路由化:model/base_url 默认值清空,由 DB ai_models/ai_capability_configs 配置驱动。
|
||||
# 环境变量仍可覆盖(兼容旧部署);无任何配置时 ai_router fallback 提供最终默认值。
|
||||
doubao_api_key: str = ""
|
||||
doubao_model: str = ""
|
||||
doubao_fast_model: str = ""
|
||||
doubao_base_url: str = ""
|
||||
doubao_timeout: int = 45
|
||||
doubao_max_retries: int = 3
|
||||
doubao_vision_model: str = ""
|
||||
doubao_vision_lite_model: str = ""
|
||||
doubao_vision_use_lite: bool = True
|
||||
doubao_embedding_model: str = ""
|
||||
doubao_video_model: str = ""
|
||||
doubao_video_timeout: int = 600
|
||||
doubao_video_poll_interval: int = 10
|
||||
doubao_image_model: str = ""
|
||||
doubao_image_size: str = "1K"
|
||||
doubao_image_timeout: int = 60
|
||||
doubao_trust_chain_enabled: bool = True
|
||||
|
||||
# ── DashScope (阿里云百炼 Wan 3.0 等) ─────────────────────────────────
|
||||
dashscope_api_key: str = ""
|
||||
dashscope_base_url: str = ""
|
||||
dashscope_video_timeout: int = 900
|
||||
dashscope_video_poll_interval: int = 10
|
||||
doubao_model: str = "doubao-seed-1-6-250615" # 推理模型(通用兜底)
|
||||
doubao_fast_model: str = "doubao-1-5-pro-32k-250115" # 快速结构化输出模型(编导脚本/意图解析/审核)
|
||||
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_video_model: str = "doubao-seedance-2-5-260628"
|
||||
doubao_video_timeout: int = 600 # 视频生成轮询总超时(秒)
|
||||
doubao_video_poll_interval: int = 10 # 轮询间隔(秒)
|
||||
|
||||
# ── MediaKit (火山引擎 AI 媒体工具) ──────────────────────────────────
|
||||
mediakit_api_key: str = ""
|
||||
mediakit_base_url: str = ""
|
||||
mediakit_base_url: str = "https://mediakit.cn-beijing.volces.com/api/v1"
|
||||
mediakit_timeout: int = 60
|
||||
mediakit_cover_enabled: bool = False
|
||||
mediakit_cover_enabled: bool = False # 封面抽帧是否走MediaKit(默认false走本地ffmpeg+cv2,<2s完成)
|
||||
|
||||
# ── 积分/会员系统 (#1895) ────────────────────────────────────────────
|
||||
# 积分系统总开关(产品要求 #1895:暂停积分系统但保留全部代码/表/接口)。
|
||||
@@ -173,73 +161,6 @@ class SharedSettings(BaseSettings):
|
||||
# 判断 Worker 可用的心跳新鲜度窗口(秒)—— last_heartbeat_at 在窗口内视为在线
|
||||
gpu_worker_stale_seconds: int = 300
|
||||
|
||||
# ── Ditto 蚂蚁数字人口型 API(#2076)─────────────────────────────────
|
||||
# 是否优先使用 Ditto(蚂蚁数字人,替代 MuseTalk)。开关开启且 base_url 配置
|
||||
# 非空时,对口型任务优先走 Ditto;失败后回退 MuseTalk/MediaKit。
|
||||
use_ditto_lipsync: bool = Field(
|
||||
default=False,
|
||||
validation_alias=AliasChoices("USE_DITTO_LIPSYNC", "use_ditto_lipsync"),
|
||||
)
|
||||
# Ditto FastAPI 内网地址(Tailscale),如 http://100.x.x.x:8000
|
||||
ditto_api_base_url: str = Field(
|
||||
default="",
|
||||
validation_alias=AliasChoices("DITTO_API_BASE_URL", "ditto_api_base_url"),
|
||||
)
|
||||
# 默认人物模板视频 URL(正面 5-10 秒循环、光线均匀、半身)。Ditto 模式下忽略
|
||||
# 用户上传的驱动视频/图片,统一用该模板;后续可扩展为多模板让用户选择。
|
||||
ditto_default_video_url: str = Field(
|
||||
default="",
|
||||
validation_alias=AliasChoices("DITTO_DEFAULT_VIDEO_URL", "ditto_default_video_url"),
|
||||
)
|
||||
# 429 GPU 繁忙时指数退避最大重试次数
|
||||
ditto_max_retries: int = Field(
|
||||
default=3,
|
||||
validation_alias=AliasChoices("DITTO_MAX_RETRIES", "ditto_max_retries"),
|
||||
)
|
||||
# Ditto 单次请求 read 超时(秒):数字人半身视频推理通常 30-120s(RTF≈2.8,40s音频约112s)
|
||||
# connect 超时固定 10s(代码硬编码,网络不通快速失败)
|
||||
ditto_request_timeout: int = Field(
|
||||
default=120,
|
||||
validation_alias=AliasChoices("DITTO_REQUEST_TIMEOUT", "ditto_request_timeout"),
|
||||
)
|
||||
# Ditto 句间过渡帧数(平滑表情/口型切换)
|
||||
ditto_blend_frames: int = Field(
|
||||
default=12,
|
||||
validation_alias=AliasChoices("DITTO_BLEND_FRAMES", "ditto_blend_frames"),
|
||||
)
|
||||
|
||||
# ── Ditto LLM 情绪分析(emo_timeline)──────────────────────────────
|
||||
# 总开关;关闭或 LLM 失败时走 GPU 端关键词匹配兜底
|
||||
ditto_emotion_enabled: bool = Field(
|
||||
default=False,
|
||||
validation_alias=AliasChoices("DITTO_EMOTION_ENABLED", "ditto_emotion_enabled"),
|
||||
)
|
||||
ditto_emotion_model: str = Field(
|
||||
default="doubao-seed-2-1-lite-250915",
|
||||
validation_alias=AliasChoices("DITTO_EMOTION_MODEL", "ditto_emotion_model"),
|
||||
)
|
||||
ditto_emotion_temperature: float = Field(
|
||||
default=0.1,
|
||||
validation_alias=AliasChoices("DITTO_EMOTION_TEMPERATURE", "ditto_emotion_temperature"),
|
||||
)
|
||||
ditto_emotion_timeout: int = Field(
|
||||
default=10,
|
||||
validation_alias=AliasChoices("DITTO_EMOTION_TIMEOUT", "ditto_emotion_timeout"),
|
||||
)
|
||||
ditto_emotion_max_tokens: int = Field(
|
||||
default=1024,
|
||||
validation_alias=AliasChoices("DITTO_EMOTION_MAX_TOKENS", "ditto_emotion_max_tokens"),
|
||||
)
|
||||
ditto_emotion_cache_size: int = Field(
|
||||
default=500,
|
||||
validation_alias=AliasChoices("DITTO_EMOTION_CACHE_SIZE", "ditto_emotion_cache_size"),
|
||||
)
|
||||
# 提示词模板:必须包含 {文案} 占位符;后台可通过环境变量覆盖
|
||||
ditto_emotion_prompt: str = Field(
|
||||
default="",
|
||||
validation_alias=AliasChoices("DITTO_EMOTION_PROMPT", "ditto_emotion_prompt"),
|
||||
)
|
||||
|
||||
# ── P4000 NVENC 硬件编码 ────────────────────────────────────────────
|
||||
# GPU 编码总开关;关闭或 endpoint 为空时始终走本机 CPU libx264
|
||||
enable_gpu_encode: bool = Field(
|
||||
|
||||
@@ -8,12 +8,7 @@ from __future__ import annotations
|
||||
|
||||
import uuid
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import datetime, timezone
|
||||
|
||||
try:
|
||||
from datetime import UTC
|
||||
except ImportError:
|
||||
UTC = timezone.utc
|
||||
from datetime import UTC, datetime
|
||||
|
||||
|
||||
@dataclass
|
||||
|
||||
@@ -60,11 +60,6 @@ class User:
|
||||
# 资料是否已完善(微信新用户首次设置昵称后置 True;邮箱注册默认 True)
|
||||
profile_completed: bool = True
|
||||
|
||||
# 会员字段 (#1895):与 users 表列对应
|
||||
is_member: bool = False
|
||||
member_type: str | None = None
|
||||
member_expires_at: datetime | None = None
|
||||
|
||||
created_at: datetime = field(default_factory=lambda: datetime.now(UTC))
|
||||
|
||||
|
||||
|
||||
@@ -1,376 +0,0 @@
|
||||
"""功能计费配置服务:从 feature_pricing_configs 读配置,300 秒 TTL 内存缓存。
|
||||
|
||||
配置表由 xiaoxia-admin 侧维护(同库 PostgreSQL),本服务只读。
|
||||
DB 不可用 / 表不存在 / 无数据时自动回落到内置兜底配置,保证业务不崩。
|
||||
|
||||
计费公式:最终积分 = (动态成本 + 固定成本) × 利润系数,price_cap 封顶。
|
||||
启用条件:全局 points_enabled 总开关 AND 功能 is_enabled 同时为 true。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import threading
|
||||
import time
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Optional
|
||||
|
||||
import sqlalchemy as sa
|
||||
|
||||
from packages.adapters.sqlalchemy_impl import session as _session_mod
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
CACHE_TTL_SECONDS = 300.0
|
||||
|
||||
# ── 爆款视频兜底模型单价(与旧硬编码表/现状一致;DB 不可用时使用) ───────
|
||||
# 结构:models[model_key][resolution]["true"/"false"] = 单价
|
||||
# token 模式:元/百万输出 tokens;per_second 模式:元/秒
|
||||
# 注意:仅 seedance-2.5 配置 true(图生视频)单价;其余模型只有 false,
|
||||
# 精确 key 缺失时由 points_rules 回落到 seedance-2.5/false(与旧现状一致)。
|
||||
_FALLBACK_VIRAL_MODEL_PRICING: dict = {
|
||||
"seedance-2.5": {
|
||||
"480p": {"false": 70.0, "true": 42.0},
|
||||
"720p": {"false": 70.0, "true": 42.0},
|
||||
"1080p": {"false": 77.0, "true": 46.0},
|
||||
},
|
||||
"seedance-2.0": {
|
||||
"480p": {"false": 46.0},
|
||||
"720p": {"false": 46.0},
|
||||
"1080p": {"false": 51.0},
|
||||
"4k": {"false": 80.0},
|
||||
},
|
||||
"seedance-2.0-fast": {
|
||||
"480p": {"false": 28.0},
|
||||
"720p": {"false": 28.0},
|
||||
},
|
||||
"seedance-2.0-mini": {
|
||||
"480p": {"false": 9.2},
|
||||
"720p": {"false": 9.2},
|
||||
},
|
||||
"wan-3.0": {
|
||||
"480p": {"false": 0.3},
|
||||
"720p": {"false": 0.6},
|
||||
"1080p": {"false": 1.2},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
@dataclass
|
||||
class FeatureConfig:
|
||||
"""功能计费配置快照。"""
|
||||
|
||||
feature_key: str
|
||||
name: str = ""
|
||||
emoji: str = ""
|
||||
is_enabled: bool = False
|
||||
fixed_cost: float = 0.0
|
||||
profit_multiplier: float = 1.0
|
||||
dynamic_unit_cost: float = 0.0
|
||||
billing_mode: str = "model_based"
|
||||
price_cap: float = 0.0
|
||||
model_pricing: dict = field(default_factory=dict)
|
||||
description: str = ""
|
||||
|
||||
|
||||
# ── 进程内缓存:(loaded_monotonic, {feature_key: FeatureConfig}) ──────────
|
||||
_lock = threading.Lock()
|
||||
_cache: Optional[tuple[float, dict[str, FeatureConfig]]] = None
|
||||
|
||||
|
||||
def _fallback_configs() -> dict[str, FeatureConfig]:
|
||||
"""内置兜底配置:爆款启用(与现状一致),其余两个关闭。"""
|
||||
return {
|
||||
"viral_video": FeatureConfig(
|
||||
feature_key="viral_video",
|
||||
name="爆款视频",
|
||||
emoji="🎬",
|
||||
is_enabled=True,
|
||||
fixed_cost=0.15,
|
||||
profit_multiplier=1.3,
|
||||
dynamic_unit_cost=0.0,
|
||||
billing_mode="model_based",
|
||||
price_cap=0.0,
|
||||
model_pricing=json.loads(json.dumps(_FALLBACK_VIRAL_MODEL_PRICING)),
|
||||
description="爆款视频动态定价(兜底配置)",
|
||||
),
|
||||
"lip_sync": FeatureConfig(
|
||||
feature_key="lip_sync",
|
||||
name="对口型",
|
||||
emoji="🎙️",
|
||||
is_enabled=False,
|
||||
fixed_cost=0.0,
|
||||
profit_multiplier=1.0,
|
||||
dynamic_unit_cost=0.0,
|
||||
billing_mode="per_second",
|
||||
price_cap=0.0,
|
||||
description="对口型计费(兜底配置,默认关闭)",
|
||||
),
|
||||
"smart_edit": FeatureConfig(
|
||||
feature_key="smart_edit",
|
||||
name="智能剪辑",
|
||||
emoji="✂️",
|
||||
is_enabled=False,
|
||||
fixed_cost=0.0,
|
||||
profit_multiplier=1.0,
|
||||
dynamic_unit_cost=0.0,
|
||||
billing_mode="model_based",
|
||||
price_cap=0.0,
|
||||
description="智能剪辑固定价计费(兜底配置,默认关闭)",
|
||||
),
|
||||
}
|
||||
|
||||
|
||||
_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://")
|
||||
if url.startswith("postgresql://"):
|
||||
url = url.replace("postgresql://", "postgresql+psycopg://")
|
||||
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 _parse_model_pricing(raw) -> dict:
|
||||
"""解析 model_pricing_json(Text JSON),空/失败 → {}。"""
|
||||
if raw is None:
|
||||
return {}
|
||||
if isinstance(raw, dict):
|
||||
return raw
|
||||
text = str(raw).strip()
|
||||
if not text:
|
||||
return {}
|
||||
try:
|
||||
data = json.loads(text)
|
||||
except (ValueError, TypeError):
|
||||
logger.warning("model_pricing_json 解析失败,按空配置处理: %r", text[:200])
|
||||
return {}
|
||||
return data if isinstance(data, dict) else {}
|
||||
|
||||
|
||||
def _to_float(value, default: float = 0.0) -> float:
|
||||
try:
|
||||
if value is None:
|
||||
return default
|
||||
return float(value)
|
||||
except (TypeError, ValueError):
|
||||
return default
|
||||
|
||||
|
||||
def _load_all() -> dict[str, FeatureConfig]:
|
||||
"""SELECT * FROM feature_pricing_configs,返回 {feature_key: FeatureConfig}。
|
||||
|
||||
表不存在 / DB 异常由调用方捕获并回落兜底配置。
|
||||
"""
|
||||
session = None
|
||||
try:
|
||||
session = _get_session()
|
||||
if session is None:
|
||||
raise RuntimeError("no db session available")
|
||||
sql = sa.text("""
|
||||
SELECT feature_key, name, emoji, is_enabled, fixed_cost,
|
||||
profit_multiplier, dynamic_unit_cost, billing_mode,
|
||||
price_cap, model_pricing_json, description
|
||||
FROM feature_pricing_configs
|
||||
""")
|
||||
rows = session.execute(sql).mappings().all()
|
||||
configs: dict[str, FeatureConfig] = {}
|
||||
for row in rows:
|
||||
key = str(row["feature_key"] or "").strip()
|
||||
if not key:
|
||||
continue
|
||||
configs[key] = FeatureConfig(
|
||||
feature_key=key,
|
||||
name=str(row["name"] or key),
|
||||
emoji=str(row["emoji"] or ""),
|
||||
is_enabled=bool(row["is_enabled"]),
|
||||
fixed_cost=_to_float(row["fixed_cost"]),
|
||||
profit_multiplier=_to_float(row["profit_multiplier"], 1.0),
|
||||
dynamic_unit_cost=_to_float(row["dynamic_unit_cost"]),
|
||||
billing_mode=str(row["billing_mode"] or "model_based"),
|
||||
price_cap=_to_float(row["price_cap"]),
|
||||
model_pricing=_parse_model_pricing(row["model_pricing_json"]),
|
||||
description=str(row["description"] or ""),
|
||||
)
|
||||
return configs
|
||||
finally:
|
||||
if session is not None:
|
||||
try:
|
||||
session.close()
|
||||
except Exception: # noqa: BLE001
|
||||
pass
|
||||
|
||||
|
||||
def _get_cache() -> dict[str, FeatureConfig]:
|
||||
"""TTL 内返回缓存,否则重新 load;DB 异常/表不存在时返回内置兜底配置。"""
|
||||
global _cache
|
||||
now = time.monotonic()
|
||||
with _lock:
|
||||
if _cache is not None and now - _cache[0] < CACHE_TTL_SECONDS:
|
||||
return _cache[1]
|
||||
|
||||
try:
|
||||
loaded = _load_all()
|
||||
except Exception: # noqa: BLE001 - 表不存在/DB 不可用时静默回落
|
||||
logger.info("feature_pricing_configs 读取失败,使用内置兜底配置", exc_info=True)
|
||||
return _fallback_configs()
|
||||
|
||||
# DB 可用但表为空:同样回落兜底(保证爆款现状不被改变)
|
||||
if not loaded:
|
||||
fallback = _fallback_configs()
|
||||
with _lock:
|
||||
_cache = (now, fallback)
|
||||
return fallback
|
||||
|
||||
# 以兜底为底(DB 未配置的 feature_key 仍有兜底),DB 行覆盖
|
||||
merged = _fallback_configs()
|
||||
merged.update(loaded)
|
||||
with _lock:
|
||||
_cache = (now, merged)
|
||||
return merged
|
||||
|
||||
|
||||
def get_feature_config(feature_key: str) -> Optional[FeatureConfig]:
|
||||
"""获取指定功能配置,未知 key 返回 None。"""
|
||||
key = str(feature_key or "").strip()
|
||||
if not key:
|
||||
return None
|
||||
return _get_cache().get(key)
|
||||
|
||||
|
||||
def _global_points_enabled() -> bool:
|
||||
"""全局积分总开关(兼容 api / worker 运行时),取不到时默认关闭。"""
|
||||
try:
|
||||
from packages.shared import get_shared_settings
|
||||
|
||||
return bool(get_shared_settings().points_enabled)
|
||||
except Exception: # noqa: BLE001
|
||||
pass
|
||||
try:
|
||||
from app.config import settings
|
||||
|
||||
return bool(getattr(settings, "points_enabled", False))
|
||||
except Exception: # noqa: BLE001
|
||||
return False
|
||||
|
||||
|
||||
def is_feature_enabled(feature_key: str) -> bool:
|
||||
"""功能是否启用并扣费:全局 points_enabled AND 功能 is_enabled。"""
|
||||
cfg = get_feature_config(feature_key)
|
||||
if cfg is None:
|
||||
return False
|
||||
return bool(cfg.is_enabled) and _global_points_enabled()
|
||||
|
||||
|
||||
def calculate_price(feature_key: str, dynamic_cost: float = 0.0) -> tuple[float, dict]:
|
||||
"""按公式计算最终积分并返回明细。
|
||||
|
||||
price = (dynamic_cost + fixed_cost) × profit_multiplier
|
||||
price_cap > 0 时封顶(取 min)。
|
||||
功能未启用 → (0.0, breakdown{is_enabled: False, charged: False})。
|
||||
"""
|
||||
cfg = get_feature_config(feature_key)
|
||||
dynamic = max(0.0, _to_float(dynamic_cost))
|
||||
if cfg is None or not cfg.is_enabled:
|
||||
return 0.0, {
|
||||
"feature_key": feature_key,
|
||||
"is_enabled": False,
|
||||
"charged": False,
|
||||
"dynamic_cost": dynamic,
|
||||
"fixed_cost": 0.0,
|
||||
"profit_multiplier": 1.0,
|
||||
"price_cap": 0.0,
|
||||
"final_price": 0.0,
|
||||
}
|
||||
|
||||
fixed = max(0.0, cfg.fixed_cost)
|
||||
multiplier = cfg.profit_multiplier if cfg.profit_multiplier > 0 else 1.0
|
||||
raw_price = (dynamic + fixed) * multiplier
|
||||
cap = cfg.price_cap if cfg.price_cap and cfg.price_cap > 0 else 0.0
|
||||
final_price = min(raw_price, cap) if cap else raw_price
|
||||
final_price = round(float(final_price), 2)
|
||||
breakdown = {
|
||||
"feature_key": cfg.feature_key,
|
||||
"is_enabled": True,
|
||||
"charged": True,
|
||||
"dynamic_cost": round(dynamic, 4),
|
||||
"fixed_cost": float(fixed),
|
||||
"profit_multiplier": float(multiplier),
|
||||
"price_cap": float(cap),
|
||||
"raw_price": round(float(raw_price), 4),
|
||||
"final_price": final_price,
|
||||
}
|
||||
return final_price, breakdown
|
||||
|
||||
|
||||
def lookup_model_price(
|
||||
model_pricing: dict,
|
||||
model_key: str,
|
||||
resolution: str,
|
||||
has_video_input: bool,
|
||||
) -> Optional[float]:
|
||||
"""从 model_pricing dict 取模型单价,兼容两种常见 JSON 结构。
|
||||
|
||||
1. 嵌套:{model: {resolution: {"true"/"false": price}}}
|
||||
(内层 bool key 也兼容直接 bool / 省略)
|
||||
2. 扁平:{"model|resolution|true_or_false": price}
|
||||
(分隔符支持 | / : / , / 空格;bool 段可省略)
|
||||
取不到返回 None。
|
||||
"""
|
||||
if not isinstance(model_pricing, dict):
|
||||
return None
|
||||
model = str(model_key or "").strip()
|
||||
res = str(resolution or "").strip()
|
||||
flag = "true" if has_video_input else "false"
|
||||
|
||||
# 1. 嵌套
|
||||
model_node = model_pricing.get(model)
|
||||
if isinstance(model_node, dict):
|
||||
res_node = model_node.get(res)
|
||||
if isinstance(res_node, dict):
|
||||
# 精确 bool key 命中才返回;不做“只有一个值就取”的模糊匹配
|
||||
# (否则缺失 true 时会错误地取到 false 价,破坏旧版回落规则)
|
||||
if flag in res_node:
|
||||
return _to_float(res_node[flag]) if res_node[flag] is not None else None
|
||||
if has_video_input in res_node:
|
||||
val = res_node[has_video_input]
|
||||
return _to_float(val) if val is not None else None
|
||||
elif isinstance(res_node, (int, float)):
|
||||
return float(res_node)
|
||||
|
||||
# 2. 扁平
|
||||
for sep in ("|", ":", ",", " "):
|
||||
for key in (
|
||||
f"{model}{sep}{res}{sep}{flag}",
|
||||
f"{model}{sep}{res}",
|
||||
):
|
||||
if key in model_pricing:
|
||||
value = model_pricing[key]
|
||||
return _to_float(value) if value is not None else None
|
||||
return None
|
||||
|
||||
|
||||
def refresh_feature_configs() -> None:
|
||||
"""清空缓存(下次读取重新 load DB;测试/admin 改配置后可手动调)。"""
|
||||
global _cache
|
||||
with _lock:
|
||||
_cache = None
|
||||
@@ -9,9 +9,9 @@ from uuid import uuid4
|
||||
class PointsAccount:
|
||||
id: str
|
||||
user_id: str
|
||||
balance: float = 0.0
|
||||
total_earned: float = 0.0
|
||||
total_spent: float = 0.0
|
||||
balance: int = 0
|
||||
total_earned: int = 0
|
||||
total_spent: int = 0
|
||||
created_at: datetime = field(default_factory=lambda: datetime.now(UTC))
|
||||
updated_at: datetime = field(default_factory=lambda: datetime.now(UTC))
|
||||
|
||||
|
||||
+54
-393
@@ -1,384 +1,32 @@
|
||||
"""积分消耗规则配置 (#1895)
|
||||
|
||||
v1.6.1: 按产品决策,智能混剪/AI数字人/AI配音/抖音解析/改写/标题/封面 全部免费,
|
||||
仅保留声音克隆合成(voice_clone_synth)的扣点逻辑;声音克隆训练保持免费。
|
||||
爆款视频(viral_video)走动态定价,计费参数 DB 化(feature_pricing_configs,
|
||||
见 feature_pricing_service),calculate_viral_video_credits 从配置读取单价/
|
||||
固定成本/利润系数/封顶,DB 不可用时回落兜底配置。
|
||||
"""
|
||||
"""积分消耗规则配置 (#1895)"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
|
||||
from packages.domain import feature_pricing_service
|
||||
|
||||
# ============ 爆款视频动态定价 ============
|
||||
# 单价/固定成本/利润系数已 DB 化(feature_pricing_configs,feature_key=viral_video),
|
||||
# 由 feature_pricing_service 读取(300s 缓存),DB 不可用时回落内置兜底配置。
|
||||
# 以下三个常量仅为向后兼容保留(旧引用方/兜底场景),值取自兜底配置。
|
||||
VIRAL_VIDEO_MODEL_PRICES: dict[tuple[str, str, bool], float] = {
|
||||
("seedance-2.5", "480p", False): 70.0,
|
||||
("seedance-2.5", "720p", False): 70.0,
|
||||
("seedance-2.5", "1080p", False): 77.0,
|
||||
("seedance-2.5", "480p", True): 42.0,
|
||||
("seedance-2.5", "720p", True): 42.0,
|
||||
("seedance-2.5", "1080p", True): 46.0,
|
||||
("seedance-2.0", "480p", False): 46.0,
|
||||
("seedance-2.0", "720p", False): 46.0,
|
||||
("seedance-2.0", "1080p", False): 51.0,
|
||||
("seedance-2.0", "4k", False): 80.0,
|
||||
("seedance-2.0-fast", "480p", False): 28.0,
|
||||
("seedance-2.0-fast", "720p", False): 28.0,
|
||||
("seedance-2.0-mini", "480p", False): 9.2,
|
||||
("seedance-2.0-mini", "720p", False): 9.2,
|
||||
("wan-3.0", "480p", False): 0.3,
|
||||
("wan-3.0", "720p", False): 0.6,
|
||||
("wan-3.0", "1080p", False): 1.2,
|
||||
}
|
||||
|
||||
# 固定成本(元):VLM 分析 + LLM 文案 + TTS + OSS + 服务器(兜底默认值)
|
||||
VIRAL_VIDEO_FIXED_COST = 0.15
|
||||
# 利润系数(兜底默认值)
|
||||
VIRAL_VIDEO_PROFIT_MULTIPLIER = 1.3
|
||||
# Seedance 输出帧率
|
||||
VIRAL_VIDEO_FPS = 24
|
||||
|
||||
# 分辨率别名映射 -> 标准 key
|
||||
_RESOLUTION_ALIASES: dict[str, str] = {
|
||||
"480p": "480p",
|
||||
"普清": "480p",
|
||||
"default": "480p",
|
||||
"low": "480p",
|
||||
"sd": "480p",
|
||||
"720p": "720p",
|
||||
"高清": "720p",
|
||||
"medium": "720p",
|
||||
"hd": "720p",
|
||||
"1080p": "1080p",
|
||||
"超清": "1080p",
|
||||
"high": "1080p",
|
||||
"ultra": "1080p",
|
||||
"全能": "1080p",
|
||||
"fhd": "1080p",
|
||||
}
|
||||
# 分辨率 -> 短边像素数(p 值代表短边,不是 height)
|
||||
_RESOLUTION_SHORT_SIDE: dict[str, int] = {"480p": 480, "720p": 720, "1080p": 1080, "4k": 2160}
|
||||
_RESOLUTION_ALIASES["4k"] = "4k"
|
||||
_RESOLUTION_ALIASES["2160p"] = "4k"
|
||||
_RESOLUTION_ALIASES["uhd"] = "4k"
|
||||
|
||||
|
||||
def resolve_video_dimensions(resolution: str, ratio: str) -> tuple[int, int]:
|
||||
"""把 (resolution, ratio) 解析为 (width, height)。
|
||||
|
||||
resolution 数字代表短边像素数(480p/720p/1080p 等):
|
||||
- 横屏 16:9:短边是 height,width = short * 16/9
|
||||
- 竖屏 9:16:短边是 width,height = short * 16/9
|
||||
- 方屏 1:1:width = height = short
|
||||
"""
|
||||
key = str(resolution or "").strip()
|
||||
key_l = key.lower()
|
||||
res_key = _RESOLUTION_ALIASES.get(key_l) or _RESOLUTION_ALIASES.get(key) or "720p"
|
||||
short = _RESOLUTION_SHORT_SIDE.get(res_key, 720)
|
||||
r = str(ratio or "").strip().lower()
|
||||
if r == "16:9":
|
||||
# 横屏:短边是 height,width 向上取整并对齐偶数
|
||||
w = math.ceil(short * 16 / 9)
|
||||
h = short
|
||||
elif r == "1:1":
|
||||
w, h = short, short
|
||||
else:
|
||||
# 9:16 竖屏(默认):短边是 width,height 向上取整并对齐偶数
|
||||
w = short
|
||||
h = math.ceil(short * 16 / 9)
|
||||
# 对齐到偶数(视频编码要求)
|
||||
w = w + (w % 2)
|
||||
h = h + (h % 2)
|
||||
return int(w), int(h)
|
||||
|
||||
|
||||
# ── 爆款视频多模型元数据 (#2159) ──────────────────────────────────────
|
||||
VIRAL_VIDEO_MODEL_CONFIG: dict[str, dict] = {
|
||||
"seedance-2.5": {
|
||||
"key": "seedance-2.5",
|
||||
"display_name": "Seedance 2.5 — 最新最强",
|
||||
"model_id": "doubao-seedance-2-5-260628",
|
||||
"provider": "doubao",
|
||||
"supports_audio": True,
|
||||
"supported_resolutions": ["480p", "720p", "1080p"],
|
||||
"max_duration": 30,
|
||||
"billing_mode": "token",
|
||||
"is_default": True,
|
||||
},
|
||||
"seedance-2.0": {
|
||||
"key": "seedance-2.0",
|
||||
"display_name": "Seedance 2.0 — 正式首选",
|
||||
"model_id": "doubao-seedance-2-0-260128",
|
||||
"provider": "doubao",
|
||||
"supports_audio": True,
|
||||
"supported_resolutions": ["480p", "720p", "1080p", "4k"],
|
||||
"max_duration": 15,
|
||||
"billing_mode": "token",
|
||||
"is_default": False,
|
||||
},
|
||||
"seedance-2.0-fast": {
|
||||
"key": "seedance-2.0-fast",
|
||||
"display_name": "Seedance 2.0 Fast — 快速低成本",
|
||||
"model_id": "doubao-seedance-2-0-fast-260128",
|
||||
"provider": "doubao",
|
||||
"supports_audio": True,
|
||||
"supported_resolutions": ["480p", "720p"],
|
||||
"max_duration": 15,
|
||||
"billing_mode": "token",
|
||||
"is_default": False,
|
||||
},
|
||||
"seedance-2.0-mini": {
|
||||
"key": "seedance-2.0-mini",
|
||||
"display_name": "Seedance 2.0 Mini — 低成本测试",
|
||||
"model_id": "doubao-seedance-2-0-mini-260615",
|
||||
"provider": "doubao",
|
||||
"supports_audio": True,
|
||||
"supported_resolutions": ["480p", "720p"],
|
||||
"max_duration": 15,
|
||||
"billing_mode": "token",
|
||||
"is_default": False,
|
||||
},
|
||||
"wan-3.0": {
|
||||
"key": "wan-3.0",
|
||||
"display_name": "Wan 3.0 — 通义万相(阿里云)",
|
||||
"model_id": "wan3.0-video",
|
||||
"provider": "dashscope",
|
||||
"supports_audio": True,
|
||||
"supported_resolutions": ["480p", "720p", "1080p"],
|
||||
"max_duration": 30,
|
||||
"billing_mode": "per_second",
|
||||
"is_default": False,
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def get_viral_video_model_config(model_key: str | None) -> dict:
|
||||
"""获取模型配置,未知 key 回落到默认 seedance-2.5。"""
|
||||
key = (model_key or "").strip().lower()
|
||||
if key and key in VIRAL_VIDEO_MODEL_CONFIG:
|
||||
return VIRAL_VIDEO_MODEL_CONFIG[key]
|
||||
return VIRAL_VIDEO_MODEL_CONFIG["seedance-2.5"]
|
||||
|
||||
|
||||
def list_viral_video_models(
|
||||
include_placeholder: bool = False,
|
||||
dashscope_available: bool = False,
|
||||
) -> list[dict]:
|
||||
"""返回前端可用的模型列表(供 GET /api/v1/viral-video/models 端点用)。"""
|
||||
out: list[dict] = []
|
||||
for _k, cfg in VIRAL_VIDEO_MODEL_CONFIG.items():
|
||||
if cfg.get("_placeholder") and not include_placeholder:
|
||||
continue
|
||||
if cfg.get("provider") == "dashscope" and not dashscope_available:
|
||||
continue
|
||||
out.append(
|
||||
{
|
||||
"key": cfg["key"],
|
||||
"display_name": cfg["display_name"],
|
||||
"supports_audio": bool(cfg.get("supports_audio", True)),
|
||||
"supported_resolutions": list(cfg.get("supported_resolutions", ["720p"])),
|
||||
"max_duration": int(cfg.get("max_duration", 15)),
|
||||
"billing_mode": cfg.get("billing_mode", "token"),
|
||||
"is_default": bool(cfg.get("is_default", False)),
|
||||
}
|
||||
)
|
||||
return out
|
||||
|
||||
|
||||
def _match_model_prefix(model: str | None) -> str:
|
||||
"""匹配 model key(支持全部内部别名,未知回落到 seedance-2.5)。
|
||||
|
||||
按 key 长度从长到短匹配,避免 "seedance-2.0-fast" 被 "seedance-2.0" 前缀命中。
|
||||
"""
|
||||
mm = (model or "").strip().lower()
|
||||
for k in sorted(VIRAL_VIDEO_MODEL_CONFIG.keys(), key=len, reverse=True):
|
||||
if mm == k or mm.startswith(k):
|
||||
return k
|
||||
return "seedance-2.5"
|
||||
|
||||
|
||||
def _infer_resolution_key(width: int, height: int) -> str:
|
||||
"""从实际 (width, height) 用短边推断 resolution key。"""
|
||||
short = min(int(width or 720), int(height or 720))
|
||||
if short >= 1900:
|
||||
return "4k"
|
||||
if short >= 1000:
|
||||
return "1080p"
|
||||
if short >= 650:
|
||||
return "720p"
|
||||
return "480p"
|
||||
|
||||
|
||||
def calculate_viral_video_credits_with_breakdown(
|
||||
duration_seconds: int,
|
||||
width: int,
|
||||
height: int,
|
||||
model: str = "seedance-2.5",
|
||||
has_video_input: bool = False,
|
||||
actual_tokens: int | None = None,
|
||||
fps: int = VIRAL_VIDEO_FPS,
|
||||
) -> tuple[float, dict]:
|
||||
"""计算爆款视频所需积分(1 积分 = 1 元),并返回计费公式明细。
|
||||
|
||||
单价/固定成本/利润系数/封顶从 feature_pricing_configs(viral_video)读取;
|
||||
DB 不可用时回落与现状一致的内置兜底配置。
|
||||
|
||||
公式:
|
||||
tokens = duration * width * height * fps / 1024
|
||||
video_cost = tokens / 1_000_000 * model_token_price
|
||||
total = round((video_cost + fixed_cost) * profit_multiplier, 2)
|
||||
price_cap > 0 时封顶取 min
|
||||
若传入 actual_tokens 则用它替代计算值。
|
||||
|
||||
Returns:
|
||||
(credits, breakdown) 二元组:
|
||||
- credits: 四舍五入保留两位小数的最终积分
|
||||
- breakdown: dict,包含 tokens / video_cost / fixed_cost / profit_multiplier /
|
||||
model_price / width / height / fps / feature_enabled / charged / price_cap
|
||||
字段,便于前端展示计费明细。功能关闭时 credits=0、charged=False。
|
||||
"""
|
||||
w = max(1, int(width or 1))
|
||||
h = max(1, int(height or 1))
|
||||
effective_fps = int(fps or VIRAL_VIDEO_FPS)
|
||||
|
||||
prefix = _match_model_prefix(model)
|
||||
cfg = get_viral_video_model_config(prefix)
|
||||
res_key = _infer_resolution_key(w, h)
|
||||
billing = cfg.get("billing_mode", "token")
|
||||
dur = max(1, int(duration_seconds or 15))
|
||||
|
||||
# ── 从 DB 配置(兜底内置)取计费参数 ──
|
||||
feature_cfg = feature_pricing_service.get_feature_config("viral_video")
|
||||
# 注意:此处 feature_enabled 只表示“功能自身开关”,不并入全局 points_enabled
|
||||
# 总开关(保持与旧版计费函数行为一致:价格照常计算)。全局总开关由业务层
|
||||
# (route/worker)通过 feature_pricing_service.is_feature_enabled 统一把关。
|
||||
feature_enabled = bool(feature_cfg.is_enabled) if feature_cfg is not None else True
|
||||
model_pricing = feature_cfg.model_pricing if feature_cfg is not None else {}
|
||||
fixed_cost = float(feature_cfg.fixed_cost) if feature_cfg is not None else float(VIRAL_VIDEO_FIXED_COST)
|
||||
multiplier = (
|
||||
float(feature_cfg.profit_multiplier)
|
||||
if feature_cfg is not None and feature_cfg.profit_multiplier > 0
|
||||
else float(VIRAL_VIDEO_PROFIT_MULTIPLIER)
|
||||
)
|
||||
price_cap = float(feature_cfg.price_cap) if feature_cfg is not None else 0.0
|
||||
|
||||
# 单价:优先配置 dict;复刻旧版回落规则——精确 key 取不到时,回落
|
||||
# seedance-2.5 同分辨率 False 单价;最终兜底 70.0。
|
||||
price = feature_pricing_service.lookup_model_price(model_pricing, prefix, res_key, bool(has_video_input))
|
||||
if price is None:
|
||||
# 配置表未命中:先尝试配置里的 seedance-2.5/False
|
||||
if prefix != "seedance-2.5" or bool(has_video_input):
|
||||
price = feature_pricing_service.lookup_model_price(model_pricing, "seedance-2.5", res_key, False)
|
||||
if price is None:
|
||||
key = (prefix, res_key, bool(has_video_input))
|
||||
price = VIRAL_VIDEO_MODEL_PRICES.get(key)
|
||||
if price is None:
|
||||
price = VIRAL_VIDEO_MODEL_PRICES.get(("seedance-2.5", res_key, False), 70.0)
|
||||
|
||||
if billing == "per_second":
|
||||
tokens = 0.0
|
||||
video_cost = dur * float(price)
|
||||
billing_unit = "second"
|
||||
else:
|
||||
if actual_tokens is not None and actual_tokens > 0:
|
||||
tokens = float(actual_tokens)
|
||||
else:
|
||||
tokens = dur * w * h * effective_fps / 1024.0
|
||||
video_cost = tokens / 1_000_000.0 * float(price)
|
||||
billing_unit = "token"
|
||||
|
||||
if not feature_enabled:
|
||||
# 功能关闭(is_enabled=false 或全局 points 关闭):不扣费,明细照旧返回
|
||||
credits = 0.0
|
||||
raw_total = (video_cost + fixed_cost) * multiplier
|
||||
breakdown = {
|
||||
"tokens": float(tokens),
|
||||
"video_cost": float(video_cost),
|
||||
"fixed_cost": float(fixed_cost),
|
||||
"profit_multiplier": float(multiplier),
|
||||
"price_cap": float(price_cap or 0.0),
|
||||
"model_price": float(price),
|
||||
"model_key": prefix,
|
||||
"billing_mode": billing,
|
||||
"billing_unit": billing_unit,
|
||||
"width": int(w),
|
||||
"height": int(h),
|
||||
"fps": int(effective_fps),
|
||||
"duration": dur,
|
||||
"feature_enabled": False,
|
||||
"charged": False,
|
||||
"raw_price": round(float(raw_total), 4),
|
||||
}
|
||||
return credits, breakdown
|
||||
|
||||
total = (video_cost + fixed_cost) * multiplier
|
||||
if price_cap and price_cap > 0:
|
||||
total = min(total, price_cap)
|
||||
credits = round(float(total), 2)
|
||||
breakdown = {
|
||||
"tokens": float(tokens),
|
||||
"video_cost": float(video_cost),
|
||||
"fixed_cost": float(fixed_cost),
|
||||
"profit_multiplier": float(multiplier),
|
||||
"price_cap": float(price_cap or 0.0),
|
||||
"model_price": float(price),
|
||||
"model_key": prefix,
|
||||
"billing_mode": billing,
|
||||
"billing_unit": billing_unit,
|
||||
"width": int(w),
|
||||
"height": int(h),
|
||||
"fps": int(effective_fps),
|
||||
"duration": dur,
|
||||
"feature_enabled": True,
|
||||
"charged": True,
|
||||
}
|
||||
return credits, breakdown
|
||||
|
||||
|
||||
def calculate_viral_video_credits(
|
||||
duration_seconds: int,
|
||||
width: int,
|
||||
height: int,
|
||||
model: str = "seedance-2.5",
|
||||
has_video_input: bool = False,
|
||||
actual_tokens: int | None = None,
|
||||
fps: int = VIRAL_VIDEO_FPS,
|
||||
) -> float:
|
||||
"""计算爆款视频所需积分(1 积分 = 1 元),仅返回积分值(向后兼容包装器)。
|
||||
|
||||
内部调用 calculate_viral_video_credits_with_breakdown,仅返回 credits 部分,
|
||||
保持旧调用方签名与返回值类型不变。
|
||||
|
||||
公式:
|
||||
tokens = duration * width * height * fps / 1024
|
||||
video_cost = tokens / 1_000_000 * model_token_price
|
||||
total = round((video_cost + fixed_cost) * profit_multiplier, 2)
|
||||
若传入 actual_tokens 则用它替代计算值。
|
||||
"""
|
||||
credits, _ = calculate_viral_video_credits_with_breakdown(
|
||||
duration_seconds=duration_seconds,
|
||||
width=width,
|
||||
height=height,
|
||||
model=model,
|
||||
has_video_input=has_video_input,
|
||||
actual_tokens=actual_tokens,
|
||||
fps=fps,
|
||||
)
|
||||
return credits
|
||||
|
||||
|
||||
# ============ 场景定义 ============
|
||||
# 每个场景: base_points(基础积分), unit(计费单位), name(显示名称), dynamic(是否动态定价)
|
||||
# 说明:爆款视频(viral_video)走动态定价(预扣→结算多退少补),因此不使用 @points_gate
|
||||
# 装饰器,base_points=0,dynamic=True;前端展示场景列表时仍可看到。
|
||||
# 每个场景: base_points(基础积分), unit(计费单位), name(显示名称)
|
||||
|
||||
POINTS_SCENES: dict[str, dict] = {
|
||||
"ai_voice": {
|
||||
"base_points": 1,
|
||||
"unit": "分钟",
|
||||
"name": "AI 配音",
|
||||
"description": "AI 配音每分钟消耗 1 积分(免费用户上浮 15%,会员 8~9 折)",
|
||||
},
|
||||
"ai_video": {
|
||||
"base_points": 3,
|
||||
"unit": "条",
|
||||
"name": "智能混剪",
|
||||
"extra_per_30s": 1,
|
||||
"description": "智能混剪每条 3 积分起,视频超过 30 秒后每 30 秒加 1 积分;免费用户每日 2 条免费额度",
|
||||
},
|
||||
"ai_digital_human": {
|
||||
"base_points": 15,
|
||||
"unit": "分钟",
|
||||
"name": "AI 数字人",
|
||||
"description": "AI 数字人每分钟消耗 15 积分",
|
||||
},
|
||||
"voice_clone_train": {
|
||||
"base_points": 0,
|
||||
"unit": "次",
|
||||
@@ -391,16 +39,23 @@ POINTS_SCENES: dict[str, dict] = {
|
||||
"name": "声音克隆合成",
|
||||
"description": "克隆音色合成每分钟消耗 1 积分",
|
||||
},
|
||||
"viral_video": {
|
||||
"base_points": 0,
|
||||
"douyin_extract": {
|
||||
"base_points": 1,
|
||||
"unit": "次",
|
||||
"name": "爆款视频",
|
||||
"dynamic": True,
|
||||
"description": "爆款视频动态定价(按视频时长/分辨率/模型计算,预扣→结算多退少补)",
|
||||
"name": "抖音链接提取",
|
||||
"description": "抖音文案提取每次 1 积分",
|
||||
},
|
||||
"ai_rewrite": {"base_points": 1, "unit": "次", "name": "AI 改写文案", "description": "AI 改写文案每次 1 积分"},
|
||||
"ai_title": {
|
||||
"base_points": 1,
|
||||
"unit": "次",
|
||||
"name": "AI 标题生成",
|
||||
"description": "AI 生成标题每次 1 积分(免费用户实际上浮后 2 积分/次)",
|
||||
},
|
||||
"ai_cover": {"base_points": 1, "unit": "张", "name": "AI 封面生成", "description": "AI 封面生成每张 1 积分"},
|
||||
}
|
||||
|
||||
# 免费用户积分消耗上浮系数(仅对 voice_clone_synth 生效)
|
||||
# 免费用户积分消耗上浮系数
|
||||
FREE_USER_MULTIPLIER = 1.15
|
||||
|
||||
# ============ 积分包定义 ============
|
||||
@@ -426,6 +81,9 @@ MEMBER_DISCOUNT: dict[str, float] = {
|
||||
"yearly": 0.8,
|
||||
}
|
||||
|
||||
# 每日免费混剪次数(免费用户)
|
||||
DAILY_FREE_CLIP_LIMIT = 2
|
||||
|
||||
|
||||
def calculate_points_cost(
|
||||
scene_key: str,
|
||||
@@ -433,45 +91,48 @@ def calculate_points_cost(
|
||||
quantity: int = 1,
|
||||
duration_minutes: float = 0,
|
||||
member_type: str | None = None,
|
||||
) -> float:
|
||||
) -> int:
|
||||
"""计算指定场景的积分消耗。
|
||||
|
||||
Args:
|
||||
scene_key: 场景标识(当前支持 voice_clone_train/voice_clone_synth/viral_video;
|
||||
viral_video 为动态定价场景,此处返回 0,由业务侧调用
|
||||
calculate_viral_video_credits 手动计算)
|
||||
scene_key: 场景标识,如 "ai_voice"、"ai_video"
|
||||
is_member: 是否付费会员
|
||||
quantity: 数量(按次计费场景)
|
||||
duration_minutes: 时长分钟数(按时长计费场景)
|
||||
member_type: 会员类型 (monthly/quarterly/yearly),用于折扣
|
||||
|
||||
Returns:
|
||||
实际消耗积分(float;已含免费用户 ×1.15 上浮或会员折扣);免费/动态/已下线场景统一返回 0。
|
||||
实际消耗积分(已含免费用户 ×1.15 上浮或会员折扣)
|
||||
|
||||
Raises:
|
||||
ValueError: 未知场景标识
|
||||
"""
|
||||
scene = POINTS_SCENES.get(scene_key)
|
||||
if not scene:
|
||||
# 已下线/未注册的场景统一返回 0(免费),保持向后兼容
|
||||
return 0.0
|
||||
|
||||
# 动态定价场景(如 viral_video)由业务侧手动计算,这里统一返回 0
|
||||
if scene.get("dynamic"):
|
||||
return 0.0
|
||||
raise ValueError(f"Unknown points scene: {scene_key}")
|
||||
|
||||
base = scene["base_points"]
|
||||
if base == 0:
|
||||
return 0.0
|
||||
return 0
|
||||
|
||||
# —— 计算基础消耗 ——
|
||||
unit = scene["unit"]
|
||||
if unit == "分钟":
|
||||
total_base = base * max(1, math.ceil(duration_minutes))
|
||||
elif unit in ("次", "张"):
|
||||
elif unit in ("条", "次", "张"):
|
||||
total_base = base * quantity
|
||||
# 混剪特殊逻辑:视频超过 30s 后每 +30s 额外加 1 积分
|
||||
if scene_key == "ai_video" and duration_minutes > 0.5:
|
||||
extra_segments = math.ceil((duration_minutes * 60 - 30) / 30)
|
||||
if extra_segments > 0:
|
||||
total_base += scene.get("extra_per_30s", 1) * extra_segments
|
||||
else:
|
||||
total_base = base
|
||||
|
||||
# —— 会员折扣 / 免费用户上浮 ——
|
||||
if is_member and member_type and member_type in MEMBER_DISCOUNT:
|
||||
total_base = max(1, math.floor(total_base * MEMBER_DISCOUNT[member_type]))
|
||||
elif not is_member:
|
||||
total_base = math.ceil(total_base * FREE_USER_MULTIPLIER)
|
||||
|
||||
return float(total_base)
|
||||
return total_base
|
||||
|
||||
@@ -13,6 +13,7 @@ from typing import Any
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from packages.domain.points_rules import (
|
||||
DAILY_FREE_CLIP_LIMIT,
|
||||
POINTS_PACKAGES,
|
||||
)
|
||||
|
||||
@@ -83,7 +84,7 @@ class PointsService:
|
||||
|
||||
# ──────────────── 余额检查 ────────────────
|
||||
|
||||
def check_balance(self, user_id: str, amount: float, db: Session) -> dict[str, Any]:
|
||||
def check_balance(self, user_id: str, amount: int, db: Session) -> dict[str, Any]:
|
||||
"""检查余额是否足够。"""
|
||||
account_data = self.get_or_create_account(user_id, db)
|
||||
balance = account_data["balance"]
|
||||
@@ -99,7 +100,7 @@ class PointsService:
|
||||
def deduct_points(
|
||||
self,
|
||||
user_id: str,
|
||||
amount: float,
|
||||
amount: int,
|
||||
source: str,
|
||||
db: Session,
|
||||
description: str = "",
|
||||
@@ -108,7 +109,7 @@ class PointsService:
|
||||
"""扣减积分(事务性:SELECT FOR UPDATE → 检查余额 → 扣减 → 流水 → 同步用户表)。
|
||||
|
||||
Returns:
|
||||
{"success": True/False, "balance": float, "transaction_id": str|None}
|
||||
{"success": True/False, "balance": int, "transaction_id": str|None}
|
||||
"""
|
||||
PointsAccountModel, PointsTransactionModel, _, _, UserModel = _get_models()
|
||||
|
||||
@@ -172,7 +173,7 @@ class PointsService:
|
||||
except Exception:
|
||||
db.rollback()
|
||||
logger.exception(
|
||||
"积分扣减失败: user_id=%s, amount=%.2f, source=%s",
|
||||
"积分扣减失败: user_id=%s, amount=%d, source=%s",
|
||||
user_id,
|
||||
amount,
|
||||
source,
|
||||
@@ -184,7 +185,7 @@ class PointsService:
|
||||
def add_points(
|
||||
self,
|
||||
user_id: str,
|
||||
amount: float,
|
||||
amount: int,
|
||||
source: str,
|
||||
db: Session,
|
||||
description: str = "",
|
||||
@@ -241,7 +242,7 @@ class PointsService:
|
||||
except Exception:
|
||||
db.rollback()
|
||||
logger.exception(
|
||||
"积分增加失败: user_id=%s, amount=%.2f, source=%s",
|
||||
"积分增加失败: user_id=%s, amount=%d, source=%s",
|
||||
user_id,
|
||||
amount,
|
||||
source,
|
||||
@@ -253,7 +254,7 @@ class PointsService:
|
||||
def refund_points(
|
||||
self,
|
||||
user_id: str,
|
||||
amount: float,
|
||||
amount: int,
|
||||
source: str,
|
||||
db: Session,
|
||||
ref_id: str = "",
|
||||
@@ -269,92 +270,6 @@ class PointsService:
|
||||
ref_id=ref_id,
|
||||
)
|
||||
|
||||
# ──────────────── 爆款视频(viral_video)动态定价 ────────────────
|
||||
|
||||
def deduct_viral_video(self, user_id: str, credits: float, job_id: str, db: Session) -> dict[str, Any]:
|
||||
"""爆款视频预扣积分(confirm-copy 阶段)。"""
|
||||
return self.deduct_points(
|
||||
user_id=user_id,
|
||||
amount=float(credits or 0),
|
||||
source="viral_video",
|
||||
db=db,
|
||||
description="爆款视频生成",
|
||||
ref_id=job_id,
|
||||
)
|
||||
|
||||
def settle_viral_video(
|
||||
self,
|
||||
user_id: str,
|
||||
estimated: float,
|
||||
actual: float,
|
||||
txn_id: str,
|
||||
db: Session,
|
||||
) -> dict[str, Any]:
|
||||
"""爆款视频完成后按实际 tokens 结算(多退少补)。
|
||||
|
||||
- actual < estimated: 退差额
|
||||
- actual > estimated: 补扣差额(余额不足时记 warning,不阻塞完成)
|
||||
- |diff| < 0.01: 不动
|
||||
"""
|
||||
diff = round(float(actual or 0) - float(estimated or 0), 2)
|
||||
if abs(diff) < 0.01:
|
||||
return {"success": True, "action": "none", "diff": 0.0}
|
||||
if diff < 0:
|
||||
refund = round(-diff, 2)
|
||||
try:
|
||||
res = self.refund_points(
|
||||
user_id=user_id,
|
||||
amount=refund,
|
||||
source="viral_video",
|
||||
db=db,
|
||||
ref_id=txn_id,
|
||||
description="爆款视频结算退费",
|
||||
)
|
||||
return {"success": bool(res.get("success")), "action": "refund", "diff": -refund, "amount": refund}
|
||||
except Exception:
|
||||
logger.exception("[viral_video] 结算退费异常 user_id=%s refund=%.2f", user_id, refund)
|
||||
return {"success": False, "action": "refund", "diff": -refund}
|
||||
else:
|
||||
extra = round(diff, 2)
|
||||
try:
|
||||
res = self.deduct_points(
|
||||
user_id=user_id,
|
||||
amount=extra,
|
||||
source="viral_video",
|
||||
db=db,
|
||||
description="爆款视频结算补扣",
|
||||
ref_id=txn_id,
|
||||
)
|
||||
if not res.get("success"):
|
||||
logger.warning(
|
||||
"[viral_video] 结算补扣余额不足 user_id=%s extra=%.2f balance=%s (不阻塞任务完成)",
|
||||
user_id,
|
||||
extra,
|
||||
res.get("balance"),
|
||||
)
|
||||
return {"success": bool(res.get("success")), "action": "deduct", "diff": extra, "amount": extra}
|
||||
except Exception:
|
||||
logger.exception("[viral_video] 结算补扣异常 user_id=%s extra=%.2f", user_id, extra)
|
||||
return {"success": False, "action": "deduct", "diff": extra}
|
||||
|
||||
def refund_viral_video(self, user_id: str, credits: float, txn_id: str, db: Session) -> dict[str, Any]:
|
||||
"""爆款视频失败全额退款。"""
|
||||
amount = float(credits or 0)
|
||||
if amount <= 0:
|
||||
return {"success": True, "action": "none", "amount": 0.0}
|
||||
try:
|
||||
return self.refund_points(
|
||||
user_id=user_id,
|
||||
amount=amount,
|
||||
source="viral_video",
|
||||
db=db,
|
||||
ref_id=txn_id,
|
||||
description="爆款视频失败退款",
|
||||
)
|
||||
except Exception:
|
||||
logger.exception("[viral_video] 失败退款异常 user_id=%s amount=%.2f", user_id, amount)
|
||||
return {"success": False, "action": "refund", "amount": amount}
|
||||
|
||||
# ──────────────── 流水查询 ────────────────
|
||||
|
||||
def get_transactions(
|
||||
@@ -409,16 +324,132 @@ class PointsService:
|
||||
"page_size": page_size,
|
||||
}
|
||||
|
||||
# ──────────────── 每日免费混剪额度(已下线:智能混剪全免费) ────────────────
|
||||
# ──────────────── 每日免费混剪额度 ────────────────
|
||||
|
||||
def _daily_key(self, user_id: str) -> str:
|
||||
"""生成 Redis 每日额度 key。格式: daily_usage:{user_id}:{YYYYMMDD}:free_clip"""
|
||||
today = datetime.now(UTC).strftime("%Y%m%d")
|
||||
return f"daily_usage:{user_id}:{today}:free_clip"
|
||||
|
||||
def check_daily_free_clip(self, user_id: str, db: Session) -> bool:
|
||||
"""检查今日是否还有免费混剪额度。
|
||||
|
||||
优先查 Redis,Redis 不可用时降级到 DB。
|
||||
"""
|
||||
redis_client = _get_redis_client()
|
||||
if redis_client:
|
||||
try:
|
||||
key = self._daily_key(user_id)
|
||||
current = redis_client.get(key)
|
||||
if current is None:
|
||||
return True
|
||||
return int(current) < DAILY_FREE_CLIP_LIMIT
|
||||
except Exception:
|
||||
logger.warning("Redis 不可用,降级到 DB 查询每日额度")
|
||||
|
||||
# 降级到 DB
|
||||
_, _, _, DailyUsageRecordModel, _ = _get_models()
|
||||
today_start = datetime.now(UTC).replace(hour=0, minute=0, second=0, microsecond=0)
|
||||
record = (
|
||||
db.query(DailyUsageRecordModel)
|
||||
.filter(
|
||||
DailyUsageRecordModel.user_id == user_id,
|
||||
DailyUsageRecordModel.usage_type == "free_clip",
|
||||
DailyUsageRecordModel.usage_date >= today_start,
|
||||
)
|
||||
.first()
|
||||
)
|
||||
if record is None:
|
||||
return True
|
||||
return record.count < DAILY_FREE_CLIP_LIMIT
|
||||
|
||||
def record_daily_free_clip(self, user_id: str, db: Session) -> bool:
|
||||
"""记录使用一次免费混剪。
|
||||
|
||||
先 INCR Redis;如果超限回退 Redis。DB 使用 upsert 语义(唯一约束)。
|
||||
"""
|
||||
redis_client = _get_redis_client()
|
||||
if redis_client:
|
||||
try:
|
||||
key = self._daily_key(user_id)
|
||||
new_count = redis_client.incr(key)
|
||||
if new_count == 1:
|
||||
redis_client.expire(key, 48 * 3600) # TTL 48h
|
||||
if new_count <= DAILY_FREE_CLIP_LIMIT:
|
||||
return True
|
||||
# 超限,回退 Redis
|
||||
redis_client.decr(key)
|
||||
except Exception:
|
||||
logger.warning("Redis 不可用,降级到 DB 记录每日额度")
|
||||
|
||||
# 降级/兜底到 DB(upsert 语义)
|
||||
_, _, _, DailyUsageRecordModel, _ = _get_models()
|
||||
today_start = datetime.now(UTC).replace(hour=0, minute=0, second=0, microsecond=0)
|
||||
|
||||
record = (
|
||||
db.query(DailyUsageRecordModel)
|
||||
.filter(
|
||||
DailyUsageRecordModel.user_id == user_id,
|
||||
DailyUsageRecordModel.usage_type == "free_clip",
|
||||
DailyUsageRecordModel.usage_date >= today_start,
|
||||
)
|
||||
.first()
|
||||
)
|
||||
|
||||
if record is None:
|
||||
if DAILY_FREE_CLIP_LIMIT <= 0:
|
||||
return False
|
||||
record = DailyUsageRecordModel(
|
||||
id=uuid.uuid4().hex,
|
||||
user_id=user_id,
|
||||
usage_type="free_clip",
|
||||
usage_date=datetime.now(UTC),
|
||||
count=1,
|
||||
)
|
||||
db.add(record)
|
||||
else:
|
||||
if record.count >= DAILY_FREE_CLIP_LIMIT:
|
||||
return False
|
||||
record.count += 1
|
||||
|
||||
db.commit()
|
||||
return True
|
||||
|
||||
def get_daily_usage(self, user_id: str, db: Session) -> dict[str, Any]:
|
||||
"""查询今日免费额度使用情况(智能混剪已全免费,返回 unlimited)。"""
|
||||
"""查询今日免费额度使用情况。"""
|
||||
redis_client = _get_redis_client()
|
||||
used = 0
|
||||
|
||||
if redis_client:
|
||||
try:
|
||||
key = self._daily_key(user_id)
|
||||
val = redis_client.get(key)
|
||||
used = int(val) if val else 0
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
if used == 0:
|
||||
# 从 DB 查
|
||||
_, _, _, DailyUsageRecordModel, _ = _get_models()
|
||||
today_start = datetime.now(UTC).replace(hour=0, minute=0, second=0, microsecond=0)
|
||||
record = (
|
||||
db.query(DailyUsageRecordModel)
|
||||
.filter(
|
||||
DailyUsageRecordModel.user_id == user_id,
|
||||
DailyUsageRecordModel.usage_type == "free_clip",
|
||||
DailyUsageRecordModel.usage_date >= today_start,
|
||||
)
|
||||
.first()
|
||||
)
|
||||
used = record.count if record else 0
|
||||
|
||||
now = datetime.now(UTC)
|
||||
tomorrow = (now + timedelta(days=1)).replace(hour=0, minute=0, second=0, microsecond=0)
|
||||
|
||||
return {
|
||||
"free_clips_used": 0,
|
||||
"free_clips_limit": -1, # -1 表示 unlimited
|
||||
"free_clips_remaining": -1,
|
||||
"free_clips_used": used,
|
||||
"free_clips_limit": DAILY_FREE_CLIP_LIMIT,
|
||||
"free_clips_remaining": max(0, DAILY_FREE_CLIP_LIMIT - used),
|
||||
"reset_at": tomorrow.isoformat(),
|
||||
}
|
||||
|
||||
|
||||
@@ -1,134 +0,0 @@
|
||||
"""系统配置领域实体 — #2246.
|
||||
|
||||
承载 setting_type(bool/int/float/string/json)及 setting_value 的
|
||||
序列化/反序列化规则;与 DB、框架无关。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import UTC, datetime
|
||||
from typing import Any
|
||||
from uuid import uuid4
|
||||
|
||||
SETTING_TYPE_BOOL = "bool"
|
||||
SETTING_TYPE_INT = "int"
|
||||
SETTING_TYPE_FLOAT = "float"
|
||||
SETTING_TYPE_STRING = "string"
|
||||
SETTING_TYPE_JSON = "json"
|
||||
|
||||
VALID_SETTING_TYPES = {
|
||||
SETTING_TYPE_BOOL,
|
||||
SETTING_TYPE_INT,
|
||||
SETTING_TYPE_FLOAT,
|
||||
SETTING_TYPE_STRING,
|
||||
SETTING_TYPE_JSON,
|
||||
}
|
||||
|
||||
|
||||
class SystemSettingError(ValueError):
|
||||
"""系统配置类型或序列化错误."""
|
||||
|
||||
|
||||
def infer_setting_type(value: Any) -> str:
|
||||
"""根据 Python 值推断 setting_type(bool 必须先于 int 判断)."""
|
||||
if isinstance(value, bool):
|
||||
return SETTING_TYPE_BOOL
|
||||
if isinstance(value, int):
|
||||
return SETTING_TYPE_INT
|
||||
if isinstance(value, float):
|
||||
return SETTING_TYPE_FLOAT
|
||||
if isinstance(value, str):
|
||||
return SETTING_TYPE_STRING
|
||||
return SETTING_TYPE_JSON
|
||||
|
||||
|
||||
def serialize_setting_value(value: Any, setting_type: str) -> str:
|
||||
"""把 Python 值按 setting_type 序列化为可入库的字符串."""
|
||||
if setting_type == SETTING_TYPE_BOOL:
|
||||
if not isinstance(value, bool):
|
||||
raise SystemSettingError(f"bool 配置值必须是布尔类型,收到 {value!r}")
|
||||
return "true" if value else "false"
|
||||
if setting_type == SETTING_TYPE_INT:
|
||||
if isinstance(value, bool) or not isinstance(value, int):
|
||||
raise SystemSettingError(f"int 配置值必须是整数,收到 {value!r}")
|
||||
return str(value)
|
||||
if setting_type == SETTING_TYPE_FLOAT:
|
||||
if isinstance(value, bool):
|
||||
raise SystemSettingError(f"float 配置值不能是布尔类型,收到 {value!r}")
|
||||
try:
|
||||
return repr(float(value))
|
||||
except (TypeError, ValueError) as exc:
|
||||
raise SystemSettingError(f"float 配置值非法:{value!r}") from exc
|
||||
if setting_type == SETTING_TYPE_STRING:
|
||||
if not isinstance(value, str):
|
||||
raise SystemSettingError(f"string 配置值必须是字符串,收到 {value!r}")
|
||||
return value
|
||||
if setting_type == SETTING_TYPE_JSON:
|
||||
try:
|
||||
return json.dumps(value, ensure_ascii=False)
|
||||
except (TypeError, ValueError) as exc:
|
||||
raise SystemSettingError(f"json 配置值无法序列化:{value!r}") from exc
|
||||
raise SystemSettingError(f"未知 setting_type: {setting_type}")
|
||||
|
||||
|
||||
def deserialize_setting_value(raw: str | None, setting_type: str) -> Any:
|
||||
"""把入库字符串按 setting_type 反序列化为 Python 值."""
|
||||
if raw is None:
|
||||
return None
|
||||
if setting_type == SETTING_TYPE_BOOL:
|
||||
return str(raw).strip().lower() in {"1", "true", "yes", "on"}
|
||||
if setting_type == SETTING_TYPE_INT:
|
||||
try:
|
||||
return int(str(raw).strip())
|
||||
except ValueError as exc:
|
||||
raise SystemSettingError(f"int 配置值损坏:{raw!r}") from exc
|
||||
if setting_type == SETTING_TYPE_FLOAT:
|
||||
try:
|
||||
return float(str(raw).strip())
|
||||
except ValueError as exc:
|
||||
raise SystemSettingError(f"float 配置值损坏:{raw!r}") from exc
|
||||
if setting_type == SETTING_TYPE_STRING:
|
||||
return raw
|
||||
if setting_type == SETTING_TYPE_JSON:
|
||||
try:
|
||||
return json.loads(raw)
|
||||
except (TypeError, ValueError) as exc:
|
||||
raise SystemSettingError(f"json 配置值损坏:{raw!r}") from exc
|
||||
raise SystemSettingError(f"未知 setting_type: {setting_type}")
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class SystemSetting:
|
||||
"""系统配置领域实体."""
|
||||
|
||||
setting_key: str
|
||||
setting_value: str | None = None
|
||||
setting_type: str = SETTING_TYPE_STRING
|
||||
description: str = ""
|
||||
is_public: bool = False
|
||||
category: str = "general"
|
||||
updated_by: str | None = None
|
||||
id: str = field(default_factory=lambda: uuid4().hex)
|
||||
created_at: datetime = field(default_factory=lambda: datetime.now(UTC))
|
||||
updated_at: datetime = field(default_factory=lambda: datetime.now(UTC))
|
||||
|
||||
def get_typed_value(self) -> Any:
|
||||
return deserialize_setting_value(self.setting_value, self.setting_type)
|
||||
|
||||
@classmethod
|
||||
def from_value(
|
||||
cls,
|
||||
setting_key: str,
|
||||
value: Any,
|
||||
setting_type: str | None = None,
|
||||
**kwargs: Any,
|
||||
) -> "SystemSetting":
|
||||
st = setting_type or infer_setting_type(value)
|
||||
return cls(
|
||||
setting_key=setting_key,
|
||||
setting_value=serialize_setting_value(value, st),
|
||||
setting_type=st,
|
||||
**kwargs,
|
||||
)
|
||||
@@ -69,6 +69,8 @@ class PromptType(StrEnum):
|
||||
STYLE_CONSTRAINT = "style_constraint"
|
||||
|
||||
|
||||
CREDITS_VIRAL_VIDEO_COST = 50
|
||||
|
||||
STAGE_LABELS = {
|
||||
ViralVideoStage.IMAGE_ANALYSIS: "图片分析",
|
||||
ViralVideoStage.VIDEO_ANALYSIS: "视频风格分析",
|
||||
@@ -87,9 +89,6 @@ 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 = ""
|
||||
@@ -122,10 +121,7 @@ class ViralVideoJob:
|
||||
phase_message: str = "" # 阶段中文提示文案,前端轮询直接展示
|
||||
heartbeat_at: datetime | None = None # worker 心跳时间,用于超时僵尸任务检测
|
||||
result_video_url: str = ""
|
||||
video_resolution: str = "720p"
|
||||
credits_prepaid: float = 0.0
|
||||
credits_transaction_id: str = ""
|
||||
credits_cost: float = 0.0
|
||||
credits_cost: int = 0
|
||||
error_msg: str = ""
|
||||
retry_count: int = 0
|
||||
started_at: datetime | None = None
|
||||
@@ -191,45 +187,13 @@ class ViralVideoJob:
|
||||
self.updated_at = datetime.now(timezone.utc)
|
||||
|
||||
def resume_from_image_analyzed(self, **kwargs) -> None:
|
||||
"""阶段2入口:允许从 IMAGE_ANALYZED/PENDING 首次进入,也允许从 COPY_GENERATED/COMPLETED/FAILED 重新生成文案。
|
||||
|
||||
重新生成时清空上一轮文案产物(copy_result/intent_result/storyboard/generated_copy_text),
|
||||
并重置 completed_at/result_video_url/error_msg,确保前端轮询能看到新的阶段2进度。
|
||||
"""
|
||||
_allowed = (
|
||||
ViralVideoStatus.IMAGE_ANALYZED,
|
||||
ViralVideoStatus.PENDING,
|
||||
ViralVideoStatus.COPY_GENERATED,
|
||||
ViralVideoStatus.COMPLETED,
|
||||
ViralVideoStatus.FAILED,
|
||||
)
|
||||
if self.status not in _allowed:
|
||||
if self.status not in (ViralVideoStatus.IMAGE_ANALYZED, ViralVideoStatus.PENDING):
|
||||
raise ValueError(f"Cannot resume from {self.status} to copy-gen")
|
||||
_is_regen = self.status in (
|
||||
ViralVideoStatus.COPY_GENERATED,
|
||||
ViralVideoStatus.COMPLETED,
|
||||
ViralVideoStatus.FAILED,
|
||||
)
|
||||
for k, v in kwargs.items():
|
||||
if hasattr(self, k) and v not in (None, "", []):
|
||||
setattr(self, k, v)
|
||||
if _is_regen:
|
||||
# 清空上一轮文案/视频产物,避免前端拿到旧数据
|
||||
self.intent_result = None
|
||||
self.copy_result = None
|
||||
self.storyboard = None
|
||||
self.generated_copy_text = ""
|
||||
self.result_video_url = ""
|
||||
self.current_stage = ""
|
||||
self.phase_message = ""
|
||||
self.error_msg = ""
|
||||
self.completed_at = None
|
||||
self.heartbeat_at = None
|
||||
_now = datetime.now(timezone.utc)
|
||||
self.started_at = _now
|
||||
self.heartbeat_at = _now
|
||||
self.status = ViralVideoStatus.RUNNING
|
||||
self.updated_at = _now
|
||||
self.updated_at = datetime.now(timezone.utc)
|
||||
|
||||
def resume_from_copy_generated(self, edited_copy: str | None = None) -> None:
|
||||
"""阶段2->阶段3:用户确认/编辑口播文案,开始跑 TTS+单次Seedance渲染。"""
|
||||
@@ -238,20 +202,14 @@ class ViralVideoJob:
|
||||
if edited_copy and isinstance(self.copy_result, dict):
|
||||
self.copy_result = {**self.copy_result, "voiceover_script": edited_copy}
|
||||
self.generated_copy_text = edited_copy
|
||||
_now = datetime.now(timezone.utc)
|
||||
self.started_at = _now
|
||||
self.heartbeat_at = _now
|
||||
self.status = ViralVideoStatus.RUNNING
|
||||
self.updated_at = _now
|
||||
self.updated_at = datetime.now(timezone.utc)
|
||||
|
||||
def resume_from_confirm(self) -> None:
|
||||
if self.status != ViralVideoStatus.WAIT_USER_CONFIRM:
|
||||
raise ValueError(f"Cannot resume from {self.status}")
|
||||
_now = datetime.now(timezone.utc)
|
||||
self.started_at = _now
|
||||
self.heartbeat_at = _now
|
||||
self.status = ViralVideoStatus.RUNNING
|
||||
self.updated_at = _now
|
||||
self.updated_at = datetime.now(timezone.utc)
|
||||
|
||||
def mark_completed(self, video_url: str) -> None:
|
||||
self.status = ViralVideoStatus.COMPLETED
|
||||
|
||||
@@ -185,6 +185,19 @@ def _execute_with_gate_impl(
|
||||
is_member = getattr(user, "is_member", False)
|
||||
member_type = getattr(user, "member_type", None)
|
||||
|
||||
if scene_key == "ai_video":
|
||||
from packages.domain.points_service import PointsService
|
||||
|
||||
svc = PointsService()
|
||||
if not is_member:
|
||||
if svc.check_daily_free_clip(user.id, db):
|
||||
svc.record_daily_free_clip(user.id, db)
|
||||
kwargs["_points_deducted"] = 0
|
||||
kwargs["_is_free_quota"] = True
|
||||
if is_async:
|
||||
return _run_async_impl(func, args, _filter_kwargs_impl(func, kwargs))
|
||||
return func(*args, **_filter_kwargs_impl(func, kwargs))
|
||||
|
||||
if per_unit is not None:
|
||||
total_points = per_unit
|
||||
else:
|
||||
|
||||
+51
-748
File diff suppressed because it is too large
Load Diff
@@ -1,62 +0,0 @@
|
||||
"""AI 配置版本号管理 — Redis 通知机制.
|
||||
|
||||
admin 后台修改 ai_models / ai_capability_configs 后调用 bump_version(),
|
||||
SaaS 端 AIRouter 每次取配置前比对版本号,变了才重新查 DB。
|
||||
|
||||
Redis key: xiaoxia:ai_config:version = 时间戳字符串
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import time
|
||||
from typing import Optional
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_REDIS_KEY = "xiaoxia:ai_config:version"
|
||||
|
||||
|
||||
def _get_redis_client():
|
||||
"""获取 Redis 客户端(复用 Celery broker 连接)."""
|
||||
try:
|
||||
import redis as _redis
|
||||
|
||||
from packages.shared.config import get_shared_settings
|
||||
|
||||
settings = get_shared_settings()
|
||||
redis_url = getattr(settings, "redis_url", None) or getattr(
|
||||
settings, "celery_broker_url", "redis://localhost:6379/0"
|
||||
)
|
||||
return _redis.Redis.from_url(redis_url, decode_responses=True, socket_timeout=2)
|
||||
except Exception as e:
|
||||
logger.warning("AI config version: Redis 客户端初始化失败: %s", e)
|
||||
return None
|
||||
|
||||
|
||||
def bump_version() -> str:
|
||||
"""写入新版本号(当前时间戳),返回版本号字符串。失败返回空串。"""
|
||||
r = _get_redis_client()
|
||||
if r is None:
|
||||
logger.warning("AI config bump_version: Redis 不可用,跳过版本号更新")
|
||||
return ""
|
||||
try:
|
||||
ver = str(int(time.time() * 1000))
|
||||
r.set(_REDIS_KEY, ver)
|
||||
logger.info("AI config version bumped to %s", ver)
|
||||
return ver
|
||||
except Exception as e:
|
||||
logger.warning("AI config bump_version 失败: %s", e)
|
||||
return ""
|
||||
|
||||
|
||||
def get_version() -> Optional[str]:
|
||||
"""读取当前版本号。Redis 不可用或异常返回 None。"""
|
||||
r = _get_redis_client()
|
||||
if r is None:
|
||||
return None
|
||||
try:
|
||||
return r.get(_REDIS_KEY)
|
||||
except Exception as e:
|
||||
logger.warning("AI config get_version 失败: %s", e)
|
||||
return None
|
||||
@@ -1,516 +0,0 @@
|
||||
"""AI 模型路由层 — 统一模型配置读取与客户端构建.
|
||||
|
||||
业务代码通过 AIRouter 获取客户端,不再硬编码 model/api_key/base_url。
|
||||
配置来源:DB ai_capability_configs JOIN ai_models → Redis 版本号缓存 → SharedSettings fallback。
|
||||
|
||||
使用方式:
|
||||
from packages.shared.ai_router import ai_router
|
||||
|
||||
client = ai_router.get_llm_client("intent_parsing")
|
||||
result = client.chat_completion(messages=[...])
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import threading
|
||||
from dataclasses import dataclass
|
||||
|
||||
from packages.shared.config import get_shared_settings
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
# ── 配置数据类 ──────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ModelConfig:
|
||||
"""单个 AI 模型配置(来自 ai_models 表)"""
|
||||
|
||||
id: str
|
||||
name: str
|
||||
provider: str
|
||||
model_key: str
|
||||
api_key: str
|
||||
api_base: str
|
||||
api_version: str | None
|
||||
status: str
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class CapabilityConfig:
|
||||
"""业务能力配置(来自 ai_capability_configs JOIN ai_models)"""
|
||||
|
||||
capability_key: str
|
||||
capability_name: str
|
||||
primary_model: ModelConfig | None
|
||||
lite_model: ModelConfig | None
|
||||
fallback_model: ModelConfig | None
|
||||
timeout_seconds: int
|
||||
max_retries: int
|
||||
max_tokens: int | None
|
||||
temperature: float | None
|
||||
concurrency: int
|
||||
extra_params: dict
|
||||
is_enabled: bool
|
||||
|
||||
|
||||
# ── 简单包装类(TTS / ImageGen / VideoGen)──────────────────────────────────
|
||||
|
||||
|
||||
class TTSClient:
|
||||
"""TTS 客户端(简单配置持有者,实际调用由 CosyVoiceService 完成)"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
provider: str,
|
||||
api_key: str,
|
||||
base_url: str,
|
||||
model: str,
|
||||
timeout: int = 60,
|
||||
extra_params: dict | None = None,
|
||||
):
|
||||
self.provider = provider
|
||||
self.api_key = api_key
|
||||
self.base_url = base_url
|
||||
self.model = model
|
||||
self.timeout = timeout
|
||||
self.extra_params = extra_params or {}
|
||||
|
||||
@property
|
||||
def is_available(self) -> bool:
|
||||
return bool(self.api_key and self.base_url and self.model)
|
||||
|
||||
|
||||
class ImageGenClient:
|
||||
"""图片生成客户端(简单配置持有者)"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
provider: str,
|
||||
api_key: str,
|
||||
base_url: str,
|
||||
model: str,
|
||||
timeout: int = 60,
|
||||
extra_params: dict | None = None,
|
||||
):
|
||||
self.provider = provider
|
||||
self.api_key = api_key
|
||||
self.base_url = base_url
|
||||
self.model = model
|
||||
self.timeout = timeout
|
||||
self.extra_params = extra_params or {}
|
||||
|
||||
@property
|
||||
def is_available(self) -> bool:
|
||||
return bool(self.api_key and self.base_url and self.model)
|
||||
|
||||
|
||||
class VideoGenClient:
|
||||
"""视频生成客户端(简单配置持有者)"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
provider: str,
|
||||
api_key: str,
|
||||
base_url: str,
|
||||
model: str,
|
||||
timeout: int = 600,
|
||||
extra_params: dict | None = None,
|
||||
):
|
||||
self.provider = provider
|
||||
self.api_key = api_key
|
||||
self.base_url = base_url
|
||||
self.model = model
|
||||
self.timeout = timeout
|
||||
self.extra_params = extra_params or {}
|
||||
|
||||
@property
|
||||
def is_available(self) -> bool:
|
||||
return bool(self.api_key and self.base_url and self.model)
|
||||
|
||||
|
||||
# ── DB Session 获取 ─────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def _get_session():
|
||||
"""获取 DB session,兼容 api / worker / 独立脚本场景"""
|
||||
from packages.adapters.sqlalchemy_impl.session import SessionLocal
|
||||
|
||||
if SessionLocal is not None:
|
||||
return SessionLocal()
|
||||
|
||||
try:
|
||||
from worker_app.db import SessionLocal as WorkerSL
|
||||
|
||||
if WorkerSL is not None:
|
||||
return WorkerSL()
|
||||
except ImportError:
|
||||
pass
|
||||
|
||||
try:
|
||||
from app.db import SessionLocal as ApiSL
|
||||
|
||||
if ApiSL is not None:
|
||||
return ApiSL()
|
||||
except ImportError:
|
||||
pass
|
||||
|
||||
return None
|
||||
|
||||
|
||||
# ── 核心路由类 ──────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class AIRouter:
|
||||
"""AI 模型路由器 — 统一配置读取与客户端构建.
|
||||
|
||||
缓存策略:
|
||||
1. 本地内存缓存 {capability_key: CapabilityConfig}
|
||||
2. 每次读取前比对 Redis 版本号,变了则清缓存重新查 DB
|
||||
3. DB 无配置 / Redis 不可用 → fallback 到 SharedSettings 环境变量
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
self._cache: dict[str, CapabilityConfig] = {}
|
||||
self._local_ver: str | None = None
|
||||
self._lock = threading.Lock()
|
||||
|
||||
def _check_version(self) -> bool:
|
||||
"""检查 Redis 版本号,变了返回 True(需要刷新缓存)"""
|
||||
from packages.shared.ai_config_version import get_version
|
||||
|
||||
current_ver = get_version()
|
||||
if current_ver is None:
|
||||
return False
|
||||
if self._local_ver != current_ver:
|
||||
return True
|
||||
return False
|
||||
|
||||
def _load_from_db(self, capability_key: str) -> CapabilityConfig | None:
|
||||
"""从 DB 加载配置(ai_capability_configs JOIN ai_models)"""
|
||||
session = _get_session()
|
||||
if session is None:
|
||||
logger.warning("AI Router: 无法获取 DB session")
|
||||
return None
|
||||
try:
|
||||
from sqlalchemy import text
|
||||
|
||||
sql = text("""
|
||||
SELECT
|
||||
cc.capability_key, cc.capability_name, cc.timeout_seconds,
|
||||
cc.max_retries, cc.max_tokens, cc.temperature,
|
||||
cc.concurrency, cc.extra_params, cc.is_enabled,
|
||||
pm.id AS pm_id, pm.name AS pm_name, pm.provider AS pm_provider,
|
||||
pm.model_key AS pm_model_key, pm.api_key AS pm_api_key,
|
||||
pm.api_base AS pm_api_base, pm.api_version AS pm_api_version,
|
||||
pm.status AS pm_status,
|
||||
lm.id AS lm_id, lm.name AS lm_name, lm.provider AS lm_provider,
|
||||
lm.model_key AS lm_model_key, lm.api_key AS lm_api_key,
|
||||
lm.api_base AS lm_api_base, lm.api_version AS lm_api_version,
|
||||
lm.status AS lm_status,
|
||||
fm.id AS fm_id, fm.name AS fm_name, fm.provider AS fm_provider,
|
||||
fm.model_key AS fm_model_key, fm.api_key AS fm_api_key,
|
||||
fm.api_base AS fm_api_base, fm.api_version AS fm_api_version,
|
||||
fm.status AS fm_status
|
||||
FROM ai_capability_configs cc
|
||||
LEFT JOIN ai_models pm ON cc.primary_model_id = pm.id AND pm.deleted_at IS NULL
|
||||
LEFT JOIN ai_models lm ON cc.lite_model_id = lm.id AND lm.deleted_at IS NULL
|
||||
LEFT JOIN ai_models fm ON cc.fallback_model_id = fm.id AND fm.deleted_at IS NULL
|
||||
WHERE cc.capability_key = :key AND cc.is_enabled = true
|
||||
""")
|
||||
row = session.execute(sql, {"key": capability_key}).first()
|
||||
if not row:
|
||||
return None
|
||||
|
||||
def _to_model(prefix: str) -> ModelConfig | None:
|
||||
mid = getattr(row, f"{prefix}_id", None)
|
||||
if not mid:
|
||||
return None
|
||||
return ModelConfig(
|
||||
id=mid,
|
||||
name=getattr(row, f"{prefix}_name", "") or "",
|
||||
provider=getattr(row, f"{prefix}_provider", "") or "",
|
||||
model_key=getattr(row, f"{prefix}_model_key", "") or "",
|
||||
api_key=getattr(row, f"{prefix}_api_key", "") or "",
|
||||
api_base=getattr(row, f"{prefix}_api_base", "") or "",
|
||||
api_version=getattr(row, f"{prefix}_api_version", None),
|
||||
status=getattr(row, f"{prefix}_status", "active") or "active",
|
||||
)
|
||||
|
||||
return CapabilityConfig(
|
||||
capability_key=row.capability_key,
|
||||
capability_name=row.capability_name,
|
||||
primary_model=_to_model("pm"),
|
||||
lite_model=_to_model("lm"),
|
||||
fallback_model=_to_model("fm"),
|
||||
timeout_seconds=row.timeout_seconds or 30,
|
||||
max_retries=row.max_retries or 1,
|
||||
max_tokens=row.max_tokens,
|
||||
temperature=row.temperature,
|
||||
concurrency=row.concurrency or 2,
|
||||
extra_params=row.extra_params or {},
|
||||
is_enabled=row.is_enabled,
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning("AI Router: DB 查询失败 (key=%s): %s", capability_key, e)
|
||||
return None
|
||||
finally:
|
||||
session.close()
|
||||
|
||||
def get_capability(self, key: str) -> CapabilityConfig | None:
|
||||
"""获取业务能力配置(带缓存)"""
|
||||
with self._lock:
|
||||
if self._check_version():
|
||||
self._cache.clear()
|
||||
from packages.shared.ai_config_version import get_version
|
||||
|
||||
self._local_ver = get_version()
|
||||
|
||||
if key in self._cache:
|
||||
return self._cache[key]
|
||||
|
||||
config = self._load_from_db(key)
|
||||
if config:
|
||||
self._cache[key] = config
|
||||
return config
|
||||
|
||||
def _get_model_or_fallback(self, cap: CapabilityConfig, variant: str = "primary") -> ModelConfig | None:
|
||||
"""按 variant 选择模型,不存在则降级。
|
||||
|
||||
- primary: primary → fallback
|
||||
- lite: lite → primary
|
||||
- fallback: fallback → primary(修复点:此前 fallback variant 被忽略,错误地使用了 primary 模型)
|
||||
"""
|
||||
if variant == "fallback":
|
||||
if cap.fallback_model:
|
||||
return cap.fallback_model
|
||||
if cap.primary_model:
|
||||
return cap.primary_model
|
||||
elif variant == "lite":
|
||||
if cap.lite_model:
|
||||
return cap.lite_model
|
||||
if cap.primary_model:
|
||||
return cap.primary_model
|
||||
else: # primary
|
||||
if cap.primary_model:
|
||||
return cap.primary_model
|
||||
if cap.fallback_model:
|
||||
return cap.fallback_model
|
||||
return None
|
||||
|
||||
# ── 构建客户端 ─────────────────────────────────────────────────────────
|
||||
|
||||
def _build_llm_client(self, model: ModelConfig, cap: CapabilityConfig):
|
||||
"""构建 LLM 客户端 — 返回 DoubaoClient 实例"""
|
||||
from packages.shared.ai_client import DoubaoClient
|
||||
|
||||
return DoubaoClient(
|
||||
api_key=model.api_key,
|
||||
base_url=model.api_base,
|
||||
model=model.model_key,
|
||||
timeout=cap.timeout_seconds,
|
||||
max_retries=cap.max_retries,
|
||||
max_tokens=cap.max_tokens,
|
||||
temperature=cap.temperature,
|
||||
extra_params=cap.extra_params,
|
||||
provider=model.provider,
|
||||
)
|
||||
|
||||
def _build_vision_client(self, model: ModelConfig, cap: CapabilityConfig):
|
||||
"""构建 VLM 客户端 — 返回 DoubaoClient 实例(DoubaoClient 已支持 vision_completion)"""
|
||||
from packages.shared.ai_client import DoubaoClient
|
||||
|
||||
return DoubaoClient(
|
||||
api_key=model.api_key,
|
||||
base_url=model.api_base,
|
||||
model=model.model_key,
|
||||
timeout=cap.timeout_seconds,
|
||||
max_retries=cap.max_retries,
|
||||
max_tokens=cap.max_tokens,
|
||||
temperature=cap.temperature,
|
||||
extra_params=cap.extra_params,
|
||||
provider=model.provider,
|
||||
)
|
||||
|
||||
def _build_tts_client(self, model: ModelConfig, cap: CapabilityConfig) -> TTSClient:
|
||||
return TTSClient(
|
||||
provider=model.provider,
|
||||
api_key=model.api_key,
|
||||
base_url=model.api_base,
|
||||
model=model.model_key,
|
||||
timeout=cap.timeout_seconds,
|
||||
extra_params=cap.extra_params,
|
||||
)
|
||||
|
||||
def _build_image_gen_client(self, model: ModelConfig, cap: CapabilityConfig) -> ImageGenClient:
|
||||
return ImageGenClient(
|
||||
provider=model.provider,
|
||||
api_key=model.api_key,
|
||||
base_url=model.api_base,
|
||||
model=model.model_key,
|
||||
timeout=cap.timeout_seconds,
|
||||
extra_params=cap.extra_params,
|
||||
)
|
||||
|
||||
def _build_video_gen_client(self, model: ModelConfig, cap: CapabilityConfig) -> VideoGenClient:
|
||||
return VideoGenClient(
|
||||
provider=model.provider,
|
||||
api_key=model.api_key,
|
||||
base_url=model.api_base,
|
||||
model=model.model_key,
|
||||
timeout=cap.timeout_seconds,
|
||||
extra_params=cap.extra_params,
|
||||
)
|
||||
|
||||
# ── 公开接口 ────────────────────────────────────────────────────────────
|
||||
|
||||
def get_llm_client(self, key: str, variant: str = "primary"):
|
||||
"""获取 LLM 客户端(返回 DoubaoClient 实例)"""
|
||||
cap = self.get_capability(key)
|
||||
if cap and cap.is_enabled:
|
||||
model = self._get_model_or_fallback(cap, variant)
|
||||
if model and model.api_key:
|
||||
return self._build_llm_client(model, cap)
|
||||
|
||||
return self._fallback_llm_client(key)
|
||||
|
||||
def get_vision_client(self, key: str, variant: str = "primary"):
|
||||
"""获取 VLM 客户端(返回 DoubaoClient 实例)"""
|
||||
cap = self.get_capability(key)
|
||||
if cap and cap.is_enabled:
|
||||
model = self._get_model_or_fallback(cap, variant)
|
||||
if model and model.api_key:
|
||||
return self._build_vision_client(model, cap)
|
||||
|
||||
return self._fallback_vision_client(key)
|
||||
|
||||
def get_tts_client(self, key: str = "tts") -> TTSClient | None:
|
||||
"""获取 TTS 客户端"""
|
||||
cap = self.get_capability(key)
|
||||
if cap and cap.is_enabled and cap.primary_model and cap.primary_model.api_key:
|
||||
return self._build_tts_client(cap.primary_model, cap)
|
||||
|
||||
return self._fallback_tts_client()
|
||||
|
||||
def get_image_gen_client(self, key: str = "image_generation") -> ImageGenClient | None:
|
||||
"""获取图片生成客户端"""
|
||||
cap = self.get_capability(key)
|
||||
if cap and cap.is_enabled and cap.primary_model and cap.primary_model.api_key:
|
||||
return self._build_image_gen_client(cap.primary_model, cap)
|
||||
|
||||
return self._fallback_image_gen_client()
|
||||
|
||||
def get_video_gen_client(self, key: str = "video_generation") -> VideoGenClient | None:
|
||||
"""获取视频生成客户端"""
|
||||
cap = self.get_capability(key)
|
||||
if cap and cap.is_enabled and cap.primary_model and cap.primary_model.api_key:
|
||||
return self._build_video_gen_client(cap.primary_model, cap)
|
||||
|
||||
return self._fallback_video_gen_client()
|
||||
|
||||
# ── Fallback 方法(读 SharedSettings 环境变量)──────────────────────────
|
||||
|
||||
def _fallback_llm_client(self, key: str):
|
||||
"""Fallback LLM 客户端 — 从 settings 读取配置,不硬编码"""
|
||||
settings = get_shared_settings()
|
||||
model_map = {
|
||||
"intent_parsing": (settings.doubao_fast_model, settings.doubao_base_url, settings.doubao_api_key),
|
||||
"copy_fusion": (settings.doubao_fast_model, settings.doubao_base_url, settings.doubao_api_key),
|
||||
"storyboard": (settings.doubao_model, settings.doubao_base_url, settings.doubao_api_key),
|
||||
"copy_review": (settings.doubao_model, settings.doubao_base_url, settings.doubao_api_key),
|
||||
"asset_classify": (settings.doubao_model, settings.doubao_base_url, settings.doubao_api_key),
|
||||
}
|
||||
if key in model_map:
|
||||
model_id, base_url, api_key = model_map[key]
|
||||
else:
|
||||
model_id = settings.doubao_model
|
||||
base_url = settings.doubao_base_url
|
||||
api_key = settings.doubao_api_key
|
||||
|
||||
if not api_key:
|
||||
return None
|
||||
|
||||
from packages.shared.ai_client import DoubaoClient
|
||||
|
||||
return DoubaoClient(
|
||||
provider="volcengine",
|
||||
api_key=api_key,
|
||||
base_url=base_url,
|
||||
model=model_id,
|
||||
timeout=settings.doubao_timeout,
|
||||
max_retries=settings.doubao_max_retries,
|
||||
)
|
||||
|
||||
def _fallback_vision_client(self, key: str):
|
||||
"""Fallback VLM 客户端 — 从 settings 读取 dashscope 配置,不硬编码"""
|
||||
settings = get_shared_settings()
|
||||
api_key = getattr(settings, "dashscope_api_key", "")
|
||||
if not api_key:
|
||||
return None
|
||||
base_url = getattr(settings, "dashscope_base_url", "") or ""
|
||||
model = getattr(settings, "dashscope_model", "") or getattr(settings, "doubao_vision_model", "")
|
||||
|
||||
from packages.shared.ai_client import DoubaoClient
|
||||
|
||||
return DoubaoClient(
|
||||
provider="dashscope",
|
||||
api_key=api_key,
|
||||
base_url=base_url,
|
||||
model=model,
|
||||
timeout=15,
|
||||
)
|
||||
|
||||
def _fallback_tts_client(self) -> TTSClient | None:
|
||||
settings = get_shared_settings()
|
||||
api_key = getattr(settings, "cosyvoice_api_key", "")
|
||||
if not api_key:
|
||||
return None
|
||||
base_url = getattr(settings, "cosyvoice_base_url", "")
|
||||
model = getattr(settings, "cosyvoice_model", "")
|
||||
|
||||
return TTSClient(provider="dashscope", api_key=api_key, base_url=base_url, model=model)
|
||||
|
||||
def _fallback_image_gen_client(self) -> ImageGenClient | None:
|
||||
settings = get_shared_settings()
|
||||
api_key = getattr(settings, "doubao_api_key", "")
|
||||
if not api_key:
|
||||
return None
|
||||
base_url = getattr(settings, "doubao_base_url", "")
|
||||
model = getattr(settings, "doubao_image_model", "")
|
||||
|
||||
return ImageGenClient(
|
||||
provider="volcengine",
|
||||
api_key=api_key,
|
||||
base_url=base_url,
|
||||
model=model,
|
||||
timeout=getattr(settings, "doubao_image_timeout", 60),
|
||||
)
|
||||
|
||||
def _fallback_video_gen_client(self) -> VideoGenClient | None:
|
||||
settings = get_shared_settings()
|
||||
api_key = getattr(settings, "doubao_api_key", "")
|
||||
if not api_key:
|
||||
return None
|
||||
base_url = getattr(settings, "doubao_base_url", "")
|
||||
model = getattr(settings, "doubao_video_model", "")
|
||||
|
||||
return VideoGenClient(
|
||||
provider="volcengine",
|
||||
api_key=api_key,
|
||||
base_url=base_url,
|
||||
model=model,
|
||||
timeout=getattr(settings, "doubao_video_timeout", 600),
|
||||
)
|
||||
|
||||
def invalidate(self):
|
||||
"""清空本地缓存"""
|
||||
with self._lock:
|
||||
self._cache.clear()
|
||||
self._local_ver = None
|
||||
|
||||
|
||||
# ── 全局单例 ──────────────────────────────────────────────────────────────
|
||||
|
||||
ai_router = AIRouter()
|
||||
@@ -502,7 +502,6 @@ 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。
|
||||
|
||||
@@ -522,7 +521,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, timeout=timeout)
|
||||
raw = client.chat_completion(messages, temperature=temperature, max_tokens=max_tokens, model=model)
|
||||
if raw is None:
|
||||
return None
|
||||
try:
|
||||
@@ -605,24 +604,6 @@ 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,
|
||||
*,
|
||||
@@ -636,25 +617,19 @@ 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 多模型版 + #2172 信任链预热)。
|
||||
) -> str | None:
|
||||
"""调用 Seedance 2.5 生成视频(v1.6.1 单次出片版),返回本地 MP4 路径;失败返回 None。
|
||||
|
||||
成功返回 {"video_path": str, "usage": dict | None}(usage 含 completion_tokens),失败返回 None。
|
||||
失败时错误详情会写入 client.last_video_error,可通过 get_last_video_error() 读取:
|
||||
{"error_code": str, "user_message": str, "status_code": int, "detail": str, ...}
|
||||
v1.6.1 关键约束(避免 20min 卡死):
|
||||
- 参考音频/视频/多图全部放进 content 数组并带 role=reference_audio/reference_video/reference_image;
|
||||
- 纯首帧无参考时(first_frame 模式),Seedance 2.5 强制 ratio=adaptive;
|
||||
传了参考音/视/多图时走 omni_reference 模式,ratio 可指定为 9:16(客户端内部自动判断)。
|
||||
- ratio 默认 9:16(竖屏),客户端会根据是否有参考自动在 first_frame/adaptive 与 omni/9:16 间切换;
|
||||
若创建任务因 ratio 报错(HTTP 400),客户端会自动回退到 adaptive 再试一次。
|
||||
"""
|
||||
client = get_doubao_client()
|
||||
if not client.is_available:
|
||||
msg = "豆包客户端未配置(DOUBAO_API_KEY 缺失),跳过视频生成"
|
||||
logger.warning("[ai_service] %s", msg)
|
||||
# 写入 last_video_error 供上层读取
|
||||
client.last_video_error = {
|
||||
"error_code": "auth_error",
|
||||
"user_message": "视频生成服务未配置,请联系管理员。",
|
||||
"status_code": 0,
|
||||
"detail": msg,
|
||||
}
|
||||
logger.warning("[ai_service] 豆包客户端未配置,跳过视频生成")
|
||||
return None
|
||||
effective_ratio = ratio or "9:16"
|
||||
try:
|
||||
@@ -670,28 +645,10 @@ 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
|
||||
return client.video_generation(**kwargs)
|
||||
except Exception as e:
|
||||
logger.error("[ai_service] call_video_generation 异常: %s", e, exc_info=True)
|
||||
client.last_video_error = {
|
||||
"error_code": "unknown",
|
||||
"user_message": f"视频生成异常:{e!s}"[:200],
|
||||
"status_code": 0,
|
||||
"detail": str(e),
|
||||
}
|
||||
return None
|
||||
|
||||
|
||||
def get_last_video_error() -> dict:
|
||||
"""读取最近一次视频生成失败的详细错误(含 error_code/user_message/status_code/detail)。
|
||||
成功或未调用过返回空 dict。
|
||||
"""
|
||||
try:
|
||||
client = get_doubao_client()
|
||||
return client.get_last_video_error() if hasattr(client, "get_last_video_error") else {}
|
||||
except Exception:
|
||||
return {}
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user