Compare commits

..

5 Commits

Author SHA1 Message Date
xiaoxia fb2cdf7cc9 fix(vision-v2): P0 全部识别失败 - timeout 过短+缺 response_format
根因:
- _FAST_TIMEOUT/_FAST_JSON_TIMEOUT 设为8s,但qwen3.8-flash关thinking后单图实测8-9s,staging网络稍慢即全部超时
- 超时cancel后走pro兜底,pro timeout=20s也偏紧
- 缺少 response_format=json_object 导致qwen偶发输出中文解释而非JSON

修复:
- fast_json: _DEFAULT_TIMEOUT 8→12s
- fast_path: _FAST_TIMEOUT/_FAST_JSON_TIMEOUT 8→12s, _PRO_TIMEOUT 20→25s
- vlm_fallback: _DEFAULT_TIMEOUT 20→25s
- fast_json + fallback 都加 response_format={'type':'json_object'}强约束JSON输出

本地3图E2E: 8.57s, 3/3 fast_json命中, portrait_prompt正常输出一位年轻男性/女性...

Refs: staging P0 3/3全返回未识别/无法判断
2026-10-05 20:35:30 +08:00
CI Bot f904256e09 style: auto-format with black + isort + ruff + prettier [skip ci-format-check] 2026-10-05 12:07:29 +00:00
xiaoxia 257921abf6 refactor(vision): 唯一后端DashScope(qwen),删除ARK/doubao和provider切换
CI/CD Pipeline / Dedup Check - skip PR tests when covered by push pipeline (pull_request) Successful in 5s
CI/CD Pipeline / Check if frontend-only change (pull_request) Successful in 2s
CI/CD Pipeline / Check push changed paths (pull_request) Has been skipped
CI/CD Pipeline / Build Staging API Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Web Image (pull_request) Has been skipped
PR Automation / Auto Approve on CI Green (pull_request) Successful in 2m59s
CI/CD Pipeline / Frontend Lint (pull_request) Has been skipped
CI/CD Pipeline / Frontend Unit Tests (pull_request) Has been skipped
CI/CD Pipeline / PR Build API Image (pull_request) Successful in 16s
CI/CD Pipeline / PR Build Web Image (pull_request) Has been skipped
Preview Deploy / Deploy Preview Environment (pull_request) Successful in 2m28s
CI/CD Pipeline / Build Staging Worker Image (pull_request) Has been skipped
CI/CD Pipeline / PR Build Worker Image (pull_request) Successful in 1m38s
AI Code Review / AI Code Review (pull_request) Successful in 7m13s
CI/CD Pipeline / Retag skipped Staging API Image (pull_request) Has been skipped
CI/CD Pipeline / Retag skipped Staging Web Image (pull_request) Has been skipped
CI/CD Pipeline / Retag skipped Staging Worker Image (pull_request) Has been skipped
PR Automation / Auto Merge on CI Green + Approved (pull_request) Successful in 3m54s
CI/CD Pipeline / Deploy Staging (Watchtower auto-deploy) (pull_request) Has been skipped
Preview Cleanup / Cleanup Preview Environment (pull_request) Successful in 1m30s
CI/CD Pipeline / Unit Tests (pull_request) Failing after 10m42s
CI/CD Pipeline / Integration Tests (pull_request) Successful in 10m38s
ACR Cleanup / ACR Image Cleanup (pull_request_target) Successful in 22s
CI/CD Pipeline / Validate - Style (pull_request) Failing after 12m20s
CI/CD Pipeline / Staging E2E Tests (pull_request) Has been skipped
CI/CD Pipeline / Staging API Integration Tests (pull_request) Has been skipped
CI/CD Pipeline / ACR Image Cleanup (pull_request) Has been skipped
CI/CD Pipeline / Validate - Python (mypy + alembic) (pull_request) Successful in 13m59s
CI/CD Pipeline / Validate - Security (pull_request) Successful in 38m41s
CI/CD Pipeline / Build Production Web Image (pull_request) Has been skipped
CI/CD Pipeline / Build Production Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Build Production API Image (pull_request) Has been skipped
CI/CD Pipeline / Deploy Production (pull_request) Has been skipped
CI/CD Pipeline / Canary Release to Production (pull_request) Has been skipped
CI/CD Pipeline / Production Browser E2E (pull_request) Has been skipped
CI/CD Pipeline / CI Gate (pull_request) Failing after 2s
按灵应指示清理冗余代码:
- 删除_provider.py双后端切换逻辑
- vlm_fast_json.py:纯httpx直连qwen3.8-flash+enable_thinking=false,无if/else
- vlm_fallback.py:纯httpx直连qwen3.7-plus精简JSON prompt(20s超时),删除ark ai_client路径、XML解析、prompt_loader依赖
- fast_path.py:固定超时(fast 8s / pro 20s),删除provider默认值逻辑
- viral_video.py:删除get_doubao_client().max_retries全局操作(现在httpx直连不经过ai_client)
- DASHSCOPE_API_KEY从环境变量读取,不入库
- 删除VISION_V2_PROVIDER/DOUBAO_VISION_*/VISION_V2_ENABLED所有相关代码
- vision/目录896行(比V1双路径+provider切换版更精简)
- viral_video.py 2014行(从2583行累计净删569行)

本地直连验证:qwen3.8-flash单图2.6s/3图并发7.3s(3/3)/8图并发10.3s(7/8)
2026-10-05 19:52:14 +08:00
xiaoxia aac785a9ed feat(vision): 新增DashScope(阿里云百炼/qwen) provider支持
CI/CD Pipeline / Check push changed paths (pull_request) Has been skipped
CI/CD Pipeline / Check if frontend-only change (pull_request) Successful in 2s
CI/CD Pipeline / Dedup Check - skip PR tests when covered by push pipeline (pull_request) Successful in 1s
CI/CD Pipeline / Build Staging API Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Web Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Frontend Lint (pull_request) Has been skipped
CI/CD Pipeline / Frontend Unit Tests (pull_request) Has been skipped
CI/CD Pipeline / PR Build Web Image (pull_request) Has been skipped
CI/CD Pipeline / Retag skipped Staging API Image (pull_request) Has been skipped
CI/CD Pipeline / Retag skipped Staging Web Image (pull_request) Has been skipped
CI/CD Pipeline / Retag skipped Staging Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Deploy Staging (Watchtower auto-deploy) (pull_request) Has been skipped
CI/CD Pipeline / Staging E2E Tests (pull_request) Has been skipped
CI/CD Pipeline / Staging API Integration Tests (pull_request) Has been skipped
CI/CD Pipeline / ACR Image Cleanup (pull_request) Has been skipped
CI/CD Pipeline / PR Build API Image (pull_request) Successful in 1m47s
CI/CD Pipeline / PR Build Worker Image (pull_request) Successful in 2m5s
Preview Deploy / Deploy Preview Environment (pull_request) Successful in 2m43s
PR Automation / Auto Approve on CI Green (pull_request) Successful in 3m1s
AI Code Review / AI Code Review (pull_request) Successful in 6m58s
CI/CD Pipeline / Integration Tests (pull_request) Successful in 10m43s
CI/CD Pipeline / Validate - Python (mypy + alembic) (pull_request) Has been cancelled
CI/CD Pipeline / Build Production API Image (pull_request) Has been cancelled
CI/CD Pipeline / Build Production Web Image (pull_request) Has been cancelled
CI/CD Pipeline / Build Production Worker Image (pull_request) Has been cancelled
CI/CD Pipeline / Deploy Production (pull_request) Has been cancelled
CI/CD Pipeline / Production Browser E2E (pull_request) Has been cancelled
CI/CD Pipeline / Canary Release to Production (pull_request) Has been cancelled
CI/CD Pipeline / CI Gate (pull_request) Has been cancelled
CI/CD Pipeline / Validate - Style (pull_request) Has been cancelled
CI/CD Pipeline / Validate - Security (pull_request) Has been cancelled
CI/CD Pipeline / Unit Tests (pull_request) Has been cancelled
PR Automation / Auto Merge on CI Green + Approved (pull_request) Has been cancelled
- 新增_provider.py统一后端切换:VISION_V2_PROVIDER=ark|dashscope(默认dashscope对比测试)
- vlm_fast_json.py:双后端httpx直连,ark用thinking={type:disabled}、dashscope用enable_thinking=false,400降级逻辑按provider处理
- vlm_fallback.py:dashscope用qwen3.7-plus+精简JSON prompt(<20s),ark保留原有ai_client+XML/JSON双解析
- fast_path.py:超时默认值从provider读取(dashscope pro 20s、ark pro 45s)
- 通过环境变量DASHSCOPE_API_KEY配置key(敏感信息不入代码)
- 本地curl验证:qwen3.8-flash关thinking单图2.6s、3图并发7.3s、8图并发10.3s(7/8可用)
2026-10-05 19:42:11 +08:00
xiaoxia cdd343131e fix(vision): #2205 thinking参数互斥修复——只传thinking=disabled,去掉reasoning_effort
CI/CD Pipeline / Check push changed paths (pull_request) Has been skipped
CI/CD Pipeline / Dedup Check - skip PR tests when covered by push pipeline (pull_request) Successful in 1s
CI/CD Pipeline / Check if frontend-only change (pull_request) Successful in 1s
CI/CD Pipeline / Frontend Lint (pull_request) Has been skipped
CI/CD Pipeline / Frontend Unit Tests (pull_request) Has been skipped
CI/CD Pipeline / PR Build Web Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Web Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging API Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Retag skipped Staging API Image (pull_request) Has been skipped
CI/CD Pipeline / Retag skipped Staging Web Image (pull_request) Has been skipped
CI/CD Pipeline / Retag skipped Staging Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Deploy Staging (Watchtower auto-deploy) (pull_request) Has been skipped
CI/CD Pipeline / Staging E2E Tests (pull_request) Has been skipped
CI/CD Pipeline / Staging API Integration Tests (pull_request) Has been skipped
CI/CD Pipeline / ACR Image Cleanup (pull_request) Has been skipped
CI/CD Pipeline / PR Build Worker Image (pull_request) Successful in 1m10s
Preview Deploy / Deploy Preview Environment (pull_request) Successful in 1m52s
CI/CD Pipeline / PR Build API Image (pull_request) Successful in 1m59s
PR Automation / Auto Approve on CI Green (pull_request) Successful in 2m58s
CI/CD Pipeline / Integration Tests (pull_request) Successful in 4m8s
CI/CD Pipeline / Build Production API Image (pull_request) Has been cancelled
CI/CD Pipeline / Build Production Web Image (pull_request) Has been cancelled
CI/CD Pipeline / Build Production Worker Image (pull_request) Has been cancelled
CI/CD Pipeline / Deploy Production (pull_request) Has been cancelled
CI/CD Pipeline / Production Browser E2E (pull_request) Has been cancelled
CI/CD Pipeline / Canary Release to Production (pull_request) Has been cancelled
CI/CD Pipeline / CI Gate (pull_request) Has been cancelled
CI/CD Pipeline / Validate - Security (pull_request) Has been cancelled
CI/CD Pipeline / Validate - Style (pull_request) Has been cancelled
CI/CD Pipeline / Validate - Python (mypy + alembic) (pull_request) Has been cancelled
CI/CD Pipeline / Unit Tests (pull_request) Has been cancelled
AI Code Review / AI Code Review (pull_request) Has been cancelled
PR Automation / Auto Merge on CI Green + Approved (pull_request) Has been cancelled
根因:上版同时传 thinking={type:disabled} + reasoning_effort=low,方舟返回400
Invalid combination of reasoning_effort and thinking type。降级重试分支
把两个参数都pop掉,模型回到默认thinking开启→响应9-12s超时,fast全败。

修复:
1. 只传 thinking={"type":"disabled"},去掉互斥的 reasoning_effort
2. 400降级只pop thinking,保留精简payload重试(不pop多个)
3. 先单独验证 lite 单次HTTP 200且reasoning_tokens=0再跑E2E,省一轮部署
2026-10-05 19:06:02 +08:00
40 changed files with 574 additions and 4563 deletions
@@ -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 -48
View File
@@ -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,7 @@ def _infer_expected_categories(script_tags: set[str] | None) -> set[str] | None:
return matched or None
logger = logging.getLogger(__name__)
router = APIRouter()
@@ -700,17 +700,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 +930,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。
+3 -13
View File
@@ -280,18 +280,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 = ""
@@ -348,9 +341,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="文案数据缺失,请先点击「生成文案」")
# 积分预扣(已扣过/重试任务跳过)
from app.config import settings as _settings
@@ -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,
)
-157
View File
@@ -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,
@@ -369,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",
@@ -418,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(
@@ -584,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(
@@ -629,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()
@@ -826,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
# 未提交的任务不轮询
@@ -853,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
@@ -872,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:
@@ -967,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
-29
View File
@@ -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:
@@ -1050,13 +1050,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,38 +1082,18 @@
/* ── 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;
@@ -1406,14 +1383,12 @@
/* 口播稿 —— 复用 vv-sb-field 样式,无额外需求 */
.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;
@@ -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,
+85 -138
View File
@@ -92,30 +92,22 @@ def _save_job(repo, job, session):
session.commit()
def _start_trust_chain_preheat(job_id: str, products: list[dict]) -> None:
"""#2172/#2174/#2220 后台启动信任链预热(Seedream t2i 文生图人像),不阻塞调用方。
def _start_trust_chain_preheat(job_id: str, portrait_descriptions: list[str]) -> None:
"""#2172/#2174 后台启动信任链预热(Seedream t2i 文生图人像),不阻塞调用方。
#2220 修复:只对 has_person=True 的图(真人照片)生成 AI 人像替换,
场景图/商品图/门店图保持原图不变,传给 Seedance 作为 reference_image 直接使用。
#2174 重要:改为 t2i 文生图模式——用 VLM 分析出的人物外貌描述做 prompt,不传 reference_images,
产物是方舟信任模型输出,Seedance 直接放行不触发肖像审核。
i2i(传用户照片做 reference)产物不被信任,实测仍被 400 portrait_intercept 拦截。
预热结果写入 job.pre_trusted_images:与 products 等长的稀疏列表,
人像位是 AI 图 URL,非人像位是 None(表示保留原图)。
预热成功后把结果写入 job.pre_trusted_images,阶段3 渲染直接使用,省掉串行等待。
预热失败静默(pre_trusted_images 保持 None),阶段3 会走 #2166 自动降级纯 t2v。
"""
# 构建人像位索引映射:person_indices[k] = products中第k个人像的位置
_person_indices: list[int] = []
_valid: list[str] = []
for _i, _p in enumerate(products or []):
if not isinstance(_p, dict):
continue
if not _p.get("has_person", False):
continue
_d = (_p.get("portrait_prompt") or "").strip()
if not _d or "无人像" in _d or len(_d) < 10:
continue
_person_indices.append(_i)
_valid.append(_d)
# 过滤有效描述:非空且不是"无人像"
_valid = [
d for d in (portrait_descriptions or []) if d and isinstance(d, str) and "无人像" not in d and len(d) >= 10
]
if not _valid:
logger.info("[trust-chain][preheat] 无有效人物描述(可能是纯商品/场景图),跳过预热 job=%s", job_id)
logger.info("[trust-chain][preheat] 无有效人物描述(可能是纯商品图),跳过预热 job=%s", job_id)
return
# 判断是否是 doubao provider(DashScope/Wan 不需要信任链)
try:
@@ -141,33 +133,21 @@ def _start_trust_chain_preheat(job_id: str, products: list[dict]) -> None:
logger.info("[trust-chain][preheat] 后台t2i预热启动 job=%s n=%d", job_id, len(_valid))
result = preheat_trust_chain(_valid, timeout=120)
if result and len(result) == len(_valid):
if result and len(result) >= 1:
sess2, repo2, job2 = _get_repo_and_job(job_id)
try:
# #2220: 构建与 products 等长的稀疏列表,人像位放AI图URL,非人像位None
_imgs2 = job2.images or []
_sparse: list[str | None] = [None] * max(len(_imgs2), len(products or []))
for _k, _url in enumerate(result):
if _k < len(_person_indices):
_sparse[_person_indices[_k]] = _url
job2.pre_trusted_images = _sparse
job2.pre_trusted_images = result
repo2.update(job2)
sess2.commit()
logger.info(
"[trust-chain][preheat] t2i预热完成并持久化 job=%s n_person=%d total=%d",
"[trust-chain][preheat] t2i预热完成并持久化 job=%s n=%d",
job_id,
len(result),
len(_sparse),
)
finally:
sess2.close()
else:
logger.info(
"[trust-chain][preheat] 预热失败或数量不匹配 job=%s got=%s expect=%d,阶段3现场跑兜底",
job_id,
len(result) if result else 0,
len(_valid),
)
logger.info("[trust-chain][preheat] 预热失败 job=%s,阶段3现场跑兜底", job_id)
except Exception as e:
logger.warning("[trust-chain][preheat] 预热异常 job=%s err=%s", job_id, e, exc_info=True)
@@ -288,7 +268,6 @@ def _recover_stale_jobs() -> int:
_DEFAULT_HARD_CONSTRAINTS = [
"无字幕、无水印、无任何自动生成文字、无 logo",
"严格还原参考图片中的真实场景、门店环境、商品陈列、人物外貌服装特征,不得凭空生成与参考图无关的人物、场景或物品",
"同一人物全程保持一致的五官、发型、服装、身材,不得换脸或变形",
"口播语音必须在指定时长内自然念完,语速自然,口型与语音同步",
"画面流畅无闪烁、无多余肢体、无扭曲变形、无穿模",
@@ -344,7 +323,6 @@ def _vision_fallback(idx: int, reason: str, extra: dict | None = None) -> dict:
"scene": "通用",
"portrait_prompt": "无人像",
"summary": "",
"has_person": False,
"_source": reason,
}
if extra:
@@ -443,14 +421,10 @@ def _step_intent_parsing(job: ViralVideoJob, image_analysis: dict) -> dict:
render_system_prompt,
render_user_prompt,
)
from packages.shared.ai_router import ai_router
from packages.shared.ai_service import call_llm
except ImportError:
return {"intent": "推广产品", "key_messages": ["产品亮点"], "tone": "专业", "suggested_title": ""}
_llm_client = ai_router.get_llm_client("intent_parsing")
if not _llm_client.is_available:
return {"intent": "推广产品", "key_messages": ["产品亮点"], "tone": "专业", "suggested_title": ""}
products_summary = ""
products = (image_analysis or {}).get("products", []) or []
for p in products:
@@ -499,22 +473,19 @@ def _step_intent_parsing(job: ViralVideoJob, image_analysis: dict) -> dict:
"suggested_title": "",
}
# #2220: 直接用 ai_router 获取 client,不再手动提取 model_key
from packages.shared.ai_router import ai_router
_client_fast = ai_router.get_llm_client("intent_parsing", variant="primary")
_client_pro = ai_router.get_llm_client("intent_parsing", variant="lite")
for _client, _lbl in [(_client_fast, "fast"), (_client_pro, "pro-fallback")]:
if not _client or not _client.is_available:
continue
_s = get_shared_settings()
_fast = _s.doubao_fast_model
_pro = _s.doubao_model
for _m, _lbl in [(_fast, "fast"), (_pro, "pro-fallback")]:
try:
logger.info("[爆款视频] 意图解析 model=%s label=%s", _client.model, _lbl)
raw = _client.chat_completion(
logger.info("[爆款视频] 意图解析 model=%s label=%s", _m, _lbl)
raw = call_llm(
[{"role": "system", "content": system}, {"role": "user", "content": user}],
temperature=0.4,
max_tokens=1024,
model=_m,
timeout=60,
)
) # #2180: 意图解析 LLM 实测需更长响应,原25s太紧
if not raw:
continue
parsed = _parse(raw)
@@ -854,14 +825,10 @@ def _step_script_generation(job: ViralVideoJob, intent: dict, image_analysis: di
GLOBAL_CONSTRAINTS,
NEGATIVE_RULES,
)
from packages.shared.ai_router import ai_router
from packages.shared.ai_service import call_llm
except ImportError:
return _fallback_script(job)
_llm_client2 = ai_router.get_llm_client("storyboard")
if not _llm_client2.is_available:
return _fallback_script(job)
products_summary = _build_products_summary(image_analysis)
dur = max(5, min(30, int(getattr(job, "duration", 15) or 15)))
@@ -895,14 +862,13 @@ def _step_script_generation(job: ViralVideoJob, intent: dict, image_analysis: di
image_analysis=products_summary,
)
def _try_gen(client, temp: float, max_tok: int, label: str, tmo: int = 25):
if not client or not client.is_available:
return None
logger.info("[爆款视频] 编导脚本生成 model=%s label=%s timeout=%d", client.model, label, tmo)
raw = client.chat_completion(
def _try_gen(model: str, temp: float, max_tok: int, label: str, tmo: int = 25):
logger.info("[爆款视频] 编导脚本生成 model=%s label=%s timeout=%d", model, label, tmo)
raw = call_llm(
[{"role": "system", "content": system_tpl}, {"role": "user", "content": user}],
temperature=temp,
max_tokens=max_tok,
model=model,
timeout=tmo,
)
if not raw:
@@ -933,22 +899,21 @@ def _step_script_generation(job: ViralVideoJob, intent: dict, image_analysis: di
)
return None if is_fallback else normalized
# #2220: 直接用 ai_router 获取 client,不再手动提取 model_key
_client_fast = ai_router.get_llm_client("storyboard", variant="primary")
_client_pro = ai_router.get_llm_client("storyboard", variant="lite")
_script_fast_tmo = int(os.environ.get("VIRAL_VIDEO_SCRIPT_FAST_TIMEOUT", "150"))
_script_pro_tmo = int(os.environ.get("VIRAL_VIDEO_SCRIPT_PRO_TIMEOUT", "150"))
_s = get_shared_settings()
_fast = _s.doubao_fast_model
_pro = getattr(_s, "doubao_model", None) or _fast
try:
# #2217: doubao-seed-2-1-pro生成长编导脚本高峰期>90s,上调到150s,支持ENV覆盖
normalized = _try_gen(_client_fast, 0.8, 2500, "fast-first", tmo=_script_fast_tmo)
# 第一次:快模型 25s
normalized = _try_gen(_fast, 0.8, 2500, "fast-first", tmo=90)
if normalized is not None:
return normalized
normalized = _try_gen(_client_fast, 0.6, 3200, "fast-retry", tmo=_script_fast_tmo)
# #2183: 实测pro 1500tok输出需75.8s,单次timeout提到90s
normalized = _try_gen(_fast, 0.6, 3200, "fast-retry", tmo=90)
if normalized is not None:
return normalized
# 第三次:用 lite/pro 模型兜底
if _client_pro and _client_pro.is_available:
normalized = _try_gen(_client_pro, 0.7, 3500, "pro-fallback", tmo=_script_pro_tmo)
# 第三次:用主力模型兜底,给 120s
if _pro and _pro != _fast:
normalized = _try_gen(_pro, 0.7, 3500, "pro-fallback", tmo=120)
if normalized is not None:
return normalized
logger.warning("[爆款视频] 编导脚本三次都未生成合格结果,使用兜底脚本")
@@ -1271,12 +1236,9 @@ def _step_render(job: ViralVideoJob, copy_result: dict, tts_audio_url: str | Non
if u not in all_portrait_urls:
all_portrait_urls.append(u)
pti = getattr(job, "pre_trusted_images", None)
# #2220: 稀疏列表模式(与images等长,None表示该位置保留原图)
_n_total = len(all_portrait_urls)
if pti and isinstance(pti, list) and len(pti) >= _n_total and _n_total > 0:
pre_trusted = list(pti[:_n_total])
_n_trusted = sum(1 for _x in pre_trusted if _x)
logger.info("[爆款视频] 使用信任链预热结果 person=%d total=%d,跳过现场Seedream AI化", _n_trusted, _n_total)
if pti and len(pti) == len(all_portrait_urls):
pre_trusted = list(pti)
logger.info("[爆款视频] 使用信任链预热结果 n=%d,跳过现场 Seedream AI 化", len(pre_trusted))
elif all_portrait_urls and _mcfg.get("provider", "doubao") == "doubao":
# #2183: 真·现场跑信任链——同步调用 Seedream t2i,拿到 AI 人像 URL 后再传 Seedance
logger.info(
@@ -1289,39 +1251,32 @@ def _step_render(job: ViralVideoJob, copy_result: dict, tts_audio_url: str | Non
_ia = getattr(job, "image_analysis", None) or {}
_prods = (_ia.get("products") if isinstance(_ia, dict) else None) or []
_live_person_idx: list[int] = []
_live_pdescs: list[str] = []
for _i, _pp in enumerate(_prods):
if not isinstance(_pp, dict) or not _pp.get("has_person", False):
continue
_d = (_pp.get("portrait_prompt") or "").strip()
if not _d or "无人像" in _d or len(_d) < 10:
continue
_live_person_idx.append(_i)
_live_pdescs.append(_d)
if _live_pdescs:
_pdescs = []
if _prods:
_pdescs = [(pp.get("portrait_prompt") or "无人像") for pp in _prods]
elif isinstance(_ia, dict):
_pp0 = _ia.get("portrait_prompt") or "无人像"
if _pp0 and _pp0 != "无人像":
_pdescs = [_pp0]
_valid = [d for d in _pdescs if d and isinstance(d, str) and "无人像" not in d and len(d) >= 10]
if _valid:
_t0 = time.time()
_live_urls = preheat_trust_chain(_live_pdescs, timeout=120)
if _live_urls and len(_live_urls) == len(_live_pdescs):
# #2220: 构建稀疏列表,人像位替换AI图,非人像位保留None(ai_client里用原图)
pre_trusted = [None] * len(all_portrait_urls)
for _k, _u in enumerate(_live_urls):
if _k < len(_live_person_idx):
pre_trusted[_live_person_idx[_k]] = _u
_live_urls = preheat_trust_chain(_valid, timeout=120)
if _live_urls and len(_live_urls) == len(all_portrait_urls):
pre_trusted = list(_live_urls)
logger.info(
"[爆款视频] 现场信任链t2i完成 %d张人像AI化 耗时%.1fs(共%d张图,其余保留原图)",
len(_live_urls),
"[爆款视频] 现场信任链t2i完成 %d张 耗时%.1fs,将用AI人像传Seedance",
len(pre_trusted),
time.time() - _t0,
len(all_portrait_urls),
)
else:
logger.warning(
"[爆款视频] 现场信任链t2i返回不匹配 urls=%s n_person=%d,人像位原图传Seedance(可能触发400拦截)",
"[爆款视频] 现场信任链t2i返回不匹配 urls=%s n_portraits=%d,回退原图+400降级纯t2v",
_live_urls,
len(_live_pdescs),
len(all_portrait_urls),
)
else:
logger.info("[爆款视频] 无有效人物描述(商品/场景图),无需AI化,直接传原图给Seedance")
logger.info("[爆款视频] 无有效人物描述(可能是商品图),无需现场跑信任链")
except Exception as _te:
logger.warning("[爆款视频] 现场跑信任链异常: %s,回退原图+400降级纯t2v", _te, exc_info=True)
@@ -1444,8 +1399,13 @@ def run_viral_video_pipeline(self: Task, job_id: str) -> dict:
if job.images:
try:
_products = (image_analysis or {}).get("products", []) or []
# #2220: 直接传 products 列表,由 _start_trust_chain_preheat 内部按 has_person 筛选
_start_trust_chain_preheat(job.id, _products)
_portrait_descs = [(p.get("portrait_prompt") or "无人像") for p in _products] if _products else []
# 兼容单图结果格式(非products列表)
if not _portrait_descs and isinstance(image_analysis, dict):
_pp = image_analysis.get("portrait_prompt") or "无人像"
if _pp and _pp != "无人像":
_portrait_descs = [_pp]
_start_trust_chain_preheat(job.id, _portrait_descs)
except Exception as _e:
logger.warning("[爆款视频][阶段1] 启动信任链t2i预热失败: %s", _e)
_save_job(repo, job, session)
@@ -1617,8 +1577,12 @@ def run_viral_video_analyze(self: Task, job_id: str) -> dict:
if job.images:
try:
_products = (image_analysis or {}).get("products", []) or []
# #2220: 直接传 products 列表,由 _start_trust_chain_preheat 内部按 has_person 筛选
_start_trust_chain_preheat(job.id, _products)
_portrait_descs = [(p.get("portrait_prompt") or "无人像") for p in _products] if _products else []
if not _portrait_descs and isinstance(image_analysis, dict):
_pp = image_analysis.get("portrait_prompt") or "无人像"
if _pp and _pp != "无人像":
_portrait_descs = [_pp]
_start_trust_chain_preheat(job.id, _portrait_descs)
except Exception as _e:
logger.warning("[爆款视频][阶段1] 启动信任链t2i预热失败: %s", _e)
_save_job(repo, job, session)
@@ -1706,7 +1670,7 @@ def run_viral_video_generate_copy(self: Task, job_id: str) -> dict:
image_analysis = job.image_analysis or {"products": []}
intent_result = _step_intent_parsing(job, image_analysis)
job.intent_result = intent_result
# #2218: 不在意图解析后单独落库,等 copy_result 生成后与 mark_copy_generated 一起原子写入
_save_job(repo, job, session)
_emit_progress(job_id, ViralVideoStage.INTENT_PARSING, 35.0, "意图解析完成")
# 阶段:编导脚本生成(核心耗时环节,已用快模型)
@@ -1764,8 +1728,7 @@ def run_viral_video_generate_copy(self: Task, job_id: str) -> dict:
except Retry:
raise
except Exception as e:
logger.error("[爆款视频][阶段2] 异常 job_id=%s: %s", job_id, e, exc_info=True)
# #2218: 阶段2任何异常都标记为 failed(由 _mark_failed_and_notify 处理),前端提示重试
logger.error("[爆款视频][阶段2] 异常: %s", e, exc_info=True)
_mark_failed_and_notify(job_id, session, None, None, str(e), ViralVideoStage.SCRIPT_GENERATION)
return {"ok": False, "job_id": job_id, "error": str(e)}
finally:
@@ -1924,26 +1887,16 @@ def _run_render_pipeline(job_id: str, session, repo, job) -> dict:
阶段2 generate-copy 已把 LLM 深度审核后置,这里在 TTS 前做最终审核(不通过则自动重写1次)。
所有阶段通过 _set_stage 持久化 current_stage/phase_message。
"""
image_analysis = job.image_analysis or {"products": []}
# #2218: render 流程严禁补生成意图+编导脚本。copy_result 必须由 generate-copy 提前准备好;
# 若缺失说明 generate-copy 未完成或数据丢失,直接报错让用户重新点「生成文案」。
# 如果没有 copy_result(旧数据/失败重试),现场补生成(意图+脚本,不走 LLM 审核,出片前会统一做)
copy_result = job.copy_result
_copy_src = "db"
if not isinstance(copy_result, dict) or not copy_result:
logger.error(
"[爆款视频][阶段3] copy_result 为空或无效,无法进入渲染流程。job_id=%s status=%s intent_len=%d,请重新触发「生成文案」",
job_id,
job.status,
len((job.intent_result or {}) if isinstance(job.intent_result, dict) else {}),
)
raise ValueError("文案数据缺失,请先点击「生成文案」完成文案生成后再生成视频")
logger.info(
"[爆款视频][阶段3] 进入渲染流程 job_id=%s copy_result_shots=%d copy_result_len=%d source=%s",
job_id,
len((copy_result.get("shots") or [])),
len(str(copy_result)),
_copy_src,
)
_set_stage(job, repo, session, ViralVideoStage.SCRIPT_GENERATION, "正在补生成编导脚本...")
intent = job.intent_result or _step_intent_parsing(job, image_analysis)
copy_result = _step_script_generation(job, intent, image_analysis)
job.mark_copy_generated(copy_result)
_save_job(repo, job, session)
# 出片前 LLM 深度合规审核(#2134 问题7:审核从阶段2后置到这里,不阻塞前端预览脚本)
_set_stage(job, repo, session, ViralVideoStage.REVIEW, "正在进行出片前合规审核...")
@@ -1956,18 +1909,12 @@ def _run_render_pipeline(job_id: str, session, repo, job) -> dict:
if isinstance(rewritten, dict) and rewritten:
copy_result = rewritten
else:
# #2218: 审核重写失败不再从意图解析重跑,直接报错让用户重新生成文案
logger.error(
"[爆款视频][阶段3] 合规审核未通过且自动重写失败 job_id=%s,终止渲染",
job_id,
)
raise ValueError("文案合规审核未通过,请修改文案后重试或重新生成文案")
intent = job.intent_result or _step_intent_parsing(job, image_analysis)
copy_result = _step_script_generation(job, intent, image_analysis)
_step_review(job, copy_result)
job.copy_result = copy_result
job.generated_copy_text = copy_result.get("voiceover_script", "") or ""
_save_job(repo, job, session)
except ValueError:
# #2218: 审核未通过/文案缺失的业务异常,不继续出片,向上抛出
raise
except Exception as e:
logger.warning("[爆款视频][阶段3] 合规审核异常,继续出片: %s", e)
_emit_progress(job_id, ViralVideoStage.REVIEW, 70.0, "合规审核完成")
@@ -1,208 +0,0 @@
# -*- coding: utf-8 -*-
"""V2 prompt 解析:优先读后台 viral_video_prompt_templates 表(prompt_type='image_analysis'
且 is_active=true),30s TTL 热加载;DB 无有效记录/异常时,fallback 到纯硬编码 JSON schema prompt。
规则(简单直接,不做字符串匹配判断):
- DB 有 is_active=true 的 image_analysis 记录(含种子版本和用户修改后的版本):
* system = DB.system_prompt(DB prompt 自带完整输出格式,不追加硬编码 schema,
避免 DB 写 XML、调用强制 json_object 造成的格式冲突)
* user = DB.user_prompt_template 渲染后使用;渲染后为空则用硬编码默认
- DB 无记录/连接异常/返回空:system/user 全部用纯硬编码 JSON schema prompt
"""
from __future__ import annotations
import logging
import threading
import time
from typing import Any
logger = logging.getLogger(__name__)
# ---- 纯硬编码 JSON schema(DB 无有效配置时全量使用) ----
_FAST_JSON_SCHEMA = (
"你是图片结构化识别器。严格按下方 JSON schema 返回一个对象,不要任何解释、"
"不要markdown、不要代码块、不要前后缀文字。字段值不确定时填 null 或空数组。\n"
"{\n"
' "has_person": true/false,\n'
' "gender": "男"/"女"/null,\n'
' "age_range": "儿童"/"青少年"/"青年"/"中年"/"老年"/null,\n'
' "upper_wear": "上装款式,如T恤/衬衫/卫衣/毛衣/西装/夹克/连衣裙/吊带/背心/外套等",\n'
' "upper_color": "上装主色",\n'
' "lower_wear": "下装款式;穿连衣裙时填null",\n'
' "lower_color": "下装主色",\n'
' "dress_color": "连衣裙主色(穿连衣裙时填)",\n'
' "accessories": ["眼镜"/"帽子"/"项链"/"耳环"/"背包"/"手表"等数组],\n'
' "hairstyle": "发型,如短发/长发/马尾/卷发/丸子头/光头等",\n'
' "expression": "表情,如微笑/严肃/酷/开心等",\n'
' "pose": "姿势,如站立/坐姿/侧身/行走等",\n'
' "scene": "场景,如室内/街拍/户外/办公室/家居/海边/雪景/森林等",\n'
' "style": "风格,如休闲/商务/运动/复古/潮流/甜美/酷飒/优雅/街头/法式等",\n'
' "has_product": true/false,\n'
' "category": "产品类目:服饰/鞋包/美妆/数码/食品/家居/配饰/母婴/非产品图",\n'
' "product_name": "产品名称,非产品图填null",\n'
' "brand": "品牌或文字标识,无则null",\n'
' "material": "材质,如棉质/牛仔/皮革/真丝/针织/涤纶等",\n'
' "pattern": "图案,如纯色/条纹/波点/格子/印花/碎花/Logo等",\n'
' "colors": ["主色数组"],\n'
' "mood": "整体氛围/情绪,如清新/活力/高级/温暖/冷峻/甜美/复古等"\n'
"}\n\n"
"你必须只返回一个合法的JSON对象,不要输出任何其他文字、解释、XML标签或markdown。"
)
DEFAULT_FAST_USER = "识别这张图片的人物穿搭与主体信息,只返回JSON对象。"
_PRO_JSON_SCHEMA = (
"你是图片分析专家。严格按下方 JSON schema 返回一个对象,不要解释、不要markdown、不要代码块、不要XML标签。\n"
"{\n"
' "has_person": true/false,\n'
' "gender": "男"/"女"/null,\n'
' "age_range": "儿童"/"青少年"/"青年"/"中年"/"老年"/null,\n'
' "outfit": "整体穿着描述(含颜色款式)",\n'
' "hair": "发型发色",\n'
' "pose": "姿势",\n'
' "expression": "表情",\n'
' "scene": "场景",\n'
' "mood": "氛围",\n'
' "has_product": true/false,\n'
' "category": "类目:服饰/鞋包/美妆/数码/食品/家居/配饰/母婴/非产品图",\n'
' "product_name": "产品名,非产品图填null",\n'
' "brand": "品牌,无则null",\n'
' "key_features": ["核心特征数组,3-6个短语"]\n'
"}\n\n"
"你必须只返回一个合法的JSON对象,不要输出任何其他文字、解释、XML标签或markdown。"
)
DEFAULT_PRO_USER = "分析这张图片,返回符合schema的JSON。"
# 保留旧 JSON schema 追加文本作为常量(DB prompt 完全控制输出格式后不再使用,
# 保留以便排查历史行为)。
_FAST_JSON_APPEND = (
"\n\n【输出格式要求】无论上文如何要求,最终你必须只返回一个合法的JSON对象,"
"严格包含以下字段(字段值不确定时填null或空数组):\n"
"{\n"
' "has_person": true/false,\n'
' "gender": "男"/"女"/null,\n'
' "age_range": "儿童"/"青少年"/"青年"/"中年"/"老年"/null,\n'
' "upper_wear": "上装款式字符串",\n'
' "upper_color": "上装主色",\n'
' "lower_wear": "下装款式(穿连衣裙时填null)",\n'
' "lower_color": "下装主色",\n'
' "dress_color": "连衣裙主色(穿连衣裙时填)",\n'
' "accessories": ["配饰数组"],\n'
' "hairstyle": "发型",\n'
' "expression": "表情",\n'
' "pose": "姿势",\n'
' "scene": "场景",\n'
' "style": "风格",\n'
' "has_product": true/false,\n'
' "category": "产品类目:服饰/鞋包/美妆/数码/食品/家居/配饰/母婴/非产品图",\n'
' "product_name": "产品名称,非产品图填null",\n'
' "brand": "品牌或文字标识,无则null",\n'
' "material": "材质",\n'
' "pattern": "图案",\n'
' "colors": ["主色数组"],\n'
' "mood": "整体氛围"\n'
"}\n"
"不要输出任何其他文字、解释、XML标签或markdown。"
)
_PRO_JSON_APPEND = (
"\n\n【输出格式要求】无论上文如何要求,最终你必须只返回一个合法的JSON对象,"
"严格包含以下字段(字段值不确定时填null或空数组):\n"
"{\n"
' "has_person": true/false,\n'
' "gender": "男"/"女"/null,\n'
' "age_range": "儿童"/"青少年"/"青年"/"中年"/"老年"/null,\n'
' "outfit": "整体穿着描述(含颜色款式)",\n'
' "hair": "发型发色",\n'
' "pose": "姿势",\n'
' "expression": "表情",\n'
' "scene": "场景",\n'
' "mood": "氛围",\n'
' "has_product": true/false,\n'
' "category": "类目:服饰/鞋包/美妆/数码/食品/家居/配饰/母婴/非产品图",\n'
' "product_name": "产品名,非产品图填null",\n'
' "brand": "品牌,无则null",\n'
' "key_features": ["核心特征3-6个短语"]\n'
"}\n"
"不要输出任何其他文字、解释、XML标签或markdown。"
)
_cache_lock = threading.Lock()
_cache: dict[str, tuple[float, Any]] = {}
_CACHE_TTL = 30.0
def _load_db_template() -> Any | None:
"""直接查DB viral_video_prompt_templates 中 is_active=true 的 image_analysis 记录;
DB不可达/无记录/异常返回None。
复用 prompt_loader._load_from_db,它只查DB不做DEFAULT_TEMPLATES fallback,
返回None表示DB无记录或异常。"""
try:
from packages.application.viral_video.prompt_loader import _load_from_db
return _load_from_db("image_analysis")
except Exception as e:
logger.warning("[vision.v2] 查询DB prompt配置失败: %s", e)
return None
def _render_user(tpl: Any | None, default_user: str) -> str:
if not tpl:
return default_user
tpl_str = getattr(tpl, "user_prompt_template", "") or ""
if not tpl_str.strip():
return default_user
rendered = tpl_str.replace("{image_count}", "1").replace("{industry}", "通用").replace("{image_urls}", "").strip()
return rendered or default_user
def resolve_fast_prompt() -> tuple[str, str]:
return _resolve("fast")
def resolve_pro_prompt() -> tuple[str, str]:
return _resolve("pro")
def _resolve(kind: 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:
return hit[1]
default_sys = _FAST_JSON_SCHEMA if kind == "fast" else _PRO_JSON_SCHEMA
default_user = DEFAULT_FAST_USER if kind == "fast" else DEFAULT_PRO_USER
sys_prompt = default_sys
usr_prompt = default_user
try:
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 prompt自带完整输出格式,不追加硬编码schema避免冲突
usr_prompt = _render_user(tpl, default_user)
logger.info(
"[vision.v2] 使用DB image_analysis prompt (kind=%s version=%s sys_len=%d)",
kind,
getattr(tpl, "version", "?"),
len(db_sys),
)
else:
logger.debug("[vision.v2] DB image_analysis system_prompt为空,使用默认JSON (kind=%s)", kind)
else:
logger.debug("[vision.v2] DB无image_analysis记录/不可达,使用默认JSON prompt (kind=%s)", kind)
except Exception as e:
logger.warning("[vision.v2] 解析DB prompt异常,使用默认: %s", e)
with _cache_lock:
_cache[cache_key] = (now, (sys_prompt, usr_prompt))
return sys_prompt, usr_prompt
def invalidate_cache() -> None:
with _cache_lock:
_cache.clear()
+100 -461
View File
@@ -1,8 +1,5 @@
# -*- coding: utf-8 -*-
"""把 fast_json VLM 输出 + OCR 文本组装为下游兼容的 product dict。
v4 schema: DB prompt完全控制输出格式,可能是v4嵌套schema(type/products/people/store_info)
或旧扁平schema(has_person/upper_wear/product_name/brand等)。assembler兼容两种格式。
"""把 fast_json VLM 输出 + OCR 文本组装为与旧 _normalize() 完全一致的 dict。
目标:下游(信任链t2i/intent_parsing/script_generation)零改动。
必出字段:name, brand, category, appearance, packaging, text_on_package,
@@ -13,16 +10,29 @@ from __future__ import annotations
from typing import Any
# ---------- portrait_prompt 模板 ----------
# 目标:60-100 字的人物穿搭描述,用于 Seedream 纯文生图。要求具体、风格化、视觉细节丰富。
# 旧 VLM 输出格式参考:"一位25岁左右的亚洲女性,身穿白色V领短袖T恤,黑色高腰阔腿裤,
# 搭配银色项链,长发披肩,表情自信,街拍风格,阳光明媚的城市街头"
def _join_parts(*parts: str | None) -> str:
return "".join(p for p in parts if p)
_AGE_PREFIX = {"青年": "年轻", "中年": "中年", "老年": "老年"}
_AGE_PREFIX = {
"青年": "年轻",
"中年": "中年",
"老年": "老年",
}
# gender 后缀
_GENDER_WORD = {"男": "男性", "女": "女性"}
def _person_subject(gender: str, age: str) -> str:
def _person_subject(fj: dict[str, Any]) -> str:
"""人物主语:年轻女性 / 中年男性 / 少女 / 小男孩 / 人物 等。"""
gender = fj.get("gender") or ""
age = fj.get("age_range") or ""
gw = _GENDER_WORD.get(gender, "")
if age == "儿童":
if gender == "女":
@@ -42,147 +52,8 @@ def _person_subject(gender: str, age: str) -> str:
return f"{prefix}人物" if prefix else "人物"
def _build_wear_from_v4(p: dict) -> str:
"""v4 person schema: upper_wear/upper_color/lower_wear/lower_color/dress_color"""
upper = p.get("upper_wear") or ""
upper_color = p.get("upper_color") or ""
lower = p.get("lower_wear") or ""
lower_color = p.get("lower_color") or ""
dress_color = p.get("dress_color") or ""
is_dress = ("连衣裙" in upper) or ("裙" in upper and not lower)
if is_dress:
c = dress_color or upper_color
return f"身穿{c}{upper}" if c else f"身穿{upper}"
parts = []
if upper:
up = f"{upper_color}{upper}" if upper_color else upper
parts.append(f"上身{up}")
if lower:
lo = f"{lower_color}{lower}" if lower_color else lower
parts.append(f"下身{lo}")
return ",".join(parts)
def _build_portrait_prompt_from_v4(p: dict) -> str:
"""v4 person: 直接用portrait_prompt字段;没有就拼"""
direct = p.get("portrait_prompt")
if direct and len(direct) >= 10:
return direct
subject = _person_subject(p.get("gender", ""), p.get("age_range", ""))
wear = _build_wear_from_v4(p)
acc = p.get("accessories") or []
if isinstance(acc, str):
acc = [acc]
acc_str = ",佩戴" + "、".join(str(a) for a in acc if a) if acc else ""
hair = p.get("hairstyle") or ""
expr = p.get("expression") or ""
pose = p.get("pose") or ""
style = p.get("outfit_style") or p.get("style") or ""
scene = p.get("scene") or ""
mood = p.get("mood") or ""
details = []
if hair:
details.append(hair)
if expr and expr not in ("自然", "平静"):
details.append(f"神情{expr}")
if pose and pose not in ("站立",):
details.append(pose)
style_parts = []
if style:
style_parts.append(style)
if mood:
style_parts.append(mood)
if scene and scene not in ("通用",):
style_parts.append(scene)
pieces = [f"一位{subject}"]
if wear:
pieces.append(wear)
if acc_str:
pieces.append(acc_str.lstrip(","))
if details:
pieces.append(",".join(details))
pieces.append(("".join(style_parts) + "风格") if style_parts else "人像写真")
full = ",".join(p for p in pieces if p)
if len(full) < 40:
full += ",自然光线下人像特写,画面清晰"
if len(full) > 120:
full = full[:120].rstrip(",") + "。"
return full
def _build_product_prompt_from_v4(prod: dict, top: dict) -> str:
"""v4 product: 拼商品视觉描述prompt(用于AI生图参考)"""
name = prod.get("product_name") or "商品"
brand = prod.get("brand") or ""
lead = f"{brand} {name}" if brand and brand not in name else name
pkg_color = prod.get("package_color") or ""
pkg_type = prod.get("package_type") or ""
cap = prod.get("cap_type") or ""
body = prod.get("body_shape") or ""
features = prod.get("product_features") or []
sell = prod.get("key_selling_points") or []
colors = top.get("colors") or []
style = top.get("style") or ""
scene = top.get("scene") or ""
mood = top.get("mood") or ""
parts = [lead]
desc = []
if pkg_color:
desc.append(pkg_color)
if pkg_type:
desc.append(pkg_type)
if cap and len(desc) < 3:
desc.append(f"配{cap}")
if body and len(desc) < 3:
desc.append(body)
if desc:
parts.append(",".join(desc))
if features:
core = [str(f) for f in features[:3] if f and len(str(f)) <= 25]
if core:
parts.append(";".join(core))
if sell:
s = [str(x) for x in sell[:2] if x]
if s:
parts.append("突出" + "、".join(s))
cnames = []
for cc in colors:
if isinstance(cc, dict) and cc.get("name"):
cnames.append(cc["name"])
elif isinstance(cc, str):
cnames.append(cc)
cnames = cnames[:3]
if cnames:
parts.append("、".join(cnames) + "主色")
if style:
parts.append(style)
if mood:
parts.append(mood)
if scene and not any(k in scene for k in ("白色背景", "纯色", "通用")):
parts.append(scene)
parts.append("产品特写,画面清晰")
prompt = ",".join(p for p in parts if p)
return prompt if len(prompt) >= 10 else "产品展示图,特写镜头"
def _is_v4_schema(fj: dict) -> bool:
"""判断是v4嵌套schema还是旧扁平schema"""
return (
isinstance(fj.get("products"), list)
or fj.get("type") in ("product", "store", "person", "other")
or isinstance(fj.get("people"), dict)
)
# ---------- 旧扁平schema兼容(保留原逻辑) ----------
def _person_subject_old(fj: dict) -> str:
return _person_subject(fj.get("gender", ""), fj.get("age_range", ""))
def _build_wear_sentence_old(fj: dict) -> str:
def _build_wear_sentence(fj: dict[str, Any]) -> str:
"""穿搭段:上装+下装/连衣裙,带颜色+材质+图案。"""
upper = fj.get("upper_wear") or ""
upper_color = fj.get("upper_color") or ""
lower = fj.get("lower_wear") or ""
@@ -190,6 +61,7 @@ def _build_wear_sentence_old(fj: dict) -> str:
dress_color = fj.get("dress_color") or ""
material = fj.get("material") or ""
pattern = fj.get("pattern") or ""
is_dress = ("连衣裙" in upper) or ("裙" in upper and not lower)
if is_dress:
c = dress_color or upper_color
@@ -199,22 +71,25 @@ def _build_wear_sentence_old(fj: dict) -> str:
if pattern and pattern not in wear and pattern != "纯色":
wear += f",{pattern}图案"
return f"身穿{wear}"
parts = []
parts: list[str] = []
if upper:
up = f"{upper_color}{upper}" if upper_color else upper
if material and material not in up:
up = f"{material}{up}"
if pattern and pattern != "纯色" and pattern not in up:
up += f"({pattern})"
parts.append(f"上身{up}")
parts.append(f"上身{up}" if up else "")
if lower:
lo = f"{lower_color}{lower}" if lower_color else lower
parts.append(f"下身{lo}")
parts.append(f"下身{lo}" if lo else "")
return ",".join(p for p in parts if p)
def _build_portrait_prompt_old(fj: dict) -> str:
def _build_portrait_prompt(fj: dict[str, Any]) -> str:
"""组装最终 portrait_prompt(目标 60-100 字,用于 Seedream 纯文生图)。"""
if not fj.get("has_person"):
# 非人像:用商品+场景+mood 拼一段
name = fj.get("product_name") or "商品"
brand = fj.get("brand") or ""
colors = fj.get("colors") or []
@@ -226,15 +101,7 @@ def _build_portrait_prompt_old(fj: dict) -> str:
pieces.append(brand)
pieces.append(name)
if colors:
cnames = []
for c in colors:
if isinstance(c, dict):
cnames.append(c.get("name", ""))
elif isinstance(c, str):
cnames.append(c)
cnames = [c for c in cnames if c][:3]
if cnames:
pieces.append("、".join(cnames) + "配色")
pieces.append("、".join(colors[:3]) + "配色")
if style:
pieces.append(style + "风格")
if mood:
@@ -244,32 +111,40 @@ def _build_portrait_prompt_old(fj: dict) -> str:
pieces.append("产品特写")
prompt = ",".join(p for p in pieces if p)
return prompt if len(prompt) >= 10 else "产品展示图,特写镜头"
subject = _person_subject_old(fj)
wear = _build_wear_sentence_old(fj)
subject = _person_subject(fj)
wear = _build_wear_sentence(fj)
accessories = fj.get("accessories") or []
if isinstance(accessories, str):
accessories = [accessories]
acc_str = ",佩戴" + "、".join(str(a) for a in accessories if a) if accessories else ""
acc_str = ""
if accessories:
acc_str = ",佩戴" + "、".join(str(a) for a in accessories if a)
hairstyle = fj.get("hairstyle") or ""
expression = fj.get("expression") or ""
pose = fj.get("pose") or ""
style = fj.get("style") or ""
scene = fj.get("scene") or ""
mood = fj.get("mood") or ""
detail_parts = []
detail_parts: list[str] = []
if hairstyle:
detail_parts.append(hairstyle)
if expression and expression not in ("自然", "平静"):
detail_parts.append(f"神情{expression}")
if pose and pose not in ("站立",):
detail_parts.append(pose)
style_parts = []
style_parts: list[str] = []
if style:
style_parts.append(style)
if mood:
style_parts.append(mood)
if scene and scene not in ("通用",):
style_parts.append(scene)
pieces = [f"一位{subject}"]
if wear:
pieces.append(wear)
@@ -277,40 +152,53 @@ def _build_portrait_prompt_old(fj: dict) -> str:
pieces.append(acc_str.lstrip(","))
if detail_parts:
pieces.append(",".join(detail_parts))
pieces.append("".join(style_parts) + "风格" if style_parts else "人像写真")
if style_parts:
# 风格词之间不用逗号,用空格紧凑
pieces.append("".join(style_parts) + "风格")
else:
pieces.append("人像写真")
full = ",".join(p for p in pieces if p)
# 过短补充镜头词
if len(full) < 40:
full += ",自然光线下人像特写,画面清晰"
# 过长截断
if len(full) > 120:
full = full[:120].rstrip(",") + "。"
return full
def _infer_name_old(fj: dict, ocr_texts: list[str]) -> str:
# ---------- 商品字段 ----------
def _infer_name(fj: dict[str, Any], ocr_texts: list[str]) -> str:
pname = fj.get("product_name")
if pname and pname != "未识别":
return str(pname)
# 人物图 → name 用穿搭主件
if fj.get("has_person"):
up = fj.get("upper_wear") or ""
if "连衣裙" in up:
return up
return up or "人物穿搭"
if ocr_texts:
# 商品名可能是 OCR 最长的一行(品牌/产品名)
return max(ocr_texts, key=len)
return "未识别"
def _infer_brand_old(fj: dict, ocr_texts: list[str]) -> str:
def _infer_brand(fj: dict[str, Any], ocr_texts: list[str]) -> str:
brand = fj.get("brand")
if brand:
return str(brand)
# OCR 里短的、纯字母/汉字短串可能是 brand
for t in ocr_texts:
if 1 < len(t) <= 12:
return t
return "无法判断"
def _infer_category_old(fj: dict) -> str:
def _infer_category(fj: dict[str, Any]) -> str:
cat = fj.get("category")
if cat:
return str(cat)
@@ -319,19 +207,27 @@ def _infer_category_old(fj: dict) -> str:
return "非产品图"
def _build_appearance_old(fj: dict) -> str:
parts = []
for key in ("upper_color", "upper_wear", "material", "pattern"):
def _build_appearance(fj: dict[str, Any]) -> str:
"""外观描述:颜色+款式+材质+图案 拼成一段。"""
parts: list[str] = []
for key, _label in [
("upper_color", "主色"),
("upper_wear", "款式"),
("material", "材质"),
("pattern", "图案"),
]:
v = fj.get(key)
if v and v not in ("无法判断", "未知", "纯色"):
parts.append(str(v))
if not parts:
return "人像穿搭整体造型" if fj.get("has_person") else "无法判断"
if fj.get("has_person"):
return "人像穿搭整体造型"
return "无法判断"
return "、".join(parts)
def _build_key_features_old(fj: dict, ocr_texts: list[str]) -> list[str]:
feats = []
def _build_key_features(fj: dict[str, Any], ocr_texts: list[str]) -> list[str]:
feats: list[str] = []
for key in (
"upper_wear",
"lower_wear",
@@ -352,7 +248,9 @@ def _build_key_features_old(fj: dict, ocr_texts: list[str]) -> list[str]:
feats.append(v)
if ocr_texts:
feats.append(f"画面文字: {'/'.join(ocr_texts[:3])}")
out, seen = [], set()
# 去重
out: list[str] = []
seen: set[str] = set()
for f in feats:
f = f.strip()
if f and f not in seen and len(f) <= 30:
@@ -361,296 +259,27 @@ def _build_key_features_old(fj: dict, ocr_texts: list[str]) -> list[str]:
return out[:6] if out else ["无法判断"]
def _flatten_colors(c) -> list[str]:
"""colors可能是字符串数组或[{hex,name,coverage}],统一返回名字数组"""
if not c:
return []
out = []
for item in c:
if isinstance(item, dict):
n = item.get("name")
if n:
out.append(n)
elif isinstance(item, str):
out.append(item)
return out
def assemble_result(idx: int, fast_json: dict | None, ocr_texts: list[str]) -> dict[str, Any]:
def assemble_result(
idx: int,
fast_json: dict[str, Any] | None,
ocr_texts: list[str],
) -> dict[str, Any]:
"""把 fast_json 结果 + OCR 文本组装成下游兼容的 product dict。"""
fj = fast_json or {}
ocr_texts = ocr_texts or []
if _is_v4_schema(fj):
return _assemble_v4(idx, fj, ocr_texts)
else:
return _assemble_old(idx, fj, ocr_texts)
def _assemble_v4(idx: int, fj: dict, ocr_texts: list[str]) -> dict[str, Any]:
"""v4嵌套schema → 下游product dict"""
vtype = fj.get("type") or "other"
products = fj.get("products") or []
portrait_prompt = _build_portrait_prompt(fj)
name = _infer_name(fj, ocr_texts)
brand = _infer_brand(fj, ocr_texts)
category = _infer_category(fj)
appearance = _build_appearance(fj)
key_features = _build_key_features(fj, ocr_texts)
scene = fj.get("scene") or "通用"
mood = fj.get("mood") or ""
colors = fj.get("colors") or []
visible_text = fj.get("visible_text") or []
color_names = _flatten_colors(colors)
# 合并OCR文字和visible_text
pkg_texts = []
for vt in visible_text:
if isinstance(vt, dict):
t = vt.get("text")
if t:
pkg_texts.append(str(t))
elif isinstance(vt, str):
pkg_texts.append(vt)
pkg_texts.extend(ocr_texts[:5])
# 去重
seen_t = set()
text_on_package = []
for t in pkg_texts:
t = str(t).strip()
if t and t not in seen_t and len(t) <= 50:
seen_t.add(t)
text_on_package.append(t)
text_on_package = text_on_package[:8]
has_person = fj.get("has_person", False)
# 人物类
if vtype == "person" or has_person:
# 取第一个人物信息(v5 schema人物信息在顶层)
person_info = fj
# 兼容people嵌套
ppl = fj.get("people")
if isinstance(ppl, dict) and ppl.get("has_person"):
person_info = {**fj, **ppl}
has_person = True
portrait_prompt = _build_portrait_prompt_from_v4(person_info)
outfit_style = person_info.get("outfit_style") or ""
upper = person_info.get("upper_wear") or ""
lower = person_info.get("lower_wear") or ""
dress = person_info.get("dress_wear") or ""
outer = person_info.get("outerwear") or ""
if dress:
name = str(dress)[:25]
elif outer and upper:
name = f"{outer}+{upper}"[:30]
elif upper:
name = (str(upper) + (f"+{lower}" if lower else ""))[:30]
else:
name = "人物穿搭"
brand = "无法判断"
category = "人物穿搭"
# appearance: 外套+上衣+下装/裙+鞋+包+发型+妆容
app_parts = []
for k in ("outerwear", "upper_wear", "lower_wear", "dress_wear", "shoes", "bag", "hairstyle", "makeup"):
v = person_info.get(k)
if v and v not in ("null", None, "无明显妆容"):
app_parts.append(str(v))
appearance = ";".join(app_parts) if app_parts else "人像穿搭整体造型"
# key_features: 服装+配饰+拍摄信息
kf = []
for k in (
"outfit_style",
"upper_wear",
"lower_wear",
"dress_wear",
"outerwear",
"shoes",
"bag",
"hairstyle",
"expression",
"pose",
):
v = person_info.get(k)
if v and v not in ("null", None, "无法判断"):
kf.append(str(v))
acc = person_info.get("accessories") or []
if isinstance(acc, list):
for a in acc:
if a and str(a) not in kf:
kf.append(str(a))
elif isinstance(acc, str) and acc:
kf.append(acc)
for k in ("shot_type", "camera_angle", "lighting", "atmosphere"):
v = person_info.get(k)
if v and v not in ("null", None):
kf.append(str(v))
if text_on_package:
kf.append(f"文字:{'/'.join(text_on_package[:3])}")
kf = kf[:8] or ["无法判断"]
summary = (outfit_style + " " if outfit_style and outfit_style not in name else "") + name[:25]
if not summary.strip():
summary = "人物穿搭"
return {
"name": name[:30],
"brand": brand,
"category": category,
"appearance": appearance[:400],
"packaging": "人物形象无包装",
"text_on_package": text_on_package,
"key_features": kf,
"scene": scene,
"mood": mood,
"portrait_prompt": portrait_prompt[:300],
"summary": summary[:50],
"_source": "v2_fast_json_v5",
"has_person": True,
}
# 商品类
if vtype == "product" and products:
# 主商品(第一个position=main或第一个)
main = products[0]
for p in products:
if p.get("position") == "main":
main = p
break
name = main.get("product_name") or "未识别"
brand = main.get("brand") or "无法判断"
category = main.get("category") or "非产品图"
# appearance: 包装外观
app_parts = []
for k in ("package_color", "package_type", "cap_type", "body_shape", "label_design"):
v = main.get(k)
if v and v not in ("null", None):
app_parts.append(str(v))
appearance = ";".join(app_parts) if app_parts else "无法判断"
# packaging: 包装信息(直接用package_type+package_color)
pkg_parts = []
if main.get("package_type"):
pkg_parts.append(str(main["package_type"]))
if main.get("package_color"):
pkg_parts.append(str(main["package_color"]))
if main.get("cap_type"):
pkg_parts.append(f"配{main['cap_type']}")
packaging = ",".join(pkg_parts) if pkg_parts else "无法判断"
# key_features: product_features字段
feats = main.get("product_features") or []
if not isinstance(feats, list):
feats = [str(feats)]
kf = [str(f) for f in feats if f and len(str(f)) <= 40][:6]
# 补充卖点
sell = main.get("key_selling_points") or []
if isinstance(sell, list):
for s in sell[:2]:
if s and len(str(s)) <= 30 and str(s) not in kf:
kf.append(f"卖点:{s}")
if text_on_package:
kf.append(f"文字: {'/'.join(text_on_package[:3])}")
kf = kf[:6] or ["无法判断"]
portrait_prompt = _build_product_prompt_from_v4(main, fj)
if brand != "无法判断" and brand not in name:
summary = f"{brand} {name}"
else:
summary = name
return {
"name": str(name)[:50],
"brand": str(brand)[:30],
"category": str(category)[:20],
"appearance": appearance[:200],
"packaging": packaging[:100],
"text_on_package": text_on_package,
"key_features": kf,
"scene": scene,
"mood": mood,
"portrait_prompt": portrait_prompt[:200],
"summary": str(summary)[:60],
"has_person": False,
"_source": "v2_fast_json_v4",
}
# 门店类或其他
if vtype == "store":
store_type = fj.get("store_type") or "店铺"
name = store_type
brand = fj.get("brand_signage") or "无法判断"
category = "门店场景"
visual = fj.get("visual_elements") or []
if isinstance(visual, str):
visual = [visual]
atmosphere = fj.get("atmosphere") or mood
appearance_parts = []
if fj.get("store_layout"):
appearance_parts.append(str(fj["store_layout"]))
if visual:
appearance_parts.append("、".join(str(v) for v in visual[:3]))
if fj.get("cleanliness"):
appearance_parts.append(str(fj["cleanliness"]))
appearance = ";".join(appearance_parts) if appearance_parts else "门店环境"
kf = []
if isinstance(visual, list):
kf.extend(str(v) for v in visual if v and len(str(v)) <= 30)
prods_vis = fj.get("product_categories_visible") or []
if isinstance(prods_vis, list):
kf.extend(str(c) for c in prods_vis[:3] if c)
promo = fj.get("promotion_elements") or []
if isinstance(promo, list) and promo:
kf.append("促销活动:" + "、".join(str(p) for p in promo[:2]))
if text_on_package:
kf.append(f"文字: {'/'.join(text_on_package[:3])}")
kf = kf[:6] or ["门店场景"]
portrait_prompt = f"{brand if brand!='无法判断' else ''}{store_type},{atmosphere},{scene}场景,{('、'.join(color_names[:3])+'配色,') if color_names else ''}产品陈列丰富,门店实拍"
portrait_prompt = portrait_prompt.strip(",")
summary = f"{store_type}场景"
return {
"name": name[:30],
"brand": str(brand)[:30],
"category": category,
"appearance": appearance[:200],
"packaging": "门店场景无包装",
"text_on_package": text_on_package,
"key_features": kf,
"scene": scene,
"mood": atmosphere or mood,
"portrait_prompt": portrait_prompt[:200],
"summary": summary[:40],
"has_person": False,
"_source": "v2_fast_json_v4",
}
# other 兜底
desc = fj.get("description") or "未识别"
return {
"name": desc[:30],
"brand": "无法判断",
"category": "非产品图",
"appearance": desc[:200],
"packaging": "无法判断",
"text_on_package": text_on_package,
"key_features": [desc[:30]] if desc != "未识别" else ["无法判断"],
"scene": scene,
"mood": mood,
"portrait_prompt": f"{scene},{mood}氛围,{desc}"[:200],
"summary": desc[:40],
"has_person": False,
"_source": "v2_fast_json_v4_other",
}
def _assemble_old(idx: int, fj: dict, ocr_texts: list[str]) -> dict[str, Any]:
"""旧扁平schema(兼容存量prompt或pro兜底输出)"""
portrait_prompt = _build_portrait_prompt_old(fj)
name = _infer_name_old(fj, ocr_texts)
brand = _infer_brand_old(fj, ocr_texts)
category = _infer_category_old(fj)
appearance = _build_appearance_old(fj)
key_features = _build_key_features_old(fj, ocr_texts)
scene = fj.get("scene") or "通用"
mood = fj.get("mood") or ""
packaging = "无法判断"
packaging = "无法判断" # 包装细节专用API无,保留占位
text_on_package = ocr_texts[:8]
if fj.get("has_person"):
up = fj.get("upper_wear") or "穿搭"
style = fj.get("style") or ""
summary = f"{style}{up}" if style and style not in up else up
elif brand != "无法判断" and name != brand:
summary = f"{brand} {name}"
else:
summary = name
summary = _build_summary(fj, name, brand, category)
return {
"name": name,
"brand": brand,
@@ -664,5 +293,15 @@ def _assemble_old(idx: int, fj: dict, ocr_texts: list[str]) -> dict[str, Any]:
"portrait_prompt": portrait_prompt,
"summary": summary,
"_source": "v2_fast_json",
"has_person": bool(fj.get("has_person", False)),
}
def _build_summary(fj: dict, name: str, brand: str, category: str) -> str:
if fj.get("has_person"):
up = fj.get("upper_wear") or "穿搭"
style = fj.get("style") or ""
base = f"{style}{up}" if style and style not in up else up
return base
if brand != "无法判断" and name != brand:
return f"{brand} {name}"
return name
@@ -23,10 +23,10 @@ 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", "15"))
_FAST_JSON_TIMEOUT = float(os.environ.get("VISION_V2_FAST_JSON_TIMEOUT", "15"))
_FAST_TIMEOUT = float(os.environ.get("VISION_V2_FAST_TIMEOUT", "12"))
_FAST_JSON_TIMEOUT = float(os.environ.get("VISION_V2_FAST_JSON_TIMEOUT", "12"))
_OCR_TIMEOUT = float(os.environ.get("VISION_V2_OCR_TIMEOUT", "6"))
_PRO_TIMEOUT = float(os.environ.get("VISION_V2_PRO_TIMEOUT", "30"))
_PRO_TIMEOUT = float(os.environ.get("VISION_V2_PRO_TIMEOUT", "25"))
_FALLBACK_RESULT = {
"name": "未识别",
@@ -1,107 +1,46 @@
# -*- coding: utf-8 -*-
"""V2 兜底路径:image_analysis(默认 qwen-vl-plus 视觉模型,fallback qwen3.7-plus / DashScope)单图调用。
"""V2 pro 兜底:qwen3.7-plus(阿里云百炼/DashScope)单次调用。
fast_json 超时/返回非 JSON/识别为空时,本路径单次调用兜底。
设计要点:
- 通过 ai_router.get_vision_client() 获取 DoubaoClient 实例,不再自己拼 httpx 请求
- enable_thinking=False + response_format=json_object
- system prompt 优先读后台 viral_video_prompt_templates 配置,DB不可用时fallback到硬编码JSON schema
- max_tokens 不传,使用 client 中 capability 的 DB 配置(避免硬编码截断 JSON)
- timeout=30s
- 返回 dict 统一走 assembler.assemble_result 组装,与 fast 路径输出格式完全一致
fast_json 结果不可用时单次调用,无竞速、无重试、无复杂超时逻辑。
直接 httpx 发精简 JSON-only prompt(比旧版 prompt_loader XML 模板短很多,降低延迟)。
"""
from __future__ import annotations
import json
import logging
import os
import time
from typing import Any
from . import _prompt, assembler
logger = logging.getLogger(__name__)
_DEFAULT_TIMEOUT = 30
_BASE_URL = "https://dashscope.aliyuncs.com/compatible-mode/v1"
_PRO_MODEL = "qwen3.7-plus"
_DEFAULT_TIMEOUT = 25
_DEFAULT_MAX_TOKENS = 800
def call_pro_vlm(
img_url: str,
idx: int,
*,
timeout: int = _DEFAULT_TIMEOUT,
max_tokens: int | None = None,
) -> dict[str, Any] | None:
"""max_tokens 默认 None:不显式传参,使用 client 内 capability 的 DB 配置。"""
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] pro vision client 不可用,跳过")
return None
except Exception as e:
logger.warning("[vision.v2] ai_router 获取失败: %s", e)
return None
system_prompt, user_prompt = _prompt.resolve_pro_prompt()
messages = [
{"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, # 图片已在 messages 中
"temperature": 0.3,
"timeout": timeout,
"enable_thinking": False,
"response_format": {"type": "json_object"},
}
if max_tokens is not None:
call_kwargs["max_tokens"] = max_tokens
raw = client.vision_completion(**call_kwargs)
elapsed = time.time() - t0
if not raw:
logger.warning("[vision.v2] pro 返回空 elapsed=%.1fs", elapsed)
return None
logger.info(
"[vision.v2] pro 完成 model=%s elapsed=%.1fs",
client.model,
elapsed,
)
s = _strip_code_fence(raw)
lpos, rr = s.find("{"), s.rfind("}")
if lpos >= 0 and rr > lpos:
s = s[lpos : rr + 1]
try:
obj = json.loads(s)
except json.JSONDecodeError:
logger.warning("[vision.v2] pro JSON 解析失败 head=%s", raw[:200])
return None
if not isinstance(obj, dict):
return None
# 通过assembler统一组装,兼容v4嵌套schema和旧扁平schema
result = assembler.assemble_result(idx, obj, [])
result["_source"] = "vlm_pro"
result["_fallback_used"] = True
return result
except Exception as e:
elapsed = time.time() - t0
logger.warning("[vision.v2] pro 异常 elapsed=%.1fs err=%s", elapsed, e, exc_info=True)
return None
_PRO_SYSTEM = (
"你是图片分析助手。仔细观察图片,严格按JSON schema返回一个对象,不要任何解释、"
"不要markdown、不要代码块、不要前后缀文字。字段值不确定时填null或空数组。\n"
"{\n"
' "has_person": true/false,\n'
' "gender": "男"/"女"/null,\n'
' "age_range": "儿童"/"青少年"/"青年"/"中年"/"老年"/null,\n'
' "outfit": "人物穿搭描述,60字以内(例:白色T恤+牛仔裤)",\n'
' "hair": "发型",\n'
' "pose": "姿态",\n'
' "expression": "表情",\n'
' "scene": "场景",\n'
' "mood": "氛围",\n'
' "has_product": true/false,\n'
' "category": "服饰/鞋包/美妆/数码/食品/家居/配饰/母婴/非产品图",\n'
' "product_name": "产品名称,非产品图填null",\n'
' "brand": "品牌或文字标识,无则null",\n'
' "key_features": ["特征数组"]\n'
"}"
)
_PRO_USER = "分析这张图片,返回符合schema的JSON。"
def _strip_code_fence(s: str) -> str:
@@ -114,3 +53,123 @@ def _strip_code_fence(s: str) -> str:
lines = lines[:-1]
s = "\n".join(lines).strip()
return s
def _assemble_pp(obj: dict[str, Any]) -> str:
if not obj.get("has_person", False):
return "无人像"
parts: list[str] = []
gender = obj.get("gender")
age = obj.get("age_range")
if gender:
parts.append(gender + ("性" if not gender.endswith("性") else ""))
if age:
parts.append(age)
parts.append("人物")
hair = obj.get("hair")
if hair:
parts.append(hair)
outfit = obj.get("outfit")
if outfit:
parts.append(f"身着{outfit}")
pose = obj.get("pose")
if pose:
parts.append(f"姿态{pose}")
expr = obj.get("expression")
if expr:
parts.append(f"表情{expr}")
return ",".join(parts) if parts else "无人像"
def call_pro_vlm(
img_url: str,
idx: int,
*,
timeout: int = _DEFAULT_TIMEOUT,
) -> dict[str, Any] | None:
"""单次调用 qwen3.7-plus,解析后返回 product dict;失败返回 None。"""
t0 = time.time()
import httpx
api_key = os.environ.get("DASHSCOPE_API_KEY")
if not api_key:
logger.warning("[vision.v2] DASHSCOPE_API_KEY 未配置,跳过 pro 兜底")
return None
url = f"{_BASE_URL}/chat/completions"
payload: dict[str, Any] = {
"model": _PRO_MODEL,
"messages": [
{"role": "system", "content": _PRO_SYSTEM},
{
"role": "user",
"content": [
{"type": "image_url", "image_url": {"url": img_url}},
{"type": "text", "text": _PRO_USER},
],
},
],
"temperature": 0.3,
"max_tokens": _DEFAULT_MAX_TOKENS,
"stream": False,
"enable_thinking": False,
"response_format": {"type": "json_object"},
}
try:
r = httpx.post(
url,
headers={"Authorization": f"Bearer {api_key}", "Content-Type": "application/json"},
json=payload,
timeout=timeout,
)
elapsed = time.time() - t0
if r.status_code != 200:
logger.warning("[vision.v2] pro HTTP %d elapsed=%.1fs body=%s", r.status_code, elapsed, r.text[:200])
return None
data = r.json()
raw = (data.get("choices") or [{}])[0].get("message", {}).get("content")
if not raw:
logger.warning("[vision.v2] pro 返回空 elapsed=%.1fs", elapsed)
return None
usage = data.get("usage") or {}
logger.info(
"[vision.v2] pro 完成 idx=%d model=%s elapsed=%.1fs in=%d out=%d",
idx,
_PRO_MODEL,
elapsed,
usage.get("prompt_tokens", 0),
usage.get("completion_tokens", 0),
)
text = _strip_code_fence(raw)
l, r_pos = text.find("{"), text.rfind("}")
if l < 0 or r_pos <= l:
logger.warning("[vision.v2] pro 无JSON elapsed=%.1fs head=%s", elapsed, raw[:200])
return None
obj = json.loads(text[l : r_pos + 1])
if not isinstance(obj, dict):
return None
scene = obj.get("scene") or "通用"
mood = obj.get("mood") or ""
pp = _assemble_pp(obj)
has_person = obj.get("has_person", False)
has_product = obj.get("has_product", False)
name = obj.get("product_name") or "未识别"
brand = obj.get("brand") or "无法判断"
category = obj.get("category") or ("非产品图" if has_person and not has_product else "无法判断")
return {
"name": name,
"brand": brand,
"category": category,
"appearance": obj.get("outfit") or "无法判断",
"packaging": "无法判断",
"text_on_package": [],
"key_features": obj.get("key_features") or ["无法判断"],
"scene": scene,
"mood": mood,
"portrait_prompt": pp,
"summary": f"{brand} {name}" if name != "未识别" else "未识别",
"_source": "vlm_pro",
}
except Exception as e:
logger.warning("[vision.v2] pro 异常 idx=%d elapsed=%.1fs err=%s", idx, time.time() - t0, e, exc_info=True)
return None
@@ -1,29 +1,67 @@
# -*- coding: utf-8 -*-
"""V2 快速路径:image_analysis capability(默认 qwen-vl-plus 视觉模型 / DashScope)强约束 JSON-only 调用。
"""V2 快速路径:qwen3.8-flash(阿里云百炼/DashScope)强约束 JSON-only 调用。
目标:替代"人体属性/商品检测/图像标签"三个火山不存在的专用云端 API。
设计要点:
- 通过 ai_router.get_vision_client() 获取 DoubaoClient 实例,不再自己拼 httpx 请求
- enable_thinking=False 关闭推理链(reasoning 是延迟主因)
- response_format=json_object 强约束JSON输出
- system prompt 优先读后台 viral_video_prompt_templates 配置,DB不可用时fallback到硬编码JSON schema
- max_tokens 不传,使用 client 中 capability 的 DB 配置(避免硬编码截断 JSON)
- temperature=0.1(稳定输出 JSON)
- timeout=15s(失败由外层走 pro 兜底)
- 直接用 httpx 发最小 payload 到 DashScope OpenAI 兼容 endpoint,不走 ai_client 包装
- enable_thinking=false 关闭推理链(reasoning 是延迟主因)
- system prompt 极致精简,只给字段 schema 和强约束(禁止自然语言、禁止 markdown)
- max_tokens=350、temperature=0.1(稳定输出 JSON)
- timeout=12s(失败由外层走 pro 兜底)
- API Key 从环境变量 DASHSCOPE_API_KEY 读取
"""
from __future__ import annotations
import json
import logging
import os
import time
from typing import Any
from . import _prompt
logger = logging.getLogger(__name__)
_DEFAULT_TIMEOUT = 15
# DashScope OpenAI 兼容 endpoint
_BASE_URL = "https://dashscope.aliyuncs.com/compatible-mode/v1"
_FAST_MODEL = "qwen3.8-flash"
_DEFAULT_TIMEOUT = 12
_DEFAULT_MAX_TOKENS = 350
# 极简 system prompt:只给字段定义 + 硬性输出要求
_FAST_SYSTEM = (
"你是图片结构化识别器。严格按下方 JSON schema 返回一个对象,不要任何解释、"
"不要markdown、不要代码块、不要前后缀文字。字段值不确定时填 null 或空数组。\n"
"{\n"
' "has_person": true/false,\n'
' "gender": "男"/"女"/null,\n'
' "age_range": "儿童"/"青少年"/"青年"/"中年"/"老年"/null,\n'
' "upper_wear": "上装款式,如T恤/衬衫/卫衣/毛衣/西装/夹克/连衣裙/吊带/背心/外套等",\n'
' "upper_color": "上装主色",\n'
' "lower_wear": "下装款式;穿连衣裙时填null",\n'
' "lower_color": "下装主色",\n'
' "dress_color": "连衣裙主色(穿连衣裙时填)",\n'
' "accessories": ["眼镜"/"帽子"/"项链"/"耳环"/"背包"/"手表"等数组],\n'
' "hairstyle": "发型,如短发/长发/马尾/卷发/丸子头/光头等",\n'
' "expression": "表情,如微笑/严肃/酷/开心等",\n'
' "pose": "姿势,如站立/坐姿/侧身/行走等",\n'
' "scene": "场景,如室内/街拍/户外/办公室/家居/海边/雪景/森林等",\n'
' "style": "风格,如休闲/商务/运动/复古/潮流/甜美/酷飒/优雅/街头/法式等",\n'
' "has_product": true/false,\n'
' "category": "产品类目:服饰/鞋包/美妆/数码/食品/家居/配饰/母婴/非产品图",\n'
' "product_name": "产品名称,非产品图填null",\n'
' "brand": "品牌或文字标识,无则null",\n'
' "material": "材质,如棉质/牛仔/皮革/真丝/针织/涤纶等",\n'
' "pattern": "图案,如纯色/条纹/波点/格子/印花/碎花/Logo等",\n'
' "colors": ["主色数组"],\n'
' "mood": "整体氛围/情绪,如清新/活力/高级/温暖/冷峻/甜美/复古等"\n'
"}"
)
_FAST_USER = "识别这张图片的人物穿搭与主体信息,只返回JSON对象。"
def _api_key() -> str | None:
return os.environ.get("DASHSCOPE_API_KEY")
def _strip_code_fence(s: str) -> str:
@@ -42,65 +80,82 @@ def call_fast_json(
img_url: str,
*,
timeout: int = _DEFAULT_TIMEOUT,
max_tokens: int | None = None,
max_tokens: int = _DEFAULT_MAX_TOKENS,
) -> dict[str, Any] | None:
"""调用 vision client 返回结构化 dict;失败/非 JSON 返回 None。
max_tokens 默认 None:不显式传参,使用 client 内 capability 的 DB 配置;
显式传入时作为覆盖。
"""
"""调用 qwen3.8-flash 返回结构化 dict;失败/非 JSON 返回 None。"""
t0 = time.time()
import httpx
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:
logger.warning("[vision.v2] ai_router 获取失败: %s", e)
api_key = _api_key()
if not api_key:
logger.warning("[vision.v2] DASHSCOPE_API_KEY 未配置,跳过 fast_json")
return None
system_prompt, user_prompt = _prompt.resolve_fast_prompt()
messages = [
{"role": "system", "content": system_prompt},
{
"role": "user",
"content": [
{"type": "image_url", "image_url": {"url": img_url}},
{"type": "text", "text": user_prompt},
],
},
]
url = f"{_BASE_URL}/chat/completions"
payload: dict[str, Any] = {
"model": _FAST_MODEL,
"messages": [
{"role": "system", "content": _FAST_SYSTEM},
{
"role": "user",
"content": [
{"type": "image_url", "image_url": {"url": img_url}},
{"type": "text", "text": _FAST_USER},
],
},
],
"temperature": 0.1,
"max_tokens": max_tokens,
"stream": False,
"enable_thinking": False,
"response_format": {"type": "json_object"},
}
try:
call_kwargs: dict[str, Any] = {
"messages": messages,
"images": None, # 图片已在 messages 中
"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
raw = client.vision_completion(**call_kwargs)
resp = httpx.post(
url,
headers={"Authorization": f"Bearer {api_key}", "Content-Type": "application/json"},
json=payload,
timeout=timeout,
)
elapsed = time.time() - t0
if resp.status_code == 400 and "enable_thinking" in resp.text[:300].lower():
# 极少数 endpoint 版本不识别 enable_thinking,重试一次不带
logger.warning("[vision.v2] fast_json HTTP 400 thinking 参数不兼容,重试 elapsed=%.1fs", elapsed)
payload.pop("enable_thinking", None)
resp = httpx.post(
url,
headers={"Authorization": f"Bearer {api_key}", "Content-Type": "application/json"},
json=payload,
timeout=timeout,
)
elapsed = time.time() - t0
if resp.status_code != 200:
logger.warning(
"[vision.v2] fast_json HTTP %d elapsed=%.1fs body=%s", resp.status_code, elapsed, resp.text[:200]
)
return None
data = resp.json()
raw = (data.get("choices") or [{}])[0].get("message", {}).get("content")
if not raw:
logger.warning("[vision.v2] fast_json 返回空 elapsed=%.1fs", elapsed)
return None
usage = data.get("usage") or {}
reasoning_tokens = usage.get("reasoning_tokens", 0)
ctd = usage.get("completion_tokens_details") or {}
if not reasoning_tokens:
reasoning_tokens = ctd.get("reasoning_tokens", 0)
logger.info(
"[vision.v2] fast_json 完成 model=%s elapsed=%.1fs",
client.model,
"[vision.v2] fast_json 完成 model=%s elapsed=%.1fs in=%d out=%d reasoning=%d",
_FAST_MODEL,
elapsed,
usage.get("prompt_tokens", 0),
usage.get("completion_tokens", 0),
reasoning_tokens,
)
text = _strip_code_fence(raw)
lpos, r = text.find("{"), text.rfind("}")
if lpos >= 0 and r > lpos:
text = text[lpos : r + 1]
l, r = text.find("{"), text.rfind("}")
if l >= 0 and r > l:
text = text[l : r + 1]
try:
obj = json.loads(text)
except json.JSONDecodeError:
@@ -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)
@@ -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 注册表 — 反向轮询模式下用于心跳与监控."""
+4 -17
View File
@@ -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 规范化:去掉末尾的路径残留(兼容旧版配置)
+2 -7
View File
@@ -55,14 +55,9 @@ _LOCATIONS = ["title", "hook", "body_points", "cta", "script_segments"]
class Reviewer:
def __init__(self, client=None):
if client is None:
try:
from packages.shared.ai_router import ai_router
from packages.shared.ai_client import get_doubao_client
client = ai_router.get_llm_client("copy_review")
except Exception:
from packages.shared.ai_client import get_doubao_client
client = get_doubao_client()
client = get_doubao_client()
self.client = client
# ── 审核 ────────────────────────────────────────────────────────────
+35 -27
View File
@@ -80,46 +80,54 @@ 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 = 1
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
doubao_model: str = "doubao-seed-2-1-pro-260915" # 推理模型(Seed 2.1 Pro,深度思考+多模态;原 seed-1-6 已下线)
doubao_fast_model: str = (
"doubao-seed-2-1-pro-260915" # #2181: lite方舟侧100%超时,默认fast_model也走pro;方舟恢复lite后通过ENV DOUBAO_FAST_MODEL切回
)
doubao_base_url: str = "https://ark.cn-beijing.volces.com/api/v3"
doubao_timeout: int = 45 # #2180: 方舟LLM高峰期响应6-8s,原30s太紧提到45s
doubao_max_retries: int = 1 # #2180: timeout调大后一次调用就够,1次重试防偶发抖动;避免6次重试叠加到351s
doubao_vision_model: str = (
"doubao-seed-2-1-pro-260915" # 高精度视觉(Seed 2.1 Pro 原生多模态;原 vision-pro-250328 已下线)
)
doubao_vision_lite_model: str = (
"doubao-seed-2-1-lite-260915" # 快速视觉(Seed 2.1 Lite 原生多模态;原 vision-lite-250315 不可用)
)
doubao_vision_use_lite: bool = True # #2188: lite恢复稳定,爆款视频默认lite-first提速(20-30s)
doubao_embedding_model: str = "doubao-embedding-vision-251215" # 多模态向量化(原 large-text-240915 已 Retiring)
doubao_video_model: str = "doubao-seedance-2-5-260628"
doubao_video_timeout: int = 600 # 视频生成轮询总超时(秒)
doubao_video_poll_interval: int = 10 # 轮询间隔(秒)
doubao_image_model: str = (
"doubao-seedream-5-0-flash-260915" # #2173: 信任链 Seedream 改 flash 模型(实测 pro 46.5s→flash 13s;pro AI化图仍被Seedance拦截)
)
doubao_image_size: str = "1K" # #2173: 1K 已足够做 Seedance 参考图,2K 在 flash 下也 22s,1K 13s
doubao_image_timeout: int = 60 # #2173: flash+1K 通常15s内,给60s余量
doubao_trust_chain_enabled: bool = (
True # #2173: 信任链总开关;若Seedream产物仍被Seedance拦截,可配 False 关闭直接t2v降级
)
# ── DashScope (阿里云百炼 Wan 3.0 等) ─────────────────────────────────
dashscope_api_key: str = ""
dashscope_base_url: str = ""
dashscope_video_timeout: int = 900
dashscope_base_url: str = "https://dashscope.aliyuncs.com/api/v1"
dashscope_video_timeout: int = 900 # Wan 视频任务轮询总超时(秒)
dashscope_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:暂停积分系统但保留全部代码/表/接口)。
-376
View File
@@ -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
+14 -77
View File
@@ -2,21 +2,17 @@
v1.6.1: 按产品决策,智能混剪/AI数字人/AI配音/抖音解析/改写/标题/封面 全部免费,
仅保留声音克隆合成(voice_clone_synth)的扣点逻辑;声音克隆训练保持免费。
爆款视频(viral_video)走动态定价,计费参数 DB 化(feature_pricing_configs,
见 feature_pricing_service),calculate_viral_video_credits 从配置读取单价/
固定成本/利润系数/封顶,DB 不可用时回落兜底配置。
爆款视频(viral_video)走动态定价,见本文件 VIRAL_VIDEO_MODEL_PRICES + calculate_viral_video_credits。
"""
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 不可用时回落内置兜底配置。
# 以下三个常量仅为向后兼容保留(旧引用方/兜底场景),值取自兜底配置。
# ============ 爆款视频动态定价 (#2151) ============
# key = (model_id, resolution, has_video_input),单位:
# - billing_mode=token: 元/百万tokens(输出)
# - billing_mode=per_second: 元/秒(视频时长)
VIRAL_VIDEO_MODEL_PRICES: dict[tuple[str, str, bool], float] = {
("seedance-2.5", "480p", False): 70.0,
("seedance-2.5", "720p", False): 70.0,
@@ -37,9 +33,9 @@ VIRAL_VIDEO_MODEL_PRICES: dict[tuple[str, str, bool], float] = {
("wan-3.0", "1080p", False): 1.2,
}
# 固定成本(元):VLM 分析 + LLM 文案 + TTS + OSS + 服务器(兜底默认值)
# 固定成本(元):VLM 分析 + LLM 文案 + TTS + OSS + 服务器
VIRAL_VIDEO_FIXED_COST = 0.15
# 利润系数(兜底默认值)
# 利润系数
VIRAL_VIDEO_PROFIT_MULTIPLIER = 1.3
# Seedance 输出帧率
VIRAL_VIDEO_FPS = 24
@@ -226,22 +222,17 @@ def calculate_viral_video_credits_with_breakdown(
) -> 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。
model_price / width / height / fps 字段,便于前端展示计费明细。
"""
w = max(1, int(width or 1))
h = max(1, int(height or 1))
@@ -251,36 +242,11 @@ def calculate_viral_video_credits_with_breakdown(
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)
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)
dur = max(1, int(duration_seconds or 15))
if billing == "per_second":
tokens = 0.0
video_cost = dur * float(price)
@@ -293,40 +259,13 @@ def calculate_viral_video_credits_with_breakdown(
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)
total = (video_cost + VIRAL_VIDEO_FIXED_COST) * VIRAL_VIDEO_PROFIT_MULTIPLIER
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),
"fixed_cost": float(VIRAL_VIDEO_FIXED_COST),
"profit_multiplier": float(VIRAL_VIDEO_PROFIT_MULTIPLIER),
"model_price": float(price),
"model_key": prefix,
"billing_mode": billing,
@@ -335,8 +274,6 @@ def calculate_viral_video_credits_with_breakdown(
"height": int(h),
"fps": int(effective_fps),
"duration": dur,
"feature_enabled": True,
"charged": True,
}
return credits, breakdown
+1 -30
View File
@@ -191,40 +191,11 @@ 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
self.status = ViralVideoStatus.RUNNING
self.updated_at = datetime.now(timezone.utc)
+25 -95
View File
@@ -170,33 +170,13 @@ class DoubaoClient:
未配置 API Key 时 is_available 为 False,调用方应降级处理。
"""
def __init__(
self,
api_key: str = "",
base_url: str = "",
model: str = "",
timeout: int = 0,
max_retries: int = 0,
max_tokens: int | None = None,
temperature: float | None = None,
extra_params: dict | None = None,
provider: str = "volcengine",
) -> None:
def __init__(self) -> None:
settings = get_shared_settings()
self.provider: str = provider
if provider == "dashscope":
self.api_key: str = api_key or getattr(settings, "dashscope_api_key", "")
self.model: str = model or getattr(settings, "dashscope_model", "")
self.base_url: str = (base_url or getattr(settings, "dashscope_base_url", "")).rstrip("/")
else: # volcengine (default)
self.api_key = api_key or settings.doubao_api_key
self.model = model or settings.doubao_model
self.base_url = (base_url or settings.doubao_base_url).rstrip("/")
self.timeout: int = timeout or settings.doubao_timeout
self.max_retries: int = max_retries or settings.doubao_max_retries
self.max_tokens: int | None = max_tokens
self.temperature: float | None = temperature
self.extra_params: dict = extra_params or {}
self.api_key: str = settings.doubao_api_key
self.model: str = settings.doubao_model
self.base_url: str = settings.doubao_base_url.rstrip("/")
self.timeout: int = settings.doubao_timeout
self.max_retries: int = settings.doubao_max_retries
self.vision_model: str = settings.doubao_vision_model
self.vision_lite_model: str = settings.doubao_vision_lite_model
self.fast_model: str = settings.doubao_fast_model
@@ -260,17 +240,16 @@ class DoubaoClient:
self,
messages: list[dict[str, str]],
temperature: float = 0.7,
max_tokens: int | None = None,
max_tokens: int = 1024,
model: str | None = None,
timeout: int | None = None,
**kwargs,
) -> Optional[str]:
"""调用 Chat Completion 接口.
Args:
messages: 对话消息列表,[{"role": "user"/"system"/"assistant", "content": "..."}]
temperature: 采样温度,0-2,默认0.7
max_tokens: 最大生成token数,默认 None(使用实例 self.max_tokens DB 配置,兜底 1024)
max_tokens: 最大生成token数,默认1024
Returns:
模型返回的文本内容,失败返回 None
@@ -283,18 +262,12 @@ class DoubaoClient:
"Authorization": f"Bearer {self.api_key}",
"Content-Type": "application/json",
}
effective_max_tokens = max_tokens if max_tokens is not None else (self.max_tokens or 1024)
payload: dict[str, Any] = {
"model": model or self.model,
"messages": messages,
"temperature": temperature,
"max_tokens": effective_max_tokens,
"max_tokens": max_tokens,
}
# 合并实例级额外参数和调用方传入的额外参数
if self.extra_params:
payload.update(self.extra_params)
if kwargs:
payload.update(kwargs)
last_error: Optional[Exception] = None
_t0 = time.time()
@@ -309,22 +282,6 @@ class DoubaoClient:
)
response.raise_for_status()
data = response.json()
finish_reason = (data.get("choices") or [{}])[0].get("finish_reason", "")
if finish_reason == "length" and attempt < self.max_retries:
# 输出被 max_tokens 截断:1.5x 扩容后重试(计入 max_retries,不额外增加)
old_max = int(payload["max_tokens"])
new_max = int(old_max * 1.5)
payload["max_tokens"] = new_max
wait = 0.5 * (2**attempt)
logger.warning(
"输出被max_tokens截断(%d),扩容到%d后重试 (第%d/%d次)",
old_max,
new_max,
attempt + 1,
self.max_retries + 1,
)
time.sleep(wait)
continue
content = data["choices"][0]["message"]["content"]
_elapsed = time.time() - _t0
logger.info(
@@ -334,7 +291,6 @@ class DoubaoClient:
data.get("usage", {}).get("completion_tokens", 0),
_elapsed,
attempt + 1,
_req_timeout,
)
return content.strip()
except Exception as e:
@@ -358,21 +314,20 @@ class DoubaoClient:
self,
messages: list[dict],
images: list[str] | None = None,
max_tokens: int | None = None,
max_tokens: int = 2048,
temperature: float = 0.3,
timeout: int | None = None,
model: str | None = None,
**kwargs,
) -> Optional[str]:
"""调用豆包视觉理解 API(OpenAI 兼容多模态格式).
将 images 附加到最后一条 user message 的 content 中,
使用构造函数传入的 self.model(DB capability 绑定的视觉模型,默认 qwen-vl-plus)。
使用 vision_model(默认 doubao-1-5-vision-pro-250328)。
Args:
messages: 对话消息列表。最后一条 user message 会被注入图片内容。
images: 图片列表,支持 base64 data URI 或 HTTP(S) URL。
max_tokens: 最大生成 token 数,默认 None(使用实例 self.max_tokens DB 配置,兜底 2048)。
max_tokens: 最大生成 token 数,默认 2048。
temperature: 采样温度,默认 0.3(视觉任务偏低更稳定)。
timeout: 单次请求超时秒数,不传则使用默认 self.timeout。
@@ -412,17 +367,12 @@ class DoubaoClient:
"Authorization": f"Bearer {self.api_key}",
"Content-Type": "application/json",
}
effective_max_tokens = max_tokens if max_tokens is not None else (self.max_tokens or 2048)
payload: dict[str, Any] = {
"model": model or self.model,
"model": model or self.vision_model,
"messages": vision_messages,
"temperature": temperature,
"max_tokens": effective_max_tokens,
"max_tokens": max_tokens,
}
if self.extra_params:
payload.update(self.extra_params)
if kwargs:
payload.update(kwargs)
req_timeout = timeout or self.timeout
last_error: Optional[Exception] = None
@@ -437,22 +387,6 @@ class DoubaoClient:
)
response.raise_for_status()
data = response.json()
finish_reason = (data.get("choices") or [{}])[0].get("finish_reason", "")
if finish_reason == "length" and attempt < self.max_retries:
# 视觉输出被 max_tokens 截断:1.5x 扩容后重试(计入 max_retries)
old_max = int(payload["max_tokens"])
new_max = int(old_max * 1.5)
payload["max_tokens"] = new_max
wait = 0.5 * (2**attempt)
logger.warning(
"视觉输出被max_tokens截断(%d),扩容到%d后重试 (第%d/%d次)",
old_max,
new_max,
attempt + 1,
self.max_retries + 1,
)
time.sleep(wait)
continue
content = data["choices"][0]["message"]["content"]
_elapsed = time.time() - _t0
logger.info(
@@ -655,32 +589,28 @@ class DoubaoClient:
# 信任链只作用于 doubao provider;DashScope(Wan) 保持原行为。
trust_chain_applied = False
if provider == "doubao" and getattr(self, "trust_chain_enabled", True) and pre_trusted_images:
# #2220: 稀疏列表模式——pre_trusted_images 与 raw_portrait_urls 等长,
# None 位保留原图,非 None 位用 AI 人像替换。
raw_portrait_urls: list[str] = []
if image_url:
raw_portrait_urls.append(image_url)
for u in ref_imgs:
if u not in raw_portrait_urls:
raw_portrait_urls.append(u)
_n_trusted = sum(1 for _x in pre_trusted_images if _x)
if _n_trusted >= 1 and len(pre_trusted_images) >= len(raw_portrait_urls):
merged: list[str] = []
for _i, _orig in enumerate(raw_portrait_urls):
_ai = pre_trusted_images[_i] if _i < len(pre_trusted_images) else None
merged.append(str(_ai) if _ai else _orig)
trusted_urls: list[str] = []
if len(pre_trusted_images) >= 1:
trusted_urls = list(pre_trusted_images)
trust_chain_applied = True
logger.info(
"[trust-chain] 稀疏替换 %d/%d 张为AI人像(场景/商品图保留原图),走reference_image模式",
_n_trusted,
"[trust-chain] 使用预热t2i结果 %d 张,替换原参考图走 reference_image 模式(原n=%d)",
len(trusted_urls),
len(raw_portrait_urls),
)
# 替换:image_url 用第一张(可能是AI或原图),ref_imgs 用剩余
if image_url and merged:
image_url = merged[0]
ref_imgs = merged[1:] if len(merged) > 1 else []
if trust_chain_applied and trusted_urls:
# 替换:原 image_url 用第一张 AI 图,ref_imgs 用剩余
if image_url and trusted_urls:
image_url = trusted_urls[0]
ref_imgs = trusted_urls[1:] if len(trusted_urls) > 1 else []
else:
ref_imgs = merged
ref_imgs = trusted_urls
# ─────────────────────────────────────────────────────────────────
# 判断任务模式:
-62
View File
@@ -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
-478
View File
@@ -1,478 +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 选择模型,不存在则 fallback"""
if variant == "lite" and cap.lite_model:
return cap.lite_model
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()
-395
View File
@@ -1,395 +0,0 @@
"""AI Router 单元测试 — 23 cases covering routing/cache/fallback/client construction."""
from __future__ import annotations
import sys
import unittest
from dataclasses import dataclass
from typing import Optional
from unittest.mock import MagicMock, patch
# ── Pre-mock heavy import chain to avoid pulling in full app ──
_mock_config = MagicMock()
_mock_settings = MagicMock()
_mock_settings.doubao_model = "doubao-seed-2-1-pro-260915"
_mock_settings.doubao_fast_model = "doubao-seed-2-1-pro-260915"
_mock_settings.doubao_base_url = "https://ark.cn-beijing.volces.com/api/v3"
_mock_settings.doubao_api_key = "test-key"
_mock_settings.doubao_timeout = 45
_mock_settings.doubao_max_retries = 1
_mock_settings.doubao_image_model = "doubao-seedream-5-0-flash-260915"
_mock_settings.doubao_image_timeout = 60
_mock_settings.doubao_video_model = "doubao-seedance-2-5-260628"
_mock_settings.doubao_video_timeout = 600
_mock_settings.dashscope_api_key = "ds-key"
_mock_settings.cosyvoice_api_key = "cv-key"
_mock_settings.cosyvoice_base_url = "https://dashscope.aliyuncs.com/api/v1"
_mock_settings.cosyvoice_model = "cosyvoice-v3-flash"
_mock_settings.redis_url = "redis://localhost:6379/0"
_mock_settings.celery_broker_url = "redis://localhost:6379/0"
_mock_config.get_shared_settings.return_value = _mock_settings
# Prevent the full packages.shared from loading
for mod_name in list(sys.modules.keys()):
if "packages.shared" in mod_name and "ai_router" not in mod_name and "ai_config_version" not in mod_name:
pass # don't remove, just prevent new imports
# Direct import of our modules (bypassing __init__.py)
import importlib.util
import os
def _load_module_from_file(name, path):
spec = importlib.util.spec_from_file_location(name, path)
mod = importlib.util.module_from_spec(spec)
sys.modules[name] = mod
spec.loader.exec_module(mod)
return mod
# Load ai_config_version
_ai_config_version = _load_module_from_file(
"packages.shared.ai_config_version",
os.path.join(os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__)))), "packages", "shared", "ai_config_version.py"),
)
# Patch get_shared_settings in the loaded module
_ai_config_version.get_shared_settings = lambda: _mock_settings
# Load ai_router - needs packages.shared.config to be available
sys.modules["packages.shared.config"] = MagicMock()
sys.modules["packages.shared.config"].get_shared_settings = lambda: _mock_settings
# Mock packages.shared.ai_client to avoid triggering packages.shared.__init__ chain
# (which fails on Python 3.10 due to datetime.UTC import in packages.domain)
_mock_ai_client = MagicMock()
class _FakeDoubaoClient:
"""Fake DoubaoClient for testing - mimics the real interface."""
def __init__(self, api_key="", base_url="", model="", timeout=0, max_retries=0,
max_tokens=None, temperature=None, extra_params=None, provider="volcengine"):
self.api_key = api_key
self.base_url = base_url
self.model = model
self.timeout = timeout
self.max_retries = max_retries
self.max_tokens = max_tokens
self.temperature = temperature
self.extra_params = extra_params or {}
self.provider = provider
self.vision_model = model
@property
def is_available(self):
return bool(self.api_key)
def chat_completion(self, messages, **kwargs):
return None
def vision_completion(self, messages, **kwargs):
return None
_mock_ai_client.DoubaoClient = _FakeDoubaoClient
sys.modules["packages.shared.ai_client"] = _mock_ai_client
_ai_router = _load_module_from_file(
"packages.shared.ai_router",
os.path.join(os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__)))), "packages", "shared", "ai_router.py"),
)
class TestAIConfigVersion(unittest.TestCase):
"""Redis 版本号机制测试"""
@patch.object(_ai_config_version, "_get_redis_client")
def test_bump_version_success(self, mock_redis_fn):
mock_r = MagicMock()
mock_r.set.return_value = True
mock_redis_fn.return_value = mock_r
ver = _ai_config_version.bump_version()
self.assertTrue(ver)
self.assertTrue(ver.isdigit())
mock_r.set.assert_called_once()
@patch.object(_ai_config_version, "_get_redis_client")
def test_bump_version_redis_unavailable(self, mock_redis_fn):
mock_redis_fn.return_value = None
ver = _ai_config_version.bump_version()
self.assertEqual(ver, "")
@patch.object(_ai_config_version, "_get_redis_client")
def test_get_version_success(self, mock_redis_fn):
mock_r = MagicMock()
mock_r.get.return_value = "1234567890"
mock_redis_fn.return_value = mock_r
ver = _ai_config_version.get_version()
self.assertEqual(ver, "1234567890")
@patch.object(_ai_config_version, "_get_redis_client")
def test_get_version_redis_down(self, mock_redis_fn):
mock_redis_fn.return_value = None
ver = _ai_config_version.get_version()
self.assertIsNone(ver)
@patch.object(_ai_config_version, "_get_redis_client")
def test_get_version_exception(self, mock_redis_fn):
mock_r = MagicMock()
mock_r.get.side_effect = Exception("connection refused")
mock_redis_fn.return_value = mock_r
ver = _ai_config_version.get_version()
self.assertIsNone(ver)
class TestAIRouter(unittest.TestCase):
"""AIRouter 路由/缓存/fallback 测试"""
def setUp(self):
self.router = _ai_router.AIRouter()
@patch.object(_ai_config_version, "get_version", return_value=None)
def test_get_capability_db_unavailable(self, mock_ver):
with patch.object(_ai_router, "_get_session", return_value=None):
cap = self.router.get_capability("intent_parsing")
self.assertIsNone(cap)
@patch.object(_ai_config_version, "get_version", return_value=None)
def test_get_capability_from_db(self, mock_ver):
mock_session = MagicMock()
mock_row = MagicMock()
mock_row.capability_key = "intent_parsing"
mock_row.capability_name = "文案意图解析"
mock_row.timeout_seconds = 45
mock_row.max_retries = 1
mock_row.max_tokens = None
mock_row.temperature = None
mock_row.concurrency = 2
mock_row.extra_params = {}
mock_row.is_enabled = True
mock_row.pm_id = "model-1"
mock_row.pm_name = "豆包"
mock_row.pm_provider = "volcengine"
mock_row.pm_model_key = "doubao-seed-1-6-250615"
mock_row.pm_api_key = "test-key"
mock_row.pm_api_base = "https://ark.test.com"
mock_row.pm_api_version = None
mock_row.pm_status = "active"
mock_row.lm_id = None
mock_row.fm_id = None
mock_session.execute.return_value.first.return_value = mock_row
with patch.object(_ai_router, "_get_session", return_value=mock_session):
cap = self.router.get_capability("intent_parsing")
self.assertIsNotNone(cap)
self.assertEqual(cap.capability_key, "intent_parsing")
self.assertEqual(cap.primary_model.model_key, "doubao-seed-1-6-250615")
@patch.object(_ai_config_version, "get_version", side_effect=[None, "v2"])
def test_cache_invalidation_on_version_change(self, mock_ver):
with patch.object(self.router, "_load_from_db", return_value=None):
self.router.get_capability("test_key")
self.router._local_ver = "v1"
self.assertTrue(self.router._check_version())
@patch.object(_ai_config_version, "get_version", return_value="same_ver")
def test_cache_hit_same_version(self, mock_ver):
model = _ai_router.ModelConfig(
id="m1", name="test", provider="volcengine", model_key="test-model",
api_key="key", api_base="https://test.com", api_version=None, status="active",
)
cap = _ai_router.CapabilityConfig(
capability_key="test", capability_name="test", primary_model=model,
lite_model=None, fallback_model=None, timeout_seconds=30,
max_retries=1, max_tokens=None, temperature=None, concurrency=2,
extra_params={}, is_enabled=True,
)
self.router._cache["test"] = cap
self.router._local_ver = "same_ver"
result = self.router.get_capability("test")
self.assertEqual(result, cap)
def test_invalidate_clears_cache(self):
self.router._cache["x"] = MagicMock()
self.router._local_ver = "v1"
self.router.invalidate()
self.assertEqual(len(self.router._cache), 0)
self.assertIsNone(self.router._local_ver)
@patch.object(_ai_router, "_get_session", return_value=None)
@patch.object(_ai_config_version, "get_version", return_value=None)
def test_get_llm_client_fallback(self, mock_ver, mock_session):
_ai_router.get_shared_settings = lambda: _mock_settings
client = self.router.get_llm_client("intent_parsing")
self.assertIsNotNone(client)
self.assertEqual(client.model, "doubao-seed-2-1-pro-260915")
self.assertEqual(client.api_key, "test-key")
@patch.object(_ai_config_version, "get_version", return_value=None)
def test_get_llm_client_from_db(self, mock_ver):
model = _ai_router.ModelConfig(
id="m1", name="test", provider="dashscope", model_key="qwen3.8-flash",
api_key="db-key", api_base="https://dashscope.test.com", api_version=None, status="active",
)
cap = _ai_router.CapabilityConfig(
capability_key="image_analysis", capability_name="图片分析",
primary_model=model, lite_model=None, fallback_model=None,
timeout_seconds=15, max_retries=1, max_tokens=350, temperature=0.1,
concurrency=2, extra_params={}, is_enabled=True,
)
with patch.object(self.router, "get_capability", return_value=cap):
client = self.router.get_llm_client("image_analysis")
self.assertIsNotNone(client)
self.assertEqual(client.model, "qwen3.8-flash")
self.assertEqual(client.provider, "dashscope")
@patch.object(_ai_config_version, "get_version", return_value=None)
def test_get_vision_client(self, mock_ver):
model = _ai_router.ModelConfig(
id="m1", name="test", provider="dashscope", model_key="qwen3.8-flash",
api_key="key", api_base="https://dashscope.test.com", api_version=None, status="active",
)
cap = _ai_router.CapabilityConfig(
capability_key="image_analysis", capability_name="图片分析",
primary_model=model, lite_model=None, fallback_model=None,
timeout_seconds=15, max_retries=1, max_tokens=None, temperature=None,
concurrency=2, extra_params={}, is_enabled=True,
)
with patch.object(self.router, "get_capability", return_value=cap):
client = self.router.get_vision_client("image_analysis")
self.assertIsNotNone(client)
# #2220: vision client is now DoubaoClient with vision_completion
self.assertTrue(hasattr(client, "vision_completion"))
@patch.object(_ai_config_version, "get_version", return_value=None)
def test_get_tts_client(self, mock_ver):
model = _ai_router.ModelConfig(
id="m1", name="test", provider="dashscope", model_key="cosyvoice-v3-flash",
api_key="key", api_base="https://dashscope.test.com", api_version=None, status="active",
)
cap = _ai_router.CapabilityConfig(
capability_key="tts", capability_name="语音合成",
primary_model=model, lite_model=None, fallback_model=None,
timeout_seconds=60, max_retries=1, max_tokens=None, temperature=None,
concurrency=2, extra_params={}, is_enabled=True,
)
with patch.object(self.router, "get_capability", return_value=cap):
client = self.router.get_tts_client()
self.assertIsNotNone(client)
self.assertEqual(client.model, "cosyvoice-v3-flash")
@patch.object(_ai_config_version, "get_version", return_value=None)
def test_get_image_gen_client(self, mock_ver):
model = _ai_router.ModelConfig(
id="m1", name="test", provider="volcengine", model_key="seedream-5.0-flash",
api_key="key", api_base="https://ark.test.com", api_version=None, status="active",
)
cap = _ai_router.CapabilityConfig(
capability_key="image_generation", capability_name="图片生成",
primary_model=model, lite_model=None, fallback_model=None,
timeout_seconds=60, max_retries=1, max_tokens=None, temperature=None,
concurrency=2, extra_params={"size": "1K"}, is_enabled=True,
)
with patch.object(self.router, "get_capability", return_value=cap):
client = self.router.get_image_gen_client()
self.assertIsNotNone(client)
self.assertEqual(client.model, "seedream-5.0-flash")
@patch.object(_ai_config_version, "get_version", return_value=None)
def test_get_video_gen_client(self, mock_ver):
model = _ai_router.ModelConfig(
id="m1", name="test", provider="volcengine", model_key="seedance-2.5",
api_key="key", api_base="https://ark.test.com", api_version=None, status="active",
)
cap = _ai_router.CapabilityConfig(
capability_key="video_generation", capability_name="视频生成",
primary_model=model, lite_model=None, fallback_model=None,
timeout_seconds=600, max_retries=1, max_tokens=None, temperature=None,
concurrency=1, extra_params={}, is_enabled=True,
)
with patch.object(self.router, "get_capability", return_value=cap):
client = self.router.get_video_gen_client()
self.assertIsNotNone(client)
self.assertEqual(client.model, "seedance-2.5")
@patch.object(_ai_config_version, "get_version", return_value=None)
def test_lite_variant_preference(self, mock_ver):
primary = _ai_router.ModelConfig(id="p1", name="pro", provider="volcengine", model_key="pro-model", api_key="k", api_base="u", api_version=None, status="active")
lite = _ai_router.ModelConfig(id="l1", name="lite", provider="volcengine", model_key="lite-model", api_key="k", api_base="u", api_version=None, status="active")
cap = _ai_router.CapabilityConfig(
capability_key="image_analysis", capability_name="图片分析",
primary_model=primary, lite_model=lite, fallback_model=None,
timeout_seconds=15, max_retries=1, max_tokens=None, temperature=None,
concurrency=2, extra_params={}, is_enabled=True,
)
model = self.router._get_model_or_fallback(cap, "lite")
self.assertEqual(model.model_key, "lite-model")
model_primary = self.router._get_model_or_fallback(cap, "primary")
self.assertEqual(model_primary.model_key, "pro-model")
@patch.object(_ai_config_version, "get_version", return_value=None)
def test_disabled_capability_returns_fallback(self, mock_ver):
cap = _ai_router.CapabilityConfig(
capability_key="test", capability_name="test",
primary_model=None, lite_model=None, fallback_model=None,
timeout_seconds=30, max_retries=1, max_tokens=None, temperature=None,
concurrency=2, extra_params={}, is_enabled=False,
)
_ai_router.get_shared_settings = lambda: _mock_settings
with patch.object(self.router, "get_capability", return_value=cap):
client = self.router.get_llm_client("test")
self.assertIsNotNone(client)
self.assertEqual(client.model, "doubao-seed-2-1-pro-260915")
@patch.object(_ai_config_version, "get_version", return_value=None)
def test_fallback_chain_primary_none(self, mock_ver):
"""primary_model 为 None 时 fallback 到 fallback_model"""
fb = _ai_router.ModelConfig(id="f1", name="fb", provider="volcengine", model_key="fb-model", api_key="k", api_base="u", api_version=None, status="active")
cap = _ai_router.CapabilityConfig(
capability_key="test", capability_name="test",
primary_model=None, lite_model=None, fallback_model=fb,
timeout_seconds=30, max_retries=1, max_tokens=None, temperature=None,
concurrency=2, extra_params={}, is_enabled=True,
)
model = self.router._get_model_or_fallback(cap, "primary")
self.assertEqual(model.model_key, "fb-model")
class TestModelConfig(unittest.TestCase):
"""数据类测试"""
def test_model_config_frozen(self):
m = _ai_router.ModelConfig(id="1", name="t", provider="p", model_key="k", api_key="a", api_base="b", api_version=None, status="active")
with self.assertRaises(AttributeError):
m.model_key = "new"
def test_capability_config_frozen(self):
c = _ai_router.CapabilityConfig(
capability_key="k", capability_name="n", primary_model=None,
lite_model=None, fallback_model=None, timeout_seconds=30,
max_retries=1, max_tokens=None, temperature=None, concurrency=2,
extra_params={}, is_enabled=True,
)
with self.assertRaises(AttributeError):
c.is_enabled = False
class TestClientAvailability(unittest.TestCase):
"""客户端可用性测试"""
def test_tts_client_available(self):
c = _ai_router.TTSClient(provider="p", api_key="k", base_url="u", model="m")
self.assertTrue(c.is_available)
def test_tts_client_unavailable_no_model(self):
c = _ai_router.TTSClient(provider="p", api_key="k", base_url="u", model="")
self.assertFalse(c.is_available)
def test_image_gen_client_unavailable_no_url(self):
c = _ai_router.ImageGenClient(provider="p", api_key="k", base_url="", model="m")
self.assertFalse(c.is_available)
def test_video_gen_client_available(self):
c = _ai_router.VideoGenClient(provider="p", api_key="k", base_url="u", model="m")
self.assertTrue(c.is_available)
if __name__ == "__main__":
unittest.main()
+3 -3
View File
@@ -73,13 +73,13 @@ class TestSharedSettingsDefaults:
def test_default_cosyvoice_settings(self):
s = SharedSettings()
assert s.cosyvoice_model == "" # 零硬编码:默认值已清空
assert s.cosyvoice_model == "cosyvoice-v3-flash"
assert s.cosyvoice_format == "mp3"
assert s.cosyvoice_sample_rate == 22050
def test_default_doubao_settings(self):
s = SharedSettings()
assert s.doubao_model == "" # 零硬编码:默认值已清空
assert "doubao" in s.doubao_model
assert s.doubao_timeout == 45 # #2180 默认提到45s
assert s.doubao_max_retries == 1
@@ -321,7 +321,7 @@ class TestWorkerSettingsDefaults:
assert s.database_url # 继承自SharedSettings
assert s.redis_url
assert s.oss_endpoint
assert s.cosyvoice_model == "" # 零硬编码:默认值已清空
assert s.cosyvoice_model == "cosyvoice-v3-flash"
class TestGetWorkerSettings:
+3 -3
View File
@@ -102,17 +102,17 @@ class TestSharedSettingsDefaults:
def test_default_cosyvoice_config(self):
"""CosyVoice 默认配置"""
s = self._make_settings()
assert s.cosyvoice_model == "" # 零硬编码:默认值已清空
assert s.cosyvoice_model == "cosyvoice-v3-flash"
assert s.cosyvoice_sample_rate == 22050
assert s.cosyvoice_format == "mp3"
assert s.cosyvoice_clone_model == "" # 零硬编码:默认值已清空
assert s.cosyvoice_clone_model == "voice-enrollment"
def test_default_doubao_config(self):
"""豆包默认配置"""
s = self._make_settings()
assert s.doubao_timeout == 45 # #2180 默认提到45s
assert s.doubao_max_retries == 1
assert s.doubao_base_url == "" # 零硬编码:默认值已清空
assert "volces.com" in s.doubao_base_url
def test_default_empty_api_keys(self):
"""API Key 默认空字符串"""
+10 -50
View File
@@ -27,10 +27,7 @@ def mock_client() -> MagicMock:
@pytest.fixture
def service(mock_client: MagicMock) -> CosyVoiceService:
"""Create CosyVoiceService with mocked HTTP client and config."""
with (
patch("packages.application.cosyvoice_service.get_shared_settings") as mock_settings,
patch("packages.shared.ai_router.ai_router") as mock_router,
):
with patch("packages.application.cosyvoice_service.get_shared_settings") as mock_settings:
settings = MagicMock()
settings.cosyvoice_api_key = "sk-test-12345678"
settings.cosyvoice_base_url = "https://dashscope.aliyuncs.com/api/v1"
@@ -40,7 +37,6 @@ def service(mock_client: MagicMock) -> CosyVoiceService:
settings.cosyvoice_voice = "longxiaochun_v3"
settings.cosyvoice_clone_model = "voice-enrollment"
mock_settings.return_value = settings
mock_router.get_tts_client.return_value = None # ai_router returns None in tests
svc = CosyVoiceService(http_client=mock_client)
svc.CLONE_POLL_INTERVAL = 0.001 # 加速测试
svc.RETRY_BACKOFF = 0.001
@@ -52,10 +48,7 @@ class TestInitConfig:
def test_base_url_with_old_text2audio_path_gets_normalized(self, mock_client: MagicMock) -> None:
"""旧版 base_url 带 text2audio 路径应自动修正."""
with (
patch("packages.application.cosyvoice_service.get_shared_settings") as mock_settings,
patch("packages.shared.ai_router.ai_router") as mock_router,
):
with patch("packages.application.cosyvoice_service.get_shared_settings") as mock_settings:
settings = MagicMock()
settings.cosyvoice_api_key = "sk-test"
settings.cosyvoice_base_url = "https://dashscope.aliyuncs.com/api/v1/services/aigc/text2audio"
@@ -65,16 +58,12 @@ class TestInitConfig:
settings.cosyvoice_voice = "test"
settings.cosyvoice_clone_model = "voice-enrollment"
mock_settings.return_value = settings
mock_router.get_tts_client.return_value = None
svc = CosyVoiceService(http_client=mock_client)
assert svc._base_url == "https://dashscope.aliyuncs.com/api/v1"
def test_custom_params_override_config(self, mock_client: MagicMock) -> None:
"""显式传入参数覆盖配置."""
with (
patch("packages.application.cosyvoice_service.get_shared_settings") as mock_settings,
patch("packages.shared.ai_router.ai_router") as mock_router,
):
with patch("packages.application.cosyvoice_service.get_shared_settings") as mock_settings:
settings = MagicMock()
settings.cosyvoice_api_key = "sk-config"
settings.cosyvoice_base_url = "https://config.example.com"
@@ -84,7 +73,6 @@ class TestInitConfig:
settings.cosyvoice_voice = "test"
settings.cosyvoice_clone_model = "voice-enrollment"
mock_settings.return_value = settings
mock_router.get_tts_client.return_value = None
svc = CosyVoiceService(
api_key="sk-custom",
base_url="https://custom.example.com/api/v1",
@@ -99,10 +87,7 @@ class TestInitConfig:
def test_context_manager(self, mock_client: MagicMock) -> None:
"""上下文管理器正常工作."""
with (
patch("packages.application.cosyvoice_service.get_shared_settings") as mock_settings,
patch("packages.shared.ai_router.ai_router") as mock_router,
):
with patch("packages.application.cosyvoice_service.get_shared_settings") as mock_settings:
settings = MagicMock()
settings.cosyvoice_api_key = "sk-test"
@@ -120,7 +105,6 @@ class TestInitConfig:
settings.cosyvoice_clone_model = "voice-enrollment"
mock_settings.return_value = settings
mock_router.get_tts_client.return_value = None
svc = CosyVoiceService(http_client=mock_client)
with svc as s:
assert s is svc
@@ -129,10 +113,7 @@ class TestInitConfig:
def test_owns_client_gets_closed(self) -> None:
"""自有client在close时被关闭."""
with (
patch("packages.application.cosyvoice_service.get_shared_settings") as mock_settings,
patch("packages.shared.ai_router.ai_router") as mock_router,
):
with patch("packages.application.cosyvoice_service.get_shared_settings") as mock_settings:
settings = MagicMock()
settings.cosyvoice_api_key = "sk-test"
@@ -150,7 +131,6 @@ class TestInitConfig:
settings.cosyvoice_clone_model = "voice-enrollment"
mock_settings.return_value = settings
mock_router.get_tts_client.return_value = None
with patch("packages.application.cosyvoice_service.httpx.Client") as mock_cls:
mock_instance = MagicMock()
mock_cls.return_value = mock_instance
@@ -193,10 +173,7 @@ class TestSubmitCloneTask:
def test_no_api_key_raises_auth_error(self, mock_client: MagicMock) -> None:
"""无API Key抛认证错误."""
with (
patch("packages.application.cosyvoice_service.get_shared_settings") as mock_settings,
patch("packages.shared.ai_router.ai_router") as mock_router,
):
with patch("packages.application.cosyvoice_service.get_shared_settings") as mock_settings:
settings = MagicMock()
settings.cosyvoice_api_key = ""
@@ -214,7 +191,6 @@ class TestSubmitCloneTask:
settings.cosyvoice_clone_model = "voice-enrollment"
mock_settings.return_value = settings
mock_router.get_tts_client.return_value = None
svc = CosyVoiceService(http_client=mock_client)
with pytest.raises(CosyVoiceAuthError, match="API Key 未配置"):
svc.submit_clone_task(audio_url="https://example.com/audio.mp3")
@@ -282,10 +258,7 @@ class TestSubmitCloneTask:
def test_audio_url_signer_is_called(self, mock_client: MagicMock) -> None:
"""配置了audio_url_signer时会被调用预签名."""
with (
patch("packages.application.cosyvoice_service.get_shared_settings") as mock_settings,
patch("packages.shared.ai_router.ai_router") as mock_router,
):
with patch("packages.application.cosyvoice_service.get_shared_settings") as mock_settings:
settings = MagicMock()
settings.cosyvoice_api_key = "sk-test"
@@ -303,7 +276,6 @@ class TestSubmitCloneTask:
settings.cosyvoice_clone_model = "voice-enrollment"
mock_settings.return_value = settings
mock_router.get_tts_client.return_value = None
signer = MagicMock(return_value="https://signed.example.com/audio.mp3?token=xxx")
svc = CosyVoiceService(http_client=mock_client, audio_url_signer=signer)
@@ -325,10 +297,7 @@ class TestSubmitCloneTask:
def test_signer_failure_falls_back_to_original_url(self, mock_client: MagicMock) -> None:
"""预签名失败时回退到原始URL."""
with (
patch("packages.application.cosyvoice_service.get_shared_settings") as mock_settings,
patch("packages.shared.ai_router.ai_router") as mock_router,
):
with patch("packages.application.cosyvoice_service.get_shared_settings") as mock_settings:
settings = MagicMock()
settings.cosyvoice_api_key = "sk-test"
@@ -346,7 +315,6 @@ class TestSubmitCloneTask:
settings.cosyvoice_clone_model = "voice-enrollment"
mock_settings.return_value = settings
mock_router.get_tts_client.return_value = None
signer = MagicMock(side_effect=RuntimeError("sign failed"))
svc = CosyVoiceService(http_client=mock_client, audio_url_signer=signer)
@@ -407,10 +375,7 @@ class TestQueryVoiceStatus:
def test_no_api_key_raises(self, mock_client: MagicMock) -> None:
"""无API Key抛认证错误."""
with (
patch("packages.application.cosyvoice_service.get_shared_settings") as mock_settings,
patch("packages.shared.ai_router.ai_router") as mock_router,
):
with patch("packages.application.cosyvoice_service.get_shared_settings") as mock_settings:
settings = MagicMock()
settings.cosyvoice_api_key = ""
@@ -428,7 +393,6 @@ class TestQueryVoiceStatus:
settings.cosyvoice_clone_model = "voice-enrollment"
mock_settings.return_value = settings
mock_router.get_tts_client.return_value = None
svc = CosyVoiceService(http_client=mock_client)
with pytest.raises(CosyVoiceAuthError):
svc.query_voice_status("v1")
@@ -625,10 +589,7 @@ class TestSubmitSynthesizeTask:
def test_no_api_key_raises(self, mock_client: MagicMock) -> None:
"""无API Key抛认证错误."""
with (
patch("packages.application.cosyvoice_service.get_shared_settings") as mock_settings,
patch("packages.shared.ai_router.ai_router") as mock_router,
):
with patch("packages.application.cosyvoice_service.get_shared_settings") as mock_settings:
settings = MagicMock()
settings.cosyvoice_api_key = ""
@@ -646,7 +607,6 @@ class TestSubmitSynthesizeTask:
settings.cosyvoice_clone_model = "voice-enrollment"
mock_settings.return_value = settings
mock_router.get_tts_client.return_value = None
svc = CosyVoiceService(http_client=mock_client)
with pytest.raises(CosyVoiceAuthError):
svc.submit_synthesize_task(text="你好", voice_id="v1")
@@ -521,81 +521,3 @@ class TestIngestJob:
storage_key="k",
)
assert job.error_message == ""
class TestViralVideoResumeForRegenerate:
"""#2222: resume_from_image_analyzed 应支持 COPY_GENERATED/COMPLETED/FAILED 重新生成文案。"""
def test_regen_from_copy_generated_clears_old_copy(self):
from datetime import datetime, timezone
from packages.domain.viral_video import ViralVideoJob, ViralVideoStatus
job = ViralVideoJob(user_id="u1", images=["img1"])
# 模拟已经生成过文案和视频
job.status = ViralVideoStatus.COPY_GENERATED
job.copy_result = {"shots": [{"x": 1}], "voiceover_script": "旧文案"}
job.intent_result = {"intent": "旧意图"}
job.storyboard = [{"x": 1}]
job.generated_copy_text = "旧文案"
job.result_video_url = "http://old.mp4"
job.completed_at = datetime(2026, 10, 6, tzinfo=timezone.utc)
job.error_msg = ""
job.current_stage = "tts_generation"
job.phase_message = "TTS完成"
# 重新生成
job.resume_from_image_analyzed()
assert job.status == ViralVideoStatus.RUNNING
assert job.copy_result is None
assert job.intent_result is None
assert job.storyboard is None
assert job.generated_copy_text == ""
assert job.result_video_url == ""
assert job.completed_at is None
assert job.error_msg == ""
assert job.current_stage == ""
assert job.phase_message == ""
def test_regen_from_completed_clears_old_copy(self):
from packages.domain.viral_video import ViralVideoJob, ViralVideoStatus
job = ViralVideoJob(user_id="u1", images=["img1"])
job.status = ViralVideoStatus.COMPLETED
job.copy_result = {"shots": [], "voiceover_script": "xx"}
job.intent_result = {"intent": "x"}
job.result_video_url = "http://v.mp4"
job.resume_from_image_analyzed()
assert job.status == ViralVideoStatus.RUNNING
assert job.copy_result is None
assert job.intent_result is None
assert job.result_video_url == ""
def test_first_call_from_image_analyzed_keeps_fields(self):
"""首次进入(IMAGE_ANALYZED)不应清空任何已有的字段。"""
from packages.domain.viral_video import ViralVideoJob, ViralVideoStatus
job = ViralVideoJob(user_id="u1", images=["img1"])
job.status = ViralVideoStatus.IMAGE_ANALYZED
job.image_analysis = {"products": []}
job.industry = "美妆"
job.resume_from_image_analyzed()
assert job.status == ViralVideoStatus.RUNNING
assert job.image_analysis == {"products": []}
assert job.industry == "美妆"
def test_wait_user_confirm_rejected(self):
"""wait_user_confirm 中间状态应被拒绝(前端正在编辑/确认文案)。"""
import pytest
from packages.domain.viral_video import ViralVideoJob, ViralVideoStatus
job = ViralVideoJob(user_id="u1", images=["img1"])
job.status = ViralVideoStatus.WAIT_USER_CONFIRM
with pytest.raises(ValueError, match="Cannot resume"):
job.resume_from_image_analyzed()
@@ -1,221 +0,0 @@
"""功能计费改造测试:爆款读配置、对口型/智能剪辑预扣逻辑。
策略:
- 爆款:通过修改缓存中的 FeatureConfig(multiplier/model_pricing)验证价格随配置变化
- lip_sync / smart_edit:直接测 LipsyncService 的预扣/结算/退款辅助方法,
PointsService 用 mock,避免依赖真实积分账户。
"""
from __future__ import annotations
from unittest.mock import MagicMock, patch
import pytest
from packages.domain import feature_pricing_service as fps
from packages.domain.feature_pricing_service import FeatureConfig, refresh_feature_configs
@pytest.fixture(autouse=True)
def _reset_cache():
refresh_feature_configs()
yield
refresh_feature_configs()
def _seed_cache(configs: dict) -> None:
import time
fps._cache = (time.monotonic(), configs)
class TestViralVideoReadsConfig:
def test_multiplier_change_changes_price(self):
"""配置里 multiplier 改大后,爆款价格随之变大(证明不再读死常量)。"""
from packages.domain.points_rules import calculate_viral_video_credits
# 基线兜底
base = calculate_viral_video_credits(15, 1280, 720)
assert base == 29.68
fallback = fps._fallback_configs()
vv = fallback["viral_video"]
vv.profit_multiplier = 2.0
_seed_cache(fallback)
changed = calculate_viral_video_credits(15, 1280, 720)
assert changed > base
# 精确校验:video_cost 相同,仅系数从 1.3 → 2.0
_, bd = __import__(
"packages.domain.points_rules", fromlist=["calculate_viral_video_credits_with_breakdown"]
).calculate_viral_video_credits_with_breakdown(15, 1280, 720)
assert bd["profit_multiplier"] == 2.0
def test_model_price_from_config(self):
"""model_pricing 改单价后,token 成本按新单价计算。"""
from packages.domain.points_rules import calculate_viral_video_credits_with_breakdown
fallback = fps._fallback_configs()
vv = fallback["viral_video"]
# seedance-2.5/720p/false 从 70 改成 100
vv.model_pricing["seedance-2.5"]["720p"]["false"] = 100.0
_seed_cache(fallback)
_, bd = calculate_viral_video_credits_with_breakdown(15, 1280, 720)
assert bd["model_price"] == 100.0
def test_price_cap_from_config(self):
from packages.domain.points_rules import calculate_viral_video_credits_with_breakdown
fallback = fps._fallback_configs()
vv = fallback["viral_video"]
vv.price_cap = 5.0
_seed_cache(fallback)
credits, bd = calculate_viral_video_credits_with_breakdown(15, 1280, 720)
assert credits == 5.0
assert bd["price_cap"] == 5.0
def test_disabled_feature_returns_zero_credits(self):
"""功能 is_enabled=false 时计费函数返回 0(纯计费层语义)。"""
from packages.domain.points_rules import calculate_viral_video_credits_with_breakdown
fallback = fps._fallback_configs()
fallback["viral_video"].is_enabled = False
_seed_cache(fallback)
credits, bd = calculate_viral_video_credits_with_breakdown(15, 1280, 720)
assert credits == 0.0
assert bd["feature_enabled"] is False
assert bd["charged"] is False
class TestLipSyncPricing:
def _make_service(self):
from app.services.lipsync_service import LipsyncService
svc = LipsyncService.__new__(LipsyncService)
svc.db = MagicMock()
return svc
def _lip_cfg(self, **kw):
base = dict(
feature_key="lip_sync",
name="对口型",
is_enabled=True,
fixed_cost=0.1,
profit_multiplier=1.0,
dynamic_unit_cost=0.05,
billing_mode="per_second",
price_cap=0.0,
model_pricing={},
description="",
)
base.update(kw)
return FeatureConfig(**base)
def test_estimate_duration_from_script(self):
svc = self._make_service()
# 10 个字 / 5 = 2 秒,下限 1
assert svc._estimate_duration(script_text="一二三四五六七八九十") == 2.0
# 无任何信息 → 默认 10 秒
assert svc._estimate_duration() == 10.0
def test_calculate_lipsync_price_per_second(self):
_seed_cache({"lip_sync": self._lip_cfg()})
price, bd = fps.calculate_price("lip_sync", dynamic_cost=20.0 * 0.05)
# dynamic 1.0 + fixed 0.1 = 1.1
assert price == 1.1
assert bd["charged"] is True
def test_settle_refunds_overcharge(self):
"""实际时长短 → 只退不补,退还差额。"""
svc = self._make_service()
_seed_cache({"lip_sync": self._lip_cfg()})
job = MagicMock()
job.credits_prepaid = 2.0
job.credits_cost = 0.0 # 未结算
job.user_id = "u1"
job.credits_transaction_id = "txn-old"
with patch("packages.domain.points_service.PointsService") as MockPS:
inst = MockPS.return_value
inst.refund_points.return_value = {"success": True}
svc._settle_lip_sync(job, actual_duration=10.0)
# final: (10*0.05 + 0.1)*1.0 = 0.6;退 2.0-0.6=1.4
assert round(job.credits_cost, 2) == 0.6
inst.refund_points.assert_called_once()
kwargs = inst.refund_points.call_args.kwargs
assert kwargs["amount"] == 1.4
def test_settle_no_refund_when_longer(self):
"""首期只退不补:实际更贵不补扣。"""
svc = self._make_service()
_seed_cache({"lip_sync": self._lip_cfg()})
job = MagicMock()
job.credits_prepaid = 0.5
job.credits_cost = 0.0
with patch("packages.domain.points_service.PointsService") as MockPS:
inst = MockPS.return_value
svc._settle_lip_sync(job, actual_duration=60.0)
assert round(job.credits_cost, 2) > 0.5
inst.refund_points.assert_not_called()
def test_refund_on_failure_full(self):
svc = self._make_service()
job = MagicMock()
job.credits_prepaid = 3.0
job.credits_cost = 0.0
job.user_id = "u1"
job.credits_transaction_id = "t1"
with patch("packages.domain.points_service.PointsService") as MockPS:
inst = MockPS.return_value
inst.refund_points.return_value = {"success": True}
svc._refund_lip_sync(job)
kwargs = inst.refund_points.call_args.kwargs
assert kwargs["amount"] == 3.0
class TestSmartEditFixedPrice:
def test_fixed_price_formula(self):
"""首期固定价:dynamic=0,price=fixed*multiplier,cap 封顶。"""
cfg = FeatureConfig(
feature_key="smart_edit",
name="智能剪辑",
is_enabled=True,
fixed_cost=2.0,
profit_multiplier=1.5,
billing_mode="model_based",
price_cap=0.0,
)
_seed_cache({"smart_edit": cfg})
price, bd = fps.calculate_price("smart_edit", dynamic_cost=0.0)
# (0+2)*1.5 = 3.0
assert price == 3.0
assert bd["dynamic_cost"] == 0.0
def test_fixed_price_with_cap(self):
cfg = FeatureConfig(
feature_key="smart_edit",
is_enabled=True,
fixed_cost=10.0,
profit_multiplier=2.0,
price_cap=8.0,
)
_seed_cache({"smart_edit": cfg})
price, _ = fps.calculate_price("smart_edit", dynamic_cost=0.0)
assert price == 8.0
def test_disabled_smart_edit_free(self):
cfg = FeatureConfig(feature_key="smart_edit", is_enabled=False, fixed_cost=2.0)
_seed_cache({"smart_edit": cfg})
price, bd = fps.calculate_price("smart_edit", dynamic_cost=0.0)
assert price == 0.0
assert bd["charged"] is False
-235
View File
@@ -1,235 +0,0 @@
"""feature_pricing_service 单元测试。
覆盖:
- 300s TTL 内存缓存(命中不重复 load / 过期重新 load / refresh 强制刷新)
- calculate_price 公式 (dynamic+fixed)*multiplier、price_cap 封顶、round
- disabled / 未知 key 返回 0
- DB 异常 / 空表 → 内置兜底配置(爆款启用且价格与现状一致)
- lookup_model_price 嵌套/扁平结构与旧版回落语义
"""
from __future__ import annotations
import time
import pytest
from packages.domain import feature_pricing_service as fps
from packages.domain.feature_pricing_service import (
CACHE_TTL_SECONDS,
FeatureConfig,
calculate_price,
get_feature_config,
is_feature_enabled,
lookup_model_price,
refresh_feature_configs,
)
@pytest.fixture(autouse=True)
def _reset_cache():
"""每个用例前后清空模块缓存,避免相互污染。"""
refresh_feature_configs()
yield
refresh_feature_configs()
def _cfg(key="x", **kw) -> FeatureConfig:
base = dict(
feature_key=key,
name=key,
is_enabled=True,
fixed_cost=0.2,
profit_multiplier=2.0,
dynamic_unit_cost=0.0,
billing_mode="per_second",
price_cap=0.0,
model_pricing={},
description="",
)
base.update(kw)
return FeatureConfig(**base)
class TestCacheTTL:
def test_cache_hit_avoids_reload(self, monkeypatch):
"""TTL 内第二次读取不再调 _load_all。"""
calls = {"n": 0}
def fake_load():
calls["n"] += 1
return {"x": _cfg()}
monkeypatch.setattr(fps, "_load_all", fake_load)
get_feature_config("x")
get_feature_config("x")
get_feature_config("x")
assert calls["n"] == 1
def test_expired_cache_reloads(self, monkeypatch):
"""超过 TTL 后重新 load。"""
calls = {"n": 0}
def fake_load():
calls["n"] += 1
return {"x": _cfg()}
monkeypatch.setattr(fps, "_load_all", fake_load)
get_feature_config("x")
assert calls["n"] == 1
# 把缓存时间戳回拨到 TTL 之前
ts, data = fps._cache
fps._cache = (ts - CACHE_TTL_SECONDS - 1, data)
get_feature_config("x")
assert calls["n"] == 2
def test_refresh_forces_reload(self, monkeypatch):
calls = {"n": 0}
def fake_load():
calls["n"] += 1
return {"x": _cfg()}
monkeypatch.setattr(fps, "_load_all", fake_load)
get_feature_config("x")
refresh_feature_configs()
get_feature_config("x")
assert calls["n"] == 2
def test_ttl_constant_is_300(self):
assert CACHE_TTL_SECONDS == 300.0
class TestCalculatePrice:
def test_basic_formula(self, monkeypatch):
# (dynamic 1.0 + fixed 0.2) * 2.0 = 2.4
monkeypatch.setattr(fps, "_load_all", lambda: {"x": _cfg(dynamic_unit_cost=1.0)})
price, bd = calculate_price("x", dynamic_cost=1.0)
assert price == 2.4
assert bd["dynamic_cost"] == 1.0
assert bd["fixed_cost"] == 0.2
assert bd["profit_multiplier"] == 2.0
assert bd["final_price"] == 2.4
assert bd["charged"] is True
def test_price_cap_clamps(self, monkeypatch):
# raw = (1+0.2)*2 = 2.4,cap=1.0 → 1.0
monkeypatch.setattr(fps, "_load_all", lambda: {"x": _cfg(price_cap=1.0)})
price, bd = calculate_price("x", dynamic_cost=1.0)
assert price == 1.0
assert bd["price_cap"] == 1.0
def test_no_cap_keeps_raw(self, monkeypatch):
# cap=0 视为不封顶
monkeypatch.setattr(fps, "_load_all", lambda: {"x": _cfg(price_cap=0.0)})
price, _ = calculate_price("x", dynamic_cost=1.0)
assert price == 2.4
def test_rounded_two_decimals(self, monkeypatch):
monkeypatch.setattr(
fps,
"_load_all",
lambda: {"x": _cfg(fixed_cost=0.1, profit_multiplier=1.0)},
)
price, _ = calculate_price("x", dynamic_cost=1.0 / 3.0)
# 0.3333... + 0.1 = 0.4333 → 0.43
assert price == 0.43
def test_negative_dynamic_treated_as_zero(self, monkeypatch):
monkeypatch.setattr(fps, "_load_all", lambda: {"x": _cfg()})
price, _ = calculate_price("x", dynamic_cost=-5.0)
# (0 + 0.2) * 2 = 0.4
assert price == 0.4
def test_disabled_returns_zero(self, monkeypatch):
monkeypatch.setattr(fps, "_load_all", lambda: {"x": _cfg(is_enabled=False)})
price, bd = calculate_price("x", dynamic_cost=1.0)
assert price == 0.0
assert bd["is_enabled"] is False
assert bd["charged"] is False
def test_unknown_key_returns_zero(self, monkeypatch):
monkeypatch.setattr(fps, "_load_all", lambda: {"x": _cfg()})
price, bd = calculate_price("nope", dynamic_cost=1.0)
assert price == 0.0
assert bd["charged"] is False
class TestDBFailureFallback:
def test_load_exception_uses_fallback(self, monkeypatch):
def boom():
raise RuntimeError("table does not exist")
monkeypatch.setattr(fps, "_load_all", boom)
cfg = get_feature_config("viral_video")
assert cfg is not None
assert cfg.is_enabled is True
assert cfg.fixed_cost == 0.15
assert cfg.profit_multiplier == 1.3
def test_empty_table_uses_fallback(self, monkeypatch):
monkeypatch.setattr(fps, "_load_all", lambda: {})
assert get_feature_config("viral_video").is_enabled is True
assert get_feature_config("lip_sync").is_enabled is False
assert get_feature_config("smart_edit").is_enabled is False
def test_fallback_viral_price_matches_current(self, monkeypatch):
"""兜底爆款价格与旧硬编码现状一致:seedance-2.5/720p/false=70。"""
monkeypatch.setattr(fps, "_load_all", lambda: {})
from packages.domain.points_rules import calculate_viral_video_credits
# 默认全局开关关闭,但纯计费函数价格照常算
assert calculate_viral_video_credits(15, 1280, 720) == 29.68
def test_db_row_overrides_fallback(self, monkeypatch):
monkeypatch.setattr(
fps,
"_load_all",
lambda: {"viral_video": _cfg("viral_video", fixed_cost=0.5, profit_multiplier=2.0, price_cap=50.0)},
)
cfg = get_feature_config("viral_video")
assert cfg.fixed_cost == 0.5
assert cfg.profit_multiplier == 2.0
assert cfg.price_cap == 50.0
class TestIsFeatureEnabled:
def test_disabled_feature(self, monkeypatch):
monkeypatch.setattr(fps, "_load_all", lambda: {"x": _cfg(is_enabled=False)})
assert is_feature_enabled("x") is False
def test_global_switch_off_blocks_enabled_feature(self, monkeypatch):
monkeypatch.setattr(fps, "_load_all", lambda: {"x": _cfg(is_enabled=True)})
monkeypatch.setattr(fps, "_global_points_enabled", lambda: False)
assert is_feature_enabled("x") is False
def test_both_switches_on(self, monkeypatch):
monkeypatch.setattr(fps, "_load_all", lambda: {"x": _cfg(is_enabled=True)})
monkeypatch.setattr(fps, "_global_points_enabled", lambda: True)
assert is_feature_enabled("x") is True
class TestLookupModelPrice:
NESTED = {
"seedance-2.5": {
"720p": {"false": 70.0, "true": 42.0},
},
"wan-3.0": {"480p": {"false": 0.3}},
}
def test_nested_exact_hit(self):
assert lookup_model_price(self.NESTED, "seedance-2.5", "720p", False) == 70.0
assert lookup_model_price(self.NESTED, "seedance-2.5", "720p", True) == 42.0
def test_missing_bool_key_returns_none(self):
# wan-3.0/480p 只有 false,请求 true → None(由调用方回落)
assert lookup_model_price(self.NESTED, "wan-3.0", "480p", True) is None
def test_unknown_model_returns_none(self):
assert lookup_model_price(self.NESTED, "nope", "720p", False) is None
def test_flat_structure(self):
flat = {"m|720p|false": 12.5}
assert lookup_model_price(flat, "m", "720p", False) == 12.5
assert lookup_model_price(flat, "m", "720p", True) is None
+2 -8
View File
@@ -3,7 +3,7 @@
from __future__ import annotations
import os
from unittest.mock import MagicMock, patch
from unittest.mock import patch
import pytest
@@ -108,14 +108,8 @@ class TestGetTtsService:
def test_empty_env_falls_back_to_auto_detect(self):
"""环境变量为空时自动检测."""
with (
patch.dict(os.environ, {"TTS_PROVIDER": ""}),
patch("packages.shared.config.get_shared_settings") as mock_settings,
):
with patch.dict(os.environ, {"TTS_PROVIDER": ""}):
# 没有 cosyvoice_api_key 时应该用 mock
settings = MagicMock()
settings.cosyvoice_api_key = ""
mock_settings.return_value = settings
service = get_tts_service(None)
assert service.provider_name == "mock"
+2 -28
View File
@@ -431,7 +431,7 @@ class TestGenerateCopy:
assert resp.id == "job-gc"
def test_generate_copy_rejects_wrong_status(self):
"""wait_user_confirm 等中间状态不允许调用 generate-copy(状态保护)。"""
"""任务在 copy_generated/completed 时不能再 generate-copy(状态保护)。"""
import pytest
from app.api.routes import viral_video as vv_mod
from app.schemas.viral_video import GenerateCopyRequest
@@ -441,8 +441,7 @@ class TestGenerateCopy:
user = _auth_user("u1")
session = MagicMock()
# wait_user_confirm 属于前端在编辑/确认文案的中间状态,应拒绝重新触发生成
job = _make_job(job_id="job-gc2", user_id="u1", status=ViralVideoStatus.WAIT_USER_CONFIRM)
job = _make_job(job_id="job-gc2", user_id="u1", status=ViralVideoStatus.COPY_GENERATED)
repo = MagicMock()
repo.get.return_value = job
@@ -451,31 +450,6 @@ class TestGenerateCopy:
vv_mod.generate_copy("job-gc2", GenerateCopyRequest(), authenticated_user=user, session=session)
assert exc.value.status_code == 409
def test_generate_copy_allows_regenerate_from_copy_generated(self):
"""#2222: COPY_GENERATED/COMPLETED 状态下点「重新生成文案」应放行入队,不返回 409。"""
from unittest.mock import patch
from app.api.routes import viral_video as vv_mod
from app.schemas.viral_video import GenerateCopyRequest
from packages.domain.viral_video import ViralVideoStatus
user = _auth_user("u1")
session = MagicMock()
for regen_status in (ViralVideoStatus.COPY_GENERATED, ViralVideoStatus.COMPLETED):
job = _make_job(job_id=f"job-regen-{regen_status}", user_id="u1", status=regen_status)
repo = MagicMock()
repo.get.return_value = job
with (
patch.object(vv_mod, "_get_job_repo", return_value=repo),
patch.object(vv_mod.celery_app, "send_task") as mock_send,
):
resp = vv_mod.generate_copy(f"job-regen-{regen_status}", GenerateCopyRequest(), authenticated_user=user, session=session)
mock_send.assert_called_once()
job.resume_from_image_analyzed.assert_called()
assert job.retry_count >= 1
assert resp.id == f"job-regen-{regen_status}"
def test_generate_copy_persists_voice_and_ratio(self):
"""generate-copy 应把 voice_id/voice_source/video_ratio 写入 job。"""
from unittest.mock import patch
+18 -51
View File
@@ -97,43 +97,21 @@ def invalidate_loader_cache():
class TestImageAnalysisWiring:
def test_step_image_analysis_uses_v2_batch_path(self, job):
"""#2200/#2207 后图片分析走 V2 批处理(OCR+lite JSON 并行),
_step_image_analysis 归一化 URL 后调用 analyze_images_v2。"""
def test_uses_loader_template_and_xml_parse(self, job):
from apps.worker.worker_app.tasks import viral_video as vv
fake_product = {
"name": "lipstick",
"brand": "品牌X",
"category": "唇部彩妆",
"key_features": ["显白", "持久"],
"text_on_package": ["品牌X", "211"],
"_source": "v2",
}
with patch.object(vv, "_normalize_image_url", side_effect=lambda raw, idx: raw):
with patch(
"worker_app.tasks.vision.analyze_images_v2",
return_value=[fake_product, fake_product],
create=True,
) as mock_v2:
result = vv._step_image_analysis(job)
with patch("packages.shared.ai_service.call_vision", return_value=IMAGE_XML) as mock_v:
result = vv._analyze_single_image(0, "https://img/1.jpg", "vlm-lite", 15)
mock_v2.assert_called_once()
# 传入的是归一化后的图片 URL 列表
assert mock_v2.call_args.args[0] == job.images
products = result["products"]
assert len(products) == 2
assert products[0]["name"] == "lipstick"
assert products[0]["brand"] == "品牌X"
assert "显白" in products[0]["key_features"]
assert products[0]["text_on_package"] == ["品牌X", "211"]
def test_step_image_analysis_empty_images(self, job):
from apps.worker.worker_app.tasks import viral_video as vv
job.images = []
result = vv._step_image_analysis(job)
assert result == {"products": []}
mock_v.assert_called_once()
# 验证调用时传入了 system_prompt(说明走了 loader 渲染的模板)
call_kwargs = mock_v.call_args.kwargs
assert "system_prompt" in call_kwargs and call_kwargs["system_prompt"]
# 结果包含从 XML 解析出的产品信息
assert result["name"] == "lipstick"
assert result["brand"] == "品牌X"
assert "显白" in result["key_features"]
assert result["text_on_package"] == ["品牌X", "211"]
# ── 2) 意图解析走模板 ───────────────────────────────────────────────
@@ -274,29 +252,18 @@ class TestEndToEndLoaderUsed:
called_types.append(prompt_type)
return real_get(prompt_type, **kwargs)
v2_product = {
"name": "lipstick",
"brand": "品牌X",
"key_features": ["显白", "持久"],
}
with (
patch.object(pl, "get_template", side_effect=spy_get),
patch.object(vv, "_normalize_image_url", side_effect=lambda raw, idx: raw),
patch(
"worker_app.tasks.vision.analyze_images_v2",
return_value=[v2_product],
create=True,
),
patch("packages.shared.ai_service.call_vision", return_value=IMAGE_XML),
patch("packages.shared.ai_service.call_llm", return_value=INTENT_XML),
):
# 1) image(V2 路径,不再经过 prompt_loader)
img_step = vv._step_image_analysis(job)
img_res = img_step["products"][0]
# 2) intent(走 loader image_analysis? 否——intent_parsing 模板)
# 1) image
img_res = vv._analyze_single_image(0, "https://img/1.jpg", "vlm", 15)
# 2) intent
intent_res = vv._step_intent_parsing(job, {"products": [img_res]})
# V2 图片分析不再调用 loader;意图解析调用 intent_parsing 模板
assert "image_analysis" not in called_types
# 前两步分别调用了 image_analysis 和 intent_parsing
assert "image_analysis" in called_types
assert "intent_parsing" in called_types
# script 和 review 单独验证(需要不同的 LLM 返回)
-291
View File
@@ -1,291 +0,0 @@
# -*- coding: utf-8 -*-
"""vision v4 prompt / assembler 单元测试:
- assembler 正确识别 v4 嵌套 schema 与旧扁平 schema
- v4 product/person/store/other 四类输出组装出下游必出字段
- 旧扁平 schema 行为不变
- _prompt._resolve:DB 有 active prompt 时原样使用(不追加硬编码 schema);
DB 无记录时回落到硬编码 JSON schema
"""
from __future__ import annotations
import types
from typing import Any
import pytest
from worker_app.tasks.vision import _prompt, assembler
REQUIRED_KEYS = {
"name",
"brand",
"category",
"appearance",
"packaging",
"text_on_package",
"key_features",
"scene",
"mood",
"portrait_prompt",
"summary",
"_source",
}
# ---------- schema 识别 ----------
def test_is_v4_schema_products_list() -> None:
assert assembler._is_v4_schema({"type": "product", "products": []})
def test_is_v4_schema_type_only() -> None:
assert assembler._is_v4_schema({"type": "person"})
def test_is_v4_schema_people_dict() -> None:
assert assembler._is_v4_schema({"people": {"has_person": True}})
def test_is_not_v4_schema_flat() -> None:
assert not assembler._is_v4_schema({"has_person": True, "upper_wear": "T恤"})
# ---------- v4 product ----------
V4_PRODUCT: dict[str, Any] = {
"type": "product",
"scene": "白色背景产品图",
"mood": "清新专业",
"style": "商业产品摄影",
"colors": [{"hex": "#E60012", "name": "亮红色", "coverage": 0.6}],
"visible_text": [{"text": "OMO奥妙除菌除螨", "position": "瓶身正面"}],
"products": [
{
"product_name": "OMO奥妙除菌除螨洗衣液",
"brand": "OMO奥妙",
"category": "洗护",
"package_type": "瓶装",
"package_color": "亮红色瓶身",
"cap_type": "透明翻盖式按压瓶口",
"body_shape": "带侧面握持把手的竖款瓶身",
"label_design": "瓶身印十字盾牌图案",
"product_features": ["亮红色瓶装", "按压式瓶口", "十字盾牌标签"],
"key_selling_points": ["天然除菌除螨"],
"position": "main",
}
],
"has_person": False,
}
def test_assemble_v4_product_fields() -> None:
r = assembler.assemble_result(0, V4_PRODUCT, ["OMO奥妙"])
assert REQUIRED_KEYS <= set(r.keys())
assert r["name"] == "OMO奥妙除菌除螨洗衣液"
assert r["brand"] == "OMO奥妙"
assert r["category"] == "洗护"
assert "瓶装" in r["packaging"]
assert isinstance(r["key_features"], list) and r["key_features"]
assert any("除菌" in str(t) for t in r["text_on_package"])
assert len(r["portrait_prompt"]) >= 10
assert r["_source"] == "v2_fast_json_v4"
def test_assemble_v4_product_multi_selects_main() -> None:
fj = {
"type": "product",
"products": [
{"product_name": "次要商品", "brand": "B"},
{"product_name": "主商品", "brand": "A", "position": "main"},
],
}
r = assembler.assemble_result(1, fj, [])
assert r["name"] == "主商品"
# ---------- v4 person ----------
V4_PERSON: dict[str, Any] = {
"type": "person",
"scene": "户外街拍",
"mood": "自信",
"style": "街拍",
"colors": [],
"visible_text": [],
"has_person": True,
"gender": "女",
"age_range": "青年",
"upper_wear": "白色V领短袖T恤",
"upper_color": "白色",
"lower_wear": "黑色高腰阔腿裤",
"lower_color": "黑色",
"dress_color": None,
"accessories": ["银色项链"],
"hairstyle": "黑色长直发",
"expression": "自信",
"pose": "侧身站立",
"outfit_style": "休闲日常",
"portrait_prompt": (
"一位年轻女性,身穿白色V领短袖T恤、黑色高腰阔腿裤,佩戴银色项链,"
"黑色长直发,神情自信,侧身站立,休闲日常风格,城市街拍场景"
),
"products": [],
}
def test_assemble_v4_person() -> None:
r = assembler.assemble_result(0, V4_PERSON, [])
assert REQUIRED_KEYS <= set(r.keys())
assert r["category"] == "人物穿搭"
assert r["_source"] == "v2_fast_json_v5"
assert "T恤" in r["name"]
assert "年轻女性" in r["portrait_prompt"]
assert "项链" in r["portrait_prompt"]
assert isinstance(r["key_features"], list) and len(r["key_features"]) <= 8
def test_assemble_v4_person_people_nested() -> None:
fj = {"type": "person", "people": {**V4_PERSON, "has_person": True}}
r = assembler.assemble_result(0, fj, [])
assert r["category"] == "人物穿搭"
assert "年轻女性" in r["portrait_prompt"]
# ---------- v4 store ----------
def test_assemble_v4_store() -> None:
fj = {
"type": "store",
"scene": "便利店内部",
"mood": "日常便民",
"style": "门店实拍",
"store_type": "社区便利店",
"store_layout": "纵深货架布局",
"brand_signage": "全家FamilyMart",
"visual_elements": ["红白主色调", "促销海报"],
"product_categories_visible": ["饮料", "零食"],
"promotion_elements": ["第二件半价海报"],
"atmosphere": "亲民生活化",
"has_person": False,
}
r = assembler.assemble_result(0, fj, [])
assert REQUIRED_KEYS <= set(r.keys())
assert r["name"] == "社区便利店"
assert r["brand"] == "全家FamilyMart"
assert r["category"] == "门店场景"
assert any("饮料" in str(f) for f in r["key_features"])
assert "门店实拍" in r["portrait_prompt"]
# ---------- v4 other ----------
def test_assemble_v4_other() -> None:
fj = {"type": "other", "description": "海边日落风景", "scene": "海边", "mood": "宁静"}
r = assembler.assemble_result(0, fj, [])
assert REQUIRED_KEYS <= set(r.keys())
assert r["name"] == "海边日落风景"
assert r["category"] == "非产品图"
# ---------- 旧扁平 schema 兼容 ----------
def test_assemble_old_flat_person() -> None:
fj = {
"has_person": True,
"gender": "男",
"age_range": "中年",
"upper_wear": "西装",
"upper_color": "深灰色",
"lower_wear": "西裤",
"lower_color": "黑色",
"accessories": ["手表"],
"hairstyle": "短发",
"expression": "严肃",
"scene": "办公室",
"style": "商务",
"mood": "专业",
}
r = assembler.assemble_result(0, fj, [])
assert REQUIRED_KEYS <= set(r.keys())
assert "中年男性" in r["portrait_prompt"]
assert r["_source"] == "v2_fast_json"
def test_assemble_old_flat_product() -> None:
fj = {
"has_person": False,
"product_name": "口红",
"brand": "Dior",
"category": "美妆",
"colors": ["红色"],
"scene": "通用",
"style": "商业",
"mood": "高级",
}
r = assembler.assemble_result(0, fj, ["Dior"])
assert r["name"] == "口红"
assert r["brand"] == "Dior"
assert r["text_on_package"] == ["Dior"]
def test_assemble_none_input() -> None:
r = assembler.assemble_result(0, None, [])
assert REQUIRED_KEYS <= set(r.keys())
# ---------- _prompt 解析 ----------
@pytest.fixture(autouse=True)
def _clear_prompt_cache() -> Any:
_prompt.invalidate_cache()
yield
_prompt.invalidate_cache()
def _fake_tpl(system_prompt: str = "v4 system prompt 只返回JSON") -> Any:
return types.SimpleNamespace(
system_prompt=system_prompt,
user_prompt_template="分析 {image_count} 张图",
version=4,
)
def test_resolve_uses_db_prompt_without_append(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setattr(_prompt, "_load_db_template", lambda: _fake_tpl("DB_V4_PROMPT_XYZ"))
sys_prompt, user_prompt = _prompt.resolve_fast_prompt()
assert sys_prompt == "DB_V4_PROMPT_XYZ"
assert "DB_V4_PROMPT_XYZ" not in _prompt._FAST_JSON_APPEND # sanity: 旧append是另一段文本
assert "分析 1 张图" in user_prompt
def test_resolve_pro_uses_db_prompt_without_append(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setattr(_prompt, "_load_db_template", lambda: _fake_tpl("DB_V4_PRO_PROMPT"))
sys_prompt, _ = _prompt.resolve_pro_prompt()
assert sys_prompt == "DB_V4_PRO_PROMPT"
assert "【输出格式要求】" not in sys_prompt
def test_resolve_falls_back_when_no_db(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setattr(_prompt, "_load_db_template", lambda: None)
sys_prompt, user_prompt = _prompt.resolve_fast_prompt()
assert sys_prompt == _prompt._FAST_JSON_SCHEMA
assert user_prompt == _prompt.DEFAULT_FAST_USER
def test_resolve_caches(monkeypatch: pytest.MonkeyPatch) -> None:
calls = {"n": 0}
def _load() -> Any:
calls["n"] += 1
return _fake_tpl("CACHED_PROMPT")
monkeypatch.setattr(_prompt, "_load_db_template", _load)
s1, _ = _prompt.resolve_fast_prompt()
s2, _ = _prompt.resolve_fast_prompt()
assert s1 == s2 == "CACHED_PROMPT"
assert calls["n"] == 1