"""#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_pre_trusted_images_replace_original_refs(self, tmp_path): """#2174: 传 pre_trusted_images(预热好的t2i信任图)时,替换原image_url/ref_imgs发给Seedance,不再现场跑Seedream。""" client = _make_client() captured_calls = [] 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")}) 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=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", pre_trusted_images=["https://ai.example.com/trusted.png"], duration=5, ratio="9:16", resolution="720p", output_dir=str(tmp_path), ) assert out is not None # 只有一次 Seedance 调用(不现场跑Seedream) assert len(captured_calls) == 1 assert "/contents/generations/tasks" in captured_calls[0]["url"] seedance_payload = captured_calls[0]["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_no_preheated_images_falls_back_to_first_frame_mode(self, tmp_path): """#2174: 无 pre_trusted_images 时原图走 first_frame 模式(ratio=adaptive),不再现场跑 Seedream。""" client = _make_client(max_retries=0) captured_calls = [] 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")}) 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 # 不现场跑 Seedream,只有一次 Seedance 调用 assert len(captured_calls) == 1 seedance_payload = captured_calls[0]["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): """#2174: 预热结果为空/None 时回退原图直传(走#2166的400→t2v自动降级路径)。""" client = _make_client(max_retries=0) # 预热结果传 None → 应该直接用原图发给 Seedance def fake_post(url, **kwargs): # 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_uses_preheated_t2i_images(self, tmp_path): """#2174: 传 pre_trusted_images(预热好的t2i信任图)时,替换原参考图发给Seedance。""" client = _make_client() captured = [] 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")}) 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"], pre_trusted_images=["https://ai.example.com/t2i-portrait.png"], duration=5, ratio="9:16", output_dir=str(tmp_path), ) assert out is not None # 只有一次 Seedance 创建任务(预热已完成,不再现场跑 Seedream) assert len(captured) == 1 seedance_payload = captured[0]["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/t2i-portrait.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"]