b5cbc3b482
CI/CD Pipeline / Check push changed paths (pull_request) Has been skipped
CI/CD Pipeline / Check if frontend-only change (pull_request) Successful in 1s
CI/CD Pipeline / Dedup Check - skip PR tests when covered by push pipeline (pull_request) Successful in 1s
CI/CD Pipeline / Frontend Lint (pull_request) Has been skipped
CI/CD Pipeline / PR Build Web Image (pull_request) Has been skipped
CI/CD Pipeline / Frontend Unit Tests (pull_request) Has been skipped
CI/CD Pipeline / Build Staging API Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Web Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Retag skipped Staging API Image (pull_request) Has been skipped
CI/CD Pipeline / Retag skipped Staging Web Image (pull_request) Has been skipped
CI/CD Pipeline / Retag skipped Staging Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Deploy Staging (Watchtower auto-deploy) (pull_request) Has been skipped
CI/CD Pipeline / Staging E2E Tests (pull_request) Has been skipped
CI/CD Pipeline / Staging API Integration Tests (pull_request) Has been skipped
CI/CD Pipeline / ACR Image Cleanup (pull_request) Has been skipped
CI/CD Pipeline / PR Build API Image (pull_request) Successful in 57s
CI/CD Pipeline / PR Build Worker Image (pull_request) Successful in 59s
Preview Deploy / Deploy Preview Environment (pull_request) Successful in 1m25s
CI/CD Pipeline / Integration Tests (pull_request) Successful in 1m47s
CI/CD Pipeline / Validate - Python (mypy + alembic) (pull_request) Successful in 1m59s
PR Automation / Auto Approve on CI Green (pull_request) Successful in 3m11s
CI/CD Pipeline / Validate - Security (pull_request) Successful in 3m31s
CI/CD Pipeline / Validate - Style (pull_request) Successful in 4m57s
AI Code Review / AI Code Review (pull_request) Successful in 6m50s
PR Automation / Auto Merge on CI Green + Approved (pull_request) Successful in 10m40s
CI/CD Pipeline / Unit Tests (pull_request) Successful in 14m55s
CI/CD Pipeline / Build Production API Image (pull_request) Has been skipped
CI/CD Pipeline / Build Production Web Image (pull_request) Has been skipped
CI/CD Pipeline / Build Production Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Deploy Production (pull_request) Has been skipped
CI/CD Pipeline / Canary Release to Production (pull_request) Has been skipped
CI/CD Pipeline / Production Browser E2E (pull_request) Has been skipped
CI/CD Pipeline / CI Gate (pull_request) Successful in 1s
ACR Cleanup / ACR Image Cleanup (pull_request_target) Successful in 4m14s
Preview Cleanup / Cleanup Preview Environment (pull_request) Successful in 4m30s
670 lines
29 KiB
Python
670 lines
29 KiB
Python
"""#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"]
|