| 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 |