test_evidence.py 7.2 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202
  1. from __future__ import annotations
  2. import tempfile
  3. import unittest
  4. from pathlib import Path
  5. from types import SimpleNamespace
  6. from langchain_core.tools import tool
  7. from PIL import Image
  8. from production_build_agents.evidence import (
  9. ValidatorEvidenceError,
  10. invoke_evidence_tool,
  11. view_image_evidence,
  12. )
  13. from production_build_agents.agents.validator.stage_context import (
  14. _stage_media_evidence,
  15. )
  16. from production_build_agents.agents.validator.task_context import (
  17. _media_evidence,
  18. )
  19. from production_build_agents.contracts.models import Artifact
  20. def _artifact(artifact_type: str, uri: str, *, ordinal: int = 1) -> Artifact:
  21. return Artifact(
  22. artifact_id=f"Task1-v1-artifact-{ordinal}",
  23. artifact_type=artifact_type,
  24. uri=uri,
  25. content_sha256="0" * 64,
  26. size_bytes=1,
  27. description="媒体测试产物",
  28. )
  29. class _MediaTools:
  30. def __init__(self, *, omit_last_view: bool = False) -> None:
  31. self.calls: list[tuple[str, object]] = []
  32. @tool("probe_media")
  33. def probe_media(source: str) -> dict:
  34. """返回测试媒体类型。"""
  35. self.calls.append(("probe_media", source))
  36. media_type = {
  37. ".mp4": "video",
  38. ".mp3": "audio",
  39. }.get(Path(source).suffix, "image")
  40. return {
  41. "success": True,
  42. "source": source,
  43. "media_type": media_type,
  44. "_duration_ms": 1,
  45. }
  46. @tool("extract_frames")
  47. def extract_frames(video: str, num_frames: int = 5) -> dict:
  48. """返回一个测试视频帧。"""
  49. self.calls.append(("extract_frames", (video, num_frames)))
  50. return {
  51. "success": True,
  52. "frames": [
  53. {"timestamp_sec": 0.5, "url": f"{video}#frame"}
  54. ],
  55. }
  56. @tool("view_images")
  57. def view_images(image_sources: list[str]) -> dict:
  58. """返回模型可见的测试图片。"""
  59. self.calls.append(("view_images", list(image_sources)))
  60. selected = image_sources[:-1] if omit_last_view else image_sources
  61. return {
  62. "success": True,
  63. "images": [
  64. {
  65. "source": source,
  66. "data_url": "data:image/png;base64,aGVsbG8=",
  67. }
  68. for source in selected
  69. ],
  70. }
  71. self.by_name = {
  72. item.name: item for item in (probe_media, extract_frames, view_images)
  73. }
  74. class ValidatorMediaEvidenceTest(unittest.TestCase):
  75. def test_task_and_stage_keep_distinct_video_payloads(self) -> None:
  76. video = "https://example.test/final.mp4"
  77. artifact = _artifact("video", video)
  78. task_tools = _MediaTools()
  79. task_evidence, task_blocks = _media_evidence(
  80. skill=SimpleNamespace(
  81. skill_id="video-validation",
  82. allowed_tools=(),
  83. ),
  84. delivery=SimpleNamespace(artifacts=[artifact]),
  85. tools=[
  86. task_tools.by_name["extract_frames"],
  87. task_tools.by_name["view_images"],
  88. ],
  89. )
  90. self.assertEqual(task_evidence[0]["video"], video)
  91. self.assertNotIn("artifact_id", task_evidence[0])
  92. self.assertEqual(task_tools.calls[0], ("extract_frames", (video, 5)))
  93. self.assertEqual(len(task_blocks), 1)
  94. stage_tools = _MediaTools()
  95. stage_evidence, stage_blocks = _stage_media_evidence(
  96. SimpleNamespace(artifacts=[artifact]),
  97. list(stage_tools.by_name.values()),
  98. )
  99. frame = next(item for item in stage_evidence if item["kind"] == "video_frame")
  100. self.assertEqual(frame["artifact_id"], artifact.artifact_id)
  101. self.assertNotIn("video", frame)
  102. probe = next(item for item in stage_evidence if item["kind"] == "media_probe")
  103. self.assertNotIn("_duration_ms", probe)
  104. self.assertEqual(len(stage_blocks), 1)
  105. def test_reference_task_probes_each_artifact_before_viewing(self) -> None:
  106. tools = _MediaTools()
  107. video = _artifact("video", "https://example.test/reference.mp4")
  108. audio = _artifact(
  109. "audio",
  110. "https://example.test/reference.mp3",
  111. ordinal=2,
  112. )
  113. evidence, blocks = _media_evidence(
  114. skill=SimpleNamespace(
  115. skill_id="reference-validation",
  116. allowed_tools=(),
  117. ),
  118. delivery=SimpleNamespace(artifacts=[video, audio]),
  119. tools=list(tools.by_name.values()),
  120. )
  121. probes = [item for item in evidence if item["kind"] == "media_probe"]
  122. frame = next(item for item in evidence if item["kind"] == "video_frame")
  123. self.assertEqual(
  124. [item["artifact_id"] for item in probes],
  125. [video.artifact_id, audio.artifact_id],
  126. )
  127. self.assertEqual(frame["artifact_id"], video.artifact_id)
  128. self.assertEqual(len(blocks), 1)
  129. def test_images_are_viewed_in_batches_of_twelve(self) -> None:
  130. tools = _MediaTools()
  131. sources = [f"https://example.test/{index}.png" for index in range(13)]
  132. blocks = view_image_evidence(
  133. tools.by_name["view_images"],
  134. sources,
  135. incomplete_prefix="缺少图片",
  136. )
  137. batches = [value for name, value in tools.calls if name == "view_images"]
  138. self.assertEqual([len(batch) for batch in batches], [12, 1])
  139. self.assertEqual(len(blocks), 13)
  140. def test_local_path_is_converted_and_every_source_is_required(self) -> None:
  141. with tempfile.TemporaryDirectory() as temp_dir:
  142. image_path = Path(temp_dir) / "frame.png"
  143. Image.new("RGB", (1, 1)).save(image_path)
  144. @tool("view_images")
  145. def local_view(image_sources: list[str]) -> dict:
  146. """只返回本地路径。"""
  147. return {
  148. "success": True,
  149. "images": [
  150. {
  151. "source": image_sources[0],
  152. "local_path": str(image_path),
  153. }
  154. ],
  155. }
  156. blocks = view_image_evidence(
  157. local_view,
  158. ["frame-source"],
  159. incomplete_prefix="缺少图片",
  160. )
  161. self.assertTrue(
  162. blocks[0]["image_url"]["url"].startswith("data:image/png;base64,")
  163. )
  164. missing = _MediaTools(omit_last_view=True)
  165. with self.assertRaisesRegex(ValidatorEvidenceError, "second.png"):
  166. view_image_evidence(
  167. missing.by_name["view_images"],
  168. ["first.png", "second.png"],
  169. incomplete_prefix="缺少图片",
  170. )
  171. def test_tool_failures_are_normalized(self) -> None:
  172. @tool("probe_media")
  173. def failed_probe(source: str) -> dict:
  174. """模拟失败。"""
  175. return {"success": False, "error": f"无法读取 {source}"}
  176. with self.assertRaisesRegex(
  177. ValidatorEvidenceError,
  178. "probe_media 未返回有效证据",
  179. ):
  180. invoke_evidence_tool(failed_probe, {"source": "broken.png"})