"""Lean ComfyUI nodes for O1Key Grok Imagine Video.""" import asyncio import math import os import re from typing import Dict from ..clients.grok_video_client import GrokVideoClient from ..utils.config import get_base_url_by_route from ..utils.image_utils import tensor_to_pil from ..utils.r2_uploader import upload_audio, upload_image, upload_video try: import folder_paths except ImportError: folder_paths = None try: from comfy.utils import ProgressBar except ImportError: ProgressBar = None try: from comfy_api.input_impl import VideoFromFile except Exception: try: from comfy_api.latest import InputImpl VideoFromFile = InputImpl.VideoFromFile except Exception: VideoFromFile = None MODEL_OPTIONS = list(GrokVideoClient.MODEL_OPTIONS) ASPECT_RATIO_OPTIONS = list(GrokVideoClient.ASPECT_RATIO_OPTIONS) RESOLUTION_OPTIONS = list(GrokVideoClient.RESOLUTION_OPTIONS) GENERATION_MODE_OPTIONS = ["文生视频", "图生视频", "参考生视频"] EDIT_MODE_OPTIONS = ["编辑视频", "续写视频"] IMAGE_INPUT_NAMES = [f"图片{i}" for i in range(1, 8)] AUDIO_INPUT_NAMES = ["音频素材", "音频素材2", "音频素材3"] def _get_output_dir() -> str: if folder_paths is not None: base_dir = folder_paths.get_temp_directory() else: plugin_dir = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) base_dir = os.path.join(os.path.dirname(os.path.dirname(plugin_dir)), "temp") output_dir = os.path.join(base_dir, "grok_video") os.makedirs(output_dir, exist_ok=True) return output_dir def _single_pil_image(image_tensor, input_name: str): if image_tensor is None: return None images = tensor_to_pil(image_tensor) if len(images) != 1: raise ValueError(f"{input_name} 只能连接 1 张图片,请拆分批次后再连接。") return images[0].convert("RGB") def _parse_voice_ids(value: object) -> list[str]: voice_ids = [item.strip() for item in re.split(r"[,,\n]", str(value or "")) if item.strip()] if len(voice_ids) > 3: raise ValueError("参考音色 ID 最多填写 3 个。") return voice_ids def _video_duration_seconds(video) -> float: getter = getattr(video, "get_duration", None) if not callable(getter): raise ValueError("无法读取输入视频时长;请连接 ComfyUI 原生 VIDEO 输出。") try: duration = float(getter()) except Exception as exc: raise ValueError("无法读取输入视频时长;请确认视频文件可以正常解码。") from exc if not math.isfinite(duration) or duration <= 0: raise ValueError("输入视频时长无效;请确认视频文件可以正常解码。") return duration def _progress_callback(): progress_bar = ProgressBar(100) if ProgressBar is not None else None progress_value = [0] def callback(progress: int, _status: str, _elapsed: float) -> None: current = max(0, min(100, int(progress or 0))) if progress_bar is not None and current > progress_value[0]: progress_bar.update(current - progress_value[0]) progress_value[0] = current return progress_bar, progress_value, callback def _finish_video(result: Dict[str, object], progress_bar, progress_value): if progress_bar is not None and progress_value[0] < 100: progress_bar.update(100 - progress_value[0]) video_path = result["video_path"] print(f"Grok Video:下载完成:{video_path}") return (VideoFromFile(video_path),) class O1keyGrokVideo: """文生、图生或参考素材生 Grok 视频。素材会自动上传为 URL。""" @classmethod def INPUT_TYPES(cls): return { "required": { "生成模式": (GENERATION_MODE_OPTIONS, {"default": "文生视频"}), "提示词": ("STRING", {"default": "", "multiline": True}), "模型": (MODEL_OPTIONS, {"default": GrokVideoClient.DEFAULT_MODEL}), "时长(秒)": ("INT", {"default": 8, "min": 1, "max": 15, "step": 1}), "宽高比": (ASPECT_RATIO_OPTIONS, {"default": "16:9"}), "分辨率": (RESOLUTION_OPTIONS, {"default": "480p"}), "参考音色ID(逗号分隔)": ("STRING", {"default": ""}), }, "optional": { **{input_name: ("IMAGE",) for input_name in IMAGE_INPUT_NAMES}, **{input_name: ("AUDIO",) for input_name in AUDIO_INPUT_NAMES}, }, } RETURN_TYPES = ("VIDEO",) RETURN_NAMES = ("视频",) FUNCTION = "generate" CATEGORY = "comfyui_o1key/Video" DESCRIPTION = ( "支持文生、图生和多参考素材生成。图生视频只连接图片 1;参考生视频最多使用 7 张图和 " "3 个参考音频(AUDIO 或 voice_id 合计)。Grok 1.5 的文生/图生可选 1080p,多参考最高 720p。" ) def generate(self, **kwargs): if VideoFromFile is None: raise RuntimeError("当前 ComfyUI 版本不支持 VideoFromFile,无法输出 VIDEO。") mode = kwargs.get("生成模式", "文生视频") model = kwargs.get("模型", GrokVideoClient.DEFAULT_MODEL) prompt = (kwargs.get("提示词") or "").strip() connected_images = [ (input_name, image) for input_name in IMAGE_INPUT_NAMES if (image := _single_pil_image(kwargs.get(input_name), input_name)) is not None ] images = [image for _, image in connected_images] audios = [kwargs.get(input_name) for input_name in AUDIO_INPUT_NAMES if kwargs.get(input_name) is not None] voice_ids = _parse_voice_ids(kwargs.get("参考音色ID(逗号分隔)")) duration = kwargs.get("时长(秒)", 8) aspect_ratio = kwargs.get("宽高比", "16:9") resolution = kwargs.get("分辨率", "480p") if mode not in GENERATION_MODE_OPTIONS: raise ValueError(f"不支持的生成模式:{mode}。") if len(audios) + len(voice_ids) > 3: raise ValueError("参考音频与参考音色 ID 合计最多 3 个。") if mode == "文生视频": if images or audios or voice_ids: raise ValueError("文生视频不需要连接图像或音频素材。") elif mode == "图生视频": if len(images) != 1 or connected_images[0][0] != "图片1": raise ValueError("图生视频需要在“图片 1”连接 1 张图片,其他图片端口请留空。") if audios or voice_ids: raise ValueError("图生视频不支持音频素材,请使用参考生视频。") else: if not images and not audios and not voice_ids: raise ValueError("参考生视频至少需要连接图像素材或音频素材。") placeholder_image = {"url": "https://example.invalid/image"} if mode == "图生视频" else None placeholder_references = ( [{"url": f"https://example.invalid/reference-{index}"} for index in range(len(images))] if mode == "参考生视频" else [] ) placeholder_audios = ( [{"url": f"https://example.invalid/audio-{index}"} for index in range(len(audios))] + [{"voice_id": voice_id} for voice_id in voice_ids] if mode == "参考生视频" else [] ) # Validate every user-controlled field before temporary uploads or paid generation calls. GrokVideoClient.build_video_body( operation="generate", prompt=prompt, model=model, duration=duration, aspect_ratio=aspect_ratio, resolution=resolution, image=placeholder_image, reference_images=placeholder_references, reference_audios=placeholder_audios, ) base_url = get_base_url_by_route() client = GrokVideoClient(base_url=base_url) async def upload_materials(): image_urls, audio_urls = await asyncio.gather( asyncio.gather(*(upload_image(image, base_url=base_url) for image in images)), asyncio.gather(*(upload_audio(audio, base_url=base_url) for audio in audios)), ) return list(image_urls), list(audio_urls) image_urls, audio_urls = client.run_async_in_thread(upload_materials()) if mode == "文生视频": image = None reference_images = [] reference_audios = [] elif mode == "图生视频": image = {"url": image_urls[0]} reference_images = [] reference_audios = [] else: image = None reference_images = [{"url": url} for url in image_urls] reference_audios = [ *({"url": url} for url in audio_urls), *({"voice_id": voice_id} for voice_id in voice_ids), ] progress_bar, progress_value, callback = _progress_callback() result = client.run_video_sync( operation="generate", prompt=prompt, model=model, duration=duration, aspect_ratio=aspect_ratio, resolution=resolution, image=image, reference_images=reference_images, reference_audios=reference_audios, output_dir=_get_output_dir(), progress_callback=callback, ) return _finish_video(result, progress_bar, progress_value) class O1keyGrokVideoEdit: """编辑或续写 Grok 视频。输入 VIDEO 会自动上传为 URL。""" @classmethod def INPUT_TYPES(cls): return { "required": { "操作": (EDIT_MODE_OPTIONS, {"default": "编辑视频"}), "提示词": ("STRING", {"default": "", "multiline": True}), "续写时长(秒)": ("INT", {"default": 6, "min": 2, "max": 10, "step": 1}), "模型": (MODEL_OPTIONS, {"default": GrokVideoClient.DEFAULT_MODEL}), }, "optional": {"视频素材": ("VIDEO",)}, } RETURN_TYPES = ("VIDEO",) RETURN_NAMES = ("视频",) FUNCTION = "generate" CATEGORY = "comfyui_o1key/Video" DESCRIPTION = ( "编辑或续写视频。编辑输入最长 8.7 秒,并保留原时长和宽高比,输出最高 720p;" "续写时长为 2–10 秒,输出总时长等于输入时长加续写时长。" ) def generate(self, **kwargs): if VideoFromFile is None: raise RuntimeError("当前 ComfyUI 版本不支持 VideoFromFile,无法输出 VIDEO。") video = kwargs.get("视频素材") if video is None: raise ValueError("请连接一个 VIDEO 类型的视频素材。") selected_operation = kwargs.get("操作", "编辑视频") if selected_operation not in EDIT_MODE_OPTIONS: raise ValueError(f"不支持的 Grok 视频操作:{selected_operation}。") operation = "edit" if selected_operation == "编辑视频" else "extend" prompt = (kwargs.get("提示词") or "").strip() model = kwargs.get("模型", GrokVideoClient.DEFAULT_MODEL) duration = kwargs.get("续写时长(秒)", 6) if operation == "edit" and _video_duration_seconds(video) > 8.7: raise ValueError("Grok 视频编辑的输入视频不能超过 8.7 秒。") GrokVideoClient.build_video_body( operation=operation, prompt=prompt, model=model, duration=duration, video={"url": "https://example.invalid/video"}, ) base_url = get_base_url_by_route() client = GrokVideoClient(base_url=base_url) video_url = client.run_async_in_thread(upload_video(video, base_url=base_url)) progress_bar, progress_value, callback = _progress_callback() result = client.run_video_sync( operation=operation, prompt=prompt, model=model, duration=duration, video={"url": video_url}, output_dir=_get_output_dir(), progress_callback=callback, ) return _finish_video(result, progress_bar, progress_value) NODE_CLASS_MAPPINGS = { "O1keyGrokVideo": O1keyGrokVideo, "O1keyGrokVideoEdit": O1keyGrokVideoEdit, } NODE_DISPLAY_NAME_MAPPINGS = { "O1keyGrokVideo": "Grok Video", "O1keyGrokVideoEdit": "Grok Video Edit", }