diff --git a/__init__.py b/__init__.py
index 2d2cca4..d1274bc 100644
--- a/__init__.py
+++ b/__init__.py
@@ -265,6 +265,10 @@ def register_routes():
from PIL import Image
+ from .repl import setup_custom_web_routes
+
+ setup_custom_web_routes(PromptServer.instance.app)
+
with contextlib.suppress(ImportError):
from cachetools import TTLCache
diff --git a/features.json b/features.json
new file mode 100644
index 0000000..9a58c45
--- /dev/null
+++ b/features.json
@@ -0,0 +1,3 @@
+{
+ "use_repl": false
+}
diff --git a/repl.py b/repl.py
new file mode 100644
index 0000000..b85391c
--- /dev/null
+++ b/repl.py
@@ -0,0 +1,637 @@
+import base64
+import code
+import io
+import re
+import sys
+from contextlib import redirect_stderr, redirect_stdout
+
+# import matplotlib.pyplot as plt
+import numpy as np
+import torch
+from aiohttp import web
+from PIL import Image
+from rich.console import Console
+from rich.traceback import Traceback
+
+from .log import log
+
+try:
+ import pyflakes.api
+ import pyflakes.reporter
+
+ _HAS_LINT = True
+except ImportError:
+ print(
+ "ComfyREPL: pyflakes not found. Linting will be disabled. Install with 'pip install pyflakes'."
+ )
+ _HAS_LINT = False
+
+# --- Linting Library ---
+# try:
+# import ruff
+# import ruff.lint
+# import ruff.lint.linter
+# import ruff.settings
+#
+# _HAS_LINT = True
+# except ImportError:
+# print(
+# "ComfyREPL: ruff not found. Linting will be disabled. Install with 'pip install ruff'."
+# )
+# _HAS_LINT = False
+
+# --- Audio/Video Libraries ---
+try:
+ import scipy.io.wavfile
+
+ _HAS_SCIPY = True
+except ImportError:
+ print(
+ "ComfyREPL: SciPy not found. Audio display will be disabled. Install with 'pip install scipy'."
+ )
+ _HAS_SCIPY = False
+
+try:
+ import imageio
+ import imageio.plugins.ffmpeg # Ensure ffmpeg plugin is available
+
+ _HAS_IMAGEIO = True
+except ImportError:
+ print(
+ "ComfyREPL: Imageio or imageio-ffmpeg not found. Video display will be disabled. Install with 'pip install imageio imageio-ffmpeg'."
+ )
+ _HAS_IMAGEIO = False
+
+
+# --- Audio Display ---
+class AudioDisplay:
+ def __init__(self, samples, sample_rate):
+ if not _HAS_SCIPY:
+ raise ImportError("Audio display requires scipy and numpy.")
+ if not isinstance(samples, (np.ndarray, torch.Tensor)):
+ raise TypeError(
+ "Audio samples must be a numpy array or torch tensor."
+ )
+ if isinstance(samples, torch.Tensor):
+ samples = samples.detach().cpu().numpy()
+
+ # Ensure samples are in a format scipy.io.wavfile can handle (e.g., int16, float32)
+ if samples.dtype == np.float64:
+ samples = samples.astype(np.float32)
+ elif samples.dtype == np.int64:
+ # Or scale to int32 if range requires
+ samples = samples.astype(np.int16)
+
+ self.samples = samples
+ self.sample_rate = sample_rate
+
+ def _to_wav_base64(self):
+ buffer = io.BytesIO()
+ try:
+ scipy.io.wavfile.write(buffer, self.sample_rate, self.samples)
+ audio_base64 = base64.b64encode(buffer.getvalue()).decode("utf-8")
+ return audio_base64
+ except Exception as e:
+ return f"
Error encoding audio: {e}
"
+
+ def _repr_html_(self):
+ base64_data = self._to_wav_base64()
+ if base64_data.startswith("'
+
+
+def render_audio(samples, sample_rate):
+ """
+ Render audio samples as an HTML audio player.
+
+ Args:
+ samples (np.ndarray or torch.Tensor): Audio samples.
+ sample_rate (int): Sample rate in Hz.
+
+ Returns
+ -------
+ AudioDisplay: An object that will render as an HTML audio player.
+ """
+ return AudioDisplay(samples, sample_rate)
+
+
+# --- Display Classes ---
+class VideoDisplay:
+ def __init__(self, frames, fps=24, options=None):
+ if not _HAS_IMAGEIO: # numpy/PIL/torch needed for frames
+ raise ImportError(
+ "Video display requires imageio, imageio-ffmpeg, and image libraries (numpy, Pillow, torch)."
+ )
+
+ self.frames = []
+ for frame in frames:
+ if isinstance(frame, Image.Image):
+ self.frames.append(np.array(frame))
+ elif isinstance(frame, np.ndarray):
+ # Ensure HWC and uint8
+ if frame.ndim == 3 and frame.shape[0] in [1, 3, 4]: # CHW
+ frame = np.transpose(frame, (1, 2, 0))
+ if frame.dtype != np.uint8:
+ frame = (
+ (frame * 255).astype(np.uint8)
+ if frame.max() <= 1.0
+ else frame.astype(np.uint8)
+ )
+ self.frames.append(frame)
+ elif isinstance(frame, torch.Tensor):
+ np_frame = frame.detach().cpu().numpy()
+ if np_frame.ndim == 3 and np_frame.shape[0] in [
+ 1,
+ 3,
+ 4,
+ ]: # CHW
+ np_frame = np.transpose(np_frame, (1, 2, 0))
+ if np_frame.dtype != np.uint8:
+ np_frame = (
+ (np_frame * 255).astype(np.uint8)
+ if np_frame.max() <= 1.0
+ else np_frame.astype(np.uint8)
+ )
+ self.frames.append(np_frame)
+ else:
+ raise TypeError(
+ f"Unsupported frame type: {type(frame)}. Must be PIL.Image, numpy.ndarray, or torch.Tensor."
+ )
+
+ self.fps = fps
+ self.options = options if options is not None else {}
+
+ def _to_mp4_base64(self):
+ buffer = io.BytesIO()
+ try:
+ # Use imageio to write frames to an in-memory MP4 file
+ imageio.mimwrite(
+ buffer,
+ self.frames,
+ format="mp4",
+ fps=self.fps,
+ codec="libx264",
+ quality=8,
+ ) # quality 1-10
+ video_base64 = base64.b64encode(buffer.getvalue()).decode("utf-8")
+ return video_base64
+ except Exception as e:
+ return f"Error encoding video: {e}
"
+
+ def _repr_html_(self):
+ base64_data = self._to_mp4_base64()
+ if base64_data.startswith("'
+
+
+def render_video(batch_tensor_or_array_of_pil_images, fps=24, options=None):
+ """
+ Render video frames as an HTML video player.
+
+ Args:
+ batch_tensor_or_array_of_pil_images (list of PIL.Image, np.ndarray, or torch.Tensor):
+ A list of frames, or a single batch tensor/array (B, H, W, C) or (B, C, H, W).
+ fps (int): Frames per second.
+ options (dict): Dictionary of HTML