fix: 💄 BatchFromHistory when "listening"

When using --listen, BatchFromHistory was trying the wrong local ip
on local remotes.
This commit is contained in:
melMass
2023-08-15 20:22:35 +02:00
parent fe8f519f88
commit 3b07984716
2 changed files with 74 additions and 15 deletions
+24 -15
View File
@@ -4,19 +4,22 @@ import urllib.request
import urllib.parse
import torch
import json
from comfy.cli_args import args
from ..utils import pil2tensor, apply_easing
from ..utils import pil2tensor, apply_easing, get_server_info
import io
import numpy as np
def get_image(filename, subfolder, folder_type):
log.debug(f"Getting image {filename} from {subfolder} of {folder_type}")
log.debug(
f"Getting image {filename} from foldertype {folder_type} {f'in subfolder: {subfolder}' if subfolder else ''}"
)
data = {"filename": filename, "subfolder": subfolder, "type": folder_type}
base_url, port = get_server_info()
url_values = urllib.parse.urlencode(data)
with urllib.request.urlopen(
f"http://{args.listen}:{args.port}/view?{url_values}"
) as response:
url = f"http://{base_url}:{port}/view?{url_values}"
log.debug(f"Fetching image from {url}")
with urllib.request.urlopen(url) as response:
return io.BytesIO(response.read())
@@ -60,10 +63,18 @@ class GetBatchFromHistory:
return (torch.zeros(0),)
frames = []
with urllib.request.urlopen(
f"http://{args.listen}:{args.port}/history"
) as response:
return self.load_batch_frames(response, offset, count, frames)
base_url, port = get_server_info()
history_url = f"http://{base_url}:{port}/history"
log.debug(f"Fetching history from {history_url}")
output = torch.zeros(0)
with urllib.request.urlopen(history_url) as response:
output = self.load_batch_frames(response, offset, count, frames)
if output.size(0) == 0:
log.warn("No output found in history")
return (output,)
def load_batch_frames(self, response, offset, count, frames):
history = json.loads(response.read())
@@ -80,7 +91,7 @@ class GetBatchFromHistory:
output_images.append(image_data)
if not output_images:
return (torch.zeros(0),)
return torch.zeros(0)
# Directly get desired range of images
start_index = max(len(output_images) - offset - count, 0)
@@ -90,13 +101,11 @@ class GetBatchFromHistory:
frames = [Image.open(image) for image in selected_images]
if not frames:
return (torch.zeros(0),)
return torch.zeros(0)
elif len(frames) != count:
log.warning(f"Expected {count} images, got {len(frames)} instead")
output = pil2tensor(frames)
return (output,)
return pil2tensor(frames)
class AnyToString:
+50
View File
@@ -11,6 +11,9 @@ import subprocess
import threading
import os
import math
import functools
import socket
import requests
try:
from .log import log
@@ -26,7 +29,54 @@ except ImportError:
log.warn("[comfy mtb] You probably called the file outside a module.")
class IPChecker:
def __init__(self):
self.ips = list(self.get_local_ips())
log.debug(f"Found {len(self.ips)} local ips")
self.checked_ips = set()
def get_working_ip(self, test_url_template):
for ip in self.ips:
if ip not in self.checked_ips:
self.checked_ips.add(ip)
test_url = test_url_template.format(ip)
if self._test_url(test_url):
return ip
return None
@staticmethod
def get_local_ips(prefix="192.168."):
hostname = socket.gethostname()
log.debug(f"Getting local ips for {hostname}")
for info in socket.getaddrinfo(hostname, None):
# Filter out IPv6 addresses if you only want IPv4
log.debug(info)
# if info[1] == socket.SOCK_STREAM and
if info[0] == socket.AF_INET and info[4][0].startswith(prefix):
yield info[4][0]
def _test_url(self, url):
try:
response = requests.get(url)
return response.status_code == 200
except Exception:
return False
# region MISC Utilities
@functools.lru_cache(maxsize=1)
def get_server_info():
from comfy.cli_args import args
ip_checker = IPChecker()
base_url = args.listen
if base_url == "0.0.0.0":
log.debug("Server set to 0.0.0.0, we will try to resolve the host IP")
base_url = ip_checker.get_working_ip(f"http://{{}}:{args.port}/history")
log.debug(f"Setting ip to {base_url}")
return (base_url, args.port)
def hex_to_rgb(hex_color):
try:
hex_color = hex_color.lstrip("#")