test_registry.py 35 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598599600601602603604605606607608609610611612613614615616617618619620621622623624625626627628629630631632633634635636637638639640641642643644645646647648649650651652653654655656657658659660661662663664665666667668669670671672673674675676677678679680681682683684685686687688689690691692693694695696697698699700701702703704705706707708709710711712713714715716717718719720721722723724725726727728729730731732733734735736737738739740741742743744745746747748749750751752753754755756757758759760761762763764765766767768769770771772773774775776777778779780781782783784785786787788789790791792793794795796797798799800801802803804805806807808809810811812813814815816817818819820821822823824825826827828829830831832833834835836837838839840841842843844845846847848849850851852853854855856857858859860861862863864865866867868869870871872873874875876877878879880881882883884885886887888889890891892893894895896897898899900901902903904905906907908909910911912913914915916917918919920921922923924925926927928929930931932933934935936937938939940941942943944945946947948949950951952953954955956957958959960961962963964965966967968969970971972973
  1. from __future__ import annotations
  2. import subprocess
  3. import sys
  4. import tempfile
  5. import time
  6. import unittest
  7. from concurrent.futures import ThreadPoolExecutor
  8. from io import BytesIO
  9. from pathlib import Path
  10. from types import SimpleNamespace
  11. from typing import Any
  12. from unittest.mock import MagicMock, patch
  13. import requests
  14. from PIL import Image
  15. from production_build_agents.run.operation_journal import (
  16. OperationOutcomeUnknownError,
  17. UncertainSideEffectError,
  18. )
  19. from production_build_agents.run.records import (
  20. VersionConflictError,
  21. )
  22. from production_build_agents.tools.discovery import RemoteToolClient
  23. from production_build_agents.tools.media import (
  24. concat_videos,
  25. crop_image,
  26. extract_frames,
  27. fit_audio_duration,
  28. grid_collage,
  29. image_as_data_url,
  30. mix_audio_tracks,
  31. mux_audio,
  32. overlay_text,
  33. probe_media,
  34. trim_audio,
  35. trim_video,
  36. )
  37. from production_build_agents.tools.publishing import publish_file
  38. from production_build_agents.tools.registry import (
  39. ToolRegistryError,
  40. create_default_tool_registry,
  41. )
  42. class FakeRemoteToolClient:
  43. def search(self, query: str, *, limit: int = 5) -> dict[str, Any]:
  44. return {
  45. "success": True,
  46. "results": [
  47. {"tool_id": "fake_generate", "summary": query},
  48. {"tool_id": "other_generate", "summary": "未授权工具"},
  49. ],
  50. "total": 2,
  51. }
  52. def inspect(self, tool_ids: list[str]) -> dict[str, Any]:
  53. return {
  54. "success": True,
  55. "tools": {
  56. tool_id: {"success": True, "input_schema": {}}
  57. for tool_id in tool_ids
  58. },
  59. }
  60. def invoke(self, tool_id: str, payload: dict[str, Any]) -> dict[str, Any]:
  61. return {
  62. "success": True,
  63. "data": {"tool_id": tool_id, "payload": payload},
  64. }
  65. class FakeResponse:
  66. def __init__(self, payload: dict[str, Any]) -> None:
  67. self.payload = payload
  68. def raise_for_status(self) -> None:
  69. return None
  70. def json(self) -> dict[str, Any]:
  71. return self.payload
  72. class FakeSession:
  73. def __init__(self) -> None:
  74. self.requests: list[tuple[str, str]] = []
  75. self.invoke_attempts = 0
  76. def post(self, url: str, **_: Any) -> FakeResponse:
  77. self.requests.append(("POST", url))
  78. return FakeResponse(
  79. {
  80. "data": {
  81. "items": [
  82. {
  83. "tool_id": "fake_generate",
  84. "description": "测试生成工具",
  85. }
  86. ]
  87. }
  88. }
  89. )
  90. def get(self, url: str, **_: Any) -> FakeResponse:
  91. self.requests.append(("GET", url))
  92. return FakeResponse(
  93. {
  94. "data": {
  95. "invoke": {
  96. "base_url": "https://invoke.test/fake_generate",
  97. "request_method": "POST",
  98. "input_schema": {"type": "object"},
  99. }
  100. }
  101. }
  102. )
  103. def request(self, method: str, url: str, **_: Any) -> FakeResponse:
  104. self.requests.append((method, url))
  105. self.invoke_attempts += 1
  106. if self.invoke_attempts == 1:
  107. raise requests.ConnectionError("temporary connection failure")
  108. return FakeResponse({"code": 0, "data": {"url": "https://result.test/a"}})
  109. class ToolRegistryTest(unittest.TestCase):
  110. def test_formal_seedance_tool_localizes_and_silences_result(self) -> None:
  111. class SeedanceClient:
  112. def invoke(self, tool_id: str, payload: dict[str, Any]) -> dict[str, Any]:
  113. self.tool_id = tool_id
  114. self.payload = payload
  115. return {
  116. "success": True,
  117. "outcome": "SUCCEEDED",
  118. "data": {"result_url": "https://result.test/shot.mp4"},
  119. }
  120. client = SeedanceClient()
  121. with tempfile.TemporaryDirectory() as temp_dir:
  122. root = Path(temp_dir)
  123. downloaded = root / "downloaded.mp4"
  124. silent = root / "silent.mp4"
  125. downloaded.touch()
  126. silent.touch()
  127. with (
  128. patch(
  129. "production_build_agents.tools.video_tools."
  130. "stable_fetch_media",
  131. return_value=downloaded,
  132. ),
  133. patch(
  134. "production_build_agents.tools.video_tools.probe_media",
  135. side_effect=[
  136. {"media_type": "video", "has_audio": True},
  137. {
  138. "media_type": "video",
  139. "has_audio": False,
  140. "duration_sec": 4.0,
  141. "width": 1080,
  142. "height": 1920,
  143. "codec": "h264",
  144. },
  145. ],
  146. ),
  147. patch(
  148. "production_build_agents.tools.video_tools."
  149. "strip_video_audio",
  150. return_value={"local_paths": [str(silent)]},
  151. ),
  152. ):
  153. tool = create_default_tool_registry(
  154. output_dir=root,
  155. remote_client=client,
  156. ).resolve(("generate_seedance_video",))[0]
  157. result = tool.invoke(
  158. {
  159. "prompt": "镜头缓慢推进",
  160. "first_frame_url": "https://published.test/anchor.png",
  161. "duration": 4,
  162. }
  163. )
  164. self.assertTrue(result["success"])
  165. self.assertEqual(client.tool_id, "seedance_generate_video")
  166. self.assertFalse(client.payload["generate_audio"])
  167. self.assertFalse(client.payload["web_search"])
  168. self.assertEqual(result["local_path"], str(silent.resolve()))
  169. self.assertFalse(result["has_audio"])
  170. def test_seedance_localization_failure_is_sealed_as_unknown(self) -> None:
  171. client = MagicMock()
  172. client.invoke.return_value = {
  173. "success": True,
  174. "outcome": "SUCCEEDED",
  175. "data": {"result_url": "https://result.test/shot.mp4"},
  176. }
  177. with tempfile.TemporaryDirectory() as temp_dir, patch(
  178. "production_build_agents.tools.video_tools.stable_fetch_media",
  179. side_effect=OSError("download interrupted"),
  180. ):
  181. root = Path(temp_dir)
  182. tool = create_default_tool_registry(
  183. output_dir=root / "outputs",
  184. remote_client=client,
  185. run_dir=root,
  186. executor_run_id="seedance-executor",
  187. ).resolve(("generate_seedance_video",))[0]
  188. params = {
  189. "prompt": "镜头缓慢推进",
  190. "first_frame_url": "https://published.test/anchor.png",
  191. "duration": 4,
  192. }
  193. with self.assertRaises(OperationOutcomeUnknownError):
  194. tool.invoke(params)
  195. client.invoke.assert_called_once()
  196. def test_remote_media_is_downloaded_once_per_run(self) -> None:
  197. image_buffer = BytesIO()
  198. Image.new("RGB", (8, 8), "navy").save(image_buffer, format="PNG")
  199. response = MagicMock()
  200. response.__enter__.return_value = response
  201. response.headers = {"content-type": "image/png"}
  202. response.iter_content.return_value = [image_buffer.getvalue()]
  203. with tempfile.TemporaryDirectory() as temp_dir, patch(
  204. "production_build_agents.tools.media._SESSION.get",
  205. return_value=response,
  206. ) as get:
  207. root = Path(temp_dir)
  208. source = "https://example.test/reference"
  209. probe_media(source, output_dir=root)
  210. data_url = image_as_data_url(source, root)
  211. self.assertTrue(data_url.startswith("data:image/png;base64,"))
  212. self.assertEqual(get.call_count, 1)
  213. def test_probe_media_identifies_image_video_and_audio(self) -> None:
  214. with tempfile.TemporaryDirectory() as temp_dir:
  215. root = Path(temp_dir)
  216. image = root / "reference.png"
  217. audio = root / "reference.wav"
  218. video = root / "reference.mp4"
  219. Image.new("RGB", (32, 18), "navy").save(image)
  220. subprocess.run(
  221. [
  222. "ffmpeg",
  223. "-y",
  224. "-f",
  225. "lavfi",
  226. "-i",
  227. "sine=frequency=440:duration=0.5",
  228. str(audio),
  229. ],
  230. check=True,
  231. capture_output=True,
  232. )
  233. subprocess.run(
  234. [
  235. "ffmpeg",
  236. "-y",
  237. "-f",
  238. "lavfi",
  239. "-i",
  240. "color=c=blue:s=64x36:d=0.5",
  241. "-pix_fmt",
  242. "yuv420p",
  243. str(video),
  244. ],
  245. check=True,
  246. capture_output=True,
  247. )
  248. image_info = probe_media(str(image), output_dir=root / "cache")
  249. audio_info = probe_media(str(audio), output_dir=root / "cache")
  250. video_info = probe_media(str(video), output_dir=root / "cache")
  251. self.assertEqual(image_info["media_type"], "image")
  252. self.assertEqual((image_info["width"], image_info["height"]), (32, 18))
  253. self.assertEqual(audio_info["media_type"], "audio")
  254. self.assertEqual(audio_info["channels"], 1)
  255. self.assertEqual(video_info["media_type"], "video")
  256. self.assertEqual((video_info["width"], video_info["height"]), (64, 36))
  257. def test_view_images_keeps_base64_out_of_tool_message_payload(self) -> None:
  258. with tempfile.TemporaryDirectory() as temp_dir:
  259. root = Path(temp_dir)
  260. image_path = root / "source.png"
  261. Image.new("RGB", (8, 8), "blue").save(image_path)
  262. view_images = create_default_tool_registry(
  263. output_dir=root / "outputs",
  264. remote_client=FakeRemoteToolClient(),
  265. ).resolve(("view_images",))[0]
  266. result = view_images.invoke(
  267. {"image_sources": [str(image_path)]}
  268. )
  269. self.assertTrue(result["success"])
  270. self.assertEqual(result["images"][0]["source"], str(image_path))
  271. self.assertNotIn("data_url", result["images"][0])
  272. self.assertTrue(
  273. Path(result["images"][0]["local_path"]).is_absolute()
  274. )
  275. def test_complete_registry_and_skill_allowlist(self) -> None:
  276. with tempfile.TemporaryDirectory() as temp_dir:
  277. registry = create_default_tool_registry(
  278. output_dir=Path(temp_dir),
  279. remote_client=FakeRemoteToolClient(),
  280. )
  281. self.assertEqual(
  282. set(registry.tool_ids),
  283. {
  284. "search_tool",
  285. "inspect_tool",
  286. "run_tool",
  287. "submit_reference_video",
  288. "fetch_generated_video",
  289. "generate_seedance_video",
  290. "probe_media",
  291. "detect_faces",
  292. "measure_audio_loudness",
  293. "view_images",
  294. "publish_media_reference",
  295. "crop_image",
  296. "grid_collage",
  297. "overlay_text",
  298. "extract_frames",
  299. "video_trim",
  300. "audio_trim",
  301. "audio_fit_duration",
  302. "mix_audio_tracks",
  303. "video_concat",
  304. "assemble_segment_media",
  305. "video_mux_audio",
  306. "synthesize_speech",
  307. "transcribe_audio",
  308. "create_ass_subtitles",
  309. "inspect_ass_subtitles",
  310. "render_ass_subtitles",
  311. },
  312. )
  313. self.assertEqual(
  314. [tool.name for tool in registry.resolve(("view_images", "run_tool"))],
  315. ["view_images", "run_tool"],
  316. )
  317. with self.assertRaisesRegex(ToolRegistryError, "未注册工具"):
  318. registry.resolve(("not_exists",))
  319. def test_tools_readme_covers_every_registered_tool(self) -> None:
  320. with tempfile.TemporaryDirectory() as temp_dir:
  321. registry = create_default_tool_registry(
  322. output_dir=Path(temp_dir),
  323. remote_client=FakeRemoteToolClient(),
  324. )
  325. readme = (
  326. Path(__file__).parents[2]
  327. / "production_build_agents"
  328. / "tools"
  329. / "README.md"
  330. ).read_text(encoding="utf-8")
  331. for tool_id, registered in registry.tools.items():
  332. self.assertIn(f"`{tool_id}`", readme)
  333. self.assertGreater(len(registered.description), 60)
  334. def test_remote_discovery_inspection_and_invocation(self) -> None:
  335. class SlowRemoteToolClient(FakeRemoteToolClient):
  336. def search(
  337. self,
  338. query: str,
  339. *,
  340. limit: int = 5,
  341. ) -> dict[str, Any]:
  342. time.sleep(0.01)
  343. return super().search(query, limit=limit)
  344. with tempfile.TemporaryDirectory() as temp_dir:
  345. registry = create_default_tool_registry(
  346. output_dir=Path(temp_dir),
  347. remote_client=SlowRemoteToolClient(),
  348. allowed_remote_tool_ids={"fake_generate"},
  349. )
  350. search, inspect, run = registry.resolve(
  351. ("search_tool", "inspect_tool", "run_tool")
  352. )
  353. searched = search.invoke({"query": "生成图片"})
  354. self.assertTrue(searched["success"])
  355. self.assertEqual(
  356. [item["tool_id"] for item in searched["results"]],
  357. ["fake_generate"],
  358. )
  359. timed = search.invoke({"query": "生成图片"})
  360. self.assertGreaterEqual(timed["_duration_ms"], 5)
  361. self.assertTrue(
  362. inspect.invoke({"tool_ids": ["fake_generate"]})["success"]
  363. )
  364. denied_inspect = inspect.invoke(
  365. {"tool_ids": ["other_generate"]}
  366. )
  367. self.assertFalse(denied_inspect["success"])
  368. result = run.invoke(
  369. {
  370. "tool_id": "fake_generate",
  371. "params": {"prompt": "test"},
  372. }
  373. )
  374. self.assertEqual(result["data"]["tool_id"], "fake_generate")
  375. denied = run.invoke(
  376. {"tool_id": "other_generate", "params": {}}
  377. )
  378. self.assertFalse(denied["success"])
  379. def test_run_tool_is_blocked_before_remote_call_without_inspect(
  380. self,
  381. ) -> None:
  382. class CountingClient(FakeRemoteToolClient):
  383. def __init__(self) -> None:
  384. self.calls = 0
  385. def invoke(self, tool_id: str, payload: dict[str, Any]) -> dict:
  386. self.calls += 1
  387. return super().invoke(tool_id, payload)
  388. client = CountingClient()
  389. with tempfile.TemporaryDirectory() as temp_dir:
  390. run_tool = create_default_tool_registry(
  391. output_dir=Path(temp_dir),
  392. remote_client=client,
  393. allowed_remote_tool_ids={"fake_generate"},
  394. ).resolve(("run_tool",))[0]
  395. result = run_tool.invoke(
  396. {
  397. "tool_id": "fake_generate",
  398. "params": {"prompt": "不得发送"},
  399. }
  400. )
  401. self.assertFalse(result["success"])
  402. self.assertIn("必须成功 inspect", result["error"])
  403. self.assertEqual(client.calls, 0)
  404. def test_run_tool_enforces_and_restores_remote_call_budget(
  405. self,
  406. ) -> None:
  407. class CountingClient(FakeRemoteToolClient):
  408. def __init__(self) -> None:
  409. self.calls = 0
  410. def invoke(
  411. self,
  412. tool_id: str,
  413. payload: dict[str, Any],
  414. ) -> dict[str, Any]:
  415. self.calls += 1
  416. return super().invoke(tool_id, payload)
  417. client = CountingClient()
  418. with tempfile.TemporaryDirectory() as temp_dir:
  419. root = Path(temp_dir)
  420. registry = create_default_tool_registry(
  421. output_dir=root / "outputs",
  422. remote_client=client,
  423. allowed_remote_tool_ids={"fake_generate"},
  424. inspected_remote_tool_ids={"fake_generate"},
  425. remote_tool_call_limits={"fake_generate": 2},
  426. )
  427. run_tool = registry.resolve(("run_tool",))[0]
  428. for index in range(2):
  429. self.assertTrue(
  430. run_tool.invoke(
  431. {
  432. "tool_id": "fake_generate",
  433. "params": {"prompt": str(index)},
  434. }
  435. )["success"]
  436. )
  437. denied = run_tool.invoke(
  438. {
  439. "tool_id": "fake_generate",
  440. "params": {"prompt": "over-budget"},
  441. }
  442. )
  443. restored = create_default_tool_registry(
  444. output_dir=root / "restored-outputs",
  445. remote_client=client,
  446. allowed_remote_tool_ids={"fake_generate"},
  447. inspected_remote_tool_ids={"fake_generate"},
  448. remote_tool_call_limits={"fake_generate": 2},
  449. remote_tool_call_counts={"fake_generate": 2},
  450. ).resolve(("run_tool",))[0]
  451. denied_after_restore = restored.invoke(
  452. {
  453. "tool_id": "fake_generate",
  454. "params": {"prompt": "after-restore"},
  455. }
  456. )
  457. self.assertFalse(denied["success"])
  458. self.assertFalse(denied_after_restore["success"])
  459. self.assertIn("调用上限 2", denied["error"])
  460. self.assertIn("调用上限 2", denied_after_restore["error"])
  461. self.assertEqual(client.calls, 2)
  462. def test_run_tool_budget_is_atomic_for_parallel_calls(self) -> None:
  463. class CountingClient(FakeRemoteToolClient):
  464. def __init__(self) -> None:
  465. self.calls: list[dict[str, Any]] = []
  466. def invoke(
  467. self,
  468. tool_id: str,
  469. payload: dict[str, Any],
  470. ) -> dict[str, Any]:
  471. self.calls.append(payload)
  472. time.sleep(0.01)
  473. return super().invoke(tool_id, payload)
  474. client = CountingClient()
  475. with tempfile.TemporaryDirectory() as temp_dir:
  476. run_tool = create_default_tool_registry(
  477. output_dir=Path(temp_dir),
  478. remote_client=client,
  479. allowed_remote_tool_ids={"fake_generate"},
  480. inspected_remote_tool_ids={"fake_generate"},
  481. remote_tool_call_limits={"fake_generate": 2},
  482. ).resolve(("run_tool",))[0]
  483. with ThreadPoolExecutor(max_workers=8) as pool:
  484. results = list(
  485. pool.map(
  486. lambda index: run_tool.invoke(
  487. {
  488. "tool_id": "fake_generate",
  489. "params": {"prompt": str(index)},
  490. }
  491. ),
  492. range(8),
  493. )
  494. )
  495. self.assertEqual(
  496. sum(result["success"] is True for result in results),
  497. 2,
  498. )
  499. self.assertEqual(len(client.calls), 2)
  500. def test_remote_client_allows_anonymous_internal_service(self) -> None:
  501. with patch.dict(
  502. "os.environ",
  503. {
  504. "TOOLS_API_BASE_URL": "https://tools.test/v1/tools",
  505. "TOOLS_API_TOKEN": "",
  506. },
  507. ):
  508. client = RemoteToolClient.from_env()
  509. self.assertEqual(client.token, "")
  510. self.assertNotIn("Authorization", client._headers())
  511. self.assertEqual(
  512. client._headers(json_body=True),
  513. {"Content-Type": "application/json"},
  514. )
  515. def test_remote_client_never_retries_unknown_connection_failure(self) -> None:
  516. session = FakeSession()
  517. client = RemoteToolClient(
  518. base_url="https://tools.test/v1/tools",
  519. token="test-token",
  520. session=session,
  521. )
  522. self.assertTrue(client.search("生成图片")["success"])
  523. inspected = client.inspect(["fake_generate"])
  524. self.assertTrue(inspected["success"])
  525. result = client.invoke("fake_generate", {"prompt": "test"})
  526. self.assertFalse(result["success"])
  527. self.assertEqual(result["outcome"], "OUTCOME_UNKNOWN")
  528. self.assertEqual(
  529. [
  530. request
  531. for request in session.requests
  532. if request[1] == "https://invoke.test/fake_generate"
  533. ],
  534. [
  535. ("POST", "https://invoke.test/fake_generate"),
  536. ],
  537. )
  538. def test_remote_client_marks_success_and_known_failure(self) -> None:
  539. session = FakeSession()
  540. session.invoke_attempts = 1
  541. client = RemoteToolClient(
  542. base_url="https://tools.test/v1/tools",
  543. session=session,
  544. )
  545. result = client.invoke("fake_generate", {"prompt": "test"})
  546. self.assertTrue(result["success"])
  547. self.assertEqual(result["outcome"], "SUCCEEDED")
  548. def test_invalid_response_after_remote_post_is_unknown(self) -> None:
  549. class InvalidJsonResponse(FakeResponse):
  550. def json(self) -> dict[str, Any]:
  551. raise ValueError("invalid json")
  552. class InvalidJsonSession(FakeSession):
  553. def request(self, method: str, url: str, **_: Any):
  554. self.requests.append((method, url))
  555. return InvalidJsonResponse({})
  556. client = RemoteToolClient(
  557. base_url="https://tools.test/v1/tools",
  558. session=InvalidJsonSession(),
  559. )
  560. result = client.invoke("fake_generate", {"prompt": "test"})
  561. self.assertFalse(result["success"])
  562. self.assertEqual(result["outcome"], "OUTCOME_UNKNOWN")
  563. def test_http_408_after_remote_post_is_unknown(self) -> None:
  564. class TimeoutResponse(FakeResponse):
  565. status_code = 408
  566. def raise_for_status(self) -> None:
  567. raise requests.HTTPError(
  568. "request timeout",
  569. response=self,
  570. )
  571. class TimeoutSession(FakeSession):
  572. def request(self, method: str, url: str, **_: Any):
  573. self.requests.append((method, url))
  574. return TimeoutResponse({})
  575. client = RemoteToolClient(
  576. base_url="https://tools.test/v1/tools",
  577. session=TimeoutSession(),
  578. )
  579. result = client.invoke("fake_generate", {"prompt": "test"})
  580. self.assertFalse(result["success"])
  581. self.assertEqual(result["outcome"], "OUTCOME_UNKNOWN")
  582. def test_run_tool_accepts_more_than_five_reference_urls(self) -> None:
  583. client = FakeRemoteToolClient()
  584. with tempfile.TemporaryDirectory() as temp_dir:
  585. registry = create_default_tool_registry(
  586. output_dir=Path(temp_dir),
  587. remote_client=client,
  588. allowed_remote_tool_ids={"fake_generate"},
  589. inspected_remote_tool_ids={"fake_generate"},
  590. )
  591. run_tool = registry.resolve(("run_tool",))[0]
  592. result = run_tool.invoke(
  593. {
  594. "tool_id": "fake_generate",
  595. "params": {
  596. "references": [
  597. f"https://example.test/{index}.png"
  598. for index in range(8)
  599. ]
  600. },
  601. }
  602. )
  603. self.assertTrue(result["success"])
  604. def test_tool_operation_journal_replays_and_rejects_drift(self) -> None:
  605. class CountingClient(FakeRemoteToolClient):
  606. def __init__(self) -> None:
  607. self.calls = 0
  608. def invoke(
  609. self,
  610. tool_id: str,
  611. payload: dict[str, Any],
  612. ) -> dict[str, Any]:
  613. self.calls += 1
  614. return super().invoke(tool_id, payload)
  615. client = CountingClient()
  616. with tempfile.TemporaryDirectory() as temp_dir:
  617. root = Path(temp_dir)
  618. kwargs = {
  619. "output_dir": root / "outputs",
  620. "run_dir": root,
  621. "executor_run_id": "executor-Task1-v1",
  622. "remote_client": client,
  623. "allowed_remote_tool_ids": {"fake_generate"},
  624. "inspected_remote_tool_ids": {"fake_generate"},
  625. }
  626. first = create_default_tool_registry(**kwargs)
  627. first_result = first.resolve(("run_tool",))[0].invoke(
  628. {
  629. "tool_id": "fake_generate",
  630. "params": {"prompt": "first"},
  631. }
  632. )
  633. replay = create_default_tool_registry(**kwargs)
  634. replay_result = replay.resolve(("run_tool",))[0].invoke(
  635. {
  636. "tool_id": "fake_generate",
  637. "params": {"prompt": "first"},
  638. }
  639. )
  640. self.assertTrue(first_result["success"])
  641. self.assertTrue(replay_result["_operation_replayed"])
  642. self.assertEqual(client.calls, 1)
  643. drifted = create_default_tool_registry(**kwargs)
  644. with self.assertRaises(VersionConflictError):
  645. drifted.resolve(("run_tool",))[0].invoke(
  646. {
  647. "tool_id": "fake_generate",
  648. "params": {"prompt": "changed"},
  649. }
  650. )
  651. def test_unknown_tool_outcome_is_not_resent_after_restart(self) -> None:
  652. session = FakeSession()
  653. client = RemoteToolClient(
  654. base_url="https://tools.test/v1/tools",
  655. session=session,
  656. )
  657. with tempfile.TemporaryDirectory() as temp_dir:
  658. root = Path(temp_dir)
  659. kwargs = {
  660. "output_dir": root / "outputs",
  661. "run_dir": root,
  662. "executor_run_id": "executor-Task1-v1",
  663. "remote_client": client,
  664. "allowed_remote_tool_ids": {"fake_generate"},
  665. "inspected_remote_tool_ids": {"fake_generate"},
  666. }
  667. first = create_default_tool_registry(**kwargs)
  668. with self.assertRaises(OperationOutcomeUnknownError):
  669. first.resolve(("run_tool",))[0].invoke(
  670. {
  671. "tool_id": "fake_generate",
  672. "params": {"prompt": "test"},
  673. }
  674. )
  675. requests_after_first = len(session.requests)
  676. restarted = create_default_tool_registry(**kwargs)
  677. with self.assertRaises(OperationOutcomeUnknownError):
  678. restarted.resolve(("run_tool",))[0].invoke(
  679. {
  680. "tool_id": "fake_generate",
  681. "params": {"prompt": "test"},
  682. }
  683. )
  684. self.assertEqual(len(session.requests), requests_after_first)
  685. def test_publish_file_local_fallback_and_configured_oss(self) -> None:
  686. with tempfile.TemporaryDirectory() as temp_dir:
  687. path = Path(temp_dir) / "artifact.txt"
  688. path.write_text("artifact", encoding="utf-8")
  689. with patch.dict(
  690. "os.environ",
  691. {
  692. "ALIYUN_OSS_ACCESS_KEY_ID": "",
  693. "ALIYUN_OSS_ACCESS_KEY_SECRET": "",
  694. },
  695. clear=False,
  696. ):
  697. self.assertEqual(
  698. publish_file(path, category="test"),
  699. str(path.resolve()),
  700. )
  701. uploaded: list[tuple[str, str]] = []
  702. class FakeBucket:
  703. def put_object_from_file(self, key: str, source: str) -> None:
  704. uploaded.append((key, source))
  705. fake_oss = SimpleNamespace(
  706. Auth=lambda *_: object(),
  707. Bucket=lambda *_: FakeBucket(),
  708. )
  709. with (
  710. patch.dict(
  711. "os.environ",
  712. {
  713. "ALIYUN_OSS_ACCESS_KEY_ID": "id",
  714. "ALIYUN_OSS_ACCESS_KEY_SECRET": "secret",
  715. "ALIYUN_OSS_CDN_BASE_URL": "https://cdn.test",
  716. },
  717. clear=False,
  718. ),
  719. patch.dict(sys.modules, {"oss2": fake_oss}),
  720. ):
  721. url = publish_file(path, category="test")
  722. self.assertTrue(url.startswith("https://cdn.test/"))
  723. self.assertEqual(uploaded[0][1], str(path))
  724. def test_oss_exception_is_treated_as_unknown_side_effect(self) -> None:
  725. with tempfile.TemporaryDirectory() as temp_dir:
  726. path = Path(temp_dir) / "artifact.txt"
  727. path.write_text("artifact", encoding="utf-8")
  728. class FailingBucket:
  729. def put_object_from_file(self, *_: str) -> None:
  730. raise TimeoutError("response lost")
  731. fake_oss = SimpleNamespace(
  732. Auth=lambda *_: object(),
  733. Bucket=lambda *_: FailingBucket(),
  734. )
  735. with (
  736. patch.dict(
  737. "os.environ",
  738. {
  739. "ALIYUN_OSS_ACCESS_KEY_ID": "id",
  740. "ALIYUN_OSS_ACCESS_KEY_SECRET": "secret",
  741. },
  742. clear=False,
  743. ),
  744. patch.dict(sys.modules, {"oss2": fake_oss}),
  745. self.assertRaises(UncertainSideEffectError),
  746. ):
  747. publish_file(path, category="test")
  748. class DeterministicImageToolTest(unittest.TestCase):
  749. def test_image_operations_write_real_files(self) -> None:
  750. with tempfile.TemporaryDirectory() as temp_dir:
  751. root = Path(temp_dir)
  752. first = root / "first.png"
  753. second = root / "second.png"
  754. Image.new("RGB", (100, 80), "red").save(first)
  755. Image.new("RGB", (100, 80), "blue").save(second)
  756. crop = crop_image(
  757. str(first),
  758. [0, 0, 0.5, 0.5],
  759. output_dir=root,
  760. scale=2,
  761. )
  762. collage = grid_collage(
  763. [str(first), str(second)],
  764. output_dir=root,
  765. rows=1,
  766. cols=2,
  767. )
  768. text = overlay_text(
  769. str(first),
  770. [{"content": "测试", "box": [0.1, 0.1, 0.9, 0.9]}],
  771. output_dir=root,
  772. )
  773. for result in (crop, collage, text):
  774. self.assertTrue(result["success"])
  775. self.assertTrue(Path(result["local_paths"][0]).is_file())
  776. class DeterministicVideoToolTest(unittest.TestCase):
  777. def _make_video(self, path: Path, color: str) -> None:
  778. subprocess.run(
  779. [
  780. "ffmpeg",
  781. "-y",
  782. "-f",
  783. "lavfi",
  784. "-i",
  785. f"color=c={color}:s=64x64:d=1:r=12",
  786. "-c:v",
  787. "libx264",
  788. "-pix_fmt",
  789. "yuv420p",
  790. str(path),
  791. ],
  792. check=True,
  793. capture_output=True,
  794. )
  795. def _make_audio(self, path: Path) -> None:
  796. subprocess.run(
  797. [
  798. "ffmpeg",
  799. "-y",
  800. "-f",
  801. "lavfi",
  802. "-i",
  803. "sine=frequency=440:duration=2",
  804. str(path),
  805. ],
  806. check=True,
  807. capture_output=True,
  808. )
  809. def test_video_operations_and_frames_write_real_files(self) -> None:
  810. with tempfile.TemporaryDirectory() as temp_dir:
  811. root = Path(temp_dir)
  812. first = root / "first.mp4"
  813. second = root / "second.mp4"
  814. audio = root / "tone.wav"
  815. self._make_video(first, "red")
  816. self._make_video(second, "blue")
  817. self._make_audio(audio)
  818. trimmed_audio = trim_audio(
  819. str(audio),
  820. output_dir=root,
  821. start_sec=0.25,
  822. duration_sec=1.25,
  823. )
  824. fitted_audio = fit_audio_duration(
  825. str(audio),
  826. target_duration_sec=1.75,
  827. output_dir=root,
  828. )
  829. mixed_audio = mix_audio_tracks(
  830. [
  831. {"audio": str(audio), "start_sec": 0, "volume": 1},
  832. {
  833. "audio": str(audio),
  834. "start_sec": 0.5,
  835. "volume": 0.158,
  836. },
  837. ],
  838. duration_sec=2,
  839. output_dir=root,
  840. )
  841. trimmed = trim_video(
  842. str(first),
  843. output_dir=root,
  844. duration_sec=0.5,
  845. )
  846. concatenated = concat_videos(
  847. [str(first), str(second)],
  848. output_dir=root,
  849. )
  850. muxed = mux_audio(
  851. str(first),
  852. str(audio),
  853. output_dir=root,
  854. )
  855. frames = extract_frames(
  856. str(first),
  857. output_dir=root,
  858. num_frames=3,
  859. )
  860. silent_info = probe_media(str(first), output_dir=root)
  861. muxed_info = probe_media(
  862. muxed["local_paths"][0],
  863. output_dir=root,
  864. )
  865. for result in (
  866. trimmed_audio,
  867. fitted_audio,
  868. mixed_audio,
  869. trimmed,
  870. concatenated,
  871. muxed,
  872. ):
  873. self.assertTrue(result["success"])
  874. self.assertTrue(Path(result["local_paths"][0]).is_file())
  875. self.assertAlmostEqual(trimmed_audio["duration_sec"], 1.25, places=2)
  876. self.assertAlmostEqual(fitted_audio["duration_sec"], 1.75, places=2)
  877. self.assertAlmostEqual(mixed_audio["duration_sec"], 2, places=2)
  878. self.assertFalse(silent_info["has_audio"])
  879. self.assertTrue(muxed_info["has_audio"])
  880. self.assertEqual(muxed_info["audio_codec"], "aac")
  881. self.assertEqual(len(frames["frames"]), 3)
  882. self.assertTrue(
  883. all(
  884. Path(frame["local_path"]).is_file()
  885. for frame in frames["frames"]
  886. )
  887. )
  888. if __name__ == "__main__":
  889. unittest.main()