Compare commits
1 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| e5cdb93c78 |
+4
-13
@@ -211,24 +211,15 @@ COSYVOICE_CLONE_MODEL=voice-enrollment
|
||||
# 用于 AI 文案生成、智能剪辑等需要大模型能力的场景
|
||||
|
||||
DOUBAO_API_KEY=your-doubao-api-key
|
||||
DOUBAO_MODEL=doubao-seed-2-1-pro-260915
|
||||
DOUBAO_FAST_MODEL=doubao-seed-2-1-lite-260915
|
||||
DOUBAO_MODEL=doubao-seed-1-6-250615
|
||||
DOUBAO_FAST_MODEL=doubao-1-5-pro-32k-250115
|
||||
DOUBAO_BASE_URL=https://ark.cn-beijing.volces.com/api/v3
|
||||
DOUBAO_TIMEOUT=30
|
||||
DOUBAO_MAX_RETRIES=2
|
||||
# 视觉模型:pro 精度高,lite 速度快(viral-video 商品识别默认用 lite 提速)
|
||||
DOUBAO_VISION_MODEL=doubao-seed-2-1-pro-260915
|
||||
DOUBAO_VISION_LITE_MODEL=doubao-seed-2-1-lite-260915
|
||||
DOUBAO_VISION_MODEL=doubao-1-5-vision-pro-250328
|
||||
DOUBAO_VISION_LITE_MODEL=doubao-1-5-vision-lite-250315
|
||||
DOUBAO_VISION_USE_LITE=true
|
||||
# Embedding 向量化模型
|
||||
DOUBAO_EMBEDDING_MODEL=doubao-embedding-vision-251215
|
||||
# 视频模型(Seedance 2.5,统一走方舟;真人参考图通过信任链自动 AI 化)
|
||||
DOUBAO_VIDEO_MODEL=doubao-seedance-2-5-260628
|
||||
DOUBAO_VIDEO_TIMEOUT=480
|
||||
DOUBAO_VIDEO_POLL_INTERVAL=10
|
||||
# 图片模型(Seedream 5.0 Pro,用于信任链真人 AI 化 + 文生图)
|
||||
DOUBAO_IMAGE_MODEL=doubao-seedream-5-0-pro-260628
|
||||
DOUBAO_IMAGE_TIMEOUT=120
|
||||
|
||||
# ==================== 积分/会员系统 (#1895) ====================
|
||||
# 积分系统总开关:默认 false(暂停积分系统)。
|
||||
|
||||
@@ -145,22 +145,19 @@ def get_rules(
|
||||
def get_packages(
|
||||
current_user: AuthenticatedUser = Depends(get_current_user),
|
||||
):
|
||||
"""查询可购买的积分包列表(读管理后台 credit_packages 表真实数据)。
|
||||
|
||||
仅返回 is_active=true;后台改价/启停后最多 30 秒生效。
|
||||
"""
|
||||
from packages.application.catalog.admin_catalog import get_points_packages
|
||||
|
||||
packages = [
|
||||
PointsPackageItem(
|
||||
code=row["code"],
|
||||
name=row["name"],
|
||||
points=row["points"],
|
||||
price_cents=row["price_cents"],
|
||||
unit_price=row["unit_price"],
|
||||
"""查询可购买的积分包列表。"""
|
||||
packages = []
|
||||
for code, pkg in POINTS_PACKAGES.items():
|
||||
unit_price = f"¥{pkg['price_cents'] / 100 / pkg['points']:.3f}/积分"
|
||||
packages.append(
|
||||
PointsPackageItem(
|
||||
code=code,
|
||||
name=pkg["name"],
|
||||
points=pkg["points"],
|
||||
price_cents=pkg["price_cents"],
|
||||
unit_price=unit_price,
|
||||
)
|
||||
)
|
||||
for row in get_points_packages()
|
||||
]
|
||||
mt = _member_type(current_user)
|
||||
discount = MEMBER_DISCOUNT.get(mt) if mt else None
|
||||
return PointsPackagesResponse(packages=packages, user_discount=discount)
|
||||
|
||||
@@ -86,13 +86,33 @@ async def get_current_subscription(
|
||||
def list_membership_plans(
|
||||
current_user: AuthenticatedUser = Depends(get_current_user),
|
||||
) -> dict[str, list[dict[str, Any]]]:
|
||||
"""查询可购买的会员套餐(读管理后台 plans 表真实数据)。
|
||||
"""查询所有会员档位(供前端会员购买页展示)。
|
||||
|
||||
仅返回 is_enabled=true 的套餐;后台启停/改价后最多 30 秒生效。
|
||||
返回 points 积分体系下的会员档位(月卡/季卡/年卡),含价格、时长、积分折扣等信息。
|
||||
"""
|
||||
from packages.application.catalog.admin_catalog import get_membership_plans
|
||||
from packages.domain.points_rules import MEMBER_DISCOUNT, MEMBERSHIP_PRICES
|
||||
|
||||
return {"plans": get_membership_plans()}
|
||||
plans: list[dict[str, Any]] = []
|
||||
for plan_id, info in MEMBERSHIP_PRICES.items():
|
||||
days = info["duration_days"]
|
||||
monthly_cents = round(info["price_cents"] * 30 / days)
|
||||
features: dict[str, Any] = {"max_resolution": "1080p"}
|
||||
if plan_id == MembershipType.MONTHLY:
|
||||
features.update({"free_clips_daily": 2})
|
||||
elif plan_id == MembershipType.QUARTERLY:
|
||||
features.update({"free_clips_daily": 5})
|
||||
elif plan_id == MembershipType.YEARLY:
|
||||
features.update({"free_clips_daily": "unlimited"})
|
||||
plans.append({
|
||||
"plan_id": plan_id,
|
||||
"name": info["name"],
|
||||
"price_cents": info["price_cents"],
|
||||
"monthly_price_cents": monthly_cents,
|
||||
"duration_days": days,
|
||||
"points_discount": MEMBER_DISCOUNT.get(plan_id, 1.0),
|
||||
"features": features,
|
||||
})
|
||||
return {"plans": plans}
|
||||
|
||||
|
||||
@router.get("/billing-records", response_model=list[BillingRecord])
|
||||
|
||||
@@ -80,8 +80,6 @@ const AiAvatarPage: React.FC = () => {
|
||||
const [finalizeLoading, setFinalizeLoading] = useState(false)
|
||||
|
||||
/* ── 对口型轮询 ── */
|
||||
/** 对口型轮询总时长上限(10分钟):超过后停止轮询并提示去历史记录查看 */
|
||||
const LIPSYNC_POLL_MAX_MS = 10 * 60 * 1000
|
||||
const lipsyncTimerRef = useRef<ReturnType<typeof setInterval> | null>(null)
|
||||
/* ── 渲染进度轮询 ── */
|
||||
const renderTimerRef = useRef<ReturnType<typeof setInterval> | null>(null)
|
||||
@@ -272,24 +270,7 @@ const AiAvatarPage: React.FC = () => {
|
||||
// 如果是预合成模式,后端会同步把状态置为 submitted(甚至可能已返回 running),
|
||||
// 但仍需轮询等 completed
|
||||
if (lipsyncTimerRef.current) clearInterval(lipsyncTimerRef.current)
|
||||
// 轮询间隔 5 秒;单请求超时 5 分钟(见 api/aiAvatar.ts);总轮询上限 10 分钟
|
||||
// 单次请求失败/超时不中断轮询,继续下一轮;超过总上限后停止并提示用户去历史记录查看
|
||||
lipsyncTimerRef.current = setInterval(async () => {
|
||||
// 总时长保护:超过 10 分钟停止轮询
|
||||
if (Date.now() - lipsyncStartAtRef.current > LIPSYNC_POLL_MAX_MS) {
|
||||
if (lipsyncTimerRef.current) {
|
||||
clearInterval(lipsyncTimerRef.current)
|
||||
lipsyncTimerRef.current = null
|
||||
}
|
||||
if (lipsyncTickRef.current) {
|
||||
clearInterval(lipsyncTickRef.current)
|
||||
lipsyncTickRef.current = null
|
||||
}
|
||||
setLipsyncStatus("failed")
|
||||
setLipsyncErrorMessage("渲染时间较长,请稍后在历史记录中查看")
|
||||
message.warning("对口型渲染时间较长,已停止自动刷新,请稍后在历史记录中查看")
|
||||
return
|
||||
}
|
||||
try {
|
||||
const updated = await getLipsyncJob(job.id)
|
||||
state.setLipsyncJob(updated)
|
||||
@@ -315,10 +296,9 @@ const AiAvatarPage: React.FC = () => {
|
||||
setLipsyncErrorMessage(updated.error_message || "对口型生成失败")
|
||||
}
|
||||
} catch (err) {
|
||||
// 单次轮询失败(含 timeout):不中断轮询,打印日志后等下一轮
|
||||
console.warn("[对口型] 轮询请求失败,将继续下一轮:", err)
|
||||
console.error("[对口型] 轮询错误:", err)
|
||||
}
|
||||
}, 5000)
|
||||
}, 3000)
|
||||
} catch (err) {
|
||||
console.error("[对口型] 创建失败:", {
|
||||
status: (err as { response?: { status?: number } })?.response?.status,
|
||||
|
||||
@@ -72,8 +72,7 @@ export const previewTts = async (data: {
|
||||
}
|
||||
|
||||
export const getLipsyncJob = async (id: string): Promise<LipsyncJob> => {
|
||||
// MuseTalk 渲染 8s 视频约 54s + 排队时间,给足 5 分钟超时避免单次轮询 AxiosError 中断
|
||||
const response = await apiClient.get<LipsyncJob>(`/lipsync/jobs/${id}`, { timeout: 300_000 })
|
||||
const response = await apiClient.get<LipsyncJob>(`/lipsync/jobs/${id}`, { timeout: 60000 })
|
||||
return response.data
|
||||
}
|
||||
|
||||
@@ -92,10 +91,7 @@ export const submitRender = async (data: {
|
||||
}
|
||||
|
||||
export const getRenderJob = async (jobId: string): Promise<RenderJob> => {
|
||||
// 渲染链路(对口型+B-roll+标题+合成+上传)耗时较长,给足 5 分钟超时
|
||||
const response = await apiClient.get<RenderJob>(`/ai-avatar/render/${jobId}`, {
|
||||
timeout: 300_000,
|
||||
})
|
||||
const response = await apiClient.get<RenderJob>(`/ai-avatar/render/${jobId}`, { timeout: 60000 })
|
||||
return response.data
|
||||
}
|
||||
|
||||
|
||||
@@ -647,7 +647,7 @@ def _build_products_summary(image_analysis: dict) -> str:
|
||||
# 优先 VLM 生成的 summary 段(自然语言,给编导模型看效果最好)
|
||||
summary = (p.get("summary") or "").strip()
|
||||
if summary and len(summary) >= 30:
|
||||
lines.append(f"- 图{i + 1} {name}:{summary}")
|
||||
lines.append(f"- 图{i+1} {name}:{summary}")
|
||||
continue
|
||||
# 结构化字段兜底
|
||||
brand = p.get("brand") or ""
|
||||
@@ -669,7 +669,7 @@ def _build_products_summary(image_analysis: dict) -> str:
|
||||
feats = p.get("key_features") or p.get("features") or []
|
||||
sellings = p.get("selling_points") or []
|
||||
scenes = p.get("suitable_scenes") or []
|
||||
parts = [f"图{i + 1} {name}"]
|
||||
parts = [f"图{i+1} {name}"]
|
||||
if brand and brand not in ("未知", "无法判断"):
|
||||
parts.append(f"品牌={brand}")
|
||||
if cat and cat not in ("无法判断", "非产品图"):
|
||||
@@ -802,7 +802,7 @@ def _validate_and_normalize_script(raw, job: ViralVideoJob) -> dict:
|
||||
continue
|
||||
shots.append(
|
||||
{
|
||||
"time_range": str(s.get("time_range") or f"{i * 3}-{(i + 1) * 3}秒"),
|
||||
"time_range": str(s.get("time_range") or f"{i*3}-{(i+1)*3}秒"),
|
||||
"shot_type_angle_movement": str(s.get("shot_type_angle_movement") or "中景平视,固定镜头"),
|
||||
"scene_and_dialogue": str(s.get("scene_and_dialogue") or ""),
|
||||
"action_details": str(s.get("action_details") or ""),
|
||||
@@ -878,8 +878,8 @@ def _step_script_generation(job: ViralVideoJob, intent: dict, image_analysis: di
|
||||
style_hint = "无"
|
||||
if isinstance(job.style_guide, dict):
|
||||
style_hint = (
|
||||
f"节奏{job.style_guide.get('cut_speed', '')}、转场{job.style_guide.get('transition', '')}、"
|
||||
f"色调{job.style_guide.get('color_grade', '')}、能量{job.style_guide.get('energy', '')}"
|
||||
f"节奏{job.style_guide.get('cut_speed','')}、转场{job.style_guide.get('transition','')}、"
|
||||
f"色调{job.style_guide.get('color_grade','')}、能量{job.style_guide.get('energy','')}"
|
||||
)
|
||||
|
||||
dur = max(5, min(30, int(getattr(job, "duration", 15) or 15)))
|
||||
@@ -1103,14 +1103,14 @@ def _assemble_seedance_prompt(copy_result: dict, job: ViralVideoJob) -> str:
|
||||
ab = s.get("audio_bgm", "")
|
||||
t = s.get("transition", "")
|
||||
ref = s.get("reference_image_index")
|
||||
lines.append(f"- 镜头{i + 1}({tr}):")
|
||||
lines.append(f"- 镜头{i+1}({tr}):")
|
||||
lines.append(f" 景别/运镜:{cam}")
|
||||
lines.append(f" 画面与对白:{sd}")
|
||||
lines.append(f" 动作细节:{act}")
|
||||
lines.append(f" 音效/BGM:{ab}")
|
||||
lines.append(f" 转场:{t}")
|
||||
if ref is not None and isinstance(ref, int):
|
||||
lines.append(f" 参考图片:第{ref + 1}张产品图")
|
||||
lines.append(f" 参考图片:第{ref+1}张产品图")
|
||||
lines.append("")
|
||||
lines.append("【硬性约束】")
|
||||
for c in hc:
|
||||
@@ -1162,7 +1162,6 @@ def _step_render(job: ViralVideoJob, copy_result: dict, tts_audio_url: str | Non
|
||||
)
|
||||
logger.info("[爆款视频] Seedance prompt (前300字): %s", prompt[:300])
|
||||
|
||||
# 第一次调用:带参考图/首帧/音频/参考视频
|
||||
result = call_video_generation(
|
||||
prompt=prompt,
|
||||
image_url=first_image,
|
||||
@@ -1176,42 +1175,15 @@ def _step_render(job: ViralVideoJob, copy_result: dict, tts_audio_url: str | Non
|
||||
reference_audios=ref_audios,
|
||||
reference_videos=ref_videos,
|
||||
)
|
||||
|
||||
# #2170: 真人/肖像拦截由 ai_client 内部信任链自动处理(Seedream AI 化后再调 Seedance);
|
||||
# 非拦截类错误直接抛错给上层展示用户友好提示。
|
||||
def _check_and_reraise(result):
|
||||
if result and isinstance(result, dict):
|
||||
return result
|
||||
from packages.shared.ai_service import get_last_video_error
|
||||
|
||||
err = get_last_video_error() or {}
|
||||
user_msg = err.get("user_message") or ""
|
||||
detail = err.get("detail") or ""
|
||||
err_code = err.get("error_code") or "unknown"
|
||||
status_code = err.get("status_code", 0)
|
||||
err_provider = err.get("provider") or _mcfg.get("provider", "doubao")
|
||||
err_msg = user_msg or f"视频生成失败({err_provider} status={status_code} code={err_code})"
|
||||
logger.error(
|
||||
"[爆款视频] 视频生成失败: provider=%s model=%s code=%s status=%s user_msg=%s detail=%s",
|
||||
err_provider,
|
||||
model or "default",
|
||||
err_code,
|
||||
status_code,
|
||||
user_msg,
|
||||
(detail or "")[:500],
|
||||
)
|
||||
raise RuntimeError(err_msg)
|
||||
|
||||
if not result or not isinstance(result, dict):
|
||||
_check_and_reraise(result)
|
||||
raise RuntimeError("Seedance 视频生成失败:返回为空")
|
||||
video_path = result.get("video_path") or ""
|
||||
usage = result.get("usage")
|
||||
if not video_path or not Path(video_path).exists() or Path(video_path).stat().st_size == 0:
|
||||
raise RuntimeError("视频生成失败:返回空文件或路径不存在")
|
||||
# #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)
|
||||
raise RuntimeError("Seedance 视频生成失败:返回空文件或路径不存在")
|
||||
logger.info(
|
||||
"[爆款视频] Seedance 单次生成完成: %s size=%d usage=%s", video_path, Path(video_path).stat().st_size, usage
|
||||
)
|
||||
return str(video_path), (usage if isinstance(usage, dict) else None)
|
||||
|
||||
|
||||
|
||||
@@ -39,9 +39,6 @@ class SQLAlchemyUserRepository(UserRepository):
|
||||
model.phone_verified = user.phone_verified
|
||||
model.binding_completed_at = user.binding_completed_at
|
||||
model.profile_completed = user.profile_completed
|
||||
model.is_member = user.is_member
|
||||
model.member_type = user.member_type
|
||||
model.member_expires_at = user.member_expires_at
|
||||
model.created_at = user.created_at
|
||||
|
||||
self.session.commit()
|
||||
@@ -118,8 +115,5 @@ class SQLAlchemyUserRepository(UserRepository):
|
||||
phone_verified=model.phone_verified or False,
|
||||
binding_completed_at=model.binding_completed_at,
|
||||
profile_completed=model.profile_completed if model.profile_completed is not None else True,
|
||||
is_member=model.is_member if model.is_member is not None else False,
|
||||
member_type=model.member_type,
|
||||
member_expires_at=model.member_expires_at,
|
||||
created_at=model.created_at,
|
||||
)
|
||||
|
||||
@@ -1 +0,0 @@
|
||||
"""应用层:对外展示目录(套餐/积分包)。"""
|
||||
@@ -1,152 +0,0 @@
|
||||
"""读取管理后台配置的会员套餐 / 积分充值包(共享库真实数据)。
|
||||
|
||||
替代旧的硬编码 MEMBERSHIP_PRICES / POINTS_PACKAGES。
|
||||
短 TTL 缓存(30 秒),后台改价/启停后用户端最多 30 秒可见。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import threading
|
||||
import time
|
||||
from typing import Any
|
||||
|
||||
_CACHE_TTL = 30.0
|
||||
_lock = threading.Lock()
|
||||
_cache: dict[str, tuple[float, Any]] = {}
|
||||
|
||||
_QUOTA_LABELS = {
|
||||
"4k": "4K 超清分辨率",
|
||||
"batch_render": "批量渲染",
|
||||
"priority_queue": "优先处理队列",
|
||||
"ai_matting": "AI 智能抠像",
|
||||
"remove_watermark": "去水印",
|
||||
}
|
||||
|
||||
|
||||
def _cached(key: str, loader):
|
||||
now = time.time()
|
||||
hit = _cache.get(key)
|
||||
if hit and now - hit[0] < _CACHE_TTL:
|
||||
return hit[1]
|
||||
with _lock:
|
||||
hit = _cache.get(key)
|
||||
if hit and time.time() - hit[0] < _CACHE_TTL:
|
||||
return hit[1]
|
||||
value = loader()
|
||||
_cache[key] = (time.time(), value)
|
||||
return value
|
||||
|
||||
|
||||
def _quota_features(quotas: dict[str, Any] | None) -> dict[str, Any]:
|
||||
quotas = quotas or {}
|
||||
features: dict[str, Any] = {}
|
||||
for k, v in quotas.items():
|
||||
if k == "credits_per_month":
|
||||
features["credits_per_month"] = v
|
||||
elif k in _QUOTA_LABELS:
|
||||
features[_QUOTA_LABELS[k]] = v
|
||||
else:
|
||||
features[k] = v
|
||||
return features
|
||||
|
||||
|
||||
def get_membership_plans() -> list[dict[str, Any]]:
|
||||
"""读取 is_enabled=true 的套餐,按年/月周期展开为用户端档位。"""
|
||||
|
||||
def _load() -> list[dict[str, Any]]:
|
||||
from sqlalchemy import text
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.session import SessionLocal
|
||||
|
||||
if SessionLocal is None:
|
||||
return []
|
||||
|
||||
session = SessionLocal()
|
||||
try:
|
||||
rows = session.execute(text("""
|
||||
SELECT plan_key, name, description, monthly_price, yearly_price,
|
||||
quotas, display_order
|
||||
FROM plans
|
||||
WHERE is_enabled = TRUE
|
||||
ORDER BY display_order NULLS LAST, created_at
|
||||
""")).fetchall()
|
||||
finally:
|
||||
session.close()
|
||||
|
||||
plans: list[dict[str, Any]] = []
|
||||
for r in rows:
|
||||
base_features = _quota_features(r.quotas if isinstance(r.quotas, dict) else None)
|
||||
if r.yearly_price and float(r.yearly_price) > 0:
|
||||
plans.append(
|
||||
{
|
||||
"plan_id": r.plan_key,
|
||||
"billing_cycle": "yearly",
|
||||
"name": r.name,
|
||||
"description": r.description,
|
||||
"price_cents": int(round(float(r.yearly_price) * 100)),
|
||||
"monthly_price_cents": int(round(float(r.yearly_price) * 100 / 12)),
|
||||
"duration_days": 365,
|
||||
"features": dict(base_features),
|
||||
}
|
||||
)
|
||||
if r.monthly_price and float(r.monthly_price) > 0:
|
||||
plans.append(
|
||||
{
|
||||
"plan_id": r.plan_key,
|
||||
"billing_cycle": "monthly",
|
||||
"name": r.name,
|
||||
"description": r.description,
|
||||
"price_cents": int(round(float(r.monthly_price) * 100)),
|
||||
"monthly_price_cents": int(round(float(r.monthly_price) * 100)),
|
||||
"duration_days": 30,
|
||||
"features": dict(base_features),
|
||||
}
|
||||
)
|
||||
return plans
|
||||
|
||||
return _cached("membership_plans", _load)
|
||||
|
||||
|
||||
def get_points_packages() -> list[dict[str, Any]]:
|
||||
"""读取 is_active=true 的积分充值包。"""
|
||||
|
||||
def _load() -> list[dict[str, Any]]:
|
||||
from sqlalchemy import text
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.session import SessionLocal
|
||||
|
||||
if SessionLocal is None:
|
||||
return []
|
||||
|
||||
session = SessionLocal()
|
||||
try:
|
||||
rows = session.execute(text("""
|
||||
SELECT package_key, name, price, credits, bonus_credits,
|
||||
is_recommended, description, sort_order
|
||||
FROM credit_packages
|
||||
WHERE is_active = TRUE
|
||||
ORDER BY sort_order NULLS LAST, price
|
||||
""")).fetchall()
|
||||
finally:
|
||||
session.close()
|
||||
|
||||
packages: list[dict[str, Any]] = []
|
||||
for r in rows:
|
||||
total_points = int(r.credits or 0) + int(r.bonus_credits or 0)
|
||||
price_cents = int(round(float(r.price) * 100))
|
||||
unit = (price_cents / 100 / total_points) if total_points else 0
|
||||
packages.append(
|
||||
{
|
||||
"code": r.package_key,
|
||||
"name": r.name,
|
||||
"points": total_points,
|
||||
"bonus_credits": int(r.bonus_credits or 0),
|
||||
"price_cents": price_cents,
|
||||
"unit_price": f"¥{unit:.3f}/积分",
|
||||
"is_recommended": bool(r.is_recommended),
|
||||
"description": r.description,
|
||||
}
|
||||
)
|
||||
return packages
|
||||
|
||||
return _cached("points_packages", _load)
|
||||
+5
-13
@@ -90,26 +90,18 @@ class SharedSettings(BaseSettings):
|
||||
|
||||
# ── 豆包大模型(火山引擎方舟) ────────────────────────────────────────
|
||||
doubao_api_key: str = ""
|
||||
doubao_model: str = "doubao-seed-2-1-pro-260915" # 推理模型(Seed 2.1 Pro,深度思考+多模态;原 seed-1-6 已下线)
|
||||
doubao_fast_model: str = (
|
||||
"doubao-seed-2-1-lite-260915" # 快速模型(Seed 2.1 Lite,高 RPM,编导/审核/VLM lite;原 1-5-pro-32k 已 Retiring)
|
||||
)
|
||||
doubao_model: str = "doubao-seed-1-6-250615" # 推理模型(通用兜底)
|
||||
doubao_fast_model: str = "doubao-1-5-pro-32k-250115" # 快速结构化输出模型(编导脚本/意图解析/审核)
|
||||
doubao_base_url: str = "https://ark.cn-beijing.volces.com/api/v3"
|
||||
doubao_timeout: int = 30
|
||||
doubao_max_retries: int = 2
|
||||
doubao_vision_model: str = (
|
||||
"doubao-seed-2-1-pro-260915" # 高精度视觉(Seed 2.1 Pro 原生多模态;原 vision-pro-250328 已下线)
|
||||
)
|
||||
doubao_vision_lite_model: str = (
|
||||
"doubao-seed-2-1-lite-260915" # 快速视觉(Seed 2.1 Lite 原生多模态;原 vision-lite-250315 不可用)
|
||||
)
|
||||
doubao_vision_model: str = "doubao-1-5-vision-pro-250328" # 高精度视觉(备用)
|
||||
doubao_vision_lite_model: str = "doubao-1-5-vision-lite-250315" # 快速视觉(商品识别默认,速度优先)
|
||||
doubao_vision_use_lite: bool = True # viral-video 图片分析默认用 lite 提速
|
||||
doubao_embedding_model: str = "doubao-embedding-vision-251215" # 多模态向量化(原 large-text-240915 已 Retiring)
|
||||
doubao_embedding_model: str = "doubao-embedding-large-text-240915"
|
||||
doubao_video_model: str = "doubao-seedance-2-5-260628"
|
||||
doubao_video_timeout: int = 600 # 视频生成轮询总超时(秒)
|
||||
doubao_video_poll_interval: int = 10 # 轮询间隔(秒)
|
||||
doubao_image_model: str = "doubao-seedream-5-0-pro-260628" # 图生图/文生图(信任链真人照片AI化)
|
||||
doubao_image_timeout: int = 120 # 图片生成超时(秒)
|
||||
|
||||
# ── DashScope (阿里云百炼 Wan 3.0 等) ─────────────────────────────────
|
||||
dashscope_api_key: str = ""
|
||||
|
||||
@@ -60,11 +60,6 @@ class User:
|
||||
# 资料是否已完善(微信新用户首次设置昵称后置 True;邮箱注册默认 True)
|
||||
profile_completed: bool = True
|
||||
|
||||
# 会员字段 (#1895):与 users 表列对应
|
||||
is_member: bool = False
|
||||
member_type: str | None = None
|
||||
member_expires_at: datetime | None = None
|
||||
|
||||
created_at: datetime = field(default_factory=lambda: datetime.now(UTC))
|
||||
|
||||
|
||||
|
||||
+33
-484
@@ -22,23 +22,14 @@ import httpx
|
||||
|
||||
from packages.shared.config import get_shared_settings
|
||||
|
||||
# 网络/超时类异常父类集合:覆盖 Timeout/Connect/Network/ReadTimeout/WriteTimeout/PoolTimeout
|
||||
_HTTP_NETWORK_ERRORS = ()
|
||||
try:
|
||||
_HTTP_NETWORK_ERRORS = (httpx.TimeoutException, httpx.NetworkError)
|
||||
except Exception:
|
||||
_HTTP_NETWORK_ERRORS = (Exception,)
|
||||
|
||||
_HTTP_STATUS_ERROR = httpx.HTTPStatusError if hasattr(httpx, "HTTPStatusError") else Exception
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
# 视频模型 ID 解析逻辑(#2159 多模型支持,#2170 方舟信任链统一走方舟)。
|
||||
# 内部使用简短别名(seedance-2.5 / seedance-2.0 / wan-3.0 等)做 PRICING key;
|
||||
# 视频模型 ID 解析逻辑(#2159 多模型支持)。
|
||||
# 内部使用简短别名(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(含信任链真人 AI 化)
|
||||
# - provider=dashscope → 阿里云 DashScope(Wan 系列,可选)
|
||||
# - provider=doubao → 火山方舟
|
||||
# - provider=dashscope → 阿里云 DashScope(Wan 系列)
|
||||
|
||||
|
||||
def _resolve_video_provider_and_id(model: str | None) -> tuple[str, str, dict]:
|
||||
@@ -76,93 +67,6 @@ def _resolve_video_model_id(model: str | None) -> str:
|
||||
return mid
|
||||
|
||||
|
||||
# ── 视频错误分类(给前端/用户展示友好提示)────────────────────────────
|
||||
|
||||
|
||||
def _classify_video_error(status_code: int, body: str, err: Exception | None) -> tuple[str, str]:
|
||||
"""根据 HTTP 状态码和响应 body 判断错误类型。
|
||||
|
||||
返回 (error_code, user_message):
|
||||
- error_code: 机器可读的错误码("portrait_intercept" / "quota_exceeded" / "model_not_found"
|
||||
/ "invalid_param" / "auth_error" / "rate_limit" / "network_error" / "task_failed" / "unknown")
|
||||
- user_message: 给用户看的中文提示
|
||||
"""
|
||||
body_lower = (body or "").lower()
|
||||
code_in_body = ""
|
||||
msg_in_body = ""
|
||||
try:
|
||||
import json as _json
|
||||
|
||||
parsed = _json.loads(body or "{}")
|
||||
if isinstance(parsed, dict):
|
||||
err_obj = parsed.get("error") or {}
|
||||
if isinstance(err_obj, dict):
|
||||
code_in_body = str(err_obj.get("code", "") or "")
|
||||
msg_in_body = str(err_obj.get("message", "") or err_obj.get("msg", "") or "")
|
||||
else:
|
||||
msg_in_body = str(parsed.get("message", "") or "")
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# 真人肖像/内容安全拦截
|
||||
if (
|
||||
status_code == 400
|
||||
and any(
|
||||
kw in body_lower
|
||||
for kw in ("portrait", "real_face", "human_face", "真人", "肖像", "人脸", "privacy", "real person", "face")
|
||||
)
|
||||
) or (
|
||||
"content" in body_lower
|
||||
and ("risk" in body_lower or "block" in body_lower or "reject" in body_lower)
|
||||
and status_code == 400
|
||||
):
|
||||
return (
|
||||
"portrait_intercept",
|
||||
"参考素材包含真人照片被安全策略拦截,AI视频模型暂不支持上传真人照片作为参考图,请移除真人图片后重试。",
|
||||
)
|
||||
|
||||
# 配额/计费问题
|
||||
if status_code in (402, 429) or any(
|
||||
kw in body_lower for kw in ("quota", "billing", "insufficient", "欠费", "余额", "限流", "rate limit")
|
||||
):
|
||||
if "rate" in body_lower or status_code == 429:
|
||||
return "rate_limit", "视频生成服务当前繁忙(限流),请稍等1-2分钟后重试。"
|
||||
return "quota_exceeded", "视频生成服务配额不足,请联系管理员充值或稍后重试。"
|
||||
|
||||
# 模型/Endpoint 不存在
|
||||
if status_code == 404 or any(
|
||||
kw in body_lower for kw in ("model not found", "endpoint not found", "不存在", "not found", "model_not_exist")
|
||||
):
|
||||
return "model_not_found", f"视频模型未开通或模型ID无效({code_in_body or ''}),请联系管理员。"
|
||||
|
||||
# 鉴权失败
|
||||
if status_code in (401, 403):
|
||||
return "auth_error", "视频生成服务鉴权失败(API Key无效或过期),请联系管理员。"
|
||||
|
||||
# 任务本身失败(轮询阶段拿到 status=failed)
|
||||
if err and "task failed" in str(err).lower():
|
||||
detail = msg_in_body or str(err)[:200]
|
||||
# 失败原因里再细分真人拦截
|
||||
if any(kw in detail.lower() for kw in ("portrait", "真人", "肖像", "人脸", "content_risk")):
|
||||
return (
|
||||
"portrait_intercept",
|
||||
"视频内容被安全策略拦截(疑似包含真人肖像),请更换参考图或调整文案后重试。",
|
||||
)
|
||||
return "task_failed", f"视频生成失败:{detail}"
|
||||
|
||||
# 参数错误
|
||||
if status_code == 400:
|
||||
return "invalid_param", f"视频生成参数错误:{msg_in_body or body[:200]}"
|
||||
|
||||
# 网络/连接问题
|
||||
if status_code == 0:
|
||||
return "network_error", "视频生成服务连接失败(网络超时),请稍后重试。"
|
||||
|
||||
# 默认
|
||||
detail = msg_in_body or (str(err) if err else "") or body[:200]
|
||||
return "unknown", f"视频生成失败(HTTP {status_code}):{detail}"
|
||||
|
||||
|
||||
class DoubaoClient:
|
||||
"""豆包大模型 API 客户端.
|
||||
|
||||
@@ -180,13 +84,6 @@ class DoubaoClient:
|
||||
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
|
||||
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。"""
|
||||
@@ -199,7 +96,7 @@ class DoubaoClient:
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
payload: dict[str, Any] = {
|
||||
"model": self.embedding_model,
|
||||
"model": getattr(self, "embedding_model", None) or "doubao-embedding-large-text-240915",
|
||||
"input": text.strip(),
|
||||
"encoding_format": "float",
|
||||
}
|
||||
@@ -410,8 +307,7 @@ class DoubaoClient:
|
||||
"""调用 Seedance 2.5 生视频(异步任务→轮询→下载)。
|
||||
|
||||
成功返回 {"video_path": str, "usage": dict | None},失败返回 None。
|
||||
失败时把详细错误信息(HTTP状态码、响应 body、分类后的用户提示)写入 self.last_video_error,
|
||||
上层可通过 get_last_video_error() 读取并展示给用户,不再笼统显示"返回为空"。
|
||||
usage 是 Seedance 返回的计费信息(含 completion_tokens)。
|
||||
|
||||
【v1.6.1 修复】严格按官方 content 数组协议构造请求:
|
||||
- 所有参考(图/音/视)必须放进 content 数组并带 role 字段,不能放顶层 reference_audios/reference_videos(非官方字段,会被忽略或导致异常)。
|
||||
@@ -419,24 +315,9 @@ class DoubaoClient:
|
||||
判定:传了参考音频/视频或 ≥1 张多参考图时,走 omni_reference(首张图 role=reference_image);纯首帧无参考时走 first_frame(ratio 强制 adaptive)。
|
||||
- 创建任务若因 ratio 报错(HTTP 400),自动回退到 ratio=adaptive 重试一次。
|
||||
"""
|
||||
# 每次调用前清空上次错误
|
||||
self.last_video_error = {}
|
||||
|
||||
if not self.is_available:
|
||||
self.last_video_error = {
|
||||
"error_code": "auth_error",
|
||||
"user_message": "视频生成服务未配置(API Key 缺失),请联系管理员。",
|
||||
"status_code": 0,
|
||||
"detail": "DoubaoClient not available (api_key empty)",
|
||||
}
|
||||
return None
|
||||
if not prompt or not prompt.strip():
|
||||
self.last_video_error = {
|
||||
"error_code": "invalid_param",
|
||||
"user_message": "视频生成提示词不能为空。",
|
||||
"status_code": 0,
|
||||
"detail": "empty prompt",
|
||||
}
|
||||
return None
|
||||
|
||||
settings = get_shared_settings()
|
||||
@@ -450,20 +331,10 @@ class DoubaoClient:
|
||||
|
||||
ds = get_dashscope_client()
|
||||
if ds is None:
|
||||
err_msg = "DashScope client 不可用(未配置 DASHSCOPE_API_KEY)"
|
||||
logger.error("%s, video_model=%s", err_msg, model)
|
||||
self.last_video_error = {
|
||||
"error_code": "auth_error",
|
||||
"user_message": "Wan 3.0 视频模型未配置 API Key,请联系管理员。",
|
||||
"status_code": 0,
|
||||
"detail": err_msg,
|
||||
}
|
||||
logger.error("DashScope client 不可用(未配置 DASHSCOPE_API_KEY),video_model=%s", model)
|
||||
return None
|
||||
try:
|
||||
# DashScope 客户端也设置 last_video_error 语义(如果它支持)
|
||||
if hasattr(ds, "last_video_error"):
|
||||
ds.last_video_error = {}
|
||||
result = ds.video_generation(
|
||||
return ds.video_generation(
|
||||
prompt=prompt,
|
||||
image_url=image_url,
|
||||
duration=duration,
|
||||
@@ -472,96 +343,25 @@ class DoubaoClient:
|
||||
output_dir=output_dir,
|
||||
model=video_model,
|
||||
)
|
||||
if not result and hasattr(ds, "last_video_error") and ds.last_video_error:
|
||||
self.last_video_error = dict(ds.last_video_error)
|
||||
return result
|
||||
except Exception as de:
|
||||
logger.error("DashScope video_generation 异常: %s", de, exc_info=True)
|
||||
self.last_video_error = {
|
||||
"error_code": "unknown",
|
||||
"user_message": f"Wan 3.0 视频生成异常:{de!s}"[:200],
|
||||
"status_code": 0,
|
||||
"detail": str(de),
|
||||
}
|
||||
return None
|
||||
|
||||
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)]
|
||||
|
||||
# ── #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)
|
||||
# 判断任务模式:有参考音/视/多图 → 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 and not trust_chain_applied
|
||||
is_first_frame_mode = bool(image_url) and not has_extra_refs
|
||||
# 最终 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 or trust_chain_applied:
|
||||
# omni_reference 或信任链模式:首张图作为 reference_image,允许指定 ratio
|
||||
if has_extra_refs:
|
||||
# omni_reference:首张图作为 reference_image,允许指定 ratio
|
||||
content.append(
|
||||
{
|
||||
"type": "image_url",
|
||||
@@ -570,7 +370,7 @@ class DoubaoClient:
|
||||
}
|
||||
)
|
||||
else:
|
||||
# 纯首帧:显式 role=first_frame
|
||||
# 纯首帧:不带 role,服务端识别为 first_frame(或显式 role=first_frame)
|
||||
content.append(
|
||||
{
|
||||
"type": "image_url",
|
||||
@@ -605,44 +405,28 @@ class DoubaoClient:
|
||||
video_model,
|
||||
duration,
|
||||
final_ratio,
|
||||
"first_frame" if is_first_frame_mode else ("omni_ref+trust_chain" if trust_chain_applied else "omni_ref"),
|
||||
"first_frame" if is_first_frame_mode else "omni_ref",
|
||||
generate_audio,
|
||||
(1 if image_url else 0) + len(ref_imgs),
|
||||
len(ref_audios),
|
||||
len(ref_videos),
|
||||
)
|
||||
# 打印完整 payload 便于排查(截断 prompt)
|
||||
debug_payload = dict(create_payload)
|
||||
if "content" in debug_payload:
|
||||
dbg_content = []
|
||||
for item in debug_payload["content"]:
|
||||
item_copy = dict(item)
|
||||
if item_copy.get("type") == "text" and isinstance(item_copy.get("text"), str):
|
||||
item_copy["text"] = item_copy["text"][:200] + ("..." if len(item_copy["text"]) > 200 else "")
|
||||
dbg_content.append(item_copy)
|
||||
debug_payload["content"] = dbg_content
|
||||
logger.info("Seedance 创建任务 payload: %s", json_safe_dumps(debug_payload))
|
||||
|
||||
def _do_create(payload: dict) -> tuple[str | None, Exception | None, int, str]:
|
||||
"""返回 (task_id, last_err, status_code, body_text)。"""
|
||||
last_err: Exception | None = None
|
||||
last_sc = 0
|
||||
last_body = ""
|
||||
for attempt in range(self.max_retries + 1):
|
||||
try:
|
||||
resp = httpx.post(create_url, headers=headers, json=payload, timeout=self.timeout)
|
||||
sc = int(getattr(resp, "status_code", 0) or 0)
|
||||
body = (getattr(resp, "text", "") or "")[:2000]
|
||||
last_sc = sc
|
||||
last_body = body
|
||||
body = (getattr(resp, "text", "") or "")[:1500]
|
||||
if sc >= 400:
|
||||
logger.error("Seedance 创建任务 HTTP %d: body=%s", sc, body)
|
||||
try:
|
||||
resp.raise_for_status()
|
||||
except Exception as ee:
|
||||
last_err = ee
|
||||
if attempt < self.max_retries and sc >= 500:
|
||||
# 仅 5xx 重试,4xx 不重试(参数/鉴权/配额错误重试无意义)
|
||||
if attempt < self.max_retries:
|
||||
time.sleep(0.5 * (2**attempt))
|
||||
continue
|
||||
return None, last_err, sc, body
|
||||
@@ -651,19 +435,9 @@ class DoubaoClient:
|
||||
if tid:
|
||||
return tid, None, sc, body
|
||||
last_err = RuntimeError(f"create ok but no id: {str(data)[:300]}")
|
||||
except _HTTP_NETWORK_ERRORS as ne:
|
||||
last_err = ne
|
||||
last_sc = 0
|
||||
last_body = f"network error: {ne}"
|
||||
logger.warning(
|
||||
"Seedance 创建网络异常(%s),重试 %d/%d", type(ne).__name__, attempt + 1, self.max_retries + 1
|
||||
)
|
||||
if attempt < self.max_retries:
|
||||
time.sleep(0.5 * (2**attempt))
|
||||
continue
|
||||
except Exception as e:
|
||||
last_err = e
|
||||
if attempt < self.max_retries and not isinstance(e, _HTTP_STATUS_ERROR):
|
||||
if attempt < self.max_retries:
|
||||
wait = 0.5 * (2**attempt)
|
||||
logger.warning(
|
||||
"Seedance 创建任务失败,%.1fs 后重试 (%d/%d): %s",
|
||||
@@ -673,7 +447,7 @@ class DoubaoClient:
|
||||
e,
|
||||
)
|
||||
time.sleep(wait)
|
||||
return None, last_err, last_sc, last_body
|
||||
return None, last_err, 0, ""
|
||||
|
||||
# 第一次尝试
|
||||
task_id, last_err, sc, body = _do_create(create_payload)
|
||||
@@ -692,30 +466,17 @@ class DoubaoClient:
|
||||
logger.warning("Seedance 创建因 ratio 失败,回退 ratio=adaptive 重试")
|
||||
create_payload["ratio"] = "adaptive"
|
||||
task_id, last_err, sc2, body2 = _do_create(create_payload)
|
||||
if task_id:
|
||||
sc, body = sc2, body2
|
||||
else:
|
||||
# 保留第二次的错误信息
|
||||
sc, body = sc2, body2
|
||||
|
||||
if not task_id:
|
||||
err_code, user_msg = _classify_video_error(sc, body, last_err)
|
||||
self.last_video_error = {
|
||||
"error_code": err_code,
|
||||
"user_message": user_msg,
|
||||
"status_code": sc,
|
||||
"detail": (body or "")[:500] or (str(last_err) if last_err else ""),
|
||||
"model": video_model,
|
||||
"base_url": self.base_url,
|
||||
}
|
||||
logger.error(
|
||||
"Seedance 创建任务最终失败: model=%s base_url=%s status=%d code=%s err=%s body=%s",
|
||||
"Seedance 创建任务最终失败: model=%s base_url=%s err=%s body=%s 【排查】"
|
||||
"1) 方舟控制台已开通 %s;2) API Key 有该模型权限;"
|
||||
"3) DOUBAO_BASE_URL=https://ark.cn-beijing.volces.com/api/v3;4) 参考素材 URL 公网可访问。",
|
||||
video_model,
|
||||
self.base_url,
|
||||
sc,
|
||||
err_code,
|
||||
last_err,
|
||||
(body or "")[:500],
|
||||
video_model,
|
||||
)
|
||||
return None
|
||||
|
||||
@@ -728,21 +489,15 @@ class DoubaoClient:
|
||||
usage: dict | None = None
|
||||
last_status: str = "queued"
|
||||
poll_count = 0
|
||||
last_poll_body: str = ""
|
||||
last_poll_sc: int = 0
|
||||
while time.time() < deadline:
|
||||
poll_count += 1
|
||||
try:
|
||||
resp = httpx.get(poll_url, headers=headers, timeout=self.timeout)
|
||||
last_poll_sc = int(getattr(resp, "status_code", 200) or 200)
|
||||
last_poll_body = (getattr(resp, "text", "") or "")[:1500]
|
||||
if last_poll_sc >= 400:
|
||||
logger.warning("Seedance 轮询 HTTP %d: %s", last_poll_sc, last_poll_body[:300])
|
||||
if poll_count < 3:
|
||||
time.sleep(poll_interval)
|
||||
continue
|
||||
last_err = RuntimeError(f"poll HTTP {last_poll_sc}: {last_poll_body[:200]}")
|
||||
break
|
||||
try:
|
||||
if int(getattr(resp, "status_code", 200)) >= 400:
|
||||
resp.raise_for_status()
|
||||
except (TypeError, ValueError):
|
||||
pass
|
||||
data = resp.json()
|
||||
status = data.get("status", "")
|
||||
last_status = status
|
||||
@@ -753,22 +508,13 @@ class DoubaoClient:
|
||||
if video_url:
|
||||
logger.info("Seedance 任务成功: task_id=%s polls=%d usage=%s", task_id, poll_count, usage)
|
||||
break
|
||||
# 成功但没 video_url:记录完整响应便于排查
|
||||
logger.error(
|
||||
"Seedance succeeded 但无 video_url: task_id=%s full_response=%s",
|
||||
task_id,
|
||||
str(data)[:1000],
|
||||
)
|
||||
last_err = RuntimeError("task succeeded but no video_url in response")
|
||||
last_poll_body = str(data)[:1000]
|
||||
last_err = RuntimeError(f"task succeeded but no video_url: {str(data)[:300]}")
|
||||
logger.error("Seedance succeeded 但无 video_url: %s", last_err)
|
||||
break
|
||||
if status == "failed":
|
||||
err = data.get("error") or {}
|
||||
err_code = str(err.get("code", "") or "")
|
||||
err_msg = str(err.get("message", "") or err.get("msg", "") or "")
|
||||
last_err = RuntimeError(f"task failed: code={err_code} msg={err_msg}")
|
||||
logger.error("Seedance 任务失败 task_id=%s code=%s msg=%s", task_id, err_code, err_msg)
|
||||
last_poll_body = str(data)[:1000]
|
||||
last_err = RuntimeError(f"task failed: code={err.get('code','')} msg={err.get('message','')}")
|
||||
logger.error("Seedance 任务失败 task_id=%s: %s", task_id, last_err)
|
||||
break
|
||||
if status in ("expired", "cancelled"):
|
||||
last_err = RuntimeError(f"task {status}")
|
||||
@@ -779,39 +525,18 @@ class DoubaoClient:
|
||||
logger.info("Seedance 轮询中: task_id=%s status=%s polls=%d", task_id, status, poll_count)
|
||||
except httpx.HTTPStatusError as e:
|
||||
last_err = e
|
||||
last_poll_sc = e.response.status_code
|
||||
last_poll_body = (e.response.text or "")[:500]
|
||||
logger.warning("Seedance 轮询 HTTP %d: %s", e.response.status_code, last_poll_body[:300])
|
||||
logger.warning("Seedance 轮询 HTTP %d: %s", e.response.status_code, (e.response.text or "")[:300])
|
||||
except Exception as e:
|
||||
last_err = e
|
||||
logger.debug("Seedance 轮询异常: %s", e)
|
||||
time.sleep(poll_interval)
|
||||
|
||||
if not video_url:
|
||||
# 区分轮询超时 vs 任务失败
|
||||
if last_status in ("queued", "running", "pending") and poll_count > 0 and time.time() >= deadline:
|
||||
err_code, user_msg = (
|
||||
"network_error",
|
||||
f"视频生成超时(>{total_timeout}s),任务仍在排队,请稍后重试或联系管理员。",
|
||||
)
|
||||
detail = f"timeout after {total_timeout}s, polls={poll_count}, last_status={last_status}"
|
||||
else:
|
||||
err_code, user_msg = _classify_video_error(last_poll_sc, last_poll_body, last_err)
|
||||
detail = (last_poll_body or "")[:500] or (str(last_err) if last_err else f"last_status={last_status}")
|
||||
self.last_video_error = {
|
||||
"error_code": err_code,
|
||||
"user_message": user_msg,
|
||||
"status_code": last_poll_sc,
|
||||
"detail": detail,
|
||||
"task_id": task_id,
|
||||
"last_status": last_status,
|
||||
}
|
||||
logger.error(
|
||||
"Seedance 任务未成功: task_id=%s last_status=%s polls=%d code=%s err=%s (总等待 %.0fs)",
|
||||
"Seedance 任务未成功: task_id=%s last_status=%s polls=%d err=%s (总等待 %.0fs)",
|
||||
task_id,
|
||||
last_status,
|
||||
poll_count,
|
||||
err_code,
|
||||
last_err,
|
||||
total_timeout,
|
||||
)
|
||||
@@ -842,188 +567,12 @@ class DoubaoClient:
|
||||
os.remove(local_path)
|
||||
except Exception:
|
||||
pass
|
||||
self.last_video_error = {
|
||||
"error_code": "unknown",
|
||||
"user_message": "视频生成成功但下载文件为空,请稍后重试。",
|
||||
"status_code": 0,
|
||||
"detail": f"downloaded 0 bytes from {video_url[:120]}",
|
||||
}
|
||||
return None
|
||||
return {"video_path": local_path, "usage": usage}
|
||||
except Exception as e:
|
||||
logger.error("Seedance 视频下载失败: %s", e, exc_info=True)
|
||||
self.last_video_error = {
|
||||
"error_code": "network_error",
|
||||
"user_message": f"视频下载失败:{e!s}"[:200],
|
||||
"status_code": 0,
|
||||
"detail": str(e),
|
||||
}
|
||||
return None
|
||||
|
||||
def image_generation(
|
||||
self,
|
||||
prompt: str,
|
||||
*,
|
||||
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:
|
||||
"""#2170: 调用方舟 Seedream 图片生成(文生图/图生图)。
|
||||
|
||||
- reference_images: 0~10 张参考图 URL;0 张 = 纯文生图;1 张 string/URL 直传;多张 list[str]。
|
||||
- 成功返回 {"url": str, "usage": dict | None};失败返回 None,错误写入 self.last_image_error。
|
||||
- 返回的 url 有时效性(通常 24h),应立即使用,不持久化存储。
|
||||
"""
|
||||
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 prompt or not prompt.strip():
|
||||
self.last_image_error = {
|
||||
"error_code": "invalid_param",
|
||||
"user_message": "图片生成提示词不能为空。",
|
||||
"detail": "empty prompt",
|
||||
}
|
||||
return None
|
||||
|
||||
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
|
||||
)
|
||||
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 表示上次成功或未调用。"""
|
||||
return dict(self.last_video_error or {})
|
||||
|
||||
|
||||
def json_safe_dumps(obj: Any, max_len: int = 2000) -> str:
|
||||
"""安全 json 序列化,失败则 fallback 到 repr,超长截断。"""
|
||||
try:
|
||||
import json as _json
|
||||
|
||||
s = _json.dumps(obj, ensure_ascii=False, default=str)
|
||||
except Exception:
|
||||
s = repr(obj)
|
||||
if len(s) > max_len:
|
||||
s = s[:max_len] + f"...(truncated, total {len(s)})"
|
||||
return s
|
||||
|
||||
|
||||
# ── 单例 ─────────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@@ -618,23 +618,20 @@ def call_video_generation(
|
||||
reference_audios: list[str] | None = None,
|
||||
reference_videos: list[str] | None = None,
|
||||
) -> dict | None:
|
||||
"""调用 Seedance / Wan 视频生成(v1.6.2 多模型版)。
|
||||
"""调用 Seedance 2.5 生成视频(v1.6.1 单次出片版)。
|
||||
|
||||
成功返回 {"video_path": str, "usage": dict | None}(usage 含 completion_tokens),失败返回 None。
|
||||
失败时错误详情会写入 client.last_video_error,可通过 get_last_video_error() 读取:
|
||||
{"error_code": str, "user_message": str, "status_code": int, "detail": str, ...}
|
||||
|
||||
v1.6.1 关键约束(避免 20min 卡死):
|
||||
- 参考音频/视频/多图全部放进 content 数组并带 role=reference_audio/reference_video/reference_image;
|
||||
- 纯首帧无参考时(first_frame 模式),Seedance 2.5 强制 ratio=adaptive;
|
||||
传了参考音/视/多图时走 omni_reference 模式,ratio 可指定为 9:16(客户端内部自动判断)。
|
||||
- ratio 默认 9:16(竖屏),客户端会根据是否有参考自动在 first_frame/adaptive 与 omni/9:16 间切换;
|
||||
若创建任务因 ratio 报错(HTTP 400),客户端会自动回退到 adaptive 再试一次。
|
||||
"""
|
||||
client = get_doubao_client()
|
||||
if not client.is_available:
|
||||
msg = "豆包客户端未配置(DOUBAO_API_KEY 缺失),跳过视频生成"
|
||||
logger.warning("[ai_service] %s", msg)
|
||||
# 写入 last_video_error 供上层读取
|
||||
client.last_video_error = {
|
||||
"error_code": "auth_error",
|
||||
"user_message": "视频生成服务未配置,请联系管理员。",
|
||||
"status_code": 0,
|
||||
"detail": msg,
|
||||
}
|
||||
logger.warning("[ai_service] 豆包客户端未配置,跳过视频生成")
|
||||
return None
|
||||
effective_ratio = ratio or "9:16"
|
||||
try:
|
||||
@@ -656,21 +653,4 @@ def call_video_generation(
|
||||
return client.video_generation(**kwargs)
|
||||
except Exception as e:
|
||||
logger.error("[ai_service] call_video_generation 异常: %s", e, exc_info=True)
|
||||
client.last_video_error = {
|
||||
"error_code": "unknown",
|
||||
"user_message": f"视频生成异常:{e!s}"[:200],
|
||||
"status_code": 0,
|
||||
"detail": str(e),
|
||||
}
|
||||
return None
|
||||
|
||||
|
||||
def get_last_video_error() -> dict:
|
||||
"""读取最近一次视频生成失败的详细错误(含 error_code/user_message/status_code/detail)。
|
||||
成功或未调用过返回空 dict。
|
||||
"""
|
||||
try:
|
||||
client = get_doubao_client()
|
||||
return client.get_last_video_error() if hasattr(client, "get_last_video_error") else {}
|
||||
except Exception:
|
||||
return {}
|
||||
|
||||
@@ -21,59 +21,11 @@ import httpx
|
||||
|
||||
from packages.shared.config import get_shared_settings
|
||||
|
||||
# 网络/超时类异常父类集合:覆盖 Timeout/Connect/Network/ReadTimeout/WriteTimeout/PoolTimeout
|
||||
_HTTP_NETWORK_ERRORS = ()
|
||||
try:
|
||||
_HTTP_NETWORK_ERRORS = (httpx.TimeoutException, httpx.NetworkError)
|
||||
except Exception:
|
||||
_HTTP_NETWORK_ERRORS = (Exception,)
|
||||
|
||||
_HTTP_STATUS_ERROR = httpx.HTTPStatusError if hasattr(httpx, "HTTPStatusError") else Exception
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_DASHSCOPE_CLIENT_SINGLETON: "DashScopeClient | None" = None
|
||||
|
||||
|
||||
def _classify_dashscope_error(status_code: int, body: str, task_msg: str = "") -> tuple[str, str]:
|
||||
"""DashScope 错误分类,返回 (error_code, user_message)。"""
|
||||
body_lower = (body or "").lower()
|
||||
msg_in_body = task_msg or ""
|
||||
try:
|
||||
import json as _json
|
||||
|
||||
parsed = _json.loads(body or "{}")
|
||||
if isinstance(parsed, dict):
|
||||
msg_in_body = msg_in_body or str(parsed.get("message", "") or "")
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
if status_code in (401, 403):
|
||||
return "auth_error", "Wan 3.0 服务鉴权失败(DASHSCOPE_API_KEY 无效或过期),请联系管理员。"
|
||||
if status_code == 429 or "rate" in body_lower or "throttl" in body_lower:
|
||||
return "rate_limit", "Wan 3.0 服务繁忙(限流),请稍等1-2分钟后重试。"
|
||||
if status_code == 400 and any(
|
||||
kw in body_lower for kw in ("portrait", "真人", "人脸", "肖像", "content_violation", "risk", "blocked")
|
||||
):
|
||||
return (
|
||||
"portrait_intercept",
|
||||
"参考素材包含真人照片或违规内容被安全策略拦截,请移除真人图片或调整文案后重试。",
|
||||
)
|
||||
if status_code == 404 or ("not found" in body_lower) or ("model" in body_lower and "not exist" in body_lower):
|
||||
return "model_not_found", "Wan 3.0 模型未开通或模型ID无效,请联系管理员。"
|
||||
if status_code in (402, 400) and ("quota" in body_lower or "billing" in body_lower or "insufficient" in body_lower):
|
||||
return "quota_exceeded", "Wan 3.0 服务配额不足,请联系管理员充值或稍后重试。"
|
||||
if status_code == 400:
|
||||
return "invalid_param", f"Wan 3.0 参数错误:{msg_in_body or body[:200]}"
|
||||
if status_code == 0:
|
||||
return "network_error", "Wan 3.0 服务连接失败(网络超时),请稍后重试。"
|
||||
# 任务内失败
|
||||
if task_msg and any(kw in task_msg.lower() for kw in ("portrait", "真人", "人脸", "violation", "blocked")):
|
||||
return "portrait_intercept", "Wan 3.0 视频内容被安全策略拦截,请调整文案或参考图后重试。"
|
||||
detail = msg_in_body or body[:200]
|
||||
return "unknown", f"Wan 3.0 视频生成失败(HTTP {status_code}):{detail}"
|
||||
|
||||
|
||||
class DashScopeClient:
|
||||
"""阿里云 DashScope 异步 API 客户端(Wan 3.0 等视频生成)。"""
|
||||
|
||||
@@ -86,24 +38,11 @@ class DashScopeClient:
|
||||
self.poll_interval: int = int(getattr(settings, "dashscope_video_poll_interval", 10) or 10)
|
||||
self.total_timeout: int = int(getattr(settings, "dashscope_video_timeout", 900) or 900)
|
||||
self.max_retries: int = 2
|
||||
self.last_video_error: dict = {}
|
||||
|
||||
@property
|
||||
def is_available(self) -> bool:
|
||||
return bool(self.api_key)
|
||||
|
||||
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 "",
|
||||
**extra,
|
||||
}
|
||||
|
||||
def video_generation(
|
||||
self,
|
||||
prompt: str,
|
||||
@@ -118,15 +57,12 @@ class DashScopeClient:
|
||||
) -> dict | None:
|
||||
"""调用 DashScope 异步视频合成接口,轮询完成后下载到本地。
|
||||
|
||||
返回 {"video_path": str, "usage": dict | None};失败返回 None,错误详情写入 self.last_video_error。
|
||||
返回 {"video_path": str, "usage": dict | None};失败返回 None。
|
||||
"""
|
||||
self.last_video_error = {}
|
||||
if not self.is_available:
|
||||
self._set_error("auth_error", "Wan 3.0 API key 未配置,请联系管理员。", detail="dashscope api_key empty")
|
||||
logger.error("[dashscope] API key 未配置,无法调用视频生成")
|
||||
return None
|
||||
if not prompt or not prompt.strip():
|
||||
self._set_error("invalid_param", "视频生成提示词不能为空。", detail="empty prompt")
|
||||
return None
|
||||
|
||||
# DashScope 分辨率参数:720P / 1080P / 480P(大写 P)
|
||||
@@ -170,27 +106,18 @@ class DashScopeClient:
|
||||
ds_res,
|
||||
bool(image_url),
|
||||
)
|
||||
logger.info("[dashscope] 创建任务 payload: model=%s params=%s", model, params)
|
||||
|
||||
# 创建任务
|
||||
task_id: str | None = None
|
||||
last_sc = 0
|
||||
last_body = ""
|
||||
last_err: Exception | None = None
|
||||
for attempt in range(self.max_retries + 1):
|
||||
try:
|
||||
resp = httpx.post(create_url, headers=headers, json=payload, timeout=60)
|
||||
sc = int(getattr(resp, "status_code", 0) or 0)
|
||||
body_text = (getattr(resp, "text", "") or "")[:2000]
|
||||
last_sc = sc
|
||||
last_body = body_text
|
||||
body_text = (getattr(resp, "text", "") or "")[:1500]
|
||||
if sc >= 400:
|
||||
logger.error("[dashscope] 创建任务 HTTP %d: %s", sc, body_text)
|
||||
if sc >= 500 and attempt < self.max_retries:
|
||||
time.sleep(0.5 * (2**attempt))
|
||||
continue
|
||||
err_code, user_msg = _classify_dashscope_error(sc, body_text)
|
||||
self._set_error(err_code, user_msg, sc, body_text, model=model)
|
||||
return None
|
||||
resp.raise_for_status()
|
||||
data = resp.json()
|
||||
tid = (data.get("output") or {}).get("task_id")
|
||||
if tid:
|
||||
@@ -199,34 +126,17 @@ class DashScopeClient:
|
||||
# 部分情况下 code != 错误
|
||||
code = data.get("code")
|
||||
if code and code != "":
|
||||
err_code, user_msg = _classify_dashscope_error(400, body_text, str(code))
|
||||
self._set_error(err_code, user_msg, sc, body_text, model=model)
|
||||
return None
|
||||
last_err = RuntimeError(f"dashscope create failed: {body_text[:300]}")
|
||||
else:
|
||||
self._set_error("unknown", "Wan 3.0 响应格式异常,未返回任务ID", sc, str(data)[:500], model=model)
|
||||
return None
|
||||
except _HTTP_NETWORK_ERRORS as ne:
|
||||
last_sc = 0
|
||||
last_body = f"network error: {ne}"
|
||||
logger.warning(
|
||||
"[dashscope] 网络异常 %s,重试 %d/%d", type(ne).__name__, attempt + 1, self.max_retries + 1
|
||||
)
|
||||
last_err = RuntimeError(f"create ok but no task_id: {str(data)[:300]}")
|
||||
except Exception as e:
|
||||
last_err = e
|
||||
if attempt < self.max_retries:
|
||||
time.sleep(0.5 * (2**attempt))
|
||||
continue
|
||||
self._set_error("network_error", "Wan 3.0 服务连接失败(网络超时),请稍后重试。", 0, str(ne))
|
||||
return None
|
||||
except Exception as _e:
|
||||
if attempt < self.max_retries:
|
||||
time.sleep(0.5 * (2**attempt))
|
||||
continue
|
||||
logger.error("[dashscope] 创建任务最终失败: %s", _e)
|
||||
self._set_error("unknown", f"Wan 3.0 创建任务异常:{_e!s}"[:200], 0, str(_e))
|
||||
logger.error("[dashscope] 创建任务最终失败: %s", last_err)
|
||||
return None
|
||||
if not task_id:
|
||||
if not self.last_video_error:
|
||||
err_code, user_msg = _classify_dashscope_error(last_sc, last_body)
|
||||
self._set_error(err_code, user_msg, last_sc, last_body, model=model)
|
||||
return None
|
||||
|
||||
# 轮询任务
|
||||
@@ -234,26 +144,16 @@ class DashScopeClient:
|
||||
deadline = time.time() + self.total_timeout
|
||||
video_url: str | None = None
|
||||
usage: dict | None = None
|
||||
poll_count = 0
|
||||
last_status = ""
|
||||
while time.time() < deadline:
|
||||
poll_count += 1
|
||||
try:
|
||||
r = httpx.get(poll_url, headers=headers, timeout=30)
|
||||
psc = int(getattr(r, "status_code", 0) or 0)
|
||||
pbody = (getattr(r, "text", "") or "")[:1500]
|
||||
if psc >= 400:
|
||||
logger.warning("[dashscope] 轮询 HTTP %d: %s", psc, pbody[:300])
|
||||
if poll_count < 3:
|
||||
time.sleep(self.poll_interval)
|
||||
continue
|
||||
err_code, user_msg = _classify_dashscope_error(psc, pbody)
|
||||
self._set_error(err_code, user_msg, psc, pbody, task_id=task_id)
|
||||
return None
|
||||
if int(getattr(r, "status_code", 0) or 0) >= 400:
|
||||
logger.warning("[dashscope] 轮询 HTTP %d", r.status_code)
|
||||
time.sleep(self.poll_interval)
|
||||
continue
|
||||
d = r.json()
|
||||
out = d.get("output") or {}
|
||||
task_status = out.get("task_status") or d.get("task_status") or ""
|
||||
last_status = task_status
|
||||
if task_status == "SUCCEEDED":
|
||||
video_url = out.get("video_url") or ""
|
||||
usage = d.get("usage")
|
||||
@@ -265,41 +165,22 @@ class DashScopeClient:
|
||||
if video_url:
|
||||
logger.info("[dashscope] 任务 %s 完成: %s", task_id, video_url[:120])
|
||||
break
|
||||
logger.error("[dashscope] 任务 %s SUCCEEDED 但无 video_url: %s", task_id, str(d)[:500])
|
||||
self._set_error(
|
||||
"unknown",
|
||||
"Wan 3.0 任务成功但未返回视频URL,请联系管理员。",
|
||||
200,
|
||||
str(d)[:500],
|
||||
task_id=task_id,
|
||||
)
|
||||
logger.error("[dashscope] 任务 %s SUCCEEDED 但无 video_url", task_id)
|
||||
return None
|
||||
if task_status in ("FAILED", "FAILED_WITH_ERROR", "ERROR"):
|
||||
msg = out.get("message") or d.get("message") or out.get("error_msg") or "unknown error"
|
||||
msg = out.get("message") or d.get("message") or "unknown error"
|
||||
logger.error("[dashscope] 任务 %s 失败: %s", task_id, msg)
|
||||
err_code, user_msg = _classify_dashscope_error(200, "", msg)
|
||||
self._set_error(err_code, user_msg, 200, msg, task_id=task_id, last_status=task_status)
|
||||
return None
|
||||
if task_status in ("CANCELED", "CANCELLED"):
|
||||
logger.warning("[dashscope] 任务 %s 被取消", task_id)
|
||||
self._set_error("unknown", "Wan 3.0 任务被取消。", 200, "task cancelled", task_id=task_id)
|
||||
return None
|
||||
# PENDING / RUNNING / SUSPENDED → 继续轮询
|
||||
if poll_count % 5 == 0:
|
||||
logger.info("[dashscope] 轮询中 task=%s status=%s polls=%d", task_id, task_status, poll_count)
|
||||
logger.debug("[dashscope] 任务 %s 状态 %s,继续轮询", task_id, task_status)
|
||||
except Exception as e:
|
||||
logger.warning("[dashscope] 轮询异常: %s", e)
|
||||
time.sleep(self.poll_interval)
|
||||
if not video_url:
|
||||
logger.error("[dashscope] 任务 %s 轮询超时(%ds)", task_id, self.total_timeout)
|
||||
self._set_error(
|
||||
"network_error",
|
||||
f"Wan 3.0 视频生成超时(>{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
|
||||
|
||||
# 下载视频
|
||||
@@ -312,23 +193,19 @@ class DashScopeClient:
|
||||
out_path = os.path.join(out_dir, f"wan_{safe_tid}{suffix}")
|
||||
try:
|
||||
with httpx.stream("GET", video_url, timeout=300, follow_redirects=True) as resp:
|
||||
dsc = int(getattr(resp, "status_code", 0) or 0)
|
||||
if dsc >= 400:
|
||||
logger.error("[dashscope] 下载 HTTP %d", dsc)
|
||||
self._set_error("network_error", "Wan 3.0 视频下载失败(HTTP错误),请稍后重试。", dsc)
|
||||
if int(getattr(resp, "status_code", 0) or 0) >= 400:
|
||||
logger.error("[dashscope] 下载 HTTP %d", resp.status_code)
|
||||
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("[dashscope] 下载视频失败: %s", e, exc_info=True)
|
||||
self._set_error("network_error", f"Wan 3.0 视频下载失败:{e!s}"[:200], 0, str(e))
|
||||
logger.error("[dashscope] 下载视频失败: %s", e)
|
||||
return None
|
||||
size = os.path.getsize(out_path) if os.path.exists(out_path) else 0
|
||||
if size < 1024:
|
||||
logger.error("[dashscope] 下载文件过小: %d bytes", size)
|
||||
self._set_error("unknown", "Wan 3.0 视频下载文件过小,请稍后重试。", 0, f"downloaded only {size} bytes")
|
||||
return None
|
||||
logger.info("[dashscope] 视频已下载: %s (%d bytes)", out_path, size)
|
||||
return {"video_path": out_path, "usage": usage}
|
||||
|
||||
@@ -1,5 +1,4 @@
|
||||
"""Additional unit tests to hit uncovered lines for diff-coverage >=60%."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
@@ -14,20 +13,15 @@ from packages.shared.ai_client import DoubaoClient
|
||||
|
||||
class _FakeSettings:
|
||||
doubao_api_key = "test-key"
|
||||
doubao_model = "doubao-seed-2-1-pro-260915"
|
||||
doubao_fast_model = "doubao-seed-2-1-lite-260915"
|
||||
doubao_model = "test-model"
|
||||
doubao_fast_model = "test-fast-model"
|
||||
doubao_base_url = "https://ark.cn-beijing.volces.com/api/v3"
|
||||
doubao_timeout = 10
|
||||
doubao_max_retries = 0
|
||||
doubao_vision_model = "doubao-seed-2-1-pro-260915"
|
||||
doubao_vision_lite_model = "doubao-seed-2-1-lite-260915"
|
||||
doubao_vision_model = "test-vision"
|
||||
doubao_vision_lite_model = "test-vision-lite"
|
||||
doubao_vision_use_lite = False
|
||||
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
|
||||
doubao_embedding_model = "test-embedding"
|
||||
|
||||
|
||||
def _make_client(api_key: str = "test-key") -> DoubaoClient:
|
||||
@@ -135,50 +129,26 @@ 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
|
||||
|
||||
@@ -189,18 +159,10 @@ 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
|
||||
|
||||
@@ -234,7 +196,9 @@ 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"]
|
||||
|
||||
|
||||
@@ -267,7 +231,6 @@ 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
|
||||
|
||||
@@ -1,669 +0,0 @@
|
||||
"""#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"]
|
||||
@@ -17,13 +17,8 @@ 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
|
||||
|
||||
|
||||
@@ -111,7 +106,7 @@ class TestVideoGenerationHappyPath:
|
||||
)
|
||||
out = client.video_generation(
|
||||
prompt=" 镜头一 ",
|
||||
# 不传 image_url:纯文生视频,不触发信任链,post 调用数为 1(创建任务)
|
||||
image_url="https://img/x.jpg",
|
||||
duration=5,
|
||||
ratio="9:16",
|
||||
resolution="720p",
|
||||
@@ -561,118 +556,3 @@ class TestResolveVideoModelId:
|
||||
# 未知 model key 会通过 get_viral_video_model_config 回落到 seedance-2.5
|
||||
with caplog.at_level(logging.WARNING, logger="shared.ai_client"):
|
||||
assert fn("some-random-model") == "doubao-seedance-2-5-260628"
|
||||
|
||||
|
||||
# ── #2165 详细错误信息和 last_video_error ─────────────────────────
|
||||
|
||||
|
||||
class TestVideoGenerationLastError:
|
||||
def test_create_400_portrait_returns_user_message(self, tmp_path):
|
||||
"""#2169: HTTP 400 + 真人拦截关键词 → 自动尝试即梦兜底;即梦未配时返回 portrait_intercept。"""
|
||||
client = _make_client(max_retries=0)
|
||||
create_resp = MagicMock()
|
||||
create_resp.status_code = 400
|
||||
create_resp.text = '{"error":{"code":"ContentRisk","message":"Real person face detected in reference image, portrait blocked"}}'
|
||||
create_resp.json.return_value = {"error": {"code": "ContentRisk", "message": "..."}}
|
||||
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", 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,
|
||||
):
|
||||
mock_s.return_value = MagicMock(
|
||||
doubao_video_poll_interval=0, doubao_video_timeout=1, doubao_video_model="seedance"
|
||||
)
|
||||
result = client.video_generation("p", output_dir=str(tmp_path), image_url="https://img/x.jpg")
|
||||
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"]
|
||||
assert err["status_code"] in (0, 400)
|
||||
|
||||
def test_create_401_returns_auth_error(self, tmp_path):
|
||||
client = _make_client(max_retries=0)
|
||||
create_resp = MagicMock()
|
||||
create_resp.status_code = 401
|
||||
create_resp.text = '{"error":{"message":"Unauthorized"}}'
|
||||
create_resp.json.return_value = {"error": {"message": "Unauthorized"}}
|
||||
create_resp.raise_for_status.side_effect = httpx.HTTPStatusError(
|
||||
"auth", request=MagicMock(), response=create_resp
|
||||
)
|
||||
with (
|
||||
patch("packages.shared.ai_client.httpx.post", return_value=create_resp),
|
||||
patch("packages.shared.ai_client.time.sleep", return_value=None),
|
||||
patch("packages.shared.ai_client.get_shared_settings") as mock_s,
|
||||
):
|
||||
mock_s.return_value = MagicMock(
|
||||
doubao_video_poll_interval=0, doubao_video_timeout=1, doubao_video_model="seedance"
|
||||
)
|
||||
result = client.video_generation("p", output_dir=str(tmp_path))
|
||||
assert result is None
|
||||
err = client.get_last_video_error()
|
||||
assert err["error_code"] == "auth_error"
|
||||
assert err["status_code"] == 401
|
||||
|
||||
def test_poll_failed_returns_task_failed_error(self, tmp_path):
|
||||
"""轮询 status=failed 时应记录 task_failed 错误并含 detail。"""
|
||||
client = _make_client(max_retries=0)
|
||||
create_resp = MagicMock()
|
||||
create_resp.status_code = 200
|
||||
create_resp.json.return_value = {"id": "t-fail"}
|
||||
create_resp.raise_for_status = MagicMock()
|
||||
poll_resp = MagicMock()
|
||||
poll_resp.status_code = 200
|
||||
poll_resp.json.return_value = {
|
||||
"status": "failed",
|
||||
"error": {"code": "InvalidParam", "message": "resolution invalid"},
|
||||
}
|
||||
poll_resp.raise_for_status = MagicMock()
|
||||
with (
|
||||
patch("packages.shared.ai_client.httpx.post", return_value=create_resp),
|
||||
patch("packages.shared.ai_client.httpx.get", return_value=poll_resp),
|
||||
patch("packages.shared.ai_client.time.sleep", return_value=None),
|
||||
patch("packages.shared.ai_client.time.time", side_effect=_fake_time_factory()),
|
||||
patch("packages.shared.ai_client.get_shared_settings") as mock_s,
|
||||
):
|
||||
mock_s.return_value = MagicMock(
|
||||
doubao_video_poll_interval=0, doubao_video_timeout=10, doubao_video_model="seedance"
|
||||
)
|
||||
result = client.video_generation("p", output_dir=str(tmp_path))
|
||||
assert result is None
|
||||
err = client.get_last_video_error()
|
||||
assert err["error_code"] == "task_failed"
|
||||
assert "InvalidParam" in err.get("detail", "") or err["status_code"] == 200
|
||||
|
||||
|
||||
class TestAiServiceLastVideoError:
|
||||
def test_call_video_generation_returns_none_sets_error(self):
|
||||
"""失败后 get_last_video_error 应返回结构化错误信息。"""
|
||||
from packages.shared import ai_service
|
||||
|
||||
mock_client = MagicMock()
|
||||
mock_client.is_available = True
|
||||
mock_client.last_video_error = {"error_code": "unknown", "user_message": "test"}
|
||||
mock_client.get_last_video_error.return_value = {"error_code": "unknown", "user_message": "test"}
|
||||
mock_client.video_generation.return_value = None
|
||||
with patch("packages.shared.ai_service.get_doubao_client", return_value=mock_client):
|
||||
assert ai_service.call_video_generation("p") is None
|
||||
err = ai_service.get_last_video_error()
|
||||
assert err["error_code"] == "unknown"
|
||||
assert "user_message" in err
|
||||
|
||||
@@ -1,198 +0,0 @@
|
||||
"""catalog 应用服务单测:会员套餐 / 积分包从共享库读取与字段映射。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _clear_cache():
|
||||
from packages.application.catalog import admin_catalog
|
||||
|
||||
admin_catalog._cache.clear()
|
||||
yield
|
||||
admin_catalog._cache.clear()
|
||||
|
||||
|
||||
def _row(**kw):
|
||||
row = MagicMock()
|
||||
for k, v in kw.items():
|
||||
setattr(row, k, v)
|
||||
return row
|
||||
|
||||
|
||||
class TestMembershipPlans:
|
||||
def test_yearly_plan_mapping(self):
|
||||
from packages.application.catalog import admin_catalog
|
||||
|
||||
row = _row(
|
||||
plan_key="premium_yearly",
|
||||
name="高级会员年卡",
|
||||
description="年度订阅",
|
||||
monthly_price=0,
|
||||
yearly_price=399,
|
||||
quotas={"4k": True, "batch_render": True, "credits_per_month": 500},
|
||||
display_order=1,
|
||||
)
|
||||
session = MagicMock()
|
||||
session.execute.return_value.fetchall.return_value = [row]
|
||||
sl = MagicMock(return_value=session)
|
||||
|
||||
with patch("packages.adapters.sqlalchemy_impl.session.SessionLocal", sl, create=True):
|
||||
plans = admin_catalog.get_membership_plans()
|
||||
|
||||
assert len(plans) == 1
|
||||
p = plans[0]
|
||||
assert p["plan_id"] == "premium_yearly"
|
||||
assert p["billing_cycle"] == "yearly"
|
||||
assert p["price_cents"] == 39900
|
||||
assert p["monthly_price_cents"] == 3325
|
||||
assert p["duration_days"] == 365
|
||||
assert p["features"]["4K 超清分辨率"] is True
|
||||
assert p["features"]["credits_per_month"] == 500
|
||||
session.close.assert_called_once()
|
||||
|
||||
def test_monthly_plan_mapping(self):
|
||||
from packages.application.catalog import admin_catalog
|
||||
|
||||
row = _row(
|
||||
plan_key="premium_monthly",
|
||||
name="高级会员月卡",
|
||||
description=None,
|
||||
monthly_price=39,
|
||||
yearly_price=0,
|
||||
quotas=None,
|
||||
display_order=2,
|
||||
)
|
||||
session = MagicMock()
|
||||
session.execute.return_value.fetchall.return_value = [row]
|
||||
sl = MagicMock(return_value=session)
|
||||
|
||||
with patch("packages.adapters.sqlalchemy_impl.session.SessionLocal", sl, create=True):
|
||||
plans = admin_catalog.get_membership_plans()
|
||||
|
||||
assert len(plans) == 1
|
||||
p = plans[0]
|
||||
assert p["billing_cycle"] == "monthly"
|
||||
assert p["price_cents"] == 3900
|
||||
assert p["monthly_price_cents"] == 3900
|
||||
assert p["duration_days"] == 30
|
||||
assert p["features"] == {}
|
||||
|
||||
def test_both_cycles_expanded(self):
|
||||
from packages.application.catalog import admin_catalog
|
||||
|
||||
row = _row(
|
||||
plan_key="premium",
|
||||
name="高级会员",
|
||||
description=None,
|
||||
monthly_price=39,
|
||||
yearly_price=399,
|
||||
quotas={},
|
||||
display_order=1,
|
||||
)
|
||||
session = MagicMock()
|
||||
session.execute.return_value.fetchall.return_value = [row]
|
||||
sl = MagicMock(return_value=session)
|
||||
|
||||
with patch("packages.adapters.sqlalchemy_impl.session.SessionLocal", sl, create=True):
|
||||
plans = admin_catalog.get_membership_plans()
|
||||
|
||||
cycles = {p["billing_cycle"] for p in plans}
|
||||
assert cycles == {"yearly", "monthly"}
|
||||
|
||||
def test_no_session_returns_empty(self):
|
||||
from packages.application.catalog import admin_catalog
|
||||
|
||||
with patch("packages.adapters.sqlalchemy_impl.session.SessionLocal", None, create=True):
|
||||
assert admin_catalog.get_membership_plans() == []
|
||||
|
||||
|
||||
class TestPointsPackages:
|
||||
def test_package_mapping_with_bonus(self):
|
||||
from packages.application.catalog import admin_catalog
|
||||
|
||||
row = _row(
|
||||
package_key="pkg_100",
|
||||
name="100元充值包",
|
||||
price=100,
|
||||
credits=1000,
|
||||
bonus_credits=100,
|
||||
is_recommended=True,
|
||||
description="推荐",
|
||||
sort_order=4,
|
||||
)
|
||||
session = MagicMock()
|
||||
session.execute.return_value.fetchall.return_value = [row]
|
||||
sl = MagicMock(return_value=session)
|
||||
|
||||
with patch("packages.adapters.sqlalchemy_impl.session.SessionLocal", sl, create=True):
|
||||
packages = admin_catalog.get_points_packages()
|
||||
|
||||
assert len(packages) == 1
|
||||
pkg = packages[0]
|
||||
assert pkg["code"] == "pkg_100"
|
||||
assert pkg["points"] == 1100
|
||||
assert pkg["price_cents"] == 10000
|
||||
assert pkg["is_recommended"] is True
|
||||
assert pkg["unit_price"] == "¥0.091/积分"
|
||||
|
||||
def test_zero_credits_unit_price_safe(self):
|
||||
from packages.application.catalog import admin_catalog
|
||||
|
||||
row = _row(
|
||||
package_key="pkg_0",
|
||||
name="空包",
|
||||
price=0,
|
||||
credits=0,
|
||||
bonus_credits=0,
|
||||
is_recommended=False,
|
||||
description=None,
|
||||
sort_order=0,
|
||||
)
|
||||
session = MagicMock()
|
||||
session.execute.return_value.fetchall.return_value = [row]
|
||||
sl = MagicMock(return_value=session)
|
||||
|
||||
with patch("packages.adapters.sqlalchemy_impl.session.SessionLocal", sl, create=True):
|
||||
packages = admin_catalog.get_points_packages()
|
||||
|
||||
assert packages[0]["points"] == 0
|
||||
assert packages[0]["price_cents"] == 0
|
||||
assert packages[0]["unit_price"] == "¥0.000/积分"
|
||||
|
||||
def test_no_session_returns_empty(self):
|
||||
from packages.application.catalog import admin_catalog
|
||||
|
||||
with patch("packages.adapters.sqlalchemy_impl.session.SessionLocal", None, create=True):
|
||||
assert admin_catalog.get_points_packages() == []
|
||||
|
||||
|
||||
class TestPackagesRoute:
|
||||
def test_get_packages_route_returns_items(self):
|
||||
from app.api.routes.points import get_packages
|
||||
|
||||
cu = MagicMock()
|
||||
cu.user.member_type = None
|
||||
rows = [
|
||||
{
|
||||
"code": "pkg_10",
|
||||
"name": "10元充值包",
|
||||
"points": 100,
|
||||
"price_cents": 1000,
|
||||
"unit_price": "¥0.100/积分",
|
||||
}
|
||||
]
|
||||
with patch(
|
||||
"packages.application.catalog.admin_catalog.get_points_packages",
|
||||
return_value=rows,
|
||||
):
|
||||
resp = get_packages(current_user=cu)
|
||||
|
||||
assert len(resp.packages) == 1
|
||||
item = resp.packages[0]
|
||||
assert item.code == "pkg_10"
|
||||
assert item.points == 100
|
||||
assert item.price_cents == 1000
|
||||
@@ -4,7 +4,6 @@ from __future__ import annotations
|
||||
|
||||
from unittest.mock import MagicMock, mock_open, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
_SINGLETON = "_DASHSCOPE_CLIENT_SINGLETON"
|
||||
@@ -158,21 +157,3 @@ class TestDashScopeVideoGeneration:
|
||||
c.video_generation(prompt=" ", duration=5, ratio="9:16", resolution="720p", output_dir="/tmp/videos")
|
||||
is None
|
||||
)
|
||||
|
||||
def test_create_400_sets_last_video_error(self, tmp_path):
|
||||
"""创建任务 HTTP 400 时应写 last_video_error。"""
|
||||
from packages.shared import dashscope_client as dc
|
||||
|
||||
dc._DASHSCOPE_CLIENT_SINGLETON = None
|
||||
with patch.dict("os.environ", {"DASHSCOPE_API_KEY": "test-key"}):
|
||||
c = dc.DashScopeClient()
|
||||
r = MagicMock()
|
||||
r.status_code = 401
|
||||
r.text = '{"code":"InvalidApiKey","message":"bad key"}'
|
||||
r.raise_for_status.side_effect = httpx.HTTPStatusError("auth", request=MagicMock(), response=r)
|
||||
with patch.object(dc.httpx, "post", return_value=r), patch.object(dc, "time"):
|
||||
out = c.video_generation("p", output_dir=str(tmp_path))
|
||||
assert out is None
|
||||
err = c.get_last_video_error()
|
||||
assert err["error_code"] == "auth_error"
|
||||
assert c.last_video_error is not None
|
||||
|
||||
@@ -173,46 +173,33 @@ class TestSubscriptionPlans:
|
||||
_spec.loader.exec_module(_mod)
|
||||
return _mod.list_membership_plans
|
||||
|
||||
def test_plans_endpoint_reads_admin_table(self):
|
||||
"""/subscription/plans 改读管理后台 plans 表:返回 catalog 服务提供的真实档位。"""
|
||||
list_membership_plans = self._import_plans_fn()
|
||||
real_plan = {
|
||||
"plan_id": "premium_yearly",
|
||||
"billing_cycle": "yearly",
|
||||
"name": "高级会员年卡",
|
||||
"description": "高级会员年度订阅,享受全部功能",
|
||||
"price_cents": 39900,
|
||||
"monthly_price_cents": 3325,
|
||||
"duration_days": 365,
|
||||
"features": {
|
||||
"4K 超清分辨率": True,
|
||||
"批量渲染": True,
|
||||
"优先处理队列": True,
|
||||
"credits_per_month": 500,
|
||||
},
|
||||
}
|
||||
with patch(
|
||||
"packages.application.catalog.admin_catalog.get_membership_plans",
|
||||
return_value=[real_plan],
|
||||
):
|
||||
resp = list_membership_plans(current_user=_make_cu())
|
||||
plans = resp["plans"]
|
||||
assert len(plans) == 1
|
||||
p0 = plans[0]
|
||||
assert p0["plan_id"] == "premium_yearly"
|
||||
assert p0["price_cents"] == 39900
|
||||
assert p0["duration_days"] == 365
|
||||
assert p0["features"]["4K 超清分辨率"] is True
|
||||
def test_plans_endpoint_returns_three_tiers(self):
|
||||
import os # noqa: F401 (used by _import_plans_fn)
|
||||
|
||||
def test_plans_endpoint_empty_when_all_disabled(self):
|
||||
"""后台停用全部套餐时,用户端返回空列表。"""
|
||||
list_membership_plans = self._import_plans_fn()
|
||||
with patch(
|
||||
"packages.application.catalog.admin_catalog.get_membership_plans",
|
||||
return_value=[],
|
||||
):
|
||||
resp = list_membership_plans(current_user=_make_cu())
|
||||
assert resp["plans"] == []
|
||||
resp = list_membership_plans(current_user=_make_cu())
|
||||
plans = resp["plans"]
|
||||
plan_ids = {p["plan_id"] for p in plans}
|
||||
assert plan_ids == {"monthly", "quarterly", "yearly"}
|
||||
for p in plans:
|
||||
assert p["price_cents"] > 0
|
||||
assert p["duration_days"] in (30, 90, 365)
|
||||
assert 0 < p["points_discount"] <= 1.0
|
||||
assert "max_resolution" in p["features"]
|
||||
|
||||
def test_longer_plans_cheaper_per_month(self):
|
||||
import os # noqa: F401
|
||||
|
||||
list_membership_plans = self._import_plans_fn()
|
||||
resp = list_membership_plans(current_user=_make_cu())
|
||||
plans = resp["plans"]
|
||||
monthly = next(p for p in plans if p["plan_id"] == "monthly")
|
||||
quarterly = next(p for p in plans if p["plan_id"] == "quarterly")
|
||||
yearly = next(p for p in plans if p["plan_id"] == "yearly")
|
||||
assert monthly["monthly_price_cents"] == 1990
|
||||
assert quarterly["monthly_price_cents"] < monthly["monthly_price_cents"]
|
||||
assert yearly["monthly_price_cents"] < quarterly["monthly_price_cents"]
|
||||
|
||||
|
||||
# ── P1-7: multiplier consistency ──────────────────────────────────────
|
||||
|
||||
|
||||
Reference in New Issue
Block a user