test_seedance.py 3.3 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114
  1. from __future__ import annotations
  2. import unittest
  3. from typing import Any
  4. from production_build_agents.tools.seedance import (
  5. HuoshanSeedanceClient,
  6. )
  7. class FakeResponse:
  8. def __init__(
  9. self,
  10. payload: dict[str, Any] | None = None,
  11. *,
  12. status_code: int = 200,
  13. ) -> None:
  14. self.payload = payload or {}
  15. self.status_code = status_code
  16. def json(self) -> dict[str, Any]:
  17. return self.payload
  18. class SeedanceSession:
  19. def __init__(self) -> None:
  20. self.submitted_json: dict[str, Any] | None = None
  21. self.post_count = 0
  22. self.get_count = 0
  23. def post(self, _: str, **kwargs: Any) -> FakeResponse:
  24. self.post_count += 1
  25. self.submitted_json = kwargs["json"]
  26. return FakeResponse({"id": "cgt-task-1"})
  27. def get(self, _: str, **__: Any) -> FakeResponse:
  28. self.get_count += 1
  29. if self.get_count == 1:
  30. return FakeResponse(
  31. {
  32. "id": "cgt-task-1",
  33. "status": "running",
  34. }
  35. )
  36. return FakeResponse(
  37. {
  38. "id": "cgt-task-1",
  39. "model": "doubao-seedance-2-0-260128",
  40. "status": "succeeded",
  41. "content": {
  42. "video_url": "https://result.test/shot.mp4"
  43. },
  44. "usage": {"completion_tokens": 42},
  45. }
  46. )
  47. def _payload() -> dict[str, Any]:
  48. return {
  49. "prompt": "人物面向镜头自然讲述",
  50. "first_frame_url": "https://input.test/first.jpg",
  51. "reference_image_urls": ["https://input.test/reference.jpg"],
  52. "reference_video_urls": ["https://input.test/motion.mp4"],
  53. "last_frame_url": "https://input.test/last.jpg",
  54. "duration": 5,
  55. "ratio": "9:16",
  56. "resolution": "1080p",
  57. "generate_audio": False,
  58. "watermark": False,
  59. "web_search": False,
  60. }
  61. class HuoshanSeedanceClientTest(unittest.TestCase):
  62. def test_maps_stable_tool_contract_and_polls_same_task(self) -> None:
  63. session = SeedanceSession()
  64. client = HuoshanSeedanceClient(
  65. api_key="secret",
  66. session=session, # type: ignore[arg-type]
  67. poll_interval=0,
  68. )
  69. result = client.invoke("seedance_generate_video", _payload())
  70. self.assertTrue(result["success"])
  71. self.assertEqual(result["data"]["task_id"], "cgt-task-1")
  72. self.assertEqual(
  73. result["data"]["result_url"],
  74. "https://result.test/shot.mp4",
  75. )
  76. self.assertEqual(session.post_count, 1)
  77. self.assertEqual(session.get_count, 2)
  78. assert session.submitted_json is not None
  79. self.assertEqual(
  80. session.submitted_json["model"],
  81. "doubao-seedance-2-0-260128",
  82. )
  83. self.assertEqual(
  84. [
  85. item.get("role")
  86. for item in session.submitted_json["content"][1:]
  87. ],
  88. [
  89. "first_frame",
  90. "reference_image",
  91. "reference_video",
  92. "last_frame",
  93. ],
  94. )
  95. self.assertFalse(session.submitted_json["generate_audio"])
  96. self.assertFalse(session.submitted_json["watermark"])
  97. if __name__ == "__main__":
  98. unittest.main()