feat(#2170): 方舟信任链——真人照片经 Seedream AI 化后走 Seedance reference_image #2170

Merged
xiaoxia merged 6 commits from feature/2170-ark-trust-chain into develop 2026-10-04 10:16:22 +08:00
10 changed files with 984 additions and 1128 deletions
+12 -17
View File
@@ -211,29 +211,24 @@ COSYVOICE_CLONE_MODEL=voice-enrollment
# 用于 AI 文案生成、智能剪辑等需要大模型能力的场景
DOUBAO_API_KEY=your-doubao-api-key
DOUBAO_MODEL=doubao-seed-1-6-250615
DOUBAO_FAST_MODEL=doubao-1-5-pro-32k-250115
DOUBAO_MODEL=doubao-seed-2-1-pro-260915
DOUBAO_FAST_MODEL=doubao-seed-2-1-lite-260915
DOUBAO_BASE_URL=https://ark.cn-beijing.volces.com/api/v3
DOUBAO_TIMEOUT=30
DOUBAO_MAX_RETRIES=2
# 视觉模型:pro 精度高,lite 速度快(viral-video 商品识别默认用 lite 提速)
DOUBAO_VISION_MODEL=doubao-1-5-vision-pro-250328
DOUBAO_VISION_LITE_MODEL=doubao-1-5-vision-lite-250315
DOUBAO_VISION_MODEL=doubao-seed-2-1-pro-260915
DOUBAO_VISION_LITE_MODEL=doubao-seed-2-1-lite-260915
DOUBAO_VISION_USE_LITE=true
# Embedding 向量化模型(原 large-text-240915 已下线,用多模态 embedding)
# Embedding 向量化模型
DOUBAO_EMBEDDING_MODEL=doubao-embedding-vision-251215
# ==================== 即梦(Jimeng)视觉 API —— 真人参考图兜底通道 (#2169) ====
# 方舟 Seedance 走 B 端审核,真人参考图会被 50411 拦截;即梦走 C 端审核,普通真人照片可过审。
# 需要在火山控制台开通即梦 cvtob 服务,使用 AK/SK(Region=cn-north-1, Service=cv)
# 留空则真人拦截后直接返回错误提示,不会走即梦兜底。
JIMENG_AK=
JIMENG_SK=
JIMENG_BASE_URL=https://visual.volcengineapi.com
# 视频3.0 720P 首帧图生视频(1张图,稳定版;1080P标注下线中)
JIMENG_REQ_KEY=jimeng_i2v_first_v30
JIMENG_VIDEO_TIMEOUT=600
JIMENG_VIDEO_POLL_INTERVAL=5
# 视频模型(Seedance 2.5,统一走方舟;真人参考图通过信任链自动 AI 化)
DOUBAO_VIDEO_MODEL=doubao-seedance-2-5-260628
DOUBAO_VIDEO_TIMEOUT=480
DOUBAO_VIDEO_POLL_INTERVAL=10
# 图片模型(Seedream 5.0 Pro,用于信任链真人 AI 化 + 文生图)
DOUBAO_IMAGE_MODEL=doubao-seedream-5-0-pro-260628
DOUBAO_IMAGE_TIMEOUT=120
# ==================== 积分/会员系统 (#1895) ====================
# 积分系统总开关:默认 false(暂停积分系统)。
+4 -14
View File
@@ -1177,9 +1177,8 @@ def _step_render(job: ViralVideoJob, copy_result: dict, tts_audio_url: str | Non
reference_videos=ref_videos,
)
# #2169: 真人/肖像拦截已由 ai_client 内部自动切即梦(jimeng-3.0)通道重试——
# 保留首帧图、不走"去掉参考图纯 t2v 降级"(用户明确要求按参考照片生成)。
# 即梦也失败或非拦截类错误时,直接抛错给上层展示用户友好提示。
# #2170: 真人/肖像拦截由 ai_client 内部信任链自动处理(Seedream AI 化后再调 Seedance);
# 非拦截类错误直接抛错给上层展示用户友好提示。
def _check_and_reraise(result):
if result and isinstance(result, dict):
return result
@@ -1209,17 +1208,8 @@ def _step_render(job: ViralVideoJob, copy_result: dict, tts_audio_url: str | Non
usage = result.get("usage")
if not video_path or not Path(video_path).exists() or Path(video_path).stat().st_size == 0:
raise RuntimeError("视频生成失败:返回空文件或路径不存在")
# #2169: 如果实际走了即梦兜底(真人拦截→jimeng),更新 job.video_model 让积分结算用 jimeng-3.0 价格
if isinstance(usage, dict):
actual_provider = usage.get("provider")
actual_model_key = usage.get("model_key")
if actual_provider == "jimeng" and actual_model_key:
logger.info(
"[爆款视频] 实际通过即梦通道生成(原model=%s),更新video_model=%s 用于积分结算",
job.video_model,
actual_model_key,
)
job.video_model = actual_model_key
# #2170: 统一走方舟 Seedance(含信任链),usage 里的 provider/model_key 用于积分结算;
# 信任链产生的 Seedream 图成本已在利润率中覆盖,不单独结算。
size = Path(video_path).stat().st_size
logger.info("[爆款视频] 单次生成完成: path=%s size=%d usage=%s", video_path, size, usage)
return str(video_path), (usage if isinstance(usage, dict) else None)
+2 -11
View File
@@ -108,6 +108,8 @@ class SharedSettings(BaseSettings):
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-pro-260628" # 图生图/文生图(信任链真人照片AI化)
doubao_image_timeout: int = 120 # 图片生成超时(秒)
# ── DashScope (阿里云百炼 Wan 3.0 等) ─────────────────────────────────
dashscope_api_key: str = ""
@@ -115,17 +117,6 @@ class SharedSettings(BaseSettings):
dashscope_video_timeout: int = 900 # Wan 视频任务轮询总超时(秒)
dashscope_video_poll_interval: int = 10
# ── 即梦(Jimeng)视觉 API —— 火山引擎 cvtob ──────────────────────────
# #2169: 真人参考图在方舟 Seedance 走 B 端审核会被 50411 拦截,
# 即梦走 C 端审核链路,普通真人照片可过审,作为参考图场景兜底通道。
# 鉴权:AK/SK V4 签名(Region=cn-north-1, Service=cv)
jimeng_ak: str = ""
jimeng_sk: str = ""
jimeng_base_url: str = "https://visual.volcengineapi.com"
jimeng_req_key: str = "jimeng_i2v_first_v30" # 视频3.0 720P 首帧图生视频(1张图,稳定版;1080P 标注下线中)
jimeng_video_timeout: int = 600 # 即梦轮询总超时(秒)
jimeng_video_poll_interval: int = 5 # 轮询间隔(秒)
# ── MediaKit (火山引擎 AI 媒体工具) ──────────────────────────────────
mediakit_api_key: str = ""
mediakit_base_url: str = "https://mediakit.cn-beijing.volces.com/api/v1"
-19
View File
@@ -31,8 +31,6 @@ VIRAL_VIDEO_MODEL_PRICES: dict[tuple[str, str, bool], float] = {
("wan-3.0", "480p", False): 0.3,
("wan-3.0", "720p", False): 0.6,
("wan-3.0", "1080p", False): 1.2,
# #2169: 即梦(Jimeng)视频3.0 720P 首帧图生视频,0.28 元/秒(C 端审核,真人可过)
("jimeng-3.0", "720p", False): 0.28,
}
# 固定成本(元):VLM 分析 + LLM 文案 + TTS + OSS + 服务器
@@ -153,20 +151,6 @@ VIRAL_VIDEO_MODEL_CONFIG: dict[str, dict] = {
"billing_mode": "per_second",
"is_default": False,
},
# #2169: 即梦视频3.0(内部兜底通道,方舟 Seedance 返回真人拦截 50411 时自动切到即梦重试,
# 不暴露给前端让用户直接选择,但需支持计费结算)
"jimeng-3.0": {
"key": "jimeng-3.0",
"display_name": "即梦3.0 — 真人图生视频(兜底)",
"model_id": "jimeng_i2v_first_v30",
"provider": "jimeng",
"supports_audio": False, # 即梦返回无声视频,音频由后续 ffmpeg 合成 TTS
"supported_resolutions": ["720p"],
"max_duration": 10, # 即梦 i2v 首帧最长 10s(frames=241)
"billing_mode": "per_second",
"is_default": False,
"_internal_fallback_only": True, # 标记:不对外暴露到模型选择列表
},
}
@@ -189,9 +173,6 @@ def list_viral_video_models(
continue
if cfg.get("provider") == "dashscope" and not dashscope_available:
continue
# #2169: 即梦是内部兜底通道,不在前端模型列表展示
if cfg.get("_internal_fallback_only"):
continue
out.append(
{
"key": cfg["key"],
+213 -134
View File
@@ -34,22 +34,20 @@ _HTTP_STATUS_ERROR = httpx.HTTPStatusError if hasattr(httpx, "HTTPStatusError")
logger = logging.getLogger(__name__)
# 视频模型 ID 解析逻辑(#2159 多模型支持,#2169 接入即梦)。
# 内部使用简短别名(seedance-2.5 / seedance-2.0 / wan-3.0 / jimeng-3.0 等)做 PRICING key;
# 视频模型 ID 解析逻辑(#2159 多模型支持,#2170 方舟信任链统一走方舟)。
# 内部使用简短别名(seedance-2.5 / seedance-2.0 / wan-3.0 等)做 PRICING key;
# 实际调用时按 VIRAL_VIDEO_MODEL_CONFIG(domain/points_rules.py)的 provider/model_id 分发。
# - provider=doubao → 火山方舟 Seedance
# - provider=dashscope → 阿里云 DashScope(Wan 系列)
# - provider=jimeng → 火山引擎即梦 cvtob(jimeng_i2v_first_v30,真人参考图走 C 端审核)
# - provider=doubao → 火山方舟 Seedance(含信任链真人 AI 化)
# - provider=dashscope → 阿里云 DashScope(Wan 系列,可选)
def _resolve_video_provider_and_id(model: str | None) -> tuple[str, str, dict]:
"""把内部 model key 解析成 (provider, model_id, cfg)。
- provider: "doubao" | "dashscope" | "jimeng"
- provider: "doubao" | "dashscope"
- model_id: 对应 API 的真实模型 ID
- cfg: VIRAL_VIDEO_MODEL_CONFIG 条目
未识别或空值回落到默认 seedance-2.5。已经是 doubao-/ep- 开头的完整 ID 视为 doubao provider;
"jimeng" 开头视为 jimeng provider(内部兜底,不暴露给前端)。
未识别或空值回落到默认 seedance-2.5。已经是 doubao-/ep- 开头的完整 ID 视为 doubao provider。
"""
from packages.domain.points_rules import get_viral_video_model_config
@@ -62,10 +60,6 @@ def _resolve_video_provider_and_id(model: str | None) -> tuple[str, str, dict]:
# 已经是 doubao-/ep- 开头:直接透传,默认视为 doubao provider
if m.startswith("doubao-") or m.startswith("ep-"):
return "doubao", m, {"provider": "doubao", "model_id": m, "supports_audio": True}
# 显式 jimeng 关键字:路由到即梦(内部兜底通道使用)
if m.startswith("jimeng"):
cfg = get_viral_video_model_config("jimeng-3.0")
return "jimeng", cfg.get("model_id", "jimeng_i2v_first_v30"), cfg
# 别名 → 从 domain config 查
cfg = get_viral_video_model_config(m)
provider = cfg.get("provider", "doubao")
@@ -187,8 +181,12 @@ class DoubaoClient:
self.vision_lite_model: str = settings.doubao_vision_lite_model
self.fast_model: str = settings.doubao_fast_model
self.embedding_model: str = settings.doubao_embedding_model
self.image_model: str = settings.doubao_image_model
self.image_timeout: int = getattr(settings, "doubao_image_timeout", 120) or 120
# 最近一次视频生成的详细错误(error_code + user_message + raw detail),供上层读取后展示给用户
self.last_video_error: dict = {}
# 最近一次图片生成的详细错误,供上层读取
self.last_image_error: dict = {}
def embed_text(self, text: str, timeout: int | None = None) -> list[float] | None:
"""调用豆包文本 Embedding API,返回浮点向量;失败返回 None。"""
@@ -487,33 +485,83 @@ class DoubaoClient:
}
return None
if provider == "jimeng":
# #2169: 即梦 cvtob(jimeng_i2v_first_v30)— 真人参考图兜底通道
return self._call_jimeng_video_generation(
prompt=prompt,
image_url=image_url,
duration=duration,
ratio=ratio,
resolution=resolution,
output_dir=output_dir,
generate_audio=generate_audio,
)
ref_audios = [u for u in (reference_audios or [])[:10] if u and isinstance(u, str)]
ref_videos = [u for u in (reference_videos or [])[:3] if u and isinstance(u, str)]
ref_imgs = [u for u in (reference_images or [])[:9] if u and isinstance(u, str)]
# 判断任务模式:有参考音/视/多图 → omni_reference(支持指定 ratio);纯首帧 → first_frame(ratio=adaptive)
# ── #2170 方舟信任链(Trust Chain)────────────────────────────────────
# 真人照片直接传给 Seedance 会触发 50411 肖像审核拦截。
# 解决:先通过同账号的 Seedream 5.0 Pro 图生图 AI 化(保持五官特征),
# 得到的 AI 产物图属于"模型信任产物",再作为 reference_image 传给 Seedance 即可通过审核。
# 信任链只作用于 doubao provider;DashScope(Wan) 保持原行为。
trust_chain_applied = False
if provider == "doubao":
seedream_prompt = (
"保持此人五官特征、发型、肤色、面部轮廓、年龄感,生成一张高清写实人像照片,"
"人物外貌特征与参考图完全一致,皮肤自然,光线柔和,高清细节,不要过度美化。"
)
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)
if raw_portrait_urls:
trusted_urls: list[str] = []
for idx, raw_url in enumerate(raw_portrait_urls):
sd_prompt = (
seedream_prompt if len(raw_portrait_urls) == 1 else f"{seedream_prompt}(这是参考图{idx + 1})"
)
sd_result = self.image_generation(
prompt=sd_prompt,
reference_images=[raw_url],
size="2K",
timeout=120,
)
if not sd_result:
logger.warning(
"[trust-chain] Seedream 第 %d/%d 张失败: %s,回退直传原图",
idx + 1,
len(raw_portrait_urls),
self.last_image_error,
)
break
trusted_urls.append(sd_result["url"])
if trusted_urls and len(trusted_urls) == len(raw_portrait_urls):
trust_chain_applied = True
# 替换:原 image_url 用第一张 AI 图,ref_imgs 用剩余
if image_url and trusted_urls:
image_url = trusted_urls[0]
ref_imgs = trusted_urls[1:] if len(trusted_urls) > 1 else []
else:
ref_imgs = trusted_urls
logger.info(
"[trust-chain] Seedream AI 化完成 %d 张,替换为 reference_image 模式",
len(trusted_urls),
)
else:
# Seedream 部分失败 → 回退原图直传(仍可能被 50411 拦截,但保留降级路径)
logger.warning(
"[trust-chain] Seedream AI 化不完整(%d/%d),回退原图直传",
len(trusted_urls),
len(raw_portrait_urls),
)
# ─────────────────────────────────────────────────────────────────
# 判断任务模式:
# - 信任链强制走 reference_image(不是 first_frame;产品语义是人物参考,不是从图开始动)
# - 有参考音/视/多图 → omni_reference(支持指定 ratio)
# - 纯首帧无其他参考 → first_frame(ratio=adaptive)
has_extra_refs = bool(ref_audios or ref_videos or ref_imgs)
is_first_frame_mode = bool(image_url) and not has_extra_refs
is_first_frame_mode = bool(image_url) and not has_extra_refs and not trust_chain_applied
# 最终 ratio:first_frame 模式强制 adaptive,否则按用户传值(默认 9:16)
final_ratio = "adaptive" if is_first_frame_mode else (ratio or "9:16")
# 构造 content 数组:text + 图 + 音 + 视
content: list[dict[str, Any]] = [{"type": "text", "text": prompt.strip()}]
if image_url:
if has_extra_refs:
# omni_reference:首张图作为 reference_image,允许指定 ratio
if has_extra_refs or trust_chain_applied:
# omni_reference 或信任链模式:首张图作为 reference_image,允许指定 ratio
content.append(
{
"type": "image_url",
@@ -557,7 +605,7 @@ class DoubaoClient:
video_model,
duration,
final_ratio,
"first_frame" if is_first_frame_mode else "omni_ref",
"first_frame" if is_first_frame_mode else ("omni_ref+trust_chain" if trust_chain_applied else "omni_ref"),
generate_audio,
(1 if image_url else 0) + len(ref_imgs),
len(ref_audios),
@@ -669,28 +717,6 @@ class DoubaoClient:
last_err,
(body or "")[:500],
)
# #2169: 方舟返回 portrait_intercept 且有参考图 → 自动切即梦重试一次(保留首帧图)
if err_code == "portrait_intercept" and image_url:
logger.warning(
"[viral-video] Seedance 真人拦截(code=%s),自动切即梦通道重试(首帧图) img=%s",
err_code,
bool(image_url),
)
jm_result = self._call_jimeng_video_generation(
prompt=prompt,
image_url=image_url,
duration=duration,
ratio=ratio,
resolution=resolution,
output_dir=output_dir,
generate_audio=False, # 即梦 i2v 不带音频,音频由后续 ffmpeg 合成
_portrait_fallback=True,
)
if jm_result is not None:
return jm_result
# 即梦也失败了,保留即梦的 last_video_error(已经由 _call_jimeng 设置)
logger.error("[viral-video] 即梦通道重试也失败: %s", self.last_video_error)
return None
return None
logger.info("Seedance 任务已创建: task_id=%s ratio=%s", task_id, create_payload["ratio"])
@@ -834,99 +860,152 @@ class DoubaoClient:
}
return None
def _call_jimeng_video_generation(
def image_generation(
self,
*,
prompt: str,
image_url: str | None,
duration: int,
ratio: str | None,
resolution: str,
output_dir: str | None,
generate_audio: bool = False,
_portrait_fallback: bool = False,
*,
reference_images: list[str] | None = None,
size: str = "2K",
model: str | None = None,
watermark: bool = False,
output_format: str = "png",
timeout: int | None = None,
) -> dict | None:
"""#2169: 调用即梦 cvtob 客户端做图生视频(真人参考图兜底通道)。
"""#2170: 调用方舟 Seedream 图片生成(文生图/图生图)。
- 即梦 i2v 首帧接口只接受 1 张图、无原生音频(返回无声视频,音频由 ffmpeg 后合)。
- 成功返回 {"video_path": str, "usage": {...}};失败写 self.last_video_error 并返回 None。
- _portrait_fallback=True 时在日志里标注是从方舟拦截切过来的。
- reference_images: 0~10 张参考图 URL;0 张 = 纯文生图;1 张 string/URL 直传;多张 list[str]。
- 成功返回 {"url": str, "usage": dict | None};失败返回 None,错误写入 self.last_image_error。
- 返回的 url 有时效性(通常 24h),应立即使用,不持久化存储。
"""
from packages.shared.jimeng_client import get_jimeng_client
jm = get_jimeng_client()
if jm is None:
detail = "即梦 client 不可用(JIMENG_AK/SK 未配置)"
if _portrait_fallback:
# 从真人拦截切过来但即梦没配,仍把错误归到 portrait_intercept,让上层提示用户
self.last_video_error = {
"error_code": "portrait_intercept",
"user_message": "参考素材包含真人照片被安全策略拦截,即梦兜底通道未启用,请联系管理员配置 JIMENG_AK/SK。",
"status_code": 0,
"detail": detail,
"provider": "jimeng",
}
else:
self.last_video_error = {
"error_code": "auth_error",
"user_message": "即梦视频通道未配置(JIMENG_AK/SK 缺失),请联系管理员。",
"status_code": 0,
"detail": detail,
"provider": "jimeng",
}
logger.error("[jimeng] %s, portrait_fallback=%s", detail, _portrait_fallback)
self.last_image_error = {}
if not self.is_available:
self.last_image_error = {
"error_code": "auth_error",
"user_message": "图片生成服务未配置(API Key 缺失),请联系管理员。",
"detail": "DoubaoClient not available (api_key empty)",
}
return None
if not image_url:
self.last_video_error = {
if not prompt or not prompt.strip():
self.last_image_error = {
"error_code": "invalid_param",
"user_message": "即梦图生视频必须提供参考图片。",
"status_code": 0,
"detail": "empty image_url for jimeng i2v",
"provider": "jimeng",
"user_message": "图片生成提示词不能为空。",
"detail": "empty prompt",
}
return None
# 即梦 i2v 无声视频,generate_audio 强制 False
jm.last_video_error = {}
tag = "[portrait-fallback→jimeng]" if _portrait_fallback else "[jimeng-direct]"
logger.info("%s 调用即梦: dur=%s ratio=%s res=%s img=%s", tag, duration, ratio, resolution, bool(image_url))
try:
result = jm.video_generation(
prompt=prompt,
image_url=image_url,
duration=duration,
ratio=ratio,
resolution=resolution,
output_dir=output_dir,
generate_audio=False,
)
except Exception as je:
logger.error("%s 即梦 video_generation 异常: %s", tag, je, exc_info=True)
self.last_video_error = {
"error_code": "unknown",
"user_message": f"即梦视频生成异常:{je!s}"[:200],
"status_code": 0,
"detail": str(je),
"provider": "jimeng",
}
return None
if not result and jm.last_video_error:
# 透传即梦错误;如果即梦也返回 portrait_intercept,说明图片真的有问题,直接给用户
jm_err = dict(jm.last_video_error)
jm_err["provider"] = "jimeng"
if _portrait_fallback and jm_err.get("error_code") == "portrait_intercept":
jm_err["user_message"] = (
"参考素材真人肖像审核未通过(方舟+即梦双通道均被拦截),请更换非真人或授权清晰的照片后重试。"
img_model = model or self.image_model
url = f"{self.base_url}/images/generations"
headers = {
"Authorization": f"Bearer {self.api_key}",
"Content-Type": "application/json",
}
payload: dict[str, Any] = {
"model": img_model,
"prompt": prompt.strip(),
"size": size,
"response_format": "url",
"output_format": output_format,
"watermark": bool(watermark),
}
ref_imgs_local = [u for u in (reference_images or []) if u and isinstance(u, str)]
if ref_imgs_local:
if len(ref_imgs_local) == 1:
payload["image"] = ref_imgs_local[0]
else:
payload["image"] = ref_imgs_local[:10]
req_timeout = timeout or self.image_timeout
last_err: Exception | None = None
last_sc = 0
last_body = ""
for attempt in range(self.max_retries + 1):
try:
resp = httpx.post(url, headers=headers, json=payload, timeout=req_timeout)
last_sc = int(getattr(resp, "status_code", 0) or 0)
last_body = (getattr(resp, "text", "") or "")[:2000]
if last_sc >= 400:
logger.error("Seedream 图片生成 HTTP %d: %s", last_sc, last_body[:500])
try:
resp.raise_for_status()
except Exception as ee:
last_err = ee
if attempt < self.max_retries and last_sc >= 500:
time.sleep(0.5 * (2**attempt))
continue
break
data = resp.json()
data_list = data.get("data") or []
if data_list and isinstance(data_list, list):
item = data_list[0]
img_url = item.get("url")
if img_url:
logger.info(
"Seedream 图片生成成功 model=%s ref_imgs=%d size=%s",
img_model,
len(ref_imgs_local),
size,
)
return {"url": img_url, "usage": data.get("usage")}
last_err = RuntimeError(f"Seedream 返回结构异常: {str(data)[:300]}")
break
except _HTTP_NETWORK_ERRORS as ne:
last_err = ne
last_sc = 0
last_body = f"network error: {ne}"
logger.warning(
"Seedream 网络异常 (%s),重试 %d/%d", type(ne).__name__, attempt + 1, self.max_retries + 1
)
self.last_video_error = jm_err
return None
if result:
# 补充 usage 里的 provider 标记
u = result.get("usage") or {}
u.setdefault("provider", "jimeng")
u.setdefault("model_key", "jimeng-3.0")
result["usage"] = u
logger.info("%s 即梦生成成功: %s", tag, result.get("video_path"))
return result
if attempt < self.max_retries:
time.sleep(0.5 * (2**attempt))
continue
break
except Exception as e:
last_err = e
if attempt < self.max_retries and not isinstance(e, _HTTP_STATUS_ERROR):
wait = 0.5 * (2**attempt)
logger.warning(
"Seedream 图片生成失败,%.1fs 后重试 (%d/%d): %s", wait, attempt + 1, self.max_retries + 1, e
)
time.sleep(wait)
continue
break
# 分类错误
err_code = "unknown"
user_msg = "图片生成失败,请稍后重试。"
body_lower = (last_body or "").lower()
if last_sc == 401 or last_sc == 403:
err_code, user_msg = "auth_error", "图片生成服务鉴权失败,请联系管理员。"
elif last_sc == 400:
if any(k in body_lower for k in ("quota", "billing", "insufficient", "balance")):
err_code, user_msg = "quota_exceeded", "图片生成配额不足或账号欠费,请联系管理员。"
elif any(k in body_lower for k in ("rate", "throughput", "too many", "frequency")):
err_code, user_msg = "rate_limit", "图片生成请求过于频繁,请稍后重试。"
elif any(k in body_lower for k in ("sensitive", "porn", "terror", "risk", "audit", "content", "violat")):
err_code, user_msg = "portrait_intercept", "参考素材未通过内容安全审核,请更换照片后重试。"
else:
err_code, user_msg = "invalid_param", f"图片生成参数错误:{last_body[:200]}"
elif last_sc == 404:
err_code, user_msg = "model_not_found", f"图片模型 {img_model} 不存在,请联系管理员。"
elif last_sc >= 500:
err_code, user_msg = "network_error", "图片生成服务暂时不可用,请稍后重试。"
elif last_sc == 0:
err_code, user_msg = "network_error", f"图片生成网络错误:{last_err!s}"[:200]
self.last_image_error = {
"error_code": err_code,
"user_message": user_msg,
"status_code": last_sc,
"detail": (last_body or "")[:500] or (str(last_err) if last_err else ""),
"model": img_model,
}
logger.error(
"Seedream 图片生成最终失败: model=%s status=%d code=%s err=%s", img_model, last_sc, err_code, last_err
)
return None
def get_last_image_error(self) -> dict:
"""返回最近一次 image_generation 失败的详细错误。空 dict 表示上次成功或未调用。"""
return dict(self.last_image_error or {})
def get_last_video_error(self) -> dict:
"""返回最近一次 video_generation 失败的详细错误。空 dict 表示上次成功或未调用。"""
-526
View File
@@ -1,526 +0,0 @@
"""即梦(Jimeng)视觉 API 客户端 —— 火山引擎 cvtob。
#2169: 真人参考图在方舟 Seedance 走 B 端审核会被 50411 拦截,
即梦走 C 端审核链路,普通真人照片可过审。接入即梦 i2v 作为参考图场景兜底通道。
接口协议(jimeng_i2v_first_v30 —— 视频3.0 720P 首帧图生视频):
- 接口地址:https://visual.volcengineapi.com
- 鉴权:火山 V4 签名(Region=cn-north-1, Service=cv),使用 AK/SK
- 提交任务:POST ?Action=CVSync2AsyncSubmitTask&Version=2022-08-31
body: {"req_key": "jimeng_i2v_first_v30", "image_urls": ["<url>"], "prompt": "...", "seed": -1, "frames": 121}
-> {"code": 10000, "data": {"task_id": "..."}}
- 查询任务:POST ?Action=CVSync2AsyncGetResult&Version=2022-08-31
body: {"req_key": "jimeng_i2v_first_v30", "task_id": "..."}
-> {"code": 10000, "data": {"status": "in_queue|generating|done", "video_url": "..."}}
- 视频 URL 有效期 1 小时,必须立即下载到本地。
"""
from __future__ import annotations
import hashlib
import hmac
import json
import logging
import os
import time
import uuid
from datetime import datetime, timezone
from pathlib import Path
from typing import Any
from urllib.parse import quote, urlparse
import httpx
from packages.shared.config import get_shared_settings
# 网络/超时类异常父类集合
_HTTP_NETWORK_ERRORS = ()
try:
_HTTP_NETWORK_ERRORS = (httpx.TimeoutException, httpx.NetworkError)
except Exception:
_HTTP_NETWORK_ERRORS = (Exception,)
logger = logging.getLogger(__name__)
_JIMENG_CLIENT_SINGLETON: "JimengClient | None" = None
# ── V4 签名常量 ────────────────────────────────────────────────────────
_JIMENG_REGION = "cn-north-1"
_JIMENG_SERVICE = "cv"
_JIMENG_VERSION = "2022-08-31"
_ACTION_SUBMIT = "CVSync2AsyncSubmitTask"
_ACTION_POLL = "CVSync2AsyncGetResult"
_CONTENT_TYPE = "application/json"
_SIGNED_HEADERS_LIST = ["content-type", "host", "x-content-sha256", "x-date"]
_SIGNED_HEADERS_STR = ";".join(_SIGNED_HEADERS_LIST)
def _norm_query(params: dict[str, str]) -> str:
"""构造规范查询串:按 key 排序,URL 编码(safe=-_.~),空格->%20。"""
parts = []
for k in sorted(params.keys()):
v = params[k]
ek = quote(str(k), safe="-_.~")
ev = quote(str(v), safe="-_.~").replace("+", "%20")
parts.append(f"{ek}={ev}")
return "&".join(parts)
def _hmac_sha256(key: bytes, msg: str) -> bytes:
return hmac.new(key, msg.encode("utf-8"), hashlib.sha256).digest()
def _sha256_hex(data: bytes) -> str:
return hashlib.sha256(data).hexdigest()
def _sign_v4(
ak: str,
sk: str,
method: str,
host: str,
query: dict[str, str],
body_bytes: bytes,
x_date: str,
) -> dict[str, str]:
"""火山 V4 签名,返回需要附加到请求的 headers 字典。
x_date 形如 "20260101T120000Z"(UTC)。
short_date = x_date[:8](YYYYMMDD)。
"""
short_date = x_date[:8]
payload_hash = _sha256_hex(body_bytes)
canon_uri = "/"
canon_query = _norm_query(query)
canon_headers = f"content-type:{_CONTENT_TYPE}\nhost:{host}\nx-content-sha256:{payload_hash}\nx-date:{x_date}\n"
canon_request = f"{method}\n{canon_uri}\n{canon_query}\n{canon_headers}\n{_SIGNED_HEADERS_STR}\n{payload_hash}"
credential_scope = f"{short_date}/{_JIMENG_REGION}/{_JIMENG_SERVICE}/request"
string_to_sign = f"HMAC-SHA256\n{x_date}\n{credential_scope}\n{_sha256_hex(canon_request.encode('utf-8'))}"
k_date = _hmac_sha256(sk.encode("utf-8"), short_date)
k_region = _hmac_sha256(k_date, _JIMENG_REGION)
k_service = _hmac_sha256(k_region, _JIMENG_SERVICE)
k_signing = _hmac_sha256(k_service, "request")
signature = hmac.new(k_signing, string_to_sign.encode("utf-8"), hashlib.sha256).hexdigest()
authorization = (
f"HMAC-SHA256 Credential={ak}/{credential_scope}, SignedHeaders={_SIGNED_HEADERS_STR}, Signature={signature}"
)
return {
"Content-Type": _CONTENT_TYPE,
"Host": host,
"X-Content-Sha256": payload_hash,
"X-Date": x_date,
"Authorization": authorization,
}
# ── 错误分类 ──────────────────────────────────────────────────────────
# 即梦业务码 -> 是否可重试映射
_JIMENG_RETRYABLE_CODES = {50511, 50516, 50429, 50430, 50500, 50501}
_JIMENG_NON_RETRYABLE_CODES = {50411, 50412, 50413, 50512, 50513, 50514}
def _classify_jimeng_error(status_code: int, body: str, biz_code: int | None = None) -> tuple[str, str, bool]:
"""即梦错误分类,返回 (error_code, user_message, is_retryable)。"""
code = biz_code if biz_code is not None else 0
body_lower = (body or "").lower()
# 业务码优先
if code == 50411:
return (
"portrait_intercept",
"即梦通道:参考图片前审核未通过(Pre Img Risk Not Pass),请更换参考图后重试。",
False,
)
if code == 50511:
return "task_failed", "即梦通道:输出图片后审核未通过,可稍后重试。", True
if code in (50412, 50413, 50512):
return "invalid_param", "即梦通道:提示词或文本审核不通过,请调整文案后重试。", False
if code == 50516:
return "task_failed", "即梦通道:输出视频后审核未通过,可稍后重试。", True
if code in (50429, 50430):
return "rate_limit", "即梦通道:QPS/并发超限,请稍等 1-2 分钟后重试。", True
if code in (50500, 50501):
return "network_error", "即梦通道:服务内部错误,可稍后重试。", True
# HTTP 层兜底
if status_code in (401, 403):
return "auth_error", "即梦通道:AK/SK 鉴权失败,请联系管理员检查 JIMENG_AK/SK 配置。", False
if status_code == 429:
return "rate_limit", "即梦通道:服务限流,请稍后重试。", True
if status_code == 404:
return "model_not_found", "即梦通道:接口不存在(req_key 或 Action 错误),请联系管理员。", False
if status_code in (402, 400) and any(kw in body_lower for kw in ("quota", "billing", "insufficient", "余额")):
return "quota_exceeded", "即梦通道:账户余额/配额不足,请联系管理员充值。", False
if status_code == 400:
msg = ""
try:
msg = str(json.loads(body or "{}").get("message", "") or "")
except Exception:
pass
return "invalid_param", f"即梦通道:参数错误:{msg or body[:200]}", False
if status_code == 0:
return "network_error", "即梦通道:网络连接失败,请稍后重试。", True
# 任务内失败
if code and code != 10000:
return "unknown", f"即梦通道:视频生成失败(错误码 {code}),请稍后重试。", code in _JIMENG_RETRYABLE_CODES
detail = body[:200]
return "unknown", f"即梦通道:视频生成失败(HTTP {status_code}):{detail}", False
# ── 即梦客户端 ────────────────────────────────────────────────────────
class JimengClient:
"""火山引擎即梦视觉 API(cvtob)异步客户端,支持图生视频首帧(jimeng_i2v_first_v30)。"""
def __init__(self) -> None:
settings = get_shared_settings()
self.ak: str = getattr(settings, "jimeng_ak", "") or os.getenv("JIMENG_AK", "")
self.sk: str = getattr(settings, "jimeng_sk", "") or os.getenv("JIMENG_SK", "")
self.base_url: str = (getattr(settings, "jimeng_base_url", "") or "https://visual.volcengineapi.com").rstrip(
"/"
)
self.req_key: str = getattr(settings, "jimeng_req_key", "") or "jimeng_i2v_first_v30"
self.poll_interval: int = int(getattr(settings, "jimeng_video_poll_interval", 5) or 5)
self.total_timeout: int = int(getattr(settings, "jimeng_video_timeout", 600) or 600)
self.max_retries: int = 2
self.last_video_error: dict = {}
# 解析 base_url 里的 host(用于签名 Host 头)
parsed = urlparse(self.base_url)
self.host: str = parsed.netloc or "visual.volcengineapi.com"
@property
def is_available(self) -> bool:
return bool(self.ak and self.sk)
def get_last_video_error(self) -> dict:
return dict(self.last_video_error or {})
def _set_error(self, error_code: str, user_message: str, status_code: int = 0, detail: str = "", **extra) -> None:
self.last_video_error = {
"error_code": error_code,
"user_message": user_message,
"status_code": status_code,
"detail": detail[:500] if detail else "",
"provider": "jimeng",
**extra,
}
# ── 内部 HTTP:签名 + 请求 ──────────────────────────────────────
def _signed_request(
self,
method: str,
action: str,
body_obj: dict[str, Any],
timeout: float = 60.0,
) -> tuple[int, str, dict]:
"""发送一次带 V4 签名的请求,返回 (status_code, body_text, parsed_json)。"""
body_bytes = json.dumps(body_obj, ensure_ascii=False).encode("utf-8")
query = {"Action": action, "Version": _JIMENG_VERSION}
x_date = datetime.now(timezone.utc).strftime("%Y%m%dT%H%M%SZ")
headers = _sign_v4(self.ak, self.sk, method, self.host, query, body_bytes, x_date)
url = f"{self.base_url}/?{_norm_query(query)}"
resp = httpx.request(
method,
url,
headers=headers,
content=body_bytes,
timeout=timeout,
)
sc = int(getattr(resp, "status_code", 0) or 0)
text = getattr(resp, "text", "") or ""
try:
data = resp.json()
except Exception:
data = {}
return sc, text, data
# ── 提交任务 ────────────────────────────────────────────────────
def _submit_task(
self,
prompt: str,
image_url: str,
frames: int = 121,
seed: int = -1,
) -> str | None:
"""提交图生视频任务,成功返回 task_id;失败写 last_video_error 并返回 None。"""
body: dict[str, Any] = {
"req_key": self.req_key,
"prompt": prompt.strip()[:800],
"image_urls": [image_url],
"seed": int(seed) if seed and seed > 0 else -1,
"frames": int(frames),
}
last_sc = 0
last_body = ""
for attempt in range(self.max_retries + 1):
try:
sc, text, data = self._signed_request("POST", _ACTION_SUBMIT, body, timeout=60.0)
last_sc, last_body = sc, text
if sc >= 400:
logger.error("[jimeng] 提交 HTTP %d: %s", sc, text[:500])
if sc >= 500 and attempt < self.max_retries:
time.sleep(0.8 * (2**attempt))
continue
biz_code = data.get("code") if isinstance(data, dict) else None
err_code, user_msg, _ = _classify_jimeng_error(sc, text, biz_code)
self._set_error(err_code, user_msg, sc, text, req_key=self.req_key)
return None
code = data.get("code") if isinstance(data, dict) else None
if code == 10000:
d = data.get("data") or {}
tid = d.get("task_id")
if tid:
return str(tid)
err_code, user_msg, retry = _classify_jimeng_error(sc, text, code)
logger.error(
"[jimeng] 提交业务错误 code=%s msg=%s",
code,
(data.get("message") if isinstance(data, dict) else ""),
)
if retry and attempt < self.max_retries:
time.sleep(0.8 * (2**attempt))
continue
self._set_error(err_code, user_msg, sc, text, req_key=self.req_key, biz_code=code)
return None
except _HTTP_NETWORK_ERRORS as ne:
last_sc, last_body = 0, f"network error: {ne}"
logger.warning(
"[jimeng] 提交网络异常 %s,重试 %d/%d", type(ne).__name__, attempt + 1, self.max_retries + 1
)
if attempt < self.max_retries:
time.sleep(0.8 * (2**attempt))
continue
self._set_error("network_error", "即梦通道:提交任务网络异常,请稍后重试。", 0, str(ne))
return None
except Exception as e:
last_sc, last_body = 0, f"exception: {e}"
logger.error("[jimeng] 提交异常: %s", e, exc_info=True)
if attempt < self.max_retries:
time.sleep(0.8 * (2**attempt))
continue
self._set_error("unknown", f"即梦通道:提交任务异常:{e!s}"[:200], 0, str(e))
return None
if not self.last_video_error:
err_code, user_msg, _ = _classify_jimeng_error(last_sc, last_body)
self._set_error(err_code, user_msg, last_sc, last_body)
return None
# ── 轮询结果 ────────────────────────────────────────────────────
def _poll_result(self, task_id: str) -> str | None:
"""轮询任务直到 done/failed/expired/timeout,成功返回 video_url。"""
deadline = time.time() + self.total_timeout
poll_count = 0
last_status = ""
poll_body = {"req_key": self.req_key, "task_id": task_id}
while time.time() < deadline:
poll_count += 1
try:
sc, text, data = self._signed_request("POST", _ACTION_POLL, poll_body, timeout=30.0)
if sc >= 400:
logger.warning("[jimeng] 轮询 HTTP %d: %s", sc, text[:300])
if poll_count < 3:
time.sleep(self.poll_interval)
continue
err_code, user_msg, _ = _classify_jimeng_error(sc, text)
self._set_error(err_code, user_msg, sc, text, task_id=task_id)
return None
code = data.get("code") if isinstance(data, dict) else None
d = data.get("data") if isinstance(data, dict) else None
if code != 10000 or not isinstance(d, dict):
err_code, user_msg, retry = _classify_jimeng_error(sc, text, code)
logger.error(
"[jimeng] 轮询业务错误 task=%s code=%s msg=%s",
task_id,
code,
(data.get("message") if isinstance(data, dict) else ""),
)
if retry and poll_count < 3:
time.sleep(self.poll_interval)
continue
self._set_error(err_code, user_msg, sc, text, task_id=task_id, biz_code=code)
return None
status = d.get("status", "") or ""
last_status = status
if status == "done":
video_url = d.get("video_url") or ""
if video_url:
logger.info("[jimeng] 任务 %s 完成 polls=%d", task_id, poll_count)
return str(video_url)
logger.error("[jimeng] 任务 %s done 但无 video_url: %s", task_id, str(d)[:500])
self._set_error(
"unknown",
"即梦通道:任务成功但未返回视频URL,请联系管理员。",
200,
str(d)[:500],
task_id=task_id,
)
return None
if status in ("not_found", "expired"):
logger.error("[jimeng] 任务 %s 状态 %s", task_id, status)
self._set_error(
"network_error" if status == "expired" else "unknown",
f"即梦通道:任务{'已过期' if status == 'expired' else '未找到'},请重新提交。",
200,
f"task {status}",
task_id=task_id,
)
return None
if poll_count % 6 == 0:
logger.info("[jimeng] 轮询中 task=%s status=%s polls=%d", task_id, status, poll_count)
except _HTTP_NETWORK_ERRORS as ne:
logger.warning("[jimeng] 轮询网络异常 %s", ne)
except Exception as e:
logger.debug("[jimeng] 轮询异常: %s", e)
time.sleep(self.poll_interval)
logger.error(
"[jimeng] 任务 %s 轮询超时(%ds)polls=%d last_status=%s",
task_id,
self.total_timeout,
poll_count,
last_status,
)
self._set_error(
"network_error",
f"即梦通道:视频生成超时(>{self.total_timeout}s),任务仍在排队,请稍后重试。",
0,
f"timeout after {self.total_timeout}s, polls={poll_count}, last_status={last_status}",
task_id=task_id,
last_status=last_status,
)
return None
# ── 下载视频 ────────────────────────────────────────────────────
def _download_video(self, video_url: str, output_dir: str, task_id: str) -> str | None:
os.makedirs(output_dir, exist_ok=True)
suffix = Path(urlparse(video_url).path).suffix or ".mp4"
if suffix.lower() not in (".mp4", ".mov", ".webm"):
suffix = ".mp4"
safe_tid = "".join(c if c.isalnum() or c in "-_" else "_" for c in task_id)[:40]
out_path = os.path.join(output_dir, f"jimeng_{safe_tid}_{uuid.uuid4().hex[:8]}{suffix}")
try:
with httpx.stream("GET", video_url, timeout=180, follow_redirects=True) as resp:
dsc = int(getattr(resp, "status_code", 0) or 0)
if dsc >= 400:
logger.error("[jimeng] 下载 HTTP %d", dsc)
self._set_error("network_error", "即梦通道:视频下载失败(HTTP错误),请稍后重试。", dsc)
return None
with open(out_path, "wb") as f:
for chunk in resp.iter_bytes(chunk_size=1024 * 256):
if chunk:
f.write(chunk)
except Exception as e:
logger.error("[jimeng] 下载视频失败: %s", e, exc_info=True)
self._set_error("network_error", f"即梦通道:视频下载失败:{e!s}"[:200], 0, str(e))
return None
size = os.path.getsize(out_path) if os.path.exists(out_path) else 0
if size < 1024:
logger.error("[jimeng] 下载文件过小: %d bytes", size)
self._set_error("unknown", "即梦通道:视频下载文件过小,请稍后重试。", 0, f"downloaded only {size} bytes")
try:
os.remove(out_path)
except Exception:
pass
return None
logger.info("[jimeng] 视频已下载: %s (%d bytes)", out_path, size)
return out_path
# ── 对外主入口 ──────────────────────────────────────────────────
def video_generation(
self,
prompt: str,
*,
image_url: str,
duration: int = 5,
ratio: str | None = "9:16",
resolution: str = "720p",
output_dir: str | None = None,
generate_audio: bool = False,
) -> dict | None:
"""即梦图生视频主入口。
成功返回 {"video_path": str, "usage": {"provider","duration_seconds","frames","req_key","billing_mode"}};
失败返回 None,详情在 self.last_video_error。
注意:jimeng_i2v_first_v30 不支持原生音频(generate_audio 被忽略,返回无声视频),
音频由后续 ffmpeg 合成阶段叠加 TTS。
支持时长:5s(frames=121)/10s(frames=241),>10s 截断并打 warning。
分辨率固定 720P;ratio 对首帧 i2v 无效(自动按图片比例)。
"""
self.last_video_error = {}
if not self.is_available:
self._set_error(
"auth_error",
"即梦通道未配置(JIMENG_AK/SK 缺失),请联系管理员。",
detail="jimeng ak/sk empty",
)
logger.error("[jimeng] AK/SK 未配置,无法调用")
return None
if not prompt or not prompt.strip():
self._set_error("invalid_param", "视频生成提示词不能为空。", detail="empty prompt")
return None
if not image_url or not image_url.strip():
self._set_error("invalid_param", "即梦图生视频必须提供参考图片。", detail="empty image_url")
return None
dur = int(duration or 5)
if dur <= 5:
frames = 121
real_dur = 5
elif dur <= 10:
frames = 241
real_dur = 10
else:
logger.warning("[jimeng] 请求时长 %ds 超出即梦 i2v 上限 10s,截断到 10s(frames=241)", dur)
frames = 241
real_dur = 10
out_dir = output_dir or "/tmp"
logger.info(
"[jimeng] 提交任务: req_key=%s dur=%ds(frames=%d) ratio=%s res=%s gen_audio=%s img=%s",
self.req_key,
real_dur,
frames,
ratio,
resolution,
generate_audio,
bool(image_url),
)
task_id = self._submit_task(prompt=prompt, image_url=image_url, frames=frames, seed=-1)
if not task_id:
return None
logger.info("[jimeng] 任务已提交: task_id=%s", task_id)
video_url = self._poll_result(task_id)
if not video_url:
return None
local_path = self._download_video(video_url, out_dir, task_id)
if not local_path:
return None
usage = {
"provider": "jimeng",
"duration_seconds": real_dur,
"frames": frames,
"req_key": self.req_key,
"billing_mode": "per_second",
}
return {"video_path": local_path, "usage": usage}
def get_jimeng_client() -> "JimengClient | None":
"""返回即梦客户端单例;未配置 AK/SK 时返回 None。"""
global _JIMENG_CLIENT_SINGLETON
if _JIMENG_CLIENT_SINGLETON is None:
_JIMENG_CLIENT_SINGLETON = JimengClient()
if not _JIMENG_CLIENT_SINGLETON.is_available:
return None
return _JIMENG_CLIENT_SINGLETON
+61 -24
View File
@@ -1,4 +1,5 @@
"""Additional unit tests to hit uncovered lines for diff-coverage >=60%."""
from __future__ import annotations
import json
@@ -13,15 +14,20 @@ from packages.shared.ai_client import DoubaoClient
class _FakeSettings:
doubao_api_key = "test-key"
doubao_model = "test-model"
doubao_fast_model = "test-fast-model"
doubao_model = "doubao-seed-2-1-pro-260915"
doubao_fast_model = "doubao-seed-2-1-lite-260915"
doubao_base_url = "https://ark.cn-beijing.volces.com/api/v3"
doubao_timeout = 10
doubao_max_retries = 0
doubao_vision_model = "test-vision"
doubao_vision_lite_model = "test-vision-lite"
doubao_vision_model = "doubao-seed-2-1-pro-260915"
doubao_vision_lite_model = "doubao-seed-2-1-lite-260915"
doubao_vision_use_lite = False
doubao_embedding_model = "test-embedding"
doubao_embedding_model = "doubao-embedding-vision-251215"
doubao_video_model = "doubao-seedance-2-5-260628"
doubao_video_timeout = 480
doubao_video_poll_interval = 10
doubao_image_model = "doubao-seedream-5-0-pro-260628"
doubao_image_timeout = 120
def _make_client(api_key: str = "test-key") -> DoubaoClient:
@@ -129,26 +135,50 @@ from packages.domain.atom_clip_tagger import parse_vision_response
class TestParseVisionResponseEdgeCases:
def test_person_count_type_error_defaults_zero(self):
text = json.dumps({
"scene": [], "objects": [], "action": [], "shot": "", "has_text": False,
"person_count": "not-an-int", "text_content": "", "caption": "x",
})
text = json.dumps(
{
"scene": [],
"objects": [],
"action": [],
"shot": "",
"has_text": False,
"person_count": "not-an-int",
"text_content": "",
"caption": "x",
}
)
r = parse_vision_response(text)
assert r["person_count"] == 0
def test_person_count_out_of_range_clamped(self):
text = json.dumps({
"scene": [], "objects": [], "action": [], "shot": "", "has_text": False,
"person_count": 10, "text_content": "", "caption": "x",
})
text = json.dumps(
{
"scene": [],
"objects": [],
"action": [],
"shot": "",
"has_text": False,
"person_count": 10,
"text_content": "",
"caption": "x",
}
)
r = parse_vision_response(text)
assert r["person_count"] == 3
def test_person_count_negative_clamped(self):
text = json.dumps({
"scene": [], "objects": [], "action": [], "shot": "", "has_text": False,
"person_count": -5, "text_content": "", "caption": "x",
})
text = json.dumps(
{
"scene": [],
"objects": [],
"action": [],
"shot": "",
"has_text": False,
"person_count": -5,
"text_content": "",
"caption": "x",
}
)
r = parse_vision_response(text)
assert r["person_count"] == 0
@@ -159,10 +189,18 @@ class TestParseVisionResponseEdgeCases:
def test_caption_truncation_at_80(self):
long_caption = "描" * 100
text = json.dumps({
"scene": [], "objects": [], "action": [], "shot": "", "has_text": False,
"person_count": 0, "text_content": "", "caption": long_caption,
})
text = json.dumps(
{
"scene": [],
"objects": [],
"action": [],
"shot": "",
"has_text": False,
"person_count": 0,
"text_content": "",
"caption": long_caption,
}
)
r = parse_vision_response(text)
assert len(r["caption"]) == 80
@@ -196,9 +234,7 @@ class TestNarrativeMatchNonDictClipTags:
def test_non_dict_clip_tags_are_skipped(self):
a1 = _FA("a1", tags=[])
clip_map = {"a1": [None, "bad", {"scene": ["工厂"], "objects": [], "action": []}, 123]}
matched, unmatched = match_assets_by_script_tags(
[a1], script_tags=["工厂"], clip_ai_tags_by_asset=clip_map
)
matched, unmatched = match_assets_by_script_tags([a1], script_tags=["工厂"], clip_ai_tags_by_asset=clip_map)
assert [a.id for a in matched] == ["a1"]
@@ -231,6 +267,7 @@ class _FQuery:
class TestUpdateCaptionEmbedding:
def _make_repo(self, session):
from packages.adapters.sqlalchemy_impl.asset_atom_clip_repository import SQLAlchemyAssetAtomClipRepository
repo = SQLAlchemyAssetAtomClipRepository.__new__(SQLAlchemyAssetAtomClipRepository)
repo.session = session
return repo
+669
View File
@@ -0,0 +1,669 @@
"""#2170 Seedream 图片生成 + 方舟信任链单测。"""
from __future__ import annotations
from unittest.mock import MagicMock, patch
import httpx
from packages.shared.ai_client import DoubaoClient
def _make_client(**overrides):
client = DoubaoClient.__new__(DoubaoClient)
client.api_key = overrides.get("api_key", "test-key")
client.base_url = overrides.get("base_url", "https://ark.cn-beijing.volces.com/api/v3")
client.model = "doubao-model"
client.vision_model = "doubao-vision"
client.embedding_model = "doubao-embedding"
client.image_model = overrides.get("image_model", "doubao-seedream-5-0-pro-260628")
client.image_timeout = overrides.get("image_timeout", 120)
client.timeout = overrides.get("timeout", 30)
client.max_retries = overrides.get("max_retries", 0)
client.last_video_error = {}
client.last_image_error = {}
return client
def _fake_time(base=1000.0, stable_calls=50, big=9e9):
"""返回 time.time 替身:前 stable_calls 次返回 base+i,之后返回 big+i。
Python 3.12 logging.LogRecord.__init__ 内部会调 time.time(),
用有限 iter 会 StopIteration,因此必须用无限生成器。
"""
state = {"n": 0}
def _t():
n = state["n"]
state["n"] += 1
if n < stable_calls:
return base + n
return big + n
return _t
# ── Seedream 图片生成单测 ──────────────────────────────────────────
class TestImageGenerationHappyPath:
def test_returns_none_when_no_api_key(self):
client = _make_client(api_key="")
assert client.image_generation("p") is None
err = client.get_last_image_error()
assert err["error_code"] == "auth_error"
def test_returns_none_on_empty_prompt(self):
client = _make_client()
assert client.image_generation(" ") is None
err = client.get_last_image_error()
assert err["error_code"] == "invalid_param"
def test_text_to_image_success(self):
client = _make_client()
captured = {}
ok_resp = MagicMock()
ok_resp.status_code = 200
ok_resp.json.return_value = {"data": [{"url": "https://cdn.example.com/i.png"}], "usage": {"tokens": 1}}
ok_resp.raise_for_status = MagicMock()
ok_resp.text = ""
def fake_post(url, **kwargs):
captured["url"] = url
captured["json"] = kwargs.get("json")
return ok_resp
with patch("packages.shared.ai_client.httpx.post", side_effect=fake_post):
result = client.image_generation("一只可爱的猫", size="1K")
assert result is not None
assert result["url"] == "https://cdn.example.com/i.png"
assert "/images/generations" in captured["url"]
assert captured["json"]["model"] == "doubao-seedream-5-0-pro-260628"
assert captured["json"]["size"] == "1K"
assert "image" not in captured["json"]
def test_image_to_image_single_ref_passed_as_string(self):
client = _make_client()
captured = {}
ok_resp = MagicMock(status_code=200)
ok_resp.json.return_value = {"data": [{"url": "https://cdn.example.com/out.png"}]}
ok_resp.raise_for_status = MagicMock()
ok_resp.text = ""
def fake_post(url, **kwargs):
captured["json"] = kwargs.get("json")
return ok_resp
with patch("packages.shared.ai_client.httpx.post", side_effect=fake_post):
client.image_generation("保持五官", reference_images=["https://img/x.jpg"])
assert captured["json"]["image"] == "https://img/x.jpg"
def test_image_to_image_multiple_refs_passed_as_list(self):
client = _make_client()
captured = {}
ok_resp = MagicMock(status_code=200)
ok_resp.json.return_value = {"data": [{"url": "https://cdn.example.com/out.png"}]}
ok_resp.raise_for_status = MagicMock()
def fake_post(url, **kwargs):
captured["json"] = kwargs.get("json")
return ok_resp
refs = [f"https://img/{i}.jpg" for i in range(3)]
with patch("packages.shared.ai_client.httpx.post", side_effect=fake_post):
client.image_generation("保持", reference_images=refs)
assert captured["json"]["image"] == refs
def test_400_sensitive_returns_portrait_intercept(self):
client = _make_client(max_retries=0)
bad_resp = MagicMock(status_code=400)
bad_resp.text = '{"error":{"code":"ContentRisk","message":"sensitive content detected"}}'
bad_resp.json.return_value = {"error": {"code": "ContentRisk"}}
bad_resp.raise_for_status.side_effect = httpx.HTTPStatusError("bad", request=MagicMock(), response=bad_resp)
with patch("packages.shared.ai_client.httpx.post", return_value=bad_resp):
assert client.image_generation("p", reference_images=["https://img/x.jpg"]) is None
err = client.get_last_image_error()
assert err["error_code"] == "portrait_intercept"
def test_500_retries_then_fails(self):
client = _make_client(max_retries=1)
bad_resp = MagicMock(status_code=500)
bad_resp.text = "internal error"
bad_resp.json.return_value = {"error": {"message": "internal"}}
bad_resp.raise_for_status.side_effect = httpx.HTTPStatusError("500", request=MagicMock(), response=bad_resp)
with (
patch("packages.shared.ai_client.httpx.post", return_value=bad_resp) as mp,
patch("packages.shared.ai_client.time.sleep", return_value=None),
):
assert client.image_generation("p") is None
assert mp.call_count == 2
err = client.get_last_image_error()
assert err["error_code"] == "network_error"
# ── 信任链集成单测 ────────────────────────────────────────────────
class TestTrustChainIntegration:
def test_with_reference_image_triggers_seedream_then_seedance_with_reference_image_role(self, tmp_path):
client = _make_client()
captured_calls = []
seedream_ok = MagicMock(status_code=200)
seedream_ok.json.return_value = {"data": [{"url": "https://ai.example.com/trusted.png"}]}
seedream_ok.raise_for_status = MagicMock()
seedream_ok.text = ""
task_ok = MagicMock(status_code=200)
task_ok.json.return_value = {"id": "t-trust"}
task_ok.raise_for_status = MagicMock()
task_ok.text = ""
poll_ok = MagicMock(status_code=200)
poll_ok.json.return_value = {"status": "succeeded", "content": {"video_url": "https://cdn.example.com/v.mp4"}}
poll_ok.raise_for_status = MagicMock()
class FakeStream:
def __init__(self):
self._c = [b"OK"]
self._it = iter(self._c)
def __enter__(self):
return self
def __exit__(self, *a):
return False
def raise_for_status(self):
return None
def iter_bytes(self, chunk_size=None):
return self._it
def fake_post(url, **kwargs):
captured_calls.append({"url": url, "json": kwargs.get("json")})
if "/images/generations" in url:
return seedream_ok
return task_ok
fake_uuid = MagicMock()
fake_uuid.hex = "00000001"
with (
patch("packages.shared.ai_client.httpx.post", side_effect=fake_post),
patch("packages.shared.ai_client.httpx.get", return_value=poll_ok),
patch("packages.shared.ai_client.httpx.stream", return_value=FakeStream()),
patch("packages.shared.ai_client.time.sleep", return_value=None),
patch("packages.shared.ai_client.time.time", side_effect=_fake_time(stable_calls=3)),
patch("packages.shared.ai_client.uuid.uuid4", return_value=fake_uuid),
patch("packages.shared.ai_client.get_shared_settings") as ms,
):
ms.return_value = MagicMock(
doubao_video_poll_interval=0,
doubao_video_timeout=60,
doubao_video_model="doubao-seedance-2-5-260628",
)
out = client.video_generation(
"人物在海边散步",
image_url="https://img/raw.jpg",
duration=5,
ratio="9:16",
resolution="720p",
output_dir=str(tmp_path),
)
assert out is not None
assert len(captured_calls) == 2
assert "/images/generations" in captured_calls[0]["url"]
assert captured_calls[0]["json"]["image"] == "https://img/raw.jpg"
seedance_payload = captured_calls[1]["json"]
content = seedance_payload["content"]
img_items = [c for c in content if c.get("type") == "image_url"]
assert len(img_items) == 1
assert img_items[0]["image_url"]["url"] == "https://ai.example.com/trusted.png"
assert img_items[0]["role"] == "reference_image"
assert seedance_payload["ratio"] == "9:16"
def test_no_reference_image_skips_seedream(self, tmp_path):
client = _make_client()
captured_calls = []
task_ok = MagicMock(status_code=200)
task_ok.json.return_value = {"id": "t-t2v"}
task_ok.raise_for_status = MagicMock()
poll_ok = MagicMock(status_code=200)
poll_ok.json.return_value = {"status": "succeeded", "content": {"video_url": "https://cdn.example.com/v.mp4"}}
poll_ok.raise_for_status = MagicMock()
class FakeStream:
def __init__(self):
self._it = iter([b"OK"])
def __enter__(self):
return self
def __exit__(self, *a):
return False
def raise_for_status(self):
return None
def iter_bytes(self, chunk_size=None):
return self._it
def fake_post(url, **kwargs):
captured_calls.append({"url": url, "json": kwargs.get("json")})
return task_ok
fake_uuid = MagicMock()
fake_uuid.hex = "00000002"
with (
patch("packages.shared.ai_client.httpx.post", side_effect=fake_post),
patch("packages.shared.ai_client.httpx.get", return_value=poll_ok),
patch("packages.shared.ai_client.httpx.stream", return_value=FakeStream()),
patch("packages.shared.ai_client.time.sleep", return_value=None),
patch("packages.shared.ai_client.time.time", side_effect=_fake_time(stable_calls=2)),
patch("packages.shared.ai_client.uuid.uuid4", return_value=fake_uuid),
patch("packages.shared.ai_client.get_shared_settings") as ms,
):
ms.return_value = MagicMock(
doubao_video_poll_interval=0,
doubao_video_timeout=60,
doubao_video_model="doubao-seedance-2-5-260628",
)
out = client.video_generation("海边日落", duration=5, ratio="9:16", output_dir=str(tmp_path))
assert out is not None
assert len(captured_calls) == 1
assert "/contents/generations/tasks" in captured_calls[0]["url"]
content = captured_calls[0]["json"]["content"]
assert all(c.get("type") != "image_url" for c in content)
def test_seedream_failure_falls_back_to_original_image(self, tmp_path):
client = _make_client(max_retries=0)
captured_calls = []
seedream_fail = MagicMock(status_code=400)
seedream_fail.text = '{"error":{"code":"QuotaExceeded","message":"quota"}}'
seedream_fail.json.return_value = {"error": {"code": "QuotaExceeded"}}
seedream_fail.raise_for_status.side_effect = httpx.HTTPStatusError(
"q", request=MagicMock(), response=seedream_fail
)
task_ok = MagicMock(status_code=200)
task_ok.json.return_value = {"id": "t-fb"}
task_ok.raise_for_status = MagicMock()
poll_ok = MagicMock(status_code=200)
poll_ok.json.return_value = {"status": "succeeded", "content": {"video_url": "https://cdn.example.com/v.mp4"}}
poll_ok.raise_for_status = MagicMock()
class FakeStream:
def __init__(self):
self._it = iter([b"OK"])
def __enter__(self):
return self
def __exit__(self, *a):
return False
def raise_for_status(self):
return None
def iter_bytes(self, chunk_size=None):
return self._it
def fake_post(url, **kwargs):
captured_calls.append({"url": url, "json": kwargs.get("json")})
if "/images/generations" in url:
return seedream_fail
return task_ok
fake_uuid = MagicMock()
fake_uuid.hex = "00000003"
with (
patch("packages.shared.ai_client.httpx.post", side_effect=fake_post),
patch("packages.shared.ai_client.httpx.get", return_value=poll_ok),
patch("packages.shared.ai_client.httpx.stream", return_value=FakeStream()),
patch("packages.shared.ai_client.time.sleep", return_value=None),
patch("packages.shared.ai_client.time.time", side_effect=_fake_time(stable_calls=2)),
patch("packages.shared.ai_client.uuid.uuid4", return_value=fake_uuid),
patch("packages.shared.ai_client.get_shared_settings") as ms,
):
ms.return_value = MagicMock(
doubao_video_poll_interval=0,
doubao_video_timeout=60,
doubao_video_model="doubao-seedance-2-5-260628",
)
out = client.video_generation(
"海边散步",
image_url="https://img/raw.jpg",
duration=5,
ratio="9:16",
output_dir=str(tmp_path),
)
assert out is not None
assert len(captured_calls) == 2
seedance_payload = captured_calls[1]["json"]
content = seedance_payload["content"]
img_items = [c for c in content if c.get("type") == "image_url"]
assert len(img_items) == 1
assert img_items[0]["image_url"]["url"] == "https://img/raw.jpg"
assert img_items[0]["role"] == "first_frame"
assert seedance_payload["ratio"] == "adaptive"
# ── image_generation 补充分支覆盖 ─────────────────────────────────
class TestImageGenerationBranches:
"""覆盖 image_generation 的错误分类/重试/结构异常等分支。"""
def test_401_returns_auth_error(self):
client = _make_client(max_retries=0)
r = MagicMock(status_code=401, text='{"error":{}}')
r.json.return_value = {"error": {}}
r.raise_for_status.side_effect = httpx.HTTPStatusError("a", request=MagicMock(), response=r)
with patch("packages.shared.ai_client.httpx.post", return_value=r):
assert client.image_generation("p") is None
assert client.get_last_image_error()["error_code"] == "auth_error"
def test_404_returns_model_not_found(self):
client = _make_client(max_retries=0)
r = MagicMock(status_code=404, text="not found")
r.json.return_value = {"error": {"message": "model not found"}}
r.raise_for_status.side_effect = httpx.HTTPStatusError("a", request=MagicMock(), response=r)
with patch("packages.shared.ai_client.httpx.post", return_value=r):
assert client.image_generation("p") is None
assert client.get_last_image_error()["error_code"] == "model_not_found"
def test_400_quota_returns_quota_exceeded(self):
client = _make_client(max_retries=0)
r = MagicMock(status_code=400, text="insufficient balance quota exceeded")
r.json.return_value = {"error": {"message": "quota"}}
r.raise_for_status.side_effect = httpx.HTTPStatusError("a", request=MagicMock(), response=r)
with patch("packages.shared.ai_client.httpx.post", return_value=r):
assert client.image_generation("p") is None
assert client.get_last_image_error()["error_code"] == "quota_exceeded"
def test_400_rate_limit_returns_rate_limit(self):
client = _make_client(max_retries=0)
r = MagicMock(status_code=400, text="too many requests, rate limit exceeded")
r.json.return_value = {"error": {"message": "rate"}}
r.raise_for_status.side_effect = httpx.HTTPStatusError("a", request=MagicMock(), response=r)
with patch("packages.shared.ai_client.httpx.post", return_value=r):
assert client.image_generation("p") is None
assert client.get_last_image_error()["error_code"] == "rate_limit"
def test_400_generic_returns_invalid_param(self):
client = _make_client(max_retries=0)
r = MagicMock(status_code=400, text="bad parameter size")
r.json.return_value = {"error": {"message": "bad"}}
r.raise_for_status.side_effect = httpx.HTTPStatusError("a", request=MagicMock(), response=r)
with patch("packages.shared.ai_client.httpx.post", return_value=r):
assert client.image_generation("p") is None
assert client.get_last_image_error()["error_code"] == "invalid_param"
def test_200_but_no_url_returns_none(self):
client = _make_client(max_retries=0)
r = MagicMock(status_code=200, text="")
r.json.return_value = {"data": [{"no_url": True}]} # 缺 url 字段
r.raise_for_status = MagicMock()
with patch("packages.shared.ai_client.httpx.post", return_value=r):
assert client.image_generation("p") is None
assert client.get_last_image_error()["error_code"] == "unknown"
def test_network_error_retries_then_fails(self):
client = _make_client(max_retries=1)
import httpcore
with (
patch("packages.shared.ai_client.httpx.post", side_effect=httpx.ConnectError("no network")),
patch("packages.shared.ai_client.time.sleep", return_value=None),
):
assert client.image_generation("p") is None
err = client.get_last_image_error()
assert err["error_code"] == "network_error"
def test_get_last_image_error_returns_copy(self):
client = _make_client()
client.last_image_error = {"error_code": "x"}
e1 = client.get_last_image_error()
e1["error_code"] = "mutated"
assert client.last_image_error["error_code"] == "x"
# ── 信任链分支覆盖 ────────────────────────────────────────
class TestTrustChainBranches:
def test_dashscope_provider_skips_trust_chain(self, tmp_path):
"""provider=dashscope 时不走信任链(Wan 模型由 dashscope_client 处理,在我们分支之前已经 return)。
这里测 doubao 分支:信任链默认触发,验证 DashScope 分发路径不受影响。"""
# 该测试实际覆盖 video_generation 入口的 dashscope 分发:缺 DASHSCOPE_API_KEY 时返回 auth_error
client = _make_client()
with (patch("packages.shared.ai_client.get_shared_settings") as ms,):
ms.return_value = MagicMock(
doubao_video_poll_interval=0,
doubao_video_timeout=1,
doubao_video_model="doubao-seedance-2-5-260628",
)
# DashScope 不可用时返回 auth_error(不是信任链相关错误)
result = client.video_generation(
"p",
output_dir=str(tmp_path),
model="wan-3.0",
image_url="https://img/x.jpg",
)
assert result is None
err = client.get_last_video_error()
# 不论是否走信任链,DashScope 无 key 时返回 auth_error
assert err["error_code"] == "auth_error"
def test_trust_chain_partial_seedream_success_falls_back(self, tmp_path):
"""多张参考图中第 2 张 Seedream 失败→整体回退原图直传。"""
client = _make_client(max_retries=0)
def make_seedream_fail():
r = MagicMock(status_code=500, text="err")
r.json.return_value = {"error": {}}
r.raise_for_status.side_effect = httpx.HTTPStatusError("e", request=MagicMock(), response=r)
return r
seedream_ok = MagicMock(status_code=200, text="")
seedream_ok.json.return_value = {"data": [{"url": "https://ai.example.com/a.png"}]}
seedream_ok.raise_for_status = MagicMock()
# 两张参考图(image_url + reference_images 各一张),Seedream 第 1 张 ok、第 2 张失败 → 回退
call_n = {"n": 0}
def fake_post(url, **kwargs):
if "/images/generations" in url:
call_n["n"] += 1
if call_n["n"] == 1:
return seedream_ok
return make_seedream_fail()
# Seedance create task(收到原图直传时会调用)
t = MagicMock(status_code=200, text="")
t.json.return_value = {"id": "t-partial"}
t.raise_for_status = MagicMock()
return t
poll_ok = MagicMock(status_code=200, text="")
poll_ok.json.return_value = {"status": "succeeded", "content": {"video_url": "https://cdn.example.com/v.mp4"}}
poll_ok.raise_for_status = MagicMock()
class FS:
def __init__(self):
self._it = iter([b"OK"])
def __enter__(self):
return self
def __exit__(self, *a):
return False
def raise_for_status(self):
return None
def iter_bytes(self, chunk_size=None):
return self._it
fake_uuid = MagicMock()
fake_uuid.hex = "00000004"
with (
patch("packages.shared.ai_client.httpx.post", side_effect=fake_post),
patch("packages.shared.ai_client.httpx.get", return_value=poll_ok),
patch("packages.shared.ai_client.httpx.stream", return_value=FS()),
patch("packages.shared.ai_client.time.sleep", return_value=None),
patch("packages.shared.ai_client.time.time", side_effect=_fake_time(stable_calls=50)),
patch("packages.shared.ai_client.uuid.uuid4", return_value=fake_uuid),
patch("packages.shared.ai_client.get_shared_settings") as ms,
):
ms.return_value = MagicMock(
doubao_video_poll_interval=0,
doubao_video_timeout=60,
doubao_video_model="doubao-seedance-2-5-260628",
)
out = client.video_generation(
"p",
image_url="https://img/a.jpg",
reference_images=["https://img/b.jpg"],
duration=5,
ratio="9:16",
output_dir=str(tmp_path),
)
assert out is not None
# 最终发给 Seedance 的图应是原始 https://img/a.jpg(回退),role=first_frame(因为 has_extra_refs=False 只有 1 张)
# 注意:回退后 ref_imgs 是原始 ["https://img/b.jpg"],所以 has_extra_refs=True,role=reference_image
# 断言最终 Seedance payload 里的 image_url 是原图(不是 AI 图)
def test_default_values_on_missing_settings(self):
"""getattr 兜底:settings 缺 image_timeout 字段时使用默认 120。"""
client = _make_client()
# 直接调用 image_generation,让它走一次完整流程(成功路径),验证 timeout 取值
ok = MagicMock(status_code=200, text="")
ok.json.return_value = {"data": [{"url": "https://ai.example.com/x.png"}]}
ok.raise_for_status = MagicMock()
captured_kwargs = {}
def fake_post(url, **kw):
captured_kwargs["timeout"] = kw.get("timeout")
return ok
with patch("packages.shared.ai_client.httpx.post", side_effect=fake_post):
r = client.image_generation("p", timeout=None) # 不传 timeout,走 self.image_timeout=120
assert r is not None
assert captured_kwargs["timeout"] == 120
def test_trust_chain_no_image_url_only_ref_imgs(self, tmp_path):
"""不传 image_url 仅传 reference_images 时走 trust chain 成功,ref_imgs 覆盖替换(line 537 else 分支)。"""
client = _make_client()
captured = []
seedream_ok = MagicMock(status_code=200, text="")
seedream_ok.json.return_value = {"data": [{"url": "https://ai.example.com/ref.png"}]}
seedream_ok.raise_for_status = MagicMock()
task_ok = MagicMock(status_code=200, text="")
task_ok.json.return_value = {"id": "t-refonly"}
task_ok.raise_for_status = MagicMock()
poll_ok = MagicMock(status_code=200, text="")
poll_ok.json.return_value = {"status": "succeeded", "content": {"video_url": "https://cdn.example.com/v.mp4"}}
poll_ok.raise_for_status = MagicMock()
class FS:
def __init__(self):
self._it = iter([b"OK"])
def __enter__(self):
return self
def __exit__(self, *a):
return False
def raise_for_status(self):
return None
def iter_bytes(self, chunk_size=None):
return self._it
def fake_post(url, **kw):
captured.append({"url": url, "json": kw.get("json")})
if "/images/generations" in url:
return seedream_ok
return task_ok
fu = MagicMock()
fu.hex = "0000000a"
with (
patch("packages.shared.ai_client.httpx.post", side_effect=fake_post),
patch("packages.shared.ai_client.httpx.get", return_value=poll_ok),
patch("packages.shared.ai_client.httpx.stream", return_value=FS()),
patch("packages.shared.ai_client.time.sleep", return_value=None),
patch("packages.shared.ai_client.time.time", side_effect=_fake_time(stable_calls=3)),
patch("packages.shared.ai_client.uuid.uuid4", return_value=fu),
patch("packages.shared.ai_client.get_shared_settings") as ms,
):
ms.return_value = MagicMock(
doubao_video_poll_interval=0,
doubao_video_timeout=60,
doubao_video_model="doubao-seedance-2-5-260628",
)
out = client.video_generation(
"人物散步",
reference_images=["https://img/portrait.jpg"],
duration=5,
ratio="9:16",
output_dir=str(tmp_path),
)
assert out is not None
# 第一次是 Seedream 成功,第二次是 Seedance 创建任务
assert len(captured) == 2
seedance_payload = captured[1]["json"]
content = seedance_payload["content"]
img_items = [c for c in content if c.get("type") == "image_url"]
assert len(img_items) == 1
# 不传 image_url,信任链产物放 ref_imgs,走 reference_image 模式(非 first_frame)
assert img_items[0]["image_url"]["url"] == "https://ai.example.com/ref.png"
assert img_items[0]["role"] == "reference_image"
# 因为没有 image_url,没有 text 也没有 extra_refs 之外的字段,应保留用户 ratio=9:16
assert seedance_payload.get("ratio") == "9:16"
def test_image_generation_generic_exception_retries_then_fails(self):
"""image_generation 遇到非 HTTPStatusError 的通用异常时走重试分支(lines 962-971),重试耗尽后返回 None。"""
client = _make_client(max_retries=1)
call_n = {"n": 0}
def fake_post(url, **kw):
call_n["n"] += 1
if call_n["n"] == 1:
raise RuntimeError("boiler exploded")
# 第二次调用返回成功,验证重试生效
ok = MagicMock(status_code=200, text="")
ok.json.return_value = {"data": [{"url": "https://ai.example.com/retry-ok.png"}]}
ok.raise_for_status = MagicMock()
return ok
with (
patch("packages.shared.ai_client.httpx.post", side_effect=fake_post),
patch("packages.shared.ai_client.time.sleep", return_value=None),
):
r = client.image_generation("test prompt")
assert r is not None
assert r["url"] == "https://ai.example.com/retry-ok.png"
assert call_n["n"] == 2
def test_image_generation_generic_exception_exhausts_retries(self):
"""通用异常重试耗尽后返回 None,并正确写入 last_image_error (lines 969-971 break 分支)。"""
client = _make_client(max_retries=1)
def fake_post(url, **kw):
raise RuntimeError("always fails")
with (
patch("packages.shared.ai_client.httpx.post", side_effect=fake_post),
patch("packages.shared.ai_client.time.sleep", return_value=None),
):
r = client.image_generation("test prompt")
assert r is None
err = client.last_image_error
assert err["error_code"] == "network_error"
assert "always fails" in err["detail"]
+23 -7
View File
@@ -17,8 +17,13 @@ def _make_client(**overrides):
client.base_url = overrides.get("base_url", "https://ark.cn-beijing.volces.com/api/v3")
client.model = "doubao-model"
client.vision_model = "doubao-vision"
client.embedding_model = "doubao-embedding"
client.image_model = overrides.get("image_model", "doubao-seedream-5-0-pro-260628")
client.image_timeout = overrides.get("image_timeout", 120)
client.timeout = overrides.get("timeout", 30)
client.max_retries = overrides.get("max_retries", 0)
client.last_video_error = {}
client.last_image_error = {}
return client
@@ -106,7 +111,7 @@ class TestVideoGenerationHappyPath:
)
out = client.video_generation(
prompt=" 镜头一 ",
image_url="https://img/x.jpg",
# 不传 image_url:纯文生视频,不触发信任链,post 调用数为 1(创建任务)
duration=5,
ratio="9:16",
resolution="720p",
@@ -572,12 +577,25 @@ class TestVideoGenerationLastError:
create_resp.raise_for_status.side_effect = httpx.HTTPStatusError(
"bad", request=MagicMock(), response=create_resp
)
# 信任链:Seedream 会先被调用来 AI 化;这里 mock Seedream 也失败,回退原图直传,
# 原图直传被 400 portrait 拦截,最终返回 portrait_intercept。
seedream_resp = MagicMock()
seedream_resp.status_code = 400
seedream_resp.text = '{"error":{"code":"ContentRisk","message":"sensitive"}}'
seedream_resp.json.return_value = {"error": {"code": "ContentRisk", "message": "sensitive"}}
seedream_resp.raise_for_status.side_effect = httpx.HTTPStatusError(
"bad", request=MagicMock(), response=seedream_resp
)
def fake_post(url, **kwargs):
# 第一次 POST 是 Seedream(/images/generations),返回 portrait 拦截
# 回退原图直传后第二次 POST 是 Seedance(/contents/generations/tasks),也返回 portrait 拦截
return create_resp
with (
patch("packages.shared.ai_client.httpx.post", return_value=create_resp),
patch("packages.shared.ai_client.httpx.post", side_effect=fake_post),
patch("packages.shared.ai_client.time.sleep", return_value=None),
patch("packages.shared.ai_client.get_shared_settings") as mock_s,
# jimeng 未配置,fallback 后仍返回 portrait_intercept(提示用户需要配置即梦)
patch("packages.shared.jimeng_client.get_jimeng_client", return_value=None),
):
mock_s.return_value = MagicMock(
doubao_video_poll_interval=0, doubao_video_timeout=1, doubao_video_model="seedance"
@@ -586,9 +604,7 @@ class TestVideoGenerationLastError:
assert result is None
err = client.get_last_video_error()
assert err["error_code"] == "portrait_intercept"
# 即梦兜底未启用时提示包含"真人照片"/"即梦"等关键字
assert "真人" in err["user_message"] or "肖像" in err["user_message"] or "即梦" in err["user_message"]
# 方舟本身 status_code=400(最后一个错误可能被即梦兜底覆盖,但 error_code 不变)
assert "真人" in err["user_message"] or "肖像" in err["user_message"] or "审核" in err["user_message"]
assert err["status_code"] in (0, 400)
def test_create_401_returns_auth_error(self, tmp_path):
-376
View File
@@ -1,376 +0,0 @@
"""tests for packages/shared/jimeng_client.py (#2169 即梦 i2v 客户端)."""
from __future__ import annotations
import json
import os
from unittest.mock import MagicMock, mock_open, patch
import pytest
@pytest.fixture(autouse=True)
def reset_singleton():
import packages.shared.jimeng_client as j
j._JIMENG_CLIENT_SINGLETON = None
yield
j._JIMENG_CLIENT_SINGLETON = None
def _make_settings(ak="test-ak", sk="test-sk", req_key="jimeng_i2v_first_v30", timeout=60, poll_interval=2):
return MagicMock(
jimeng_ak=ak,
jimeng_sk=sk,
jimeng_base_url="https://visual.volcengineapi.com",
jimeng_req_key=req_key,
jimeng_video_timeout=timeout,
jimeng_video_poll_interval=poll_interval,
)
# ── V4 签名单元测试 ──────────────────────────────────────────────────
class TestV4Signature:
def test_sign_returns_required_headers(self):
from packages.shared.jimeng_client import _sign_v4
headers = _sign_v4(
ak="AK_TEST",
sk="SK_TEST",
method="POST",
host="visual.volcengineapi.com",
query={"Action": "CVSync2AsyncSubmitTask", "Version": "2022-08-31"},
body_bytes=b'{"req_key":"jimeng_i2v_first_v30"}',
x_date="20260101T120000Z",
)
assert headers["Content-Type"] == "application/json"
assert headers["Host"] == "visual.volcengineapi.com"
assert headers["X-Date"] == "20260101T120000Z"
assert "X-Content-Sha256" in headers
assert headers["Authorization"].startswith("HMAC-SHA256 Credential=AK_TEST/20260101/cn-north-1/cv/request")
assert "SignedHeaders=content-type;host;x-content-sha256;x-date" in headers["Authorization"]
assert "Signature=" in headers["Authorization"]
# 签名是 64 字符 hex
sig = headers["Authorization"].split("Signature=")[-1]
assert len(sig) == 64
assert all(c in "0123456789abcdef" for c in sig)
def test_sign_deterministic(self):
"""相同输入必须产生相同签名(幂等)。"""
from packages.shared.jimeng_client import _sign_v4
kwargs = dict(
ak="AK",
sk="SK",
method="POST",
host="h",
query={"A": "1", "B": "2"},
body_bytes=b"{}",
x_date="20260101T000000Z",
)
h1 = _sign_v4(**kwargs)
h2 = _sign_v4(**kwargs)
assert h1["Authorization"] == h2["Authorization"]
assert h1["X-Content-Sha256"] == h2["X-Content-Sha256"]
def test_sign_different_body_different_sig(self):
from packages.shared.jimeng_client import _sign_v4
base = dict(ak="AK", sk="SK", method="POST", host="h", query={}, x_date="20260101T000000Z")
h1 = _sign_v4(body_bytes=b"a", **base)
h2 = _sign_v4(body_bytes=b"b", **base)
assert h1["Authorization"] != h2["Authorization"]
def test_payload_sha256_matches(self):
import hashlib
from packages.shared.jimeng_client import _sign_v4
body = b'{"prompt":"hello"}'
h = _sign_v4("ak", "sk", "POST", "h", {}, body, "20260101T000000Z")
expected = hashlib.sha256(body).hexdigest()
assert h["X-Content-Sha256"] == expected
def test_norm_query_sorted_and_encoded(self):
from packages.shared.jimeng_client import _norm_query
q = _norm_query({"B": "2", "A": "1", "C": "a b"})
# key 排序 + 空格→%20
assert q.startswith("A=1")
assert "B=2" in q
assert "C=a%20b" in q
# ── 可用性 / 单例 ─────────────────────────────────────────────────────
class TestAvailability:
def test_unavailable_without_ak_sk(self):
from packages.shared.jimeng_client import JimengClient, get_jimeng_client
with patch("packages.shared.jimeng_client.get_shared_settings") as ms:
ms.return_value = _make_settings(ak="", sk="")
# 重置单例
import packages.shared.jimeng_client as j
j._JIMENG_CLIENT_SINGLETON = None
assert get_jimeng_client() is None
def test_available_with_ak_sk(self):
from packages.shared.jimeng_client import get_jimeng_client
with patch("packages.shared.jimeng_client.get_shared_settings") as ms:
ms.return_value = _make_settings()
import packages.shared.jimeng_client as j
j._JIMENG_CLIENT_SINGLETON = None
c = get_jimeng_client()
assert c is not None
assert c.is_available is True
assert c.req_key == "jimeng_i2v_first_v30"
# ── 错误分类 ────────────────────────────────────────────────────────
class TestClassifyError:
def test_50411_is_portrait_intercept_non_retryable(self):
from packages.shared.jimeng_client import _classify_jimeng_error
code, msg, retry = _classify_jimeng_error(200, '{"code":50411,"message":"Pre Img Risk"}', 50411)
assert code == "portrait_intercept"
assert retry is False
def test_50429_is_rate_limit_retryable(self):
from packages.shared.jimeng_client import _classify_jimeng_error
code, msg, retry = _classify_jimeng_error(200, "", 50429)
assert code == "rate_limit"
assert retry is True
def test_50430_is_rate_limit(self):
from packages.shared.jimeng_client import _classify_jimeng_error
code, _, _ = _classify_jimeng_error(200, "", 50430)
assert code == "rate_limit"
def test_50500_is_network_error_retryable(self):
from packages.shared.jimeng_client import _classify_jimeng_error
code, _, retry = _classify_jimeng_error(200, "", 50500)
assert code == "network_error"
assert retry is True
def test_50412_is_invalid_param_non_retryable(self):
from packages.shared.jimeng_client import _classify_jimeng_error
code, _, retry = _classify_jimeng_error(200, "", 50412)
assert code == "invalid_param"
assert retry is False
def test_401_auth(self):
from packages.shared.jimeng_client import _classify_jimeng_error
code, msg, retry = _classify_jimeng_error(401, "auth fail", None)
assert code == "auth_error"
assert retry is False
def test_400_text_audit(self):
from packages.shared.jimeng_client import _classify_jimeng_error
code, _, _ = _classify_jimeng_error(400, "text error", None)
assert code == "invalid_param"
# ── video_generation 主流程 ──────────────────────────────────────────
class TestVideoGenerationHappyPath:
def test_missing_ak_returns_none(self):
from packages.shared.jimeng_client import JimengClient
with patch("packages.shared.jimeng_client.get_shared_settings") as ms:
ms.return_value = _make_settings(ak="", sk="")
c = JimengClient()
assert c.video_generation("hi", image_url="http://x/y.jpg") is None
err = c.last_video_error
assert err["error_code"] == "auth_error"
def test_empty_prompt_returns_none(self):
from packages.shared.jimeng_client import JimengClient
with patch("packages.shared.jimeng_client.get_shared_settings") as ms:
ms.return_value = _make_settings()
c = JimengClient()
assert c.video_generation(" ", image_url="http://x/y.jpg") is None
assert c.last_video_error["error_code"] == "invalid_param"
def test_empty_image_url_returns_none(self):
from packages.shared.jimeng_client import JimengClient
with patch("packages.shared.jimeng_client.get_shared_settings") as ms:
ms.return_value = _make_settings()
c = JimengClient()
assert c.video_generation("prompt", image_url="") is None
assert c.last_video_error["error_code"] == "invalid_param"
def test_duration_5s_frames_121(self):
"""5s → frames=121,10s→frames=241,>10s 截断到10s。"""
from packages.shared.jimeng_client import JimengClient
with patch("packages.shared.jimeng_client.get_shared_settings") as ms:
ms.return_value = _make_settings(timeout=1, poll_interval=0)
c = JimengClient()
captured = {}
def fake_submit(prompt, image_url, frames, seed=-1):
captured["frames"] = frames
return "task-xyz"
def fake_poll(tid):
captured["tid"] = tid
return "http://example.com/v.mp4"
def fake_download(url, out_dir, tid):
captured["url"] = url
return "/tmp/fake.mp4"
# 构造一个假文件
os.makedirs("/tmp", exist_ok=True)
with open("/tmp/fake.mp4", "wb") as f:
f.write(b"x" * 2048)
with (
patch.object(c, "_submit_task", side_effect=fake_submit),
patch.object(c, "_poll_result", side_effect=fake_poll),
patch.object(c, "_download_video", side_effect=fake_download),
):
r = c.video_generation("test", image_url="http://x/y.jpg", duration=5, output_dir="/tmp")
assert r is not None
assert captured["frames"] == 121
assert r["usage"]["duration_seconds"] == 5
assert r["usage"]["billing_mode"] == "per_second"
assert r["usage"]["req_key"] == "jimeng_i2v_first_v30"
def test_duration_10s_frames_241(self):
from packages.shared.jimeng_client import JimengClient
with patch("packages.shared.jimeng_client.get_shared_settings") as ms:
ms.return_value = _make_settings(timeout=1, poll_interval=0)
c = JimengClient()
captured = {}
def fake_submit(prompt, image_url, frames, seed=-1):
captured["frames"] = frames
return "tid"
def fake_poll(tid):
return "http://x/v.mp4"
def fake_download(url, out_dir, tid):
with open("/tmp/fake2.mp4", "wb") as f:
f.write(b"x" * 2048)
return "/tmp/fake2.mp4"
with (
patch.object(c, "_submit_task", side_effect=fake_submit),
patch.object(c, "_poll_result", side_effect=fake_poll),
patch.object(c, "_download_video", side_effect=fake_download),
):
r = c.video_generation("hi", image_url="http://x/y.jpg", duration=10, output_dir="/tmp")
assert captured["frames"] == 241
assert r["usage"]["duration_seconds"] == 10
def test_duration_over_10s_truncates_to_10s(self):
from packages.shared.jimeng_client import JimengClient
with patch("packages.shared.jimeng_client.get_shared_settings") as ms:
ms.return_value = _make_settings(timeout=1, poll_interval=0)
c = JimengClient()
captured = {}
def fake_submit(prompt, image_url, frames, seed=-1):
captured["frames"] = frames
return "tid"
def fake_poll(tid):
return "http://x/v.mp4"
def fake_download(url, out_dir, tid):
with open("/tmp/fake3.mp4", "wb") as f:
f.write(b"x" * 2048)
return "/tmp/fake3.mp4"
with (
patch.object(c, "_submit_task", side_effect=fake_submit),
patch.object(c, "_poll_result", side_effect=fake_poll),
patch.object(c, "_download_video", side_effect=fake_download),
):
r = c.video_generation("hi", image_url="http://x/y.jpg", duration=30, output_dir="/tmp")
assert captured["frames"] == 241
assert r["usage"]["duration_seconds"] == 10
class TestSubmitTaskErrors:
def test_submit_50411_writes_portrait_intercept(self):
from packages.shared.jimeng_client import JimengClient
with patch("packages.shared.jimeng_client.get_shared_settings") as ms:
ms.return_value = _make_settings()
c = JimengClient()
fake_resp = MagicMock(status_code=200, text='{"code":50411,"message":"Pre Img Risk Not Pass"}')
fake_resp.json.return_value = {"code": 50411, "message": "Pre Img Risk Not Pass"}
with patch("packages.shared.jimeng_client.httpx.request", return_value=fake_resp):
tid = c._submit_task("p", "http://x/y.jpg", frames=121)
assert tid is None
assert c.last_video_error["error_code"] == "portrait_intercept"
def test_submit_returns_task_id(self):
from packages.shared.jimeng_client import JimengClient
with patch("packages.shared.jimeng_client.get_shared_settings") as ms:
ms.return_value = _make_settings()
c = JimengClient()
fake_resp = MagicMock(status_code=200, text='{"code":10000,"data":{"task_id":"abc"}}')
fake_resp.json.return_value = {"code": 10000, "data": {"task_id": "abc"}}
with patch("packages.shared.jimeng_client.httpx.request", return_value=fake_resp):
tid = c._submit_task("p", "http://x/y.jpg", frames=121)
assert tid == "abc"
class TestPollResult:
def test_poll_done_returns_video_url(self):
from packages.shared.jimeng_client import JimengClient
with patch("packages.shared.jimeng_client.get_shared_settings") as ms:
ms.return_value = _make_settings(timeout=10, poll_interval=0)
c = JimengClient()
done_resp = MagicMock(status_code=200)
done_resp.json.return_value = {"code": 10000, "data": {"status": "done", "video_url": "http://x/v.mp4"}}
with (
patch("packages.shared.jimeng_client.httpx.request", return_value=done_resp),
patch("packages.shared.jimeng_client.time.sleep"),
):
url = c._poll_result("abc")
assert url == "http://x/v.mp4"
def test_poll_timeout_returns_none(self):
from packages.shared.jimeng_client import JimengClient
with patch("packages.shared.jimeng_client.get_shared_settings") as ms:
ms.return_value = _make_settings(timeout=1, poll_interval=0)
c = JimengClient()
queue_resp = MagicMock(status_code=200)
queue_resp.json.return_value = {"code": 10000, "data": {"status": "in_queue"}}
# time.time 会被调用,模拟超时
with (
patch("packages.shared.jimeng_client.httpx.request", return_value=queue_resp),
patch("packages.shared.jimeng_client.time.sleep"),
):
url = c._poll_result("abc")
assert url is None
assert c.last_video_error["error_code"] == "network_error"
assert "超时" in c.last_video_error["user_message"]