Compare commits
8 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 9cbae5cf60 | |||
| c025b8b6ac | |||
| c8e0756d05 | |||
| 0c9839bcb7 | |||
| 6737ff6ef7 | |||
| 1d1245551f | |||
| dbc684db2f | |||
| 2d811fccf1 |
@@ -0,0 +1,124 @@
|
||||
# -*- 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()
|
||||
|
||||
# 检查表是否存在
|
||||
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
|
||||
|
||||
# 查询目标模型的 ID(使用 model_key 查询,不硬编码 UUID)
|
||||
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]
|
||||
|
||||
# 查询旧模型 ID
|
||||
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},
|
||||
)
|
||||
|
||||
# 更新 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:
|
||||
conn = op.get_bind()
|
||||
table_check = conn.execute(sa.text("SELECT to_regclass(\public.ai_models\)")).scalar()
|
||||
if not table_check:
|
||||
return
|
||||
|
||||
old_model_row = conn.execute(
|
||||
sa.text("SELECT id FROM ai_models WHERE model_key = \doubao-seed-1-6-250615\ LIMIT 1")
|
||||
).first()
|
||||
if not old_model_row:
|
||||
return
|
||||
old_model_id = old_model_row[0]
|
||||
|
||||
for cap_key in ["intent_parsing", "copy_fusion", "storyboard", "copy_review", "asset_classify"]:
|
||||
conn.execute(
|
||||
sa.text(
|
||||
"UPDATE ai_capability_configs SET primary_model_id = :old_id, updated_at = NOW() WHERE capability_key = :cap_key"
|
||||
),
|
||||
{"old_id": old_model_id, "cap_key": cap_key},
|
||||
)
|
||||
@@ -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
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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:
|
||||
@@ -499,22 +477,30 @@ 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
|
||||
_s = get_shared_settings()
|
||||
try:
|
||||
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
|
||||
_cap = ai_router.get_capability("intent_parsing")
|
||||
_fast = (_cap.primary_model.model_key if _cap and _cap.primary_model else None) or _s.doubao_fast_model
|
||||
_pro = (
|
||||
(_cap.lite_model.model_key if _cap and _cap.lite_model else None)
|
||||
or (_cap.primary_model.model_key if _cap and _cap.primary_model else None)
|
||||
or _s.doubao_model
|
||||
)
|
||||
except Exception:
|
||||
_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 = _llm_client.chat_completion(
|
||||
[{"role": "system", "content": system}, {"role": "user", "content": user}],
|
||||
temperature=0.4,
|
||||
max_tokens=1024,
|
||||
model=_m,
|
||||
timeout=60,
|
||||
)
|
||||
) # #2180/#2215: 直接用 client.chat_completion 传 messages list,不再走 call_llm 字符串包装
|
||||
if not raw:
|
||||
continue
|
||||
parsed = _parse(raw)
|
||||
@@ -895,14 +881,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 = _llm_client2.chat_completion(
|
||||
[{"role": "system", "content": system_tpl}, {"role": "user", "content": user}],
|
||||
temperature=temp,
|
||||
max_tokens=max_tok,
|
||||
model=model,
|
||||
timeout=tmo,
|
||||
)
|
||||
if not raw:
|
||||
@@ -933,22 +918,34 @@ 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")
|
||||
_s = get_shared_settings()
|
||||
try:
|
||||
from packages.shared.ai_router import ai_router
|
||||
|
||||
_cap = ai_router.get_capability("storyboard")
|
||||
_fast = (_cap.primary_model.model_key if _cap and _cap.primary_model else None) or _s.doubao_fast_model
|
||||
_pro = (
|
||||
(_cap.lite_model.model_key if _cap and _cap.lite_model else None)
|
||||
or (_cap.primary_model.model_key if _cap and _cap.primary_model else None)
|
||||
or getattr(_s, "doubao_model", None)
|
||||
or _fast
|
||||
)
|
||||
except Exception:
|
||||
_fast = _s.doubao_fast_model
|
||||
_pro = getattr(_s, "doubao_model", None) or _fast
|
||||
_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"))
|
||||
try:
|
||||
# #2217: doubao-seed-2-1-pro生成长编导脚本高峰期>90s,上调到150s,支持ENV覆盖
|
||||
normalized = _try_gen(_client_fast, 0.8, 2500, "fast-first", tmo=_script_fast_tmo)
|
||||
normalized = _try_gen(_fast, 0.8, 2500, "fast-first", tmo=_script_fast_tmo)
|
||||
if normalized is not None:
|
||||
return normalized
|
||||
normalized = _try_gen(_client_fast, 0.6, 3200, "fast-retry", tmo=_script_fast_tmo)
|
||||
normalized = _try_gen(_fast, 0.6, 3200, "fast-retry", tmo=_script_fast_tmo)
|
||||
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)
|
||||
# 第三次:用主力模型兜底
|
||||
if _pro and _pro != _fast:
|
||||
normalized = _try_gen(_pro, 0.7, 3500, "pro-fallback", tmo=_script_pro_tmo)
|
||||
if normalized is not None:
|
||||
return normalized
|
||||
logger.warning("[爆款视频] 编导脚本三次都未生成合格结果,使用兜底脚本")
|
||||
@@ -1271,12 +1268,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 +1283,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 +1431,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 +1609,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 +1702,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 +1760,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 +1919,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 +1941,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, "合规审核完成")
|
||||
|
||||
@@ -498,7 +498,6 @@ def _assemble_v4(idx: int, fj: dict, ocr_texts: list[str]) -> dict[str, Any]:
|
||||
"portrait_prompt": portrait_prompt[:300],
|
||||
"summary": summary[:50],
|
||||
"_source": "v2_fast_json_v5",
|
||||
"has_person": True,
|
||||
}
|
||||
|
||||
# 商品类
|
||||
@@ -559,7 +558,6 @@ def _assemble_v4(idx: int, fj: dict, ocr_texts: list[str]) -> dict[str, Any]:
|
||||
"mood": mood,
|
||||
"portrait_prompt": portrait_prompt[:200],
|
||||
"summary": str(summary)[:60],
|
||||
"has_person": False,
|
||||
"_source": "v2_fast_json_v4",
|
||||
}
|
||||
|
||||
@@ -608,7 +606,6 @@ def _assemble_v4(idx: int, fj: dict, ocr_texts: list[str]) -> dict[str, Any]:
|
||||
"mood": atmosphere or mood,
|
||||
"portrait_prompt": portrait_prompt[:200],
|
||||
"summary": summary[:40],
|
||||
"has_person": False,
|
||||
"_source": "v2_fast_json_v4",
|
||||
}
|
||||
|
||||
@@ -626,7 +623,6 @@ def _assemble_v4(idx: int, fj: dict, ocr_texts: list[str]) -> dict[str, Any]:
|
||||
"mood": mood,
|
||||
"portrait_prompt": f"{scene},{mood}氛围,{desc}"[:200],
|
||||
"summary": desc[:40],
|
||||
"has_person": False,
|
||||
"_source": "v2_fast_json_v4_other",
|
||||
}
|
||||
|
||||
@@ -664,5 +660,4 @@ 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)),
|
||||
}
|
||||
|
||||
@@ -3,8 +3,8 @@
|
||||
|
||||
fast_json 超时/返回非 JSON/识别为空时,本路径单次调用兜底。
|
||||
设计要点:
|
||||
- 通过 ai_router.get_vision_client() 获取 DoubaoClient 实例,不再自己拼 httpx 请求
|
||||
- enable_thinking=False + response_format=json_object
|
||||
- 通过 ai_router 动态获取 model/api_key/base_url,不再硬编码
|
||||
- enable_thinking=false + response_format=json_object
|
||||
- system prompt 优先读后台 viral_video_prompt_templates 配置,DB不可用时fallback到硬编码JSON schema
|
||||
- timeout=25s
|
||||
- 返回 dict 统一走 assembler.assemble_result 组装,与 fast 路径输出格式完全一致
|
||||
@@ -25,6 +25,25 @@ _DEFAULT_TIMEOUT = 30
|
||||
_DEFAULT_MAX_TOKENS = 800
|
||||
|
||||
|
||||
def _get_vision_config(variant: str = "primary") -> tuple[str, str, str]:
|
||||
"""从 ai_router 获取 image_analysis 配置,返回 (api_key, base_url, model)。"""
|
||||
try:
|
||||
from packages.shared.ai_router import ai_router
|
||||
|
||||
# 先尝试 lite,再 fallback
|
||||
client = ai_router.get_vision_client("image_analysis", variant=variant)
|
||||
if client and client.is_available:
|
||||
return client.api_key, client.base_url, client.model
|
||||
except Exception as e:
|
||||
logger.warning("[vision.v2] ai_router 获取失败 (%s),fallback 环境变量: %s", variant, e)
|
||||
|
||||
# Fallback: 环境变量
|
||||
import os
|
||||
|
||||
api_key = os.environ.get("DASHSCOPE_API_KEY", "")
|
||||
return api_key, "https://dashscope.aliyuncs.com/compatible-mode/v1", "qwen3.7-plus"
|
||||
|
||||
|
||||
def call_pro_vlm(
|
||||
img_url: str,
|
||||
idx: int,
|
||||
@@ -32,52 +51,70 @@ def call_pro_vlm(
|
||||
timeout: int = _DEFAULT_TIMEOUT,
|
||||
) -> dict[str, Any] | 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] pro vision client 不可用,跳过")
|
||||
return None
|
||||
except Exception as e:
|
||||
logger.warning("[vision.v2] ai_router 获取失败: %s", e)
|
||||
api_key, base_url, model = _get_vision_config("primary")
|
||||
if not api_key:
|
||||
logger.warning("[vision.v2] pro DASHSCOPE_API_KEY 未配置,跳过")
|
||||
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},
|
||||
],
|
||||
},
|
||||
]
|
||||
|
||||
payload: dict[str, Any] = {
|
||||
"model": model,
|
||||
"messages": [
|
||||
{"role": "system", "content": system_prompt},
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "image_url", "image_url": {"url": img_url}},
|
||||
{"type": "text", "text": user_prompt},
|
||||
],
|
||||
},
|
||||
],
|
||||
"temperature": 0.3,
|
||||
"max_tokens": _DEFAULT_MAX_TOKENS,
|
||||
"stream": False,
|
||||
"enable_thinking": False,
|
||||
"response_format": {"type": "json_object"},
|
||||
}
|
||||
try:
|
||||
raw = client.vision_completion(
|
||||
messages=messages,
|
||||
images=None, # 图片已在 messages 中
|
||||
temperature=0.3,
|
||||
max_tokens=_DEFAULT_MAX_TOKENS,
|
||||
r = httpx.post(
|
||||
f"{base_url.rstrip('/')}/chat/completions",
|
||||
headers={"Authorization": f"Bearer {api_key}", "Content-Type": "application/json"},
|
||||
json=payload,
|
||||
timeout=timeout,
|
||||
enable_thinking=False,
|
||||
response_format={"type": "json_object"},
|
||||
)
|
||||
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 {}
|
||||
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] pro 完成 model=%s elapsed=%.1fs",
|
||||
client.model,
|
||||
"[vision.v2] pro 完成 model=%s elapsed=%.1fs in=%d out=%d reasoning=%d",
|
||||
model,
|
||||
elapsed,
|
||||
usage.get("prompt_tokens", 0),
|
||||
usage.get("completion_tokens", 0),
|
||||
reasoning_tokens,
|
||||
)
|
||||
s = _strip_code_fence(raw)
|
||||
s = raw.strip()
|
||||
if s.startswith("```"):
|
||||
lines = s.split("\n")
|
||||
if lines and lines[0].startswith("```"):
|
||||
lines = lines[1:]
|
||||
if lines and lines[-1].strip().startswith("```"):
|
||||
lines = lines[:-1]
|
||||
s = "\n".join(lines).strip()
|
||||
lpos, rr = s.find("{"), s.rfind("}")
|
||||
if lpos >= 0 and rr > lpos:
|
||||
s = s[lpos : rr + 1]
|
||||
@@ -98,15 +135,3 @@ def call_pro_vlm(
|
||||
elapsed = time.time() - t0
|
||||
logger.warning("[vision.v2] pro 异常 elapsed=%.1fs err=%s", elapsed, e, exc_info=True)
|
||||
return None
|
||||
|
||||
|
||||
def _strip_code_fence(s: str) -> str:
|
||||
s = s.strip()
|
||||
if s.startswith("```"):
|
||||
lines = s.split("\n")
|
||||
if lines and lines[0].startswith("```"):
|
||||
lines = lines[1:]
|
||||
if lines and lines[-1].strip().startswith("```"):
|
||||
lines = lines[:-1]
|
||||
s = "\n".join(lines).strip()
|
||||
return s
|
||||
|
||||
@@ -3,8 +3,8 @@
|
||||
|
||||
目标:替代"人体属性/商品检测/图像标签"三个火山不存在的专用云端 API。
|
||||
设计要点:
|
||||
- 通过 ai_router.get_vision_client() 获取 DoubaoClient 实例,不再自己拼 httpx 请求
|
||||
- enable_thinking=False 关闭推理链(reasoning 是延迟主因)
|
||||
- 通过 ai_router 动态获取 model/api_key/base_url,不再硬编码
|
||||
- enable_thinking=false 关闭推理链(reasoning 是延迟主因)
|
||||
- response_format=json_object 强约束JSON输出
|
||||
- system prompt 优先读后台 viral_video_prompt_templates 配置,DB不可用时fallback到硬编码JSON schema
|
||||
- max_tokens=350、temperature=0.1(稳定输出 JSON)
|
||||
@@ -26,6 +26,24 @@ _DEFAULT_TIMEOUT = 15
|
||||
_DEFAULT_MAX_TOKENS = 350
|
||||
|
||||
|
||||
def _get_vision_config() -> tuple[str, str, str]:
|
||||
"""从 ai_router 获取 image_analysis 配置,返回 (api_key, base_url, model)。"""
|
||||
try:
|
||||
from packages.shared.ai_router import ai_router
|
||||
|
||||
client = ai_router.get_vision_client("image_analysis", variant="primary")
|
||||
if client and client.is_available:
|
||||
return client.api_key, client.base_url, client.model
|
||||
except Exception as e:
|
||||
logger.warning("[vision.v2] ai_router 获取失败,fallback 环境变量: %s", e)
|
||||
|
||||
# Fallback: 环境变量
|
||||
import os
|
||||
|
||||
api_key = os.environ.get("DASHSCOPE_API_KEY", "")
|
||||
return api_key, "https://dashscope.aliyuncs.com/compatible-mode/v1", "qwen3.8-flash"
|
||||
|
||||
|
||||
def _strip_code_fence(s: str) -> str:
|
||||
s = s.strip()
|
||||
if s.startswith("```"):
|
||||
@@ -44,52 +62,76 @@ def call_fast_json(
|
||||
timeout: int = _DEFAULT_TIMEOUT,
|
||||
max_tokens: int = _DEFAULT_MAX_TOKENS,
|
||||
) -> dict[str, Any] | None:
|
||||
"""调用 vision client 返回结构化 dict;失败/非 JSON 返回 None。"""
|
||||
"""调用 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, base_url, model = _get_vision_config()
|
||||
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.rstrip('/')}/chat/completions"
|
||||
payload: dict[str, Any] = {
|
||||
"model": model,
|
||||
"messages": [
|
||||
{"role": "system", "content": system_prompt},
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{"type": "image_url", "image_url": {"url": img_url}},
|
||||
{"type": "text", "text": user_prompt},
|
||||
],
|
||||
},
|
||||
],
|
||||
"temperature": 0.1,
|
||||
"max_tokens": max_tokens,
|
||||
"stream": False,
|
||||
"enable_thinking": False,
|
||||
"response_format": {"type": "json_object"},
|
||||
}
|
||||
try:
|
||||
raw = client.vision_completion(
|
||||
messages=messages,
|
||||
images=None, # 图片已在 messages 中
|
||||
temperature=0.1,
|
||||
max_tokens=max_tokens,
|
||||
resp = httpx.post(
|
||||
url,
|
||||
headers={"Authorization": f"Bearer {api_key}", "Content-Type": "application/json"},
|
||||
json=payload,
|
||||
timeout=timeout,
|
||||
enable_thinking=False,
|
||||
response_format={"type": "json_object"},
|
||||
)
|
||||
elapsed = time.time() - t0
|
||||
if resp.status_code == 400 and "enable_thinking" in resp.text[:300].lower():
|
||||
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",
|
||||
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("}")
|
||||
|
||||
@@ -369,7 +369,7 @@ class CosyVoiceService:
|
||||
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", "voice-enrollment")
|
||||
self._clone_model = clone_model or getattr(settings, "cosyvoice_clone_model", "")
|
||||
self._audio_url_signer = audio_url_signer
|
||||
|
||||
# base_url 规范化:去掉末尾的路径残留(兼容旧版配置)
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -170,28 +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.api_key: str = api_key or settings.doubao_api_key
|
||||
self.model: str = model or settings.doubao_model
|
||||
self.base_url: str = (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.provider: str = provider
|
||||
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
|
||||
@@ -258,7 +243,6 @@ class DoubaoClient:
|
||||
max_tokens: int = 1024,
|
||||
model: str | None = None,
|
||||
timeout: int | None = None,
|
||||
**kwargs,
|
||||
) -> Optional[str]:
|
||||
"""调用 Chat Completion 接口.
|
||||
|
||||
@@ -284,11 +268,6 @@ class DoubaoClient:
|
||||
"temperature": temperature,
|
||||
"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()
|
||||
@@ -340,7 +319,6 @@ class DoubaoClient:
|
||||
temperature: float = 0.3,
|
||||
timeout: int | None = None,
|
||||
model: str | None = None,
|
||||
**kwargs,
|
||||
) -> Optional[str]:
|
||||
"""调用豆包视觉理解 API(OpenAI 兼容多模态格式).
|
||||
|
||||
@@ -396,10 +374,6 @@ class DoubaoClient:
|
||||
"temperature": temperature,
|
||||
"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
|
||||
@@ -616,32 +590,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
|
||||
# ─────────────────────────────────────────────────────────────────
|
||||
|
||||
# 判断任务模式:
|
||||
|
||||
@@ -56,11 +56,80 @@ class CapabilityConfig:
|
||||
is_enabled: bool
|
||||
|
||||
|
||||
# ── 简单包装类(TTS / ImageGen / VideoGen)──────────────────────────────────
|
||||
# ── 客户端包装 ──────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class LLMClient:
|
||||
"""统一 LLM 客户端接口"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
provider: str,
|
||||
api_key: str,
|
||||
base_url: str,
|
||||
model: str,
|
||||
timeout: int = 45,
|
||||
max_retries: int = 1,
|
||||
max_tokens: int | None = None,
|
||||
temperature: float | None = None,
|
||||
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.max_retries = max_retries
|
||||
self.max_tokens = max_tokens
|
||||
self.temperature = temperature
|
||||
self.extra_params = extra_params or {}
|
||||
|
||||
def chat_completion(self, messages: list[dict], **kwargs) -> dict:
|
||||
"""调用 LLM chat completion API"""
|
||||
import httpx
|
||||
|
||||
url = f"{self.base_url.rstrip('/')}/chat/completions"
|
||||
headers = {
|
||||
"Authorization": f"Bearer {self.api_key}",
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
payload: dict = {
|
||||
"model": self.model,
|
||||
"messages": messages,
|
||||
}
|
||||
if self.max_tokens is not None:
|
||||
payload["max_tokens"] = self.max_tokens
|
||||
if self.temperature is not None:
|
||||
payload["temperature"] = self.temperature
|
||||
payload.update(self.extra_params)
|
||||
payload.update(kwargs)
|
||||
|
||||
resp = httpx.post(url, json=payload, headers=headers, timeout=self.timeout)
|
||||
resp.raise_for_status()
|
||||
return resp.json()
|
||||
|
||||
@property
|
||||
def is_available(self) -> bool:
|
||||
return bool(self.api_key and self.base_url and self.model)
|
||||
|
||||
|
||||
class VisionClient(LLMClient):
|
||||
"""VLM 多模态客户端(继承 LLM,增加图片支持)"""
|
||||
|
||||
def call_with_images(self, image_urls: list[str], system_prompt: str, user_prompt: str, **kwargs) -> dict:
|
||||
"""VLM 多图片调用"""
|
||||
content: list[dict] = [{"type": "text", "text": user_prompt}]
|
||||
for url in image_urls:
|
||||
content.append({"type": "image_url", "image_url": {"url": url}})
|
||||
messages = [
|
||||
{"role": "system", "content": system_prompt},
|
||||
{"role": "user", "content": content},
|
||||
]
|
||||
return self.chat_completion(messages, **kwargs)
|
||||
|
||||
|
||||
class TTSClient:
|
||||
"""TTS 客户端(简单配置持有者,实际调用由 CosyVoiceService 完成)"""
|
||||
"""TTS 客户端"""
|
||||
|
||||
def __init__(self, provider: str, api_key: str, base_url: str, model: str, timeout: int = 60, extra_params: dict | None = None):
|
||||
self.provider = provider
|
||||
@@ -76,7 +145,7 @@ class TTSClient:
|
||||
|
||||
|
||||
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
|
||||
@@ -92,7 +161,7 @@ class ImageGenClient:
|
||||
|
||||
|
||||
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
|
||||
@@ -112,11 +181,13 @@ class VideoGenClient:
|
||||
|
||||
def _get_session():
|
||||
"""获取 DB session,兼容 api / worker / 独立脚本场景"""
|
||||
# 方式1:全局 SessionLocal(worker/api 启动时通过 build_session_factory 设置)
|
||||
from packages.adapters.sqlalchemy_impl.session import SessionLocal
|
||||
|
||||
if SessionLocal is not None:
|
||||
return SessionLocal()
|
||||
|
||||
# 方式2:尝试 worker_app.db
|
||||
try:
|
||||
from worker_app.db import SessionLocal as WorkerSL
|
||||
|
||||
@@ -125,6 +196,7 @@ def _get_session():
|
||||
except ImportError:
|
||||
pass
|
||||
|
||||
# 方式3:尝试 api 的 db 模块
|
||||
try:
|
||||
from app.db import SessionLocal as ApiSL
|
||||
|
||||
@@ -262,13 +334,9 @@ class AIRouter:
|
||||
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(
|
||||
def _build_llm_client(self, model: ModelConfig, cap: CapabilityConfig) -> LLMClient:
|
||||
return LLMClient(
|
||||
provider=model.provider,
|
||||
api_key=model.api_key,
|
||||
base_url=model.api_base,
|
||||
model=model.model_key,
|
||||
@@ -277,14 +345,11 @@ class AIRouter:
|
||||
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(
|
||||
def _build_vision_client(self, model: ModelConfig, cap: CapabilityConfig) -> VisionClient:
|
||||
return VisionClient(
|
||||
provider=model.provider,
|
||||
api_key=model.api_key,
|
||||
base_url=model.api_base,
|
||||
model=model.model_key,
|
||||
@@ -293,7 +358,6 @@ class AIRouter:
|
||||
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:
|
||||
@@ -326,10 +390,8 @@ class AIRouter:
|
||||
extra_params=cap.extra_params,
|
||||
)
|
||||
|
||||
# ── 公开接口 ────────────────────────────────────────────────────────────
|
||||
|
||||
def get_llm_client(self, key: str, variant: str = "primary"):
|
||||
"""获取 LLM 客户端(返回 DoubaoClient 实例)"""
|
||||
def get_llm_client(self, key: str, variant: str = "primary") -> LLMClient | None:
|
||||
"""获取 LLM 客户端"""
|
||||
cap = self.get_capability(key)
|
||||
if cap and cap.is_enabled:
|
||||
model = self._get_model_or_fallback(cap, variant)
|
||||
@@ -338,8 +400,8 @@ class AIRouter:
|
||||
|
||||
return self._fallback_llm_client(key)
|
||||
|
||||
def get_vision_client(self, key: str, variant: str = "primary"):
|
||||
"""获取 VLM 客户端(返回 DoubaoClient 实例)"""
|
||||
def get_vision_client(self, key: str, variant: str = "primary") -> VisionClient | None:
|
||||
"""获取 VLM 客户端"""
|
||||
cap = self.get_capability(key)
|
||||
if cap and cap.is_enabled:
|
||||
model = self._get_model_or_fallback(cap, variant)
|
||||
@@ -374,8 +436,7 @@ class AIRouter:
|
||||
|
||||
# ── Fallback 方法(读 SharedSettings 环境变量)──────────────────────────
|
||||
|
||||
def _fallback_llm_client(self, key: str):
|
||||
"""Fallback LLM 客户端 — 从 settings 读取配置,不硬编码"""
|
||||
def _fallback_llm_client(self, key: str) -> LLMClient | None:
|
||||
settings = get_shared_settings()
|
||||
model_map = {
|
||||
"intent_parsing": (settings.doubao_fast_model, settings.doubao_base_url, settings.doubao_api_key),
|
||||
@@ -394,9 +455,7 @@ class AIRouter:
|
||||
if not api_key:
|
||||
return None
|
||||
|
||||
from packages.shared.ai_client import DoubaoClient
|
||||
|
||||
return DoubaoClient(
|
||||
return LLMClient(
|
||||
provider="volcengine",
|
||||
api_key=api_key,
|
||||
base_url=base_url,
|
||||
@@ -405,18 +464,15 @@ class AIRouter:
|
||||
max_retries=settings.doubao_max_retries,
|
||||
)
|
||||
|
||||
def _fallback_vision_client(self, key: str):
|
||||
"""Fallback VLM 客户端 — 从 settings 读取 dashscope 配置,不硬编码"""
|
||||
def _fallback_vision_client(self, key: str) -> VisionClient | None:
|
||||
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", "")
|
||||
base_url = "https://dashscope.aliyuncs.com/compatible-mode/v1"
|
||||
model = "qwen3.8-flash"
|
||||
|
||||
from packages.shared.ai_client import DoubaoClient
|
||||
|
||||
return DoubaoClient(
|
||||
return VisionClient(
|
||||
provider="dashscope",
|
||||
api_key=api_key,
|
||||
base_url=base_url,
|
||||
@@ -429,8 +485,8 @@ class AIRouter:
|
||||
api_key = getattr(settings, "cosyvoice_api_key", "")
|
||||
if not api_key:
|
||||
return None
|
||||
base_url = getattr(settings, "cosyvoice_base_url", "https://dashscope.aliyuncs.com/api/v1")
|
||||
model = getattr(settings, "cosyvoice_model", "cosyvoice-v3-flash")
|
||||
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)
|
||||
|
||||
@@ -439,8 +495,8 @@ class AIRouter:
|
||||
api_key = getattr(settings, "doubao_api_key", "")
|
||||
if not api_key:
|
||||
return None
|
||||
base_url = getattr(settings, "doubao_base_url", "https://ark.cn-beijing.volces.com/api/v3")
|
||||
model = getattr(settings, "doubao_image_model", "doubao-seedream-5-0-flash-260915")
|
||||
base_url = getattr(settings, "doubao_base_url", "")
|
||||
model = getattr(settings, "doubao_image_model", "")
|
||||
|
||||
return ImageGenClient(
|
||||
provider="volcengine",
|
||||
@@ -455,8 +511,8 @@ class AIRouter:
|
||||
api_key = getattr(settings, "doubao_api_key", "")
|
||||
if not api_key:
|
||||
return None
|
||||
base_url = getattr(settings, "doubao_base_url", "https://ark.cn-beijing.volces.com/api/v3")
|
||||
model = getattr(settings, "doubao_video_model", "doubao-seedance-2-5-260628")
|
||||
base_url = getattr(settings, "doubao_base_url", "")
|
||||
model = getattr(settings, "doubao_video_model", "")
|
||||
|
||||
return VideoGenClient(
|
||||
provider="volcengine",
|
||||
|
||||
@@ -59,38 +59,6 @@ _ai_config_version.get_shared_settings = lambda: _mock_settings
|
||||
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"),
|
||||
@@ -255,8 +223,7 @@ class TestAIRouter(unittest.TestCase):
|
||||
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"))
|
||||
self.assertTrue(hasattr(client, "call_with_images"))
|
||||
|
||||
@patch.object(_ai_config_version, "get_version", return_value=None)
|
||||
def test_get_tts_client(self, mock_ver):
|
||||
@@ -374,10 +341,14 @@ class TestModelConfig(unittest.TestCase):
|
||||
class TestClientAvailability(unittest.TestCase):
|
||||
"""客户端可用性测试"""
|
||||
|
||||
def test_tts_client_available(self):
|
||||
c = _ai_router.TTSClient(provider="p", api_key="k", base_url="u", model="m")
|
||||
def test_llm_client_available(self):
|
||||
c = _ai_router.LLMClient(provider="p", api_key="k", base_url="u", model="m")
|
||||
self.assertTrue(c.is_available)
|
||||
|
||||
def test_llm_client_unavailable_no_key(self):
|
||||
c = _ai_router.LLMClient(provider="p", api_key="", base_url="u", model="m")
|
||||
self.assertFalse(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)
|
||||
|
||||
@@ -65,7 +65,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(http_client=mock_client)
|
||||
assert svc._base_url == "https://dashscope.aliyuncs.com/api/v1"
|
||||
|
||||
@@ -84,7 +83,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",
|
||||
@@ -120,7 +118,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
|
||||
@@ -150,7 +147,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
|
||||
@@ -214,7 +210,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")
|
||||
@@ -303,7 +298,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)
|
||||
|
||||
@@ -346,7 +340,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)
|
||||
|
||||
@@ -428,7 +421,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")
|
||||
@@ -646,7 +638,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()
|
||||
|
||||
@@ -3,7 +3,7 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from unittest.mock import MagicMock, patch
|
||||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user