| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260 |
- from __future__ import annotations
- import tempfile
- import unittest
- from pathlib import Path
- from typing import Any
- from unittest.mock import MagicMock, patch
- import requests
- from production_build_agents.run.operation_journal import (
- UncertainSideEffectError,
- )
- from production_build_agents.tools.video_generation import (
- BailianVideoClient,
- )
- from production_build_agents.tools.registry import (
- create_default_tool_registry,
- )
- class FakeResponse:
- def __init__(
- self,
- payload: dict[str, Any] | None = None,
- *,
- content: bytes = b"",
- status_code: int = 200,
- ) -> None:
- self.payload = payload or {}
- self.content = content
- self.status_code = status_code
- def __enter__(self) -> FakeResponse:
- return self
- def __exit__(self, *_: Any) -> None:
- return None
- def json(self) -> dict[str, Any]:
- return self.payload
- def raise_for_status(self) -> None:
- if self.status_code >= 400:
- raise requests.HTTPError(f"HTTP {self.status_code}")
- def iter_content(self, _: int):
- yield self.content
- class FakeSession:
- def __init__(self) -> None:
- self.submitted_json: dict[str, Any] | None = None
- self.get_count = 0
- def post(self, _: str, **kwargs: Any) -> FakeResponse:
- self.submitted_json = kwargs["json"]
- return FakeResponse(
- {
- "request_id": "request-1",
- "output": {
- "task_id": "task-1",
- "task_status": "PENDING",
- },
- }
- )
- def get(self, url: str, **_: Any) -> FakeResponse:
- self.get_count += 1
- if "/api/v1/tasks/" in url:
- return FakeResponse(
- {
- "request_id": "request-2",
- "output": {
- "task_id": "task-1",
- "task_status": "SUCCEEDED",
- "video_url": "https://result.test/video.mp4",
- },
- "usage": {"output_video_duration": 2},
- }
- )
- return FakeResponse(content=b"fake-mp4")
- class BailianVideoClientTest(unittest.TestCase):
- def test_submit_and_fetch_preserve_task_identity(self) -> None:
- session = FakeSession()
- client = BailianVideoClient(
- api_key="secret",
- session=session, # type: ignore[arg-type]
- poll_interval=0,
- )
- submitted = client.submit_reference_video(
- prompt="Image 1 follows Video 1",
- reference_images=["https://input.test/anchor.jpg"],
- reference_videos=["https://input.test/motion.mp4"],
- duration_sec=2,
- resolution="720P",
- ratio="9:16",
- watermark=False,
- seed=7,
- )
- with tempfile.TemporaryDirectory() as temp_dir:
- fetched = client.fetch_generated_video(
- submitted["task_id"],
- output_dir=Path(temp_dir),
- )
- output = Path(fetched["local_path"])
- self.assertEqual(output.read_bytes(), b"fake-mp4")
- self.assertEqual(submitted["task_id"], "task-1")
- assert session.submitted_json is not None
- self.assertEqual(
- [item["type"] for item in session.submitted_json["input"]["media"]],
- ["reference_image", "reference_video"],
- )
- self.assertEqual(fetched["usage"]["output_video_duration"], 2)
- def test_video_reference_limits_output_to_ten_seconds(self) -> None:
- client = BailianVideoClient(
- api_key="secret",
- session=FakeSession(), # type: ignore[arg-type]
- )
- with self.assertRaisesRegex(ValueError, "2~10"):
- client.submit_reference_video(
- prompt="test",
- reference_images=["https://input.test/anchor.jpg"],
- reference_videos=["https://input.test/motion.mp4"],
- duration_sec=11,
- resolution="720P",
- ratio="9:16",
- watermark=False,
- seed=None,
- )
- def test_uncertain_submit_is_not_reported_as_explicit_failure(self) -> None:
- class UncertainSession(FakeSession):
- def post(self, _: str, **__: Any) -> FakeResponse:
- raise requests.Timeout("unknown")
- client = BailianVideoClient(
- api_key="secret",
- session=UncertainSession(), # type: ignore[arg-type]
- )
- with self.assertRaises(UncertainSideEffectError):
- client.submit_reference_video(
- prompt="test",
- reference_images=["https://input.test/anchor.jpg"],
- reference_videos=[],
- duration_sec=2,
- resolution="720P",
- ratio="9:16",
- watermark=False,
- seed=None,
- )
- def test_registry_enforces_submit_budget_before_provider_call(self) -> None:
- client = MagicMock()
- client.submit_reference_video.return_value = {
- "success": True,
- "task_id": "task-1",
- "task_status": "PENDING",
- }
- with tempfile.TemporaryDirectory() as temp_dir:
- root = Path(temp_dir)
- submit = create_default_tool_registry(
- output_dir=root / "outputs",
- video_client=client,
- run_dir=root,
- executor_run_id="segment-executor",
- video_submit_call_limit=1,
- ).resolve(("submit_reference_video",))[0]
- params = {
- "prompt": "test",
- "reference_images": ["https://input.test/anchor.jpg"],
- "reference_videos": [],
- "duration_sec": 2,
- "resolution": "720P",
- "ratio": "9:16",
- "watermark": False,
- "seed": None,
- }
- first = submit.invoke(params)
- second = submit.invoke(params)
- self.assertTrue(first["success"])
- self.assertFalse(second["success"])
- self.assertIn("提交上限", second["error"])
- self.assertEqual(client.submit_reference_video.call_count, 1)
- def test_registry_rejects_submit_after_final_render_on_recovery(
- self,
- ) -> None:
- client = MagicMock()
- with tempfile.TemporaryDirectory() as temp_dir:
- submit = create_default_tool_registry(
- output_dir=Path(temp_dir) / "outputs",
- video_client=client,
- video_submission_closed=True,
- ).resolve(("submit_reference_video",))[0]
- result = submit.invoke(
- {
- "prompt": "test",
- "reference_images": [
- "https://input.test/anchor.jpg"
- ],
- "reference_videos": [],
- "duration_sec": 2,
- }
- )
- self.assertFalse(result["success"])
- self.assertIn("最终字幕视频已生成", result["error"])
- client.submit_reference_video.assert_not_called()
- def test_successful_final_render_closes_video_submission(self) -> None:
- client = MagicMock()
- with (
- tempfile.TemporaryDirectory() as temp_dir,
- patch(
- "production_build_agents.tools.subtitle_tools."
- "render_ass_subtitles",
- return_value={
- "success": True,
- "local_paths": ["/tmp/final.mp4"],
- },
- ),
- ):
- registry = create_default_tool_registry(
- output_dir=Path(temp_dir) / "outputs",
- video_client=client,
- )
- rendered = registry.resolve(
- ("render_ass_subtitles",)
- )[0].invoke(
- {"video": "video.mp4", "subtitle": "subtitle.ass"}
- )
- submitted = registry.resolve(
- ("submit_reference_video",)
- )[0].invoke(
- {
- "prompt": "test",
- "reference_images": [
- "https://input.test/anchor.jpg"
- ],
- "reference_videos": [],
- "duration_sec": 2,
- }
- )
- self.assertTrue(rendered["success"])
- self.assertFalse(submitted["success"])
- client.submit_reference_video.assert_not_called()
- if __name__ == "__main__":
- unittest.main()
|