Replace the prior release tree with the current plugin, frontend, tests, and documentation. Document retired node IDs and the public Gitea update source.
497 lines
22 KiB
Python
497 lines
22 KiB
Python
"""Offline regression tests for the panel-driven o1key video generator."""
|
||
|
||
from __future__ import annotations
|
||
|
||
import asyncio
|
||
import os
|
||
from pathlib import Path
|
||
import sys
|
||
import tempfile
|
||
import time
|
||
import types
|
||
import unittest
|
||
import uuid
|
||
from unittest.mock import AsyncMock, patch
|
||
|
||
from PIL import Image
|
||
|
||
|
||
ROOT = Path(__file__).resolve().parents[1]
|
||
CUSTOM_NODES_ROOT = ROOT.parent
|
||
COMFY_ROOT = CUSTOM_NODES_ROOT.parent
|
||
sys.path.insert(0, str(COMFY_ROOT))
|
||
sys.path.insert(0, str(CUSTOM_NODES_ROOT))
|
||
|
||
# Import the focused modules without executing the plugin-wide registry. The
|
||
# portable test environment intentionally has a CPU-only torch build, while
|
||
# unrelated legacy nodes probe CUDA during package import.
|
||
package = types.ModuleType("comfyui_o1key")
|
||
package.__path__ = [str(ROOT)]
|
||
sys.modules["comfyui_o1key"] = package
|
||
for child in ("nodes", "utils", "clients"):
|
||
module = types.ModuleType(f"comfyui_o1key.{child}")
|
||
module.__path__ = [str(ROOT / child)]
|
||
sys.modules[f"comfyui_o1key.{child}"] = module
|
||
|
||
from comfyui_o1key.nodes.o1key_video_generator import ( # noqa: E402
|
||
O1keyVideoGenerator,
|
||
O1keyVideoResult,
|
||
_parse_result_descriptor,
|
||
)
|
||
from comfyui_o1key.utils import o1key_video_jobs as JOBS # noqa: E402
|
||
from comfyui_o1key.utils.o1key_video_catalog import ( # noqa: E402
|
||
normalize_seedance_parameters,
|
||
public_video_capabilities,
|
||
resolve_seedance_model,
|
||
)
|
||
|
||
|
||
def payload(**overrides):
|
||
value = {
|
||
"batch_id": f"video_{uuid.uuid4().hex}",
|
||
"generator_node_id": 1,
|
||
"result_node_id": 2,
|
||
"provider": "seedance",
|
||
"model": "seedance-2.0",
|
||
"route": "domestic",
|
||
"generation_mode": "text",
|
||
"asset_creation_mode": "auto",
|
||
"prompt": "一只猫穿过有雾的森林",
|
||
"resolution": "720p",
|
||
"aspect_ratio": "16:9",
|
||
"duration": "5",
|
||
"generate_audio": False,
|
||
"return_last_frame": False,
|
||
"seed": 42,
|
||
"filename_prefix": "video",
|
||
"save_location": "video",
|
||
"media": {},
|
||
"assets": {},
|
||
}
|
||
value.update(overrides)
|
||
return value
|
||
|
||
|
||
class CatalogTests(unittest.TestCase):
|
||
def test_overseas_route_uses_short_display_label(self):
|
||
routes = public_video_capabilities()["providers"][0]["routes"]
|
||
self.assertIn({"value": "overseas_hc", "label": "海外"}, routes)
|
||
|
||
def test_model_route_matrix_and_duration_caps(self):
|
||
self.assertEqual(
|
||
resolve_seedance_model("seedance-2.0", "overseas_hc"),
|
||
"dreamina-seedance-2-0-hc",
|
||
)
|
||
with self.assertRaises(ValueError):
|
||
normalize_seedance_parameters(payload(model="seedance-2.0-fast", resolution="1080p"))
|
||
with self.assertRaises(ValueError):
|
||
normalize_seedance_parameters(payload(model="seedance-2.0", duration="16"))
|
||
result = normalize_seedance_parameters(payload(model="seedance-2.5", duration="30", resolution="4k"))
|
||
self.assertEqual(result["duration"], 30)
|
||
|
||
def test_asset_creation_mode_defaults_and_options(self):
|
||
normalized = normalize_seedance_parameters({
|
||
key: value for key, value in payload().items()
|
||
if key != "asset_creation_mode"
|
||
})
|
||
self.assertEqual(normalized["asset_creation_mode"], "auto")
|
||
modes = public_video_capabilities()["providers"][0]["asset_creation_modes"]
|
||
self.assertEqual(modes, [
|
||
{"value": "auto", "label": "自动创建"},
|
||
{"value": "manual", "label": "手动"},
|
||
])
|
||
with self.assertRaisesRegex(ValueError, "素材创建模式"):
|
||
normalize_seedance_parameters(payload(asset_creation_mode="invalid"))
|
||
|
||
|
||
class NodeSchemaTests(unittest.TestCase):
|
||
def test_generator_exposes_completed_media_without_queue_generation_side_effects(self):
|
||
schema = O1keyVideoGenerator.define_schema()
|
||
schema.validate()
|
||
self.assertEqual(schema.node_id, "O1keyVideoGenerator")
|
||
self.assertEqual([item.id for item in schema.outputs], ["VIDEO", "LAST_FRAME"])
|
||
self.assertTrue(schema.not_idempotent)
|
||
generation_mode = next(item for item in schema.inputs if item.id == "generation_mode")
|
||
self.assertEqual(generation_mode.default, "multimodal")
|
||
self.assertEqual(
|
||
[item.id for item in schema.inputs],
|
||
[
|
||
"prompt", "provider", "model", "route", "generation_mode",
|
||
"resolution", "aspect_ratio", "duration", "generate_audio",
|
||
"return_last_frame", "seed", "media_manifest", "asset_manifest",
|
||
"provider_options", "filename_prefix", "save_location",
|
||
"asset_creation_mode", "video_manifest", "last_frame_manifest",
|
||
],
|
||
)
|
||
output = O1keyVideoGenerator.execute(prompt="不会发起请求")
|
||
self.assertIsNone(output.result)
|
||
|
||
with patch(
|
||
"comfyui_o1key.nodes.o1key_video_generator._result_values",
|
||
return_value=("native-video", "last-frame"),
|
||
):
|
||
completed = O1keyVideoGenerator.execute(
|
||
video_manifest='{"filename":"video.mp4"}',
|
||
last_frame_manifest='{"filename":"frame.png"}',
|
||
)
|
||
self.assertEqual(completed.result, ("native-video", "last-frame"))
|
||
|
||
def test_result_node_exposes_native_video_and_last_frame(self):
|
||
schema = O1keyVideoResult.define_schema()
|
||
schema.validate()
|
||
self.assertEqual(schema.node_id, "O1keyVideoResult")
|
||
self.assertTrue(schema.is_deprecated)
|
||
self.assertEqual([item.id for item in schema.outputs], ["VIDEO", "LAST_FRAME"])
|
||
|
||
def test_result_descriptor_rejects_traversal(self):
|
||
with self.assertRaises(ValueError):
|
||
_parse_result_descriptor({"filename": "video.mp4", "subfolder": "../private", "type": "output"})
|
||
self.assertEqual(
|
||
_parse_result_descriptor({"filename": "video.mp4", "subfolder": "clips", "type": "output"}),
|
||
{"filename": "video.mp4", "subfolder": "clips", "type": "output"},
|
||
)
|
||
|
||
|
||
class PayloadAndBodyTests(unittest.TestCase):
|
||
def test_modes_are_validated_before_media_upload(self):
|
||
with self.assertRaisesRegex(ValueError, "首帧"):
|
||
JOBS.normalize_video_job_payload(payload(generation_mode="first_frame", prompt=""))
|
||
with self.assertRaisesRegex(ValueError, "不能携带"):
|
||
JOBS.normalize_video_job_payload(payload(media={"first_frame": {"name": "a.png", "type": "input"}}))
|
||
|
||
def test_first_last_body_preserves_roles_and_has_no_search_option(self):
|
||
job = JOBS.normalize_video_job_payload(payload(
|
||
generation_mode="first_last_frame",
|
||
media={
|
||
"first_frame": {"name": "first.png", "type": "input"},
|
||
"last_frame": {"name": "last.png", "type": "input"},
|
||
},
|
||
))
|
||
body = JOBS.build_seedance_video_body(job, {
|
||
"first_frame": "https://example.invalid/first.png",
|
||
"last_frame": "https://example.invalid/last.png",
|
||
"reference_images": [],
|
||
"reference_videos": [],
|
||
"reference_audios": [],
|
||
})
|
||
self.assertEqual([item.get("role") for item in body["content"][1:]], ["first_frame", "last_frame"])
|
||
self.assertNotIn("web_search", body)
|
||
self.assertNotIn("online_search", body)
|
||
|
||
def test_multimodal_manual_asset_ids_respect_model_limits(self):
|
||
with self.assertRaisesRegex(ValueError, "最多支持 9"):
|
||
JOBS.normalize_video_job_payload(payload(
|
||
generation_mode="multimodal",
|
||
asset_creation_mode="manual",
|
||
assets={"images": [f"image-{index}" for index in range(10)]},
|
||
))
|
||
|
||
def test_legacy_mixed_multimodal_request_keeps_combined_limit(self):
|
||
references = [{"name": f"{index}.png", "type": "input"} for index in range(9)]
|
||
legacy = payload(
|
||
generation_mode="multimodal",
|
||
media={"reference_images": references},
|
||
assets={"persons": ["legacy-image"]},
|
||
)
|
||
legacy.pop("asset_creation_mode")
|
||
with self.assertRaisesRegex(ValueError, "最多支持 9"):
|
||
JOBS.normalize_video_job_payload(legacy)
|
||
|
||
def test_asset_ids_cannot_hide_temporary_urls(self):
|
||
with self.assertRaisesRegex(ValueError, "无效素材 ID"):
|
||
JOBS.normalize_video_job_payload(payload(
|
||
generation_mode="multimodal",
|
||
asset_creation_mode="manual",
|
||
assets={"images": ["https://signed.example.invalid/image?id=secret"]},
|
||
))
|
||
|
||
def test_manual_asset_mode_drives_frame_and_multimodal_requests(self):
|
||
first = JOBS.normalize_video_job_payload(payload(
|
||
generation_mode="first_frame",
|
||
asset_creation_mode="manual",
|
||
assets={"images": ["image-first"]},
|
||
))
|
||
first_body = JOBS.build_seedance_video_body(first, {
|
||
"first_frame": None,
|
||
"last_frame": None,
|
||
"reference_images": [],
|
||
"reference_videos": [],
|
||
"reference_audios": [],
|
||
})
|
||
self.assertEqual(first_body["content"][1]["image_url"]["url"], "asset://image-first")
|
||
|
||
multimodal = JOBS.normalize_video_job_payload(payload(
|
||
generation_mode="multimodal",
|
||
asset_creation_mode="manual",
|
||
assets={"images": ["image-1"], "videos": ["video-1"], "audios": ["audio-1"]},
|
||
))
|
||
self.assertEqual(multimodal["assets"]["images"], ["image-1"])
|
||
body = JOBS.build_seedance_video_body(multimodal, {
|
||
"first_frame": None,
|
||
"last_frame": None,
|
||
"reference_images": [],
|
||
"reference_videos": [],
|
||
"reference_audios": [],
|
||
})
|
||
self.assertEqual(
|
||
[item.get("type") for item in body["content"][1:]],
|
||
["image_url", "video_url", "audio_url"],
|
||
)
|
||
|
||
def test_manual_asset_mode_rejects_uploads_and_missing_ids(self):
|
||
with self.assertRaisesRegex(ValueError, "文生视频模式不需要"):
|
||
JOBS.normalize_video_job_payload(payload(asset_creation_mode="manual"))
|
||
with self.assertRaisesRegex(ValueError, "只能填写素材 ID"):
|
||
JOBS.normalize_video_job_payload(payload(
|
||
generation_mode="multimodal",
|
||
asset_creation_mode="manual",
|
||
media={"reference_images": [{"name": "image.png", "type": "input"}]},
|
||
assets={"images": ["image-1"]},
|
||
))
|
||
with self.assertRaisesRegex(ValueError, "至少需要填写一个素材 ID"):
|
||
JOBS.normalize_video_job_payload(payload(
|
||
generation_mode="multimodal",
|
||
asset_creation_mode="manual",
|
||
prompt="prompt alone is not enough in manual mode",
|
||
))
|
||
|
||
def test_legacy_person_assets_infer_manual_mode(self):
|
||
legacy = payload(generation_mode="multimodal", assets={"persons": ["legacy-image"]})
|
||
legacy.pop("asset_creation_mode")
|
||
normalized = JOBS.normalize_video_job_payload(legacy)
|
||
self.assertEqual(normalized["asset_creation_mode"], "manual")
|
||
self.assertEqual(normalized["assets"]["images"], ["legacy-image"])
|
||
|
||
|
||
class OutputAllocationTests(unittest.TestCase):
|
||
def test_parallel_output_names_are_reserved_before_copy(self):
|
||
with tempfile.TemporaryDirectory() as output_dir:
|
||
job = {"filename_prefix": "clip", "save_location": "video"}
|
||
first = JOBS._allocate_output_path(job, output_dir)
|
||
second = JOBS._allocate_output_path(job, output_dir)
|
||
self.assertNotEqual(first[0], second[0])
|
||
self.assertEqual(os.path.getsize(first[0]), 0)
|
||
self.assertEqual(os.path.getsize(second[0]), 0)
|
||
|
||
|
||
class ReferenceMediaResolutionTests(unittest.TestCase):
|
||
def test_all_reference_image_roles_use_official_dimensions_and_ratio(self):
|
||
with tempfile.TemporaryDirectory() as temp_dir:
|
||
too_narrow = os.path.join(temp_dir, "too-narrow.png")
|
||
valid_small_total = os.path.join(temp_dir, "valid-small-total.png")
|
||
bad_ratio = os.path.join(temp_dir, "bad-ratio.png")
|
||
Image.new("RGB", (299, 750)).save(too_narrow)
|
||
Image.new("RGB", (300, 300)).save(valid_small_total)
|
||
Image.new("RGB", (300, 751)).save(bad_ratio)
|
||
|
||
for role in ("first_frame", "last_frame", "reference_images"):
|
||
with self.subTest(role=role):
|
||
with self.assertRaisesRegex(ValueError, "宽高必须分别在 300~6000px"):
|
||
JOBS._validate_media_file(too_narrow, role)
|
||
with self.assertRaisesRegex(ValueError, "宽高比必须在 0.4~2.5"):
|
||
JOBS._validate_media_file(bad_ratio, role)
|
||
JOBS._validate_media_file(valid_small_total, role)
|
||
|
||
def test_reference_video_uses_official_total_pixel_range(self):
|
||
JOBS._validate_reference_dimensions(614, 664, "参考视频", require_video_pixel_range=True)
|
||
JOBS._validate_reference_dimensions(3326, 2494, "参考视频", require_video_pixel_range=True)
|
||
with self.assertRaisesRegex(ValueError, "407,696~8,295,044"):
|
||
JOBS._validate_reference_dimensions(613, 664, "参考视频", require_video_pixel_range=True)
|
||
with self.assertRaisesRegex(ValueError, "407,696~8,295,044"):
|
||
JOBS._validate_reference_dimensions(3327, 2494, "参考视频", require_video_pixel_range=True)
|
||
|
||
def test_reference_video_path_is_probed_before_snapshot(self):
|
||
with tempfile.TemporaryDirectory() as temp_dir:
|
||
path = os.path.join(temp_dir, "reference.mp4")
|
||
with open(path, "wb") as handle:
|
||
handle.write(b"not-a-real-video")
|
||
with patch.object(JOBS, "_probe_video_dimensions", return_value=(854, 480)) as probe:
|
||
JOBS._validate_media_file(path, "reference_videos")
|
||
probe.assert_called_once_with(path)
|
||
|
||
|
||
class AutomaticAssetPreparationTests(unittest.IsolatedAsyncioTestCase):
|
||
async def test_both_routes_create_assets_with_bounded_ordered_concurrency(self):
|
||
for route, request_type in (("overseas_hc", "hc"), ("domestic", "doubao")):
|
||
with self.subTest(route=route), tempfile.TemporaryDirectory() as temp_dir:
|
||
paths = []
|
||
for index in range(5):
|
||
path = os.path.join(temp_dir, f"image-{index}.png")
|
||
Image.new("RGB", (2, 2), (index, 0, 0)).save(path)
|
||
paths.append(path)
|
||
|
||
active = 0
|
||
maximum = 0
|
||
calls = []
|
||
|
||
class FakeClient:
|
||
async def create_hc_asset_and_wait(self, **kwargs):
|
||
nonlocal active, maximum
|
||
active += 1
|
||
maximum = max(maximum, active)
|
||
calls.append(kwargs)
|
||
await asyncio.sleep((6 - int(kwargs["name"].rsplit("-", 1)[1])) * 0.002)
|
||
active -= 1
|
||
return {"Id": f"asset-{kwargs['name']}"}
|
||
|
||
job = {
|
||
"route": route,
|
||
"asset_creation_mode": "auto",
|
||
"snapshot_media": {
|
||
"first_frame": None,
|
||
"last_frame": None,
|
||
"reference_images": paths,
|
||
"reference_videos": [],
|
||
"reference_audios": [],
|
||
},
|
||
}
|
||
with (
|
||
patch.object(JOBS, "SeedanceElementClient", return_value=FakeClient()),
|
||
patch.object(
|
||
JOBS,
|
||
"upload_image",
|
||
new=AsyncMock(side_effect=lambda *_args, **_kwargs: "https://upload.invalid/image"),
|
||
),
|
||
patch.object(JOBS, "get_base_url_by_route", return_value="https://api.invalid"),
|
||
):
|
||
prepared = await JOBS._prepare_seedance_media(job, lambda **_values: None)
|
||
|
||
self.assertLessEqual(maximum, 3)
|
||
self.assertEqual([call["request_type"] for call in calls], [request_type] * 5)
|
||
self.assertEqual(
|
||
prepared["reference_images"],
|
||
[f"asset://asset-o1key-image-{index}" for index in range(1, 6)],
|
||
)
|
||
self.assertEqual(
|
||
job["resolved_assets"]["images"],
|
||
[f"asset-o1key-image-{index}" for index in range(1, 6)],
|
||
)
|
||
|
||
|
||
class ParallelManagerTests(unittest.IsolatedAsyncioTestCase):
|
||
def test_video_error_formatter_covers_review_subjects_and_preserves_other_errors(self):
|
||
cases = {
|
||
"The request failed because the output audio may be related to copyright restrictions":
|
||
"请求失败,输出视频中音频触发版权限制!",
|
||
"upstream rejected: field=video, reason=copyright":
|
||
"请求失败,输出视频触发版权限制!",
|
||
"field=content; reason=Copyright policy":
|
||
"请求失败,提示词触发版权限制!",
|
||
"field=real; reason=COPYRIGHT restriction":
|
||
"请求失败,真人内容触发版权限制!",
|
||
'{"field":"content","reason":"blocked by safety policy"}':
|
||
"请求失败,提示词触发审查!",
|
||
"moderation rejected; field: real":
|
||
"请求失败,真人内容触发审查!",
|
||
"OutputVideoSensitiveContentDetected.PolicyViolation":
|
||
"请求失败,输出视频触发审查!",
|
||
}
|
||
for raw, expected in cases.items():
|
||
with self.subTest(raw=raw):
|
||
self.assertEqual(JOBS.format_o1key_video_error(raw), expected)
|
||
self.assertEqual(
|
||
JOBS.format_o1key_video_error("provider connection timed out"),
|
||
"provider connection timed out",
|
||
)
|
||
self.assertEqual(
|
||
JOBS.format_o1key_video_error("request rejected: invalid API parameter"),
|
||
"request rejected: invalid API parameter",
|
||
)
|
||
|
||
async def test_manager_formats_errors_for_future_video_providers(self):
|
||
async def executor(_job, _update):
|
||
raise RuntimeError(
|
||
"The request failed because the output audio may be related "
|
||
"to copyright restrictions"
|
||
)
|
||
|
||
async def sender(_event, _payload):
|
||
return None
|
||
|
||
manager = JOBS.ParallelVideoJobManager(executor, sender)
|
||
with tempfile.TemporaryDirectory() as temp_dir:
|
||
job = {
|
||
"batch_id": "video_future_provider_error",
|
||
"generator_node_id": 1,
|
||
"result_node_id": 2,
|
||
"provider": "future-provider",
|
||
"model": "future-video-model",
|
||
"submitted_at": time.time(),
|
||
"temp_directory": temp_dir,
|
||
}
|
||
await manager.submit(job)
|
||
await manager.jobs[job["batch_id"]].task
|
||
self.assertEqual(
|
||
manager.status(job["batch_id"])["error"],
|
||
"请求失败,输出视频中音频触发版权限制!",
|
||
)
|
||
|
||
async def test_failed_video_job_keeps_resolved_asset_ids_for_retry(self):
|
||
async def executor(job, _update):
|
||
job["resolved_assets"] = {"images": ["safe-image-id"], "videos": [], "audios": []}
|
||
raise RuntimeError("provider failed")
|
||
|
||
async def sender(_event, _payload):
|
||
return None
|
||
|
||
manager = JOBS.ParallelVideoJobManager(executor, sender)
|
||
with tempfile.TemporaryDirectory() as temp_dir:
|
||
job = {
|
||
"batch_id": "video_failed_assets",
|
||
"generator_node_id": 1,
|
||
"result_node_id": 2,
|
||
"provider": "seedance",
|
||
"model": "seedance-2.0",
|
||
"submitted_at": time.time(),
|
||
"temp_directory": temp_dir,
|
||
}
|
||
await manager.submit(job)
|
||
await manager.jobs[job["batch_id"]].task
|
||
self.assertEqual(
|
||
manager.status(job["batch_id"])["resolved_assets"],
|
||
{"images": ["safe-image-id"], "videos": [], "audios": []},
|
||
)
|
||
|
||
async def test_repeated_submissions_start_immediately_without_concurrency_limit(self):
|
||
active = 0
|
||
maximum = 0
|
||
release = asyncio.Event()
|
||
|
||
async def executor(job, update):
|
||
nonlocal active, maximum
|
||
active += 1
|
||
maximum = max(maximum, active)
|
||
update(stage="polling", progress=0.5, provider_task_id=f"remote-{job['batch_id']}")
|
||
await release.wait()
|
||
active -= 1
|
||
return {"video": {"filename": f"{job['batch_id']}.mp4", "subfolder": "video", "type": "output"}}
|
||
|
||
async def sender(_event, _payload):
|
||
return None
|
||
|
||
manager = JOBS.ParallelVideoJobManager(executor, sender)
|
||
with tempfile.TemporaryDirectory() as temp_dir:
|
||
submitted = []
|
||
for index in range(3):
|
||
job = {
|
||
"batch_id": f"video_parallel_{index}",
|
||
"generator_node_id": 1,
|
||
"result_node_id": index + 2,
|
||
"provider": "seedance",
|
||
"model": "seedance-2.0",
|
||
"submitted_at": time.time(),
|
||
"temp_directory": temp_dir,
|
||
}
|
||
submitted.append(await manager.submit(job))
|
||
await asyncio.sleep(0.05)
|
||
self.assertEqual(maximum, 3)
|
||
self.assertTrue(all(manager.status(item["batch_id"])["state"] == "running" for item in submitted))
|
||
self.assertTrue(all("max_concurrent_jobs" not in manager.status(item["batch_id"]) for item in submitted))
|
||
release.set()
|
||
await asyncio.gather(*(record.task for record in manager.jobs.values()))
|
||
self.assertTrue(all(manager.status(item["batch_id"])["state"] == "completed" for item in submitted))
|
||
|
||
|
||
if __name__ == "__main__":
|
||
unittest.main()
|