test_video_generation.py 8.5 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260
  1. from __future__ import annotations
  2. import tempfile
  3. import unittest
  4. from pathlib import Path
  5. from typing import Any
  6. from unittest.mock import MagicMock, patch
  7. import requests
  8. from production_build_agents.run.operation_journal import (
  9. UncertainSideEffectError,
  10. )
  11. from production_build_agents.tools.video_generation import (
  12. BailianVideoClient,
  13. )
  14. from production_build_agents.tools.registry import (
  15. create_default_tool_registry,
  16. )
  17. class FakeResponse:
  18. def __init__(
  19. self,
  20. payload: dict[str, Any] | None = None,
  21. *,
  22. content: bytes = b"",
  23. status_code: int = 200,
  24. ) -> None:
  25. self.payload = payload or {}
  26. self.content = content
  27. self.status_code = status_code
  28. def __enter__(self) -> FakeResponse:
  29. return self
  30. def __exit__(self, *_: Any) -> None:
  31. return None
  32. def json(self) -> dict[str, Any]:
  33. return self.payload
  34. def raise_for_status(self) -> None:
  35. if self.status_code >= 400:
  36. raise requests.HTTPError(f"HTTP {self.status_code}")
  37. def iter_content(self, _: int):
  38. yield self.content
  39. class FakeSession:
  40. def __init__(self) -> None:
  41. self.submitted_json: dict[str, Any] | None = None
  42. self.get_count = 0
  43. def post(self, _: str, **kwargs: Any) -> FakeResponse:
  44. self.submitted_json = kwargs["json"]
  45. return FakeResponse(
  46. {
  47. "request_id": "request-1",
  48. "output": {
  49. "task_id": "task-1",
  50. "task_status": "PENDING",
  51. },
  52. }
  53. )
  54. def get(self, url: str, **_: Any) -> FakeResponse:
  55. self.get_count += 1
  56. if "/api/v1/tasks/" in url:
  57. return FakeResponse(
  58. {
  59. "request_id": "request-2",
  60. "output": {
  61. "task_id": "task-1",
  62. "task_status": "SUCCEEDED",
  63. "video_url": "https://result.test/video.mp4",
  64. },
  65. "usage": {"output_video_duration": 2},
  66. }
  67. )
  68. return FakeResponse(content=b"fake-mp4")
  69. class BailianVideoClientTest(unittest.TestCase):
  70. def test_submit_and_fetch_preserve_task_identity(self) -> None:
  71. session = FakeSession()
  72. client = BailianVideoClient(
  73. api_key="secret",
  74. session=session, # type: ignore[arg-type]
  75. poll_interval=0,
  76. )
  77. submitted = client.submit_reference_video(
  78. prompt="Image 1 follows Video 1",
  79. reference_images=["https://input.test/anchor.jpg"],
  80. reference_videos=["https://input.test/motion.mp4"],
  81. duration_sec=2,
  82. resolution="720P",
  83. ratio="9:16",
  84. watermark=False,
  85. seed=7,
  86. )
  87. with tempfile.TemporaryDirectory() as temp_dir:
  88. fetched = client.fetch_generated_video(
  89. submitted["task_id"],
  90. output_dir=Path(temp_dir),
  91. )
  92. output = Path(fetched["local_path"])
  93. self.assertEqual(output.read_bytes(), b"fake-mp4")
  94. self.assertEqual(submitted["task_id"], "task-1")
  95. assert session.submitted_json is not None
  96. self.assertEqual(
  97. [item["type"] for item in session.submitted_json["input"]["media"]],
  98. ["reference_image", "reference_video"],
  99. )
  100. self.assertEqual(fetched["usage"]["output_video_duration"], 2)
  101. def test_video_reference_limits_output_to_ten_seconds(self) -> None:
  102. client = BailianVideoClient(
  103. api_key="secret",
  104. session=FakeSession(), # type: ignore[arg-type]
  105. )
  106. with self.assertRaisesRegex(ValueError, "2~10"):
  107. client.submit_reference_video(
  108. prompt="test",
  109. reference_images=["https://input.test/anchor.jpg"],
  110. reference_videos=["https://input.test/motion.mp4"],
  111. duration_sec=11,
  112. resolution="720P",
  113. ratio="9:16",
  114. watermark=False,
  115. seed=None,
  116. )
  117. def test_uncertain_submit_is_not_reported_as_explicit_failure(self) -> None:
  118. class UncertainSession(FakeSession):
  119. def post(self, _: str, **__: Any) -> FakeResponse:
  120. raise requests.Timeout("unknown")
  121. client = BailianVideoClient(
  122. api_key="secret",
  123. session=UncertainSession(), # type: ignore[arg-type]
  124. )
  125. with self.assertRaises(UncertainSideEffectError):
  126. client.submit_reference_video(
  127. prompt="test",
  128. reference_images=["https://input.test/anchor.jpg"],
  129. reference_videos=[],
  130. duration_sec=2,
  131. resolution="720P",
  132. ratio="9:16",
  133. watermark=False,
  134. seed=None,
  135. )
  136. def test_registry_enforces_submit_budget_before_provider_call(self) -> None:
  137. client = MagicMock()
  138. client.submit_reference_video.return_value = {
  139. "success": True,
  140. "task_id": "task-1",
  141. "task_status": "PENDING",
  142. }
  143. with tempfile.TemporaryDirectory() as temp_dir:
  144. root = Path(temp_dir)
  145. submit = create_default_tool_registry(
  146. output_dir=root / "outputs",
  147. video_client=client,
  148. run_dir=root,
  149. executor_run_id="segment-executor",
  150. video_submit_call_limit=1,
  151. ).resolve(("submit_reference_video",))[0]
  152. params = {
  153. "prompt": "test",
  154. "reference_images": ["https://input.test/anchor.jpg"],
  155. "reference_videos": [],
  156. "duration_sec": 2,
  157. "resolution": "720P",
  158. "ratio": "9:16",
  159. "watermark": False,
  160. "seed": None,
  161. }
  162. first = submit.invoke(params)
  163. second = submit.invoke(params)
  164. self.assertTrue(first["success"])
  165. self.assertFalse(second["success"])
  166. self.assertIn("提交上限", second["error"])
  167. self.assertEqual(client.submit_reference_video.call_count, 1)
  168. def test_registry_rejects_submit_after_final_render_on_recovery(
  169. self,
  170. ) -> None:
  171. client = MagicMock()
  172. with tempfile.TemporaryDirectory() as temp_dir:
  173. submit = create_default_tool_registry(
  174. output_dir=Path(temp_dir) / "outputs",
  175. video_client=client,
  176. video_submission_closed=True,
  177. ).resolve(("submit_reference_video",))[0]
  178. result = submit.invoke(
  179. {
  180. "prompt": "test",
  181. "reference_images": [
  182. "https://input.test/anchor.jpg"
  183. ],
  184. "reference_videos": [],
  185. "duration_sec": 2,
  186. }
  187. )
  188. self.assertFalse(result["success"])
  189. self.assertIn("最终字幕视频已生成", result["error"])
  190. client.submit_reference_video.assert_not_called()
  191. def test_successful_final_render_closes_video_submission(self) -> None:
  192. client = MagicMock()
  193. with (
  194. tempfile.TemporaryDirectory() as temp_dir,
  195. patch(
  196. "production_build_agents.tools.subtitle_tools."
  197. "render_ass_subtitles",
  198. return_value={
  199. "success": True,
  200. "local_paths": ["/tmp/final.mp4"],
  201. },
  202. ),
  203. ):
  204. registry = create_default_tool_registry(
  205. output_dir=Path(temp_dir) / "outputs",
  206. video_client=client,
  207. )
  208. rendered = registry.resolve(
  209. ("render_ass_subtitles",)
  210. )[0].invoke(
  211. {"video": "video.mp4", "subtitle": "subtitle.ass"}
  212. )
  213. submitted = registry.resolve(
  214. ("submit_reference_video",)
  215. )[0].invoke(
  216. {
  217. "prompt": "test",
  218. "reference_images": [
  219. "https://input.test/anchor.jpg"
  220. ],
  221. "reference_videos": [],
  222. "duration_sec": 2,
  223. }
  224. )
  225. self.assertTrue(rendered["success"])
  226. self.assertFalse(submitted["success"])
  227. client.submit_reference_video.assert_not_called()
  228. if __name__ == "__main__":
  229. unittest.main()