返回 ViMax
test_image_selection.py
根目录 / tests / test_image_selection.py
1 import asyncio
2 import json
3 import os
4 import tempfile
5 import unittest
6 from pathlib import Path
7 from types import SimpleNamespace
8 from unittest.mock import AsyncMock, MagicMock, patch
9
10 from PIL import Image
11 from langchain_core.messages import AIMessage
12 from tenacity import stop_after_attempt, wait_none
13
14 from agents.best_image_selector import BestImageResponse, BestImageSelector
15 from agent_runtime.session_index import SessionIndex
16 from agent_runtime.tools import ToolRuntimeContext
17 from agent_runtime.vimax_adapters import ViMaxAdapters
18 from interfaces import Camera, ImageOutput, ShotDescription
19 from pipelines.idea2video_pipeline import Idea2VideoPipeline
20 from pipelines.script2video_pipeline import Script2VideoPipeline
21 from utils.image_selection import image_candidate_count_from_config, validate_image_candidate_count
22
23
24 class CandidateGenerator:
25 def __init__(self, failures=(), portrait_indices=(), blocked=False):
26 self.calls = []
27 self.failures = set(failures)
28 self.portrait_indices = set(portrait_indices)
29 self.started = asyncio.Event()
30 self.release = asyncio.Event()
31 if not blocked:
32 self.release.set()
33 self.running = 0
34 self.peak = 0
35
36 async def generate_single_image(self, **kwargs):
37 index = len(self.calls)
38 self.calls.append(kwargs)
39 self.running += 1
40 self.peak = max(self.peak, self.running)
41 if self.running == 2:
42 self.started.set()
43 try:
44 await self.release.wait()
45 if index in self.failures:
46 raise RuntimeError(f"candidate {index} unavailable")
47 size = (9, 16) if index in self.portrait_indices else (16, 9)
48 return ImageOutput(fmt="pil", ext="png", data=Image.new("RGB", size, (index, 0, 0)))
49 finally:
50 self.running -= 1
51
52
53 def build_pipeline(root, generator, count=2, chosen=1):
54 pipeline = Script2VideoPipeline(
55 chat_model=object(), image_generator=generator, video_generator=object(),
56 working_dir=str(root), num_image_candidates=count,
57 )
58 pipeline.best_image_selector = SimpleNamespace(select=AsyncMock(
59 return_value=BestImageResponse(best_image_index=chosen, reason="Matches the target composition."),
60 ))
61 return pipeline
62
63
64 async def generate(pipeline, directory, progress=None, prompt="two references", **kwargs):
65 return await pipeline.generate_and_select_best_image(
66 prompt=prompt, reference_image_paths=[], reference_image_path_and_text_pairs=[],
67 target_description="The target composition.", candidates_save_dir=str(directory),
68 progress=progress, **kwargs,
69 )
70
71
72 class ImageCandidateTests(unittest.IsolatedAsyncioTestCase):
73 async def test_candidates_are_concurrent_and_selected_result_is_persisted(self):
74 with tempfile.TemporaryDirectory() as tmp:
75 directory = Path(tmp) / "shots" / "0" / "first_frame_candidates"
76 generator = CandidateGenerator(blocked=True)
77 pipeline = build_pipeline(tmp, generator)
78 events = []
79 task = asyncio.create_task(generate(pipeline, directory, lambda stage, message, metadata: events.append(stage)))
80 try:
81 await asyncio.wait_for(generator.started.wait(), 1)
82 self.assertEqual(generator.peak, 2)
83 self.assertEqual(generator.calls[0], generator.calls[1])
84 generator.release.set()
85 output = await asyncio.wait_for(task, 2)
86 finally:
87 generator.release.set()
88 if not task.done():
89 task.cancel()
90 await asyncio.gather(task, return_exceptions=True)
91 self.assertEqual(output.data.getpixel((0, 0)), (1, 0, 0))
92 self.assertTrue((directory / "candidate_0.png").is_file())
93 self.assertTrue((directory / "candidate_1.png").is_file())
94 record = json.loads((directory / "selection.json").read_text())
95 self.assertEqual(record["status"], "selected")
96 self.assertEqual(record["selected_candidate_index"], 1)
97 self.assertEqual(record["selection_method"], "vlm")
98 self.assertEqual(record["reason"], "Matches the target composition.")
99 self.assertEqual(events.count("image_candidate_done"), 2)
100 self.assertIn("image_selection_start", events)
101 self.assertEqual(events[-1], "image_selection_done")
102
103 async def test_single_candidate_disables_vlm(self):
104 with tempfile.TemporaryDirectory() as tmp:
105 generator = CandidateGenerator()
106 pipeline = build_pipeline(tmp, generator, count=1)
107 await generate(pipeline, Path(tmp) / "first_frame_candidates")
108 self.assertEqual(len(generator.calls), 1)
109 pipeline.best_image_selector.select.assert_not_awaited()
110
111 async def test_failed_or_portrait_candidate_is_excluded(self):
112 for kwargs in ({"failures": [0]}, {"portrait_indices": [0]}):
113 with self.subTest(kwargs=kwargs), tempfile.TemporaryDirectory() as tmp:
114 directory = Path(tmp) / "candidates"
115 generator = CandidateGenerator(**kwargs)
116 pipeline = build_pipeline(tmp, generator)
117 output = await generate(pipeline, directory)
118 self.assertEqual(output.data.getpixel((0, 0)), (1, 0, 0))
119 pipeline.best_image_selector.select.assert_not_awaited()
120 record = json.loads((directory / "selection.json").read_text())
121 self.assertEqual(record["selection_method"], "single_valid_candidate")
122 self.assertEqual(record["selected_candidate_index"], 1)
123
124 async def test_filtered_indices_map_back_to_original_candidate(self):
125 with tempfile.TemporaryDirectory() as tmp:
126 generator = CandidateGenerator(failures=[0])
127 pipeline = build_pipeline(tmp, generator, count=3, chosen=1)
128 directory = Path(tmp) / "candidates"
129 output = await generate(pipeline, directory)
130 self.assertEqual(output.data.getpixel((0, 0)), (2, 0, 0))
131 self.assertEqual(json.loads((directory / "selection.json").read_text())["selected_candidate_index"], 2)
132
133 async def test_all_candidate_failures_remain_errors(self):
134 with tempfile.TemporaryDirectory() as tmp:
135 directory = Path(tmp) / "candidates"
136 pipeline = build_pipeline(tmp, CandidateGenerator(failures=[0, 1]))
137 with self.assertRaisesRegex(RuntimeError, "All 2 image candidates failed"):
138 await generate(pipeline, directory)
139 pipeline.best_image_selector.select.assert_not_awaited()
140 self.assertEqual(json.loads((directory / "selection.json").read_text())["status"], "generation_failed")
141
142 async def test_selection_failure_keeps_candidates_and_resume_reuses_them(self):
143 with tempfile.TemporaryDirectory() as tmp:
144 directory = Path(tmp) / "candidates"
145 generator = CandidateGenerator()
146 pipeline = build_pipeline(tmp, generator)
147 pipeline.best_image_selector.select.side_effect = [
148 RuntimeError("VLM unavailable"),
149 BestImageResponse(best_image_index=1, reason="Second candidate is better."),
150 ]
151 with self.assertRaisesRegex(RuntimeError, "VLM unavailable"):
152 await generate(pipeline, directory)
153 self.assertEqual(json.loads((directory / "selection.json").read_text())["status"], "selection_failed")
154 self.assertTrue((directory / "candidate_0.png").exists())
155 await generate(pipeline, directory)
156 self.assertEqual(len(generator.calls), 2)
157 record = json.loads((directory / "selection.json").read_text())
158 self.assertTrue(all(item["reused"] for item in record["candidates"]))
159
160 async def test_changed_prompt_does_not_reuse_old_candidates(self):
161 with tempfile.TemporaryDirectory() as tmp:
162 directory = Path(tmp) / "candidates"
163 generator = CandidateGenerator()
164 pipeline = build_pipeline(tmp, generator)
165 await generate(pipeline, directory)
166 await generate(pipeline, directory, prompt="a revised composition")
167 self.assertEqual(len(generator.calls), 4)
168
169 async def test_cancellation_stops_candidate_tasks_and_persists_status(self):
170 with tempfile.TemporaryDirectory() as tmp:
171 directory = Path(tmp) / "candidates"
172 generator = CandidateGenerator(blocked=True)
173 pipeline = build_pipeline(tmp, generator)
174 task = asyncio.create_task(generate(pipeline, directory))
175 await asyncio.wait_for(generator.started.wait(), 1)
176 task.cancel()
177 with self.assertRaises(asyncio.CancelledError):
178 await task
179 self.assertEqual(generator.running, 0)
180 pipeline.best_image_selector.select.assert_not_awaited()
181 self.assertEqual(json.loads((directory / "selection.json").read_text())["status"], "cancelled")
182
183 async def test_first_and_last_frame_candidates_are_separate_and_cached(self):
184 with tempfile.TemporaryDirectory() as tmp:
185 generator = CandidateGenerator()
186 pipeline = build_pipeline(tmp, generator)
187 pipeline.reference_image_selector.select_reference_images_and_generate_prompt = AsyncMock(
188 return_value={"reference_image_path_and_text_pairs": [], "text_prompt": "frame"},
189 )
190 pipeline.frame_events[0] = {"first_frame": asyncio.Event(), "last_frame": asyncio.Event()}
191 (Path(tmp) / "shots" / "0").mkdir(parents=True)
192 for frame in ("first_frame", "last_frame", "first_frame"):
193 await pipeline.generate_frame_for_single_shot(0, frame, ("reference.png", "reference"), "a scene", [], {})
194 self.assertEqual(len(generator.calls), 4)
195 shot = Path(tmp) / "shots" / "0"
196 self.assertTrue((shot / "first_frame_candidates" / "selection.json").is_file())
197 self.assertTrue((shot / "last_frame_candidates" / "selection.json").is_file())
198 self.assertTrue((shot / "first_frame.png").is_file())
199 self.assertTrue((shot / "last_frame.png").is_file())
200
201 async def test_camera_first_frame_uses_candidate_selection(self):
202 with tempfile.TemporaryDirectory() as tmp:
203 generator = CandidateGenerator()
204 pipeline = build_pipeline(tmp, generator)
205 pipeline.reference_image_selector.select_reference_images_and_generate_prompt = AsyncMock(
206 return_value={"reference_image_path_and_text_pairs": [], "text_prompt": "frame"},
207 )
208 pipeline.frame_events[0] = {"first_frame": asyncio.Event(), "last_frame": asyncio.Event()}
209 (Path(tmp) / "shots" / "0").mkdir(parents=True)
210 shot = ShotDescription(idx=0, is_last=True, cam_idx=0, visual_desc="scene", variation_type="small", variation_reason="still", ff_desc="scene", ff_vis_char_idxs=[], lf_desc="scene", lf_vis_char_idxs=[], motion_desc="still", audio_desc="silent")
211 await pipeline.generate_frames_for_single_camera(Camera(idx=0, active_shot_idxs=[0]), [shot], [], {}, [])
212 self.assertEqual(len(generator.calls), 2)
213 pipeline.best_image_selector.select.assert_awaited_once()
214 self.assertTrue(pipeline.frame_events[0]["first_frame"].is_set())
215
216
217 class BestImageSelectorTests(unittest.IsolatedAsyncioTestCase):
218 async def test_reuses_configured_model_and_sends_reference_and_candidates(self):
219 with tempfile.TemporaryDirectory() as tmp:
220 files = []
221 for name in ("reference", "candidate_0", "candidate_1"):
222 path = Path(tmp) / f"{name}.png"
223 Image.new("RGB", (16, 9)).save(path)
224 files.append(str(path))
225 model = SimpleNamespace(ainvoke=AsyncMock(return_value=AIMessage(content='{"best_image_index":1,"reason":"Better alignment"}')))
226 selector = BestImageSelector(chat_model=model)
227 self.assertIs(selector.chat_model, model)
228 selected = await selector([(files[0], "The character")], "Target scene", files[1:])
229 self.assertEqual(selected, files[2])
230 messages = model.ainvoke.await_args.args[0]
231 self.assertIn("exactly 2 candidate", messages[0].content)
232 self.assertEqual(len([block for block in messages[1].content if block["type"] == "image_url"]), 3)
233
234 async def test_schema_wrapped_response_parses_without_resampling(self):
235 model = SimpleNamespace(ainvoke=AsyncMock(return_value=AIMessage(content='{"properties":{"best_image_index":1,"reason":"Aligned"}}')))
236 selector = BestImageSelector(chat_model=model)
237 with patch("agents.best_image_selector.image_path_to_b64", return_value="data:image/png;base64,AA=="):
238 response = await selector.select([], "scene", ["0.png", "1.png"])
239 self.assertEqual(response.best_image_index, 1)
240 self.assertEqual(model.ainvoke.await_count, 1)
241
242 async def test_invalid_index_fails_after_bounded_retries(self):
243 model = SimpleNamespace(ainvoke=AsyncMock(return_value=AIMessage(content='{"best_image_index":9,"reason":"Invalid"}')))
244 selector = BestImageSelector(chat_model=model)
245 select = selector.select.retry_with(stop=stop_after_attempt(3), wait=wait_none())
246 with patch("agents.best_image_selector.image_path_to_b64", return_value="data:image/png;base64,AA=="):
247 with self.assertRaisesRegex(ValueError, "invalid candidate index"):
248 await select(selector, [], "scene", ["0.png", "1.png"])
249 self.assertEqual(model.ainvoke.await_count, 3)
250
251
252 class ImageSelectionIntegrationTests(unittest.IsolatedAsyncioTestCase):
253 async def test_idea_scenes_receive_count_and_scoped_progress(self):
254 with tempfile.TemporaryDirectory() as tmp:
255 pipeline = Idea2VideoPipeline(
256 chat_model=object(), image_generator=object(), video_generator=object(),
257 working_dir=tmp, num_image_candidates=3,
258 )
259 pipeline.develop_story = AsyncMock(return_value="story")
260 pipeline.extract_characters = AsyncMock(return_value=[])
261 pipeline.generate_character_portraits = AsyncMock(return_value={})
262 pipeline.write_script_based_on_story = AsyncMock(return_value=["scene one", "scene two"])
263 events = []
264
265 async def render_scene(**kwargs):
266 kwargs["progress"]("image_selection_done", "selected", {"shot_idx": 0})
267 return "scene.mp4"
268
269 renderer = AsyncMock(side_effect=render_scene)
270 with patch("pipelines.idea2video_pipeline.Script2VideoPipeline", return_value=renderer) as factory, \
271 patch("pipelines.idea2video_pipeline.concatenate_video_files") as concatenate:
272 await pipeline("idea", "short", "noir", quiet=True, progress=lambda stage, message, metadata: events.append(metadata))
273 self.assertEqual(factory.call_count, 2)
274 self.assertTrue(all(call.kwargs["num_image_candidates"] == 3 for call in factory.call_args_list))
275 self.assertEqual([event["scene_idx"] for event in events], [0, 1])
276 self.assertTrue(all(event["shot_idx"] == 0 for event in events))
277 concatenate.assert_called_once()
278
279 async def test_adapter_reads_workspace_count_and_rejects_invalid_count(self):
280 for workflow in ("idea2video", "script2video"):
281 for count in (3, 0):
282 with self.subTest(workflow=workflow, count=count), tempfile.TemporaryDirectory() as tmp, \
283 patch.dict(os.environ, {}, clear=True):
284 workspace = Path(tmp)
285 config = workspace / "configs" / "agent.local.yaml"
286 config.parent.mkdir()
287 config.write_text(f"image_selection:\n num_candidates: {count}\n", encoding="utf-8")
288 index = SessionIndex(tmp)
289 session = index.create(idea="short scene")
290 root = workspace / session["working_dir"] / workflow
291 scene = root / "scene_0" if workflow == "idea2video" else root
292 (scene / "shots" / "0").mkdir(parents=True)
293 (root / "characters.json").write_text("[]", encoding="utf-8")
294 (scene / "storyboard.json").write_text("[]", encoding="utf-8")
295 (scene / "camera_tree.json").write_text("[]", encoding="utf-8")
296 (scene / "shots" / "0" / "shot_description.json").write_text("{}", encoding="utf-8")
297 if workflow == "idea2video":
298 (root / "story.txt").write_text("story", encoding="utf-8")
299 (root / "script.json").write_text("[]", encoding="utf-8")
300 else:
301 (root / "script.txt").write_text("script", encoding="utf-8")
302 adapter = ViMaxAdapters(workspace, index)
303 events = []
304 runtime = ToolRuntimeContext("vimax_render_video", "vimax_render_video", progress_callback=events.append)
305
306 async def render(**kwargs):
307 kwargs["progress"]("image_selection_done", "selected", {"index": 1})
308 output = root / "final_video.mp4"
309 output.write_bytes(b"unit-test video")
310 return str(output)
311
312 factory_name = "Idea2VideoPipeline" if workflow == "idea2video" else "Script2VideoPipeline"
313 with patch("agent_runtime.vimax_adapters._build_chat_model", return_value=object()), \
314 patch("agent_runtime.vimax_adapters._build_image_generator", return_value=object()), \
315 patch("agent_runtime.vimax_adapters._build_video_generator", return_value=object()), \
316 patch(f"agent_runtime.vimax_adapters.{factory_name}", return_value=AsyncMock(side_effect=render)) as factory:
317 result = await adapter.vimax_render_video({}, runtime)
318 if count == 0:
319 self.assertFalse(result.ok)
320 factory.assert_not_called()
321 self.assertEqual(index.get(session["session_id"])["stage"], "error")
322 else:
323 self.assertTrue(result.ok)
324 self.assertEqual(factory.call_args.kwargs["num_image_candidates"], 3)
325 stages = [event["progress"]["stage"] for event in events if event.get("type") == "tool_progress"]
326 self.assertIn("image_selection_done", stages)
327
328
329 class ImageSelectionConfigTests(unittest.TestCase):
330 def test_count_validation(self):
331 for value in (False, True, 0, -1, 1.5, None):
332 with self.subTest(value=value), self.assertRaises(ValueError):
333 validate_image_candidate_count(value)
334 self.assertEqual(validate_image_candidate_count(3), 3)
335
336 def test_pipeline_yaml_wires_count_to_both_workflows(self):
337 with tempfile.TemporaryDirectory() as tmp, patch.dict(os.environ, {}, clear=True):
338 config = Path(tmp) / "config.yaml"
339 config.write_text(f"chat_model:\n init_args:\n model: vision-model\nimage_selection:\n num_candidates: 3\nworking_dir: {tmp}\n", encoding="utf-8")
340 backend = SimpleNamespace(image_generator=object(), video_generator=object())
341 for cls, module in ((Idea2VideoPipeline, "idea2video"), (Script2VideoPipeline, "script2video")):
342 with self.subTest(workflow=module), patch(f"pipelines.{module}_pipeline.init_chat_model", return_value=object()), patch(f"pipelines.{module}_pipeline.RenderBackend.from_config", return_value=backend):
343 pipeline = cls.init_from_config(str(config))
344 self.assertEqual(pipeline.num_image_candidates, 3)
345
346 def test_malformed_selection_section_is_rejected(self):
347 with patch.dict(os.environ, {}, clear=True), self.assertRaisesRegex(ValueError, "YAML mapping"):
348 image_candidate_count_from_config({"image_selection": False})
349
350
351 if __name__ == "__main__":
352 unittest.main()
353
353 lines PYTHON