482 lines
19 KiB
Python
482 lines
19 KiB
Python
"""#2106 DoubaoClient.video_generation 单测,覆盖 submit/poll/download 主路径和失败分支。"""
|
||
|
||
from __future__ import annotations
|
||
|
||
from pathlib import Path
|
||
from unittest.mock import MagicMock, patch
|
||
|
||
import httpx
|
||
import pytest
|
||
|
||
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.timeout = overrides.get("timeout", 30)
|
||
client.max_retries = overrides.get("max_retries", 0)
|
||
return client
|
||
|
||
|
||
def _fake_time_factory(base=1000.0, jump_after=2, jump=1e9):
|
||
"""返回一个 time.time() 替身:前 jump_after 次返回 base+offset,之后返回巨大值让 deadline 立即触发。
|
||
|
||
避免 Python logging 内部也调 time.time() 导致 StopIteration。
|
||
"""
|
||
state = {"n": 0}
|
||
|
||
def _t():
|
||
n = state["n"]
|
||
state["n"] += 1
|
||
if n < jump_after:
|
||
return base + n
|
||
return base + jump + n
|
||
|
||
return _t
|
||
|
||
|
||
class TestVideoGenerationHappyPath:
|
||
def test_happy_path_generates_and_downloads(self, tmp_path):
|
||
client = _make_client()
|
||
|
||
fake_task_resp = MagicMock()
|
||
fake_task_resp.json.return_value = {"id": "task-001"}
|
||
fake_task_resp.raise_for_status = MagicMock()
|
||
|
||
fake_poll_resp = MagicMock()
|
||
fake_poll_resp.json.return_value = {
|
||
"status": "succeeded",
|
||
"content": {"video_url": "https://cdn.example.com/v.mp4"},
|
||
}
|
||
fake_poll_resp.raise_for_status = MagicMock()
|
||
|
||
class FakeStreamResponse:
|
||
def __init__(self):
|
||
self._chunks = [b"FAKE", b"MP4", b"DATA"]
|
||
self._it = iter(self._chunks)
|
||
|
||
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
|
||
|
||
calls = {"post": 0, "get": 0}
|
||
|
||
def fake_post(url, **kwargs):
|
||
calls["post"] += 1
|
||
return fake_task_resp
|
||
|
||
def fake_get(url, **kwargs):
|
||
calls["get"] += 1
|
||
if "/tasks/task-001" in url:
|
||
return fake_poll_resp
|
||
raise AssertionError(f"unexpected GET (not stream): {url}")
|
||
|
||
fake_uuid = MagicMock()
|
||
fake_uuid.hex = "abcd1234"
|
||
|
||
with (
|
||
patch("packages.shared.ai_client.httpx.post", side_effect=fake_post),
|
||
patch("packages.shared.ai_client.httpx.get", side_effect=fake_get),
|
||
patch("packages.shared.ai_client.httpx.stream", return_value=FakeStreamResponse()),
|
||
patch("packages.shared.ai_client.time.sleep", return_value=None),
|
||
patch("packages.shared.ai_client.time.time", side_effect=_fake_time_factory(jump_after=2)),
|
||
patch("packages.shared.ai_client.uuid.uuid4", return_value=fake_uuid),
|
||
patch("packages.shared.ai_client.get_shared_settings") as mock_settings,
|
||
):
|
||
mock_settings.return_value = MagicMock(
|
||
doubao_video_poll_interval=0,
|
||
doubao_video_timeout=60,
|
||
doubao_video_model="doubao-seedance-2-5-260628",
|
||
)
|
||
out = client.video_generation(
|
||
prompt=" 镜头一 ",
|
||
image_url="https://img/x.jpg",
|
||
duration=5,
|
||
ratio="9:16",
|
||
resolution="720p",
|
||
output_dir=str(tmp_path),
|
||
)
|
||
assert out is not None
|
||
assert Path(out).exists()
|
||
assert Path(out).name == "seedance_task-001_abcd1234.mp4"
|
||
assert Path(out).read_bytes() == b"FAKEMP4DATA"
|
||
assert calls["post"] == 1
|
||
assert calls["get"] == 1
|
||
|
||
|
||
class TestVideoGenerationFailures:
|
||
def test_returns_none_when_unavailable(self, tmp_path):
|
||
client = _make_client(api_key="")
|
||
assert client.video_generation("p", output_dir=str(tmp_path)) is None
|
||
|
||
def test_returns_none_on_empty_prompt(self, tmp_path):
|
||
client = _make_client()
|
||
assert client.video_generation(" ", output_dir=str(tmp_path)) is None
|
||
|
||
def test_returns_none_when_create_returns_no_id(self, tmp_path):
|
||
client = _make_client(max_retries=0)
|
||
fake_resp = MagicMock()
|
||
fake_resp.json.return_value = {"error": "bad"}
|
||
fake_resp.raise_for_status = MagicMock()
|
||
with (
|
||
patch("packages.shared.ai_client.httpx.post", return_value=fake_resp),
|
||
patch("packages.shared.ai_client.get_shared_settings") as mock_s,
|
||
patch("packages.shared.ai_client.time.sleep", return_value=None),
|
||
patch("packages.shared.ai_client.time.time", side_effect=_fake_time_factory()),
|
||
):
|
||
mock_s.return_value = MagicMock(
|
||
doubao_video_poll_interval=1, doubao_video_timeout=60, doubao_video_model="seedance"
|
||
)
|
||
assert client.video_generation("p", output_dir=str(tmp_path)) is None
|
||
|
||
def test_returns_none_when_poll_returns_failed(self, tmp_path):
|
||
client = _make_client(max_retries=0)
|
||
create_resp = MagicMock()
|
||
create_resp.json.return_value = {"id": "t2"}
|
||
create_resp.raise_for_status = MagicMock()
|
||
poll_resp = MagicMock()
|
||
poll_resp.json.return_value = {"status": "failed", "error": {"code": "C1", "message": "bad"}}
|
||
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.get_shared_settings") as mock_s,
|
||
patch("packages.shared.ai_client.time.sleep", return_value=None),
|
||
patch("packages.shared.ai_client.time.time", side_effect=_fake_time_factory()),
|
||
):
|
||
mock_s.return_value = MagicMock(
|
||
doubao_video_poll_interval=0, doubao_video_timeout=10, doubao_video_model="seedance"
|
||
)
|
||
assert client.video_generation("p", output_dir=str(tmp_path)) is None
|
||
|
||
def test_returns_none_when_download_raises(self, tmp_path):
|
||
client = _make_client(max_retries=0)
|
||
create_resp = MagicMock()
|
||
create_resp.json.return_value = {"id": "t3"}
|
||
create_resp.raise_for_status = MagicMock()
|
||
poll_resp = MagicMock()
|
||
poll_resp.json.return_value = {"status": "succeeded", "content": {"video_url": "https://cdn/v.mp4"}}
|
||
poll_resp.raise_for_status = MagicMock()
|
||
|
||
class BadStream:
|
||
def __enter__(self):
|
||
return self
|
||
|
||
def __exit__(self, *a):
|
||
return False
|
||
|
||
def raise_for_status(self):
|
||
raise RuntimeError("network down")
|
||
|
||
def iter_bytes(self, **kw):
|
||
return iter([])
|
||
|
||
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.httpx.stream", return_value=BadStream()),
|
||
patch("packages.shared.ai_client.get_shared_settings") as mock_s,
|
||
patch("packages.shared.ai_client.time.sleep", return_value=None),
|
||
patch("packages.shared.ai_client.time.time", side_effect=_fake_time_factory()),
|
||
):
|
||
mock_s.return_value = MagicMock(
|
||
doubao_video_poll_interval=0, doubao_video_timeout=10, doubao_video_model="seedance"
|
||
)
|
||
assert client.video_generation("p", output_dir=str(tmp_path)) is None
|
||
|
||
|
||
class TestVideoGenerationRetryAndPoll:
|
||
def test_create_retries_then_succeeds(self, tmp_path):
|
||
client = _make_client(max_retries=1)
|
||
|
||
ok_resp = MagicMock()
|
||
ok_resp.json.return_value = {"id": "t-retry"}
|
||
ok_resp.raise_for_status = MagicMock()
|
||
poll_resp = MagicMock()
|
||
poll_resp.json.return_value = {"status": "expired"}
|
||
poll_resp.raise_for_status = MagicMock()
|
||
|
||
calls = {"post": 0}
|
||
|
||
def fake_post(url, **kwargs):
|
||
calls["post"] += 1
|
||
if calls["post"] == 1:
|
||
raise httpx.HTTPError("network")
|
||
return ok_resp
|
||
|
||
with (
|
||
patch("packages.shared.ai_client.httpx") as mock_httpx,
|
||
patch("packages.shared.ai_client.time.sleep", return_value=None),
|
||
patch("packages.shared.ai_client.time.time", side_effect=_fake_time_factory(jump_after=3)),
|
||
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"
|
||
)
|
||
mock_httpx.HTTPError = httpx.HTTPError
|
||
mock_httpx.post.side_effect = fake_post
|
||
mock_httpx.get.return_value = poll_resp
|
||
assert client.video_generation("p", output_dir=str(tmp_path)) is None
|
||
assert calls["post"] == 2
|
||
|
||
def test_succeeded_but_no_video_url_returns_none(self, tmp_path):
|
||
client = _make_client(max_retries=0)
|
||
create_resp = MagicMock()
|
||
create_resp.json.return_value = {"id": "t-nourl"}
|
||
create_resp.raise_for_status = MagicMock()
|
||
poll_resp = MagicMock()
|
||
poll_resp.json.return_value = {"status": "succeeded", "content": {}}
|
||
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"
|
||
)
|
||
assert client.video_generation("p", output_dir=str(tmp_path)) is None
|
||
|
||
|
||
class TestAiServiceCallVideoGeneration:
|
||
def test_returns_none_on_exception(self):
|
||
from packages.shared import ai_service
|
||
|
||
with patch("packages.shared.ai_service.get_doubao_client") as mock_get:
|
||
mock_client = MagicMock()
|
||
mock_client.is_available = True
|
||
mock_client.video_generation.side_effect = RuntimeError("boom")
|
||
mock_get.return_value = mock_client
|
||
assert ai_service.call_video_generation("p") is None
|
||
|
||
|
||
class TestVideoGenerationPollLoop:
|
||
def test_poll_queued_then_running_then_succeeded(self, tmp_path):
|
||
client = _make_client(max_retries=0)
|
||
create_resp = MagicMock()
|
||
create_resp.json.return_value = {"id": "t-wait"}
|
||
create_resp.raise_for_status = MagicMock()
|
||
|
||
queued = MagicMock(json=MagicMock(return_value={"status": "queued"}))
|
||
queued.raise_for_status = MagicMock()
|
||
running = MagicMock(json=MagicMock(return_value={"status": "running"}))
|
||
running.raise_for_status = MagicMock()
|
||
ok = MagicMock(
|
||
json=MagicMock(return_value={"status": "succeeded", "content": {"video_url": "https://cdn/x.mp4"}})
|
||
)
|
||
ok.raise_for_status = MagicMock()
|
||
poll_seq = [queued, running, ok]
|
||
|
||
class EmptyChunkStream:
|
||
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):
|
||
yield b""
|
||
yield b"D"
|
||
yield b""
|
||
yield b"ATA"
|
||
|
||
get_calls = {"n": 0}
|
||
|
||
def fake_get(url, **kw):
|
||
if "/tasks/t-wait" in url:
|
||
resp = poll_seq[min(get_calls["n"], len(poll_seq) - 1)]
|
||
get_calls["n"] += 1
|
||
return resp
|
||
raise AssertionError(url)
|
||
|
||
sleeps = []
|
||
# jump_after 要足够大:deadline 计算一次 + 3次 while 条件判断 = 4 次
|
||
with (
|
||
patch("packages.shared.ai_client.httpx.post", return_value=create_resp),
|
||
patch("packages.shared.ai_client.httpx.get", side_effect=fake_get),
|
||
patch("packages.shared.ai_client.httpx.stream", return_value=EmptyChunkStream()),
|
||
patch("packages.shared.ai_client.time.sleep", side_effect=lambda s: sleeps.append(s)),
|
||
patch("packages.shared.ai_client.time.time", side_effect=_fake_time_factory(jump_after=5, jump=1)),
|
||
patch("packages.shared.ai_client.uuid.uuid4", return_value=MagicMock(hex="ef012345")),
|
||
patch("packages.shared.ai_client.get_shared_settings") as mock_s,
|
||
):
|
||
mock_s.return_value = MagicMock(
|
||
doubao_video_poll_interval=0, doubao_video_timeout=200, doubao_video_model="seedance"
|
||
)
|
||
out = client.video_generation("p", output_dir=str(tmp_path))
|
||
assert out is not None
|
||
assert Path(out).read_bytes() == b"DATA"
|
||
# queued 和 running 各 sleep 一次
|
||
assert len(sleeps) >= 2
|
||
|
||
def test_poll_exception_does_not_crash(self, tmp_path):
|
||
client = _make_client(max_retries=0)
|
||
create_resp = MagicMock(json=MagicMock(return_value={"id": "t-err"}))
|
||
create_resp.raise_for_status = MagicMock()
|
||
ok = MagicMock(
|
||
json=MagicMock(return_value={"status": "succeeded", "content": {"video_url": "https://cdn/e.mp4"}})
|
||
)
|
||
ok.raise_for_status = MagicMock()
|
||
|
||
class OkStream:
|
||
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):
|
||
yield b"OK"
|
||
|
||
poll_calls = {"n": 0}
|
||
|
||
def fake_get(url, **kw):
|
||
poll_calls["n"] += 1
|
||
if poll_calls["n"] == 1:
|
||
raise httpx.HTTPError("transient")
|
||
return ok
|
||
|
||
with (
|
||
patch("packages.shared.ai_client.httpx.post", return_value=create_resp),
|
||
patch("packages.shared.ai_client.httpx.get", side_effect=fake_get),
|
||
patch("packages.shared.ai_client.httpx.stream", return_value=OkStream()),
|
||
patch("packages.shared.ai_client.time.sleep", return_value=None),
|
||
patch("packages.shared.ai_client.time.time", side_effect=_fake_time_factory(jump_after=3)),
|
||
patch("packages.shared.ai_client.uuid.uuid4", return_value=MagicMock(hex="11111111")),
|
||
patch("packages.shared.ai_client.get_shared_settings") as mock_s,
|
||
):
|
||
mock_s.return_value = MagicMock(
|
||
doubao_video_poll_interval=0, doubao_video_timeout=100, doubao_video_model="seedance"
|
||
)
|
||
out = client.video_generation("p", output_dir=str(tmp_path))
|
||
assert out is not None
|
||
assert Path(out).exists()
|
||
assert poll_calls["n"] == 2
|
||
|
||
def test_default_output_dir_and_audio_watermark(self, tmp_path, monkeypatch):
|
||
"""不传 output_dir 时落到 /tmp;generate_audio/watermark=True 也能正常提交。"""
|
||
client = _make_client()
|
||
|
||
create_resp = MagicMock(json=MagicMock(return_value={"id": "t-default"}))
|
||
create_resp.raise_for_status = MagicMock()
|
||
poll_resp = MagicMock(
|
||
json=MagicMock(return_value={"status": "succeeded", "content": {"video_url": "https://cdn/d.mp4"}})
|
||
)
|
||
poll_resp.raise_for_status = MagicMock()
|
||
|
||
# 用 tmp_path 伪造 /tmp 避免污染真 /tmp
|
||
monkeypatch.setattr("packages.shared.ai_client.os.makedirs", lambda d, exist_ok=True: None)
|
||
|
||
class S:
|
||
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):
|
||
yield b"D"
|
||
|
||
# 捕获 POST payload 断言
|
||
captured = {}
|
||
|
||
def fake_post(url, **kw):
|
||
captured["json"] = kw.get("json")
|
||
return create_resp
|
||
|
||
def fake_get(url, **kw):
|
||
return poll_resp
|
||
|
||
def fake_open(path, mode):
|
||
# 返回一个 MagicMock file,模拟写入
|
||
f = MagicMock()
|
||
f.__enter__ = MagicMock(return_value=f)
|
||
f.__exit__ = MagicMock(return_value=False)
|
||
captured["path"] = path
|
||
return f
|
||
|
||
monkeypatch.setattr("packages.shared.ai_client.httpx.post", fake_post)
|
||
monkeypatch.setattr("packages.shared.ai_client.httpx.get", fake_get)
|
||
monkeypatch.setattr("packages.shared.ai_client.httpx.stream", lambda *a, **kw: S())
|
||
monkeypatch.setattr("builtins.open", fake_open)
|
||
monkeypatch.setattr("packages.shared.ai_client.os.path.getsize", lambda p: 99)
|
||
|
||
with (
|
||
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.uuid.uuid4", return_value=MagicMock(hex="00000001")),
|
||
patch("packages.shared.ai_client.get_shared_settings") as mock_s,
|
||
):
|
||
mock_s.return_value = MagicMock(
|
||
doubao_video_poll_interval=0,
|
||
doubao_video_timeout=60,
|
||
doubao_video_model="seedance",
|
||
)
|
||
out = client.video_generation(
|
||
"p", duration=3, ratio="1:1", resolution="480p", generate_audio=True, watermark=True
|
||
)
|
||
assert out is not None
|
||
assert "/tmp/seedance_t-default_00000001.mp4" in out
|
||
assert captured["json"]["generate_audio"] is True
|
||
assert captured["json"]["watermark"] is True
|
||
assert captured["json"]["ratio"] == "1:1"
|
||
assert captured["json"]["resolution"] == "480p"
|
||
|
||
|
||
class TestGetDoubaoClientSingleton:
|
||
def test_singleton_lazy_init(self):
|
||
from packages.shared import ai_client
|
||
|
||
prev = ai_client._client
|
||
try:
|
||
ai_client._client = None
|
||
c1 = ai_client.get_doubao_client()
|
||
c2 = ai_client.get_doubao_client()
|
||
assert c1 is c2
|
||
assert isinstance(c1, ai_client.DoubaoClient)
|
||
finally:
|
||
ai_client._client = prev
|
||
|
||
|
||
class TestVideoGenerationCancelled:
|
||
def test_poll_cancelled_returns_none(self, tmp_path):
|
||
client = _make_client(max_retries=0)
|
||
create_resp = MagicMock(json=MagicMock(return_value={"id": "t-can"}))
|
||
create_resp.raise_for_status = MagicMock()
|
||
poll_resp = MagicMock(json=MagicMock(return_value={"status": "cancelled"}))
|
||
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"
|
||
)
|
||
assert client.video_generation("p", output_dir=str(tmp_path)) is None
|