WAI-4030
This commit is contained in:
@@ -1,7 +1,7 @@
|
||||
import requests
|
||||
import torch
|
||||
|
||||
from .common import preprocess_image, image_to_base64, poll_status_until_completed
|
||||
from .common import deserialize_and_get_comfy_key, preprocess_image, image_to_base64, poll_status_until_completed
|
||||
|
||||
class AttributionByImageNode():
|
||||
@classmethod
|
||||
@@ -26,6 +26,7 @@ class AttributionByImageNode():
|
||||
def execute(self, image, model_version, api_key):
|
||||
if api_key.strip() == "" or api_key.strip() == "BRIA_API_TOKEN":
|
||||
raise Exception("Please insert a valid API key.")
|
||||
api_key = deserialize_and_get_comfy_key(api_key)
|
||||
|
||||
# Check if image is tensor, if so, convert to NumPy array
|
||||
if isinstance(image, torch.Tensor):
|
||||
|
||||
@@ -6,6 +6,7 @@ import base64
|
||||
from torchvision.transforms import ToPILImage
|
||||
import requests
|
||||
import time
|
||||
import json
|
||||
|
||||
def postprocess_image(image):
|
||||
result_image = Image.open(io.BytesIO(image))
|
||||
@@ -48,6 +49,7 @@ def preprocess_mask(mask):
|
||||
def process_request(api_url, image, mask, api_key, visual_input_content_moderation, visual_output_content_moderation):
|
||||
if api_key.strip() == "" or api_key.strip() == "BRIA_API_TOKEN":
|
||||
raise Exception("Please insert a valid API key.")
|
||||
api_key = deserialize_and_get_comfy_key(api_key)
|
||||
|
||||
# Check if image and mask are tensors, if so, convert to NumPy arrays
|
||||
if isinstance(image, torch.Tensor):
|
||||
@@ -145,3 +147,15 @@ def poll_status_until_completed(status_url, api_key, timeout=360, check_interval
|
||||
raise Exception(f"Error checking status: {e}")
|
||||
|
||||
raise Exception(f"Timeout reached after {timeout} seconds")
|
||||
|
||||
def deserialize_and_get_comfy_key(encoded: str):
|
||||
try:
|
||||
decoded = base64.b64decode(encoded).decode("utf-8")
|
||||
payload = json.loads(decoded)
|
||||
if (payload['type'] == "comfy"):
|
||||
return payload['apiKey']
|
||||
else:
|
||||
raise Exception(f"Invalid token type")
|
||||
except Exception as e:
|
||||
raise Exception(f"Invalid token")
|
||||
|
||||
|
||||
@@ -2,6 +2,7 @@ import requests
|
||||
import torch
|
||||
|
||||
from .common import (
|
||||
deserialize_and_get_comfy_key,
|
||||
postprocess_image,
|
||||
preprocess_image,
|
||||
image_to_base64,
|
||||
@@ -96,6 +97,7 @@ class _BaseGenerateImageNodeV2:
|
||||
seed,
|
||||
images,
|
||||
)
|
||||
api_token = deserialize_and_get_comfy_key(api_token)
|
||||
|
||||
headers = {"Content-Type": "application/json", "api_token": api_token}
|
||||
|
||||
|
||||
@@ -4,7 +4,7 @@ from PIL import Image
|
||||
import io
|
||||
import torch
|
||||
|
||||
from .common import preprocess_image, preprocess_mask, image_to_base64, poll_status_until_completed
|
||||
from .common import deserialize_and_get_comfy_key, preprocess_image, preprocess_mask, image_to_base64, poll_status_until_completed
|
||||
|
||||
|
||||
class GenFillNode():
|
||||
@@ -39,6 +39,7 @@ class GenFillNode():
|
||||
def execute(self, image, mask, prompt, api_key, seed, prompt_content_moderation, visual_input_content_moderation, visual_output_content_moderation):
|
||||
if api_key.strip() == "" or api_key.strip() == "BRIA_API_TOKEN":
|
||||
raise Exception("Please insert a valid API key.")
|
||||
api_key = deserialize_and_get_comfy_key(api_key)
|
||||
|
||||
# Check if image and mask are tensors, if so, convert to NumPy arrays
|
||||
if isinstance(image, torch.Tensor):
|
||||
|
||||
@@ -4,7 +4,7 @@ from PIL import Image
|
||||
import io
|
||||
import torch
|
||||
|
||||
from .common import image_to_base64, preprocess_image, poll_status_until_completed
|
||||
from .common import deserialize_and_get_comfy_key, image_to_base64, preprocess_image, poll_status_until_completed
|
||||
|
||||
|
||||
class ImageExpansionNode():
|
||||
@@ -55,6 +55,7 @@ class ImageExpansionNode():
|
||||
api_key):
|
||||
if api_key.strip() == "" or api_key.strip() == "BRIA_API_TOKEN":
|
||||
raise Exception("Please insert a valid API key.")
|
||||
api_key = deserialize_and_get_comfy_key(api_key)
|
||||
original_image_size = [int(x.strip()) for x in original_image_size.split(",")] if original_image_size else ()
|
||||
original_image_location = [int(x.strip()) for x in original_image_location.split(",")] if original_image_location else ()
|
||||
canvas_size = [int(x.strip()) for x in canvas_size.split(",")] if canvas_size else ()
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
import requests
|
||||
from .common import poll_status_until_completed, postprocess_image
|
||||
from .common import deserialize_and_get_comfy_key, poll_status_until_completed, postprocess_image
|
||||
|
||||
|
||||
class _BaseRefineImageNodeV2:
|
||||
@@ -83,7 +83,7 @@ class _BaseRefineImageNodeV2:
|
||||
guidance_scale,
|
||||
seed,
|
||||
)
|
||||
|
||||
api_token = deserialize_and_get_comfy_key(api_token)
|
||||
headers = {"Content-Type": "application/json", "api_token": api_token}
|
||||
|
||||
try:
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
import requests
|
||||
|
||||
from .common import postprocess_image, preprocess_image, image_to_base64
|
||||
from .common import deserialize_and_get_comfy_key, postprocess_image, preprocess_image, image_to_base64
|
||||
|
||||
|
||||
class ReimagineNode():
|
||||
@@ -37,7 +37,8 @@ class ReimagineNode():
|
||||
steps_num, fast, structure_ref_influence, structure_image=None,
|
||||
tailored_model_id=None, tailored_model_influence=None, tailored_generation_prefix=None,
|
||||
content_moderation=0,
|
||||
):
|
||||
):
|
||||
api_key = deserialize_and_get_comfy_key(api_key)
|
||||
payload = {
|
||||
"prompt": tailored_generation_prefix + prompt,
|
||||
"num_results": 1,
|
||||
|
||||
@@ -4,7 +4,7 @@ from PIL import Image
|
||||
import io
|
||||
import torch
|
||||
|
||||
from .common import preprocess_image, image_to_base64, poll_status_until_completed
|
||||
from .common import deserialize_and_get_comfy_key, preprocess_image, image_to_base64, poll_status_until_completed
|
||||
|
||||
|
||||
class RemoveForegroundNode():
|
||||
@@ -34,6 +34,7 @@ class RemoveForegroundNode():
|
||||
def execute(self, image, visual_input_content_moderation, visual_output_content_moderation, preserve_alpha, api_key):
|
||||
if api_key.strip() == "" or api_key.strip() == "BRIA_API_TOKEN":
|
||||
raise Exception("Please insert a valid API key.")
|
||||
api_key = deserialize_and_get_comfy_key(api_key)
|
||||
|
||||
# Check if image is tensor, if so, convert to NumPy array
|
||||
if isinstance(image, torch.Tensor):
|
||||
|
||||
@@ -4,7 +4,7 @@ from PIL import Image
|
||||
import io
|
||||
import torch
|
||||
|
||||
from .common import image_to_base64, preprocess_image, preprocess_mask, poll_status_until_completed
|
||||
from .common import deserialize_and_get_comfy_key, image_to_base64, preprocess_image, preprocess_mask, poll_status_until_completed
|
||||
|
||||
|
||||
class ReplaceBgNode():
|
||||
@@ -53,6 +53,7 @@ class ReplaceBgNode():
|
||||
ref_images=None,):
|
||||
if api_key.strip() == "" or api_key.strip() == "BRIA_API_TOKEN":
|
||||
raise Exception("Please insert a valid API key.")
|
||||
api_key = deserialize_and_get_comfy_key(api_key)
|
||||
|
||||
# Check if image and mask are tensors, if so, convert to NumPy arrays
|
||||
if isinstance(image, torch.Tensor):
|
||||
|
||||
+2
-2
@@ -4,7 +4,7 @@ from PIL import Image
|
||||
import io
|
||||
import torch
|
||||
|
||||
from .common import preprocess_image, image_to_base64, poll_status_until_completed
|
||||
from .common import deserialize_and_get_comfy_key, preprocess_image, image_to_base64, poll_status_until_completed
|
||||
|
||||
class RmbgNode():
|
||||
@classmethod
|
||||
@@ -34,7 +34,7 @@ class RmbgNode():
|
||||
def execute(self, image, visual_input_content_moderation, visual_output_content_moderation, preserve_alpha, api_key):
|
||||
if api_key.strip() == "" or api_key.strip() == "BRIA_API_TOKEN":
|
||||
raise Exception("Please insert a valid API key.")
|
||||
|
||||
api_key = deserialize_and_get_comfy_key(api_key)
|
||||
# Check if image is tensor, if so, convert to NumPy array
|
||||
if isinstance(image, torch.Tensor):
|
||||
image = preprocess_image(image)
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
import requests
|
||||
|
||||
from .common import postprocess_image, preprocess_image, image_to_base64
|
||||
from .common import deserialize_and_get_comfy_key, postprocess_image, preprocess_image, image_to_base64
|
||||
|
||||
|
||||
class TailoredGenNode():
|
||||
@@ -45,6 +45,7 @@ class TailoredGenNode():
|
||||
guidance_method_2=None, guidance_method_2_scale=None, guidance_method_2_image=None,
|
||||
content_moderation=0,
|
||||
):
|
||||
api_key = deserialize_and_get_comfy_key(api_key)
|
||||
payload = {
|
||||
"prompt": generation_prefix + prompt,
|
||||
"num_results": 1,
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
import requests
|
||||
|
||||
from .common import deserialize_and_get_comfy_key
|
||||
|
||||
class TailoredModelInfoNode():
|
||||
@classmethod
|
||||
@@ -21,6 +21,7 @@ class TailoredModelInfoNode():
|
||||
|
||||
# Define the execute method as expected by ComfyUI
|
||||
def execute(self, model_id, api_key):
|
||||
api_key = deserialize_and_get_comfy_key(api_key)
|
||||
response = requests.get(
|
||||
self.api_url + model_id,
|
||||
headers={"api_token": api_key}
|
||||
|
||||
@@ -4,7 +4,7 @@ from PIL import Image
|
||||
import io
|
||||
import torch
|
||||
|
||||
from .common import image_to_base64, preprocess_image
|
||||
from .common import deserialize_and_get_comfy_key, image_to_base64, preprocess_image
|
||||
|
||||
class TailoredPortraitNode():
|
||||
@classmethod
|
||||
@@ -12,7 +12,7 @@ class TailoredPortraitNode():
|
||||
return {
|
||||
"required": {
|
||||
"image": ("IMAGE",), # Input image from another node
|
||||
"tailored_model_id": ("INT",),
|
||||
"tailored_model_id": ("STRING",), # API Key input with a default value
|
||||
"api_key": ("STRING", {"default": "BRIA_API_TOKEN"}), # API Key input with a default value
|
||||
},
|
||||
"optional": {
|
||||
@@ -34,6 +34,7 @@ class TailoredPortraitNode():
|
||||
def execute(self, image, tailored_model_id, api_key, seed, tailored_model_influence, id_strength):
|
||||
if api_key.strip() == "" or api_key.strip() == "BRIA_API_TOKEN":
|
||||
raise Exception("Please insert a valid API key.")
|
||||
api_key = deserialize_and_get_comfy_key(api_key)
|
||||
|
||||
# Convert the image and mask directly to if isinstance(image, torch.Tensor):
|
||||
if isinstance(image, torch.Tensor):
|
||||
@@ -44,7 +45,7 @@ class TailoredPortraitNode():
|
||||
# Prepare the API request payload
|
||||
payload = {
|
||||
"id_image_file": f"{image_base64}",
|
||||
"tailored_model_id": tailored_model_id,
|
||||
"tailored_model_id": int(tailored_model_id),
|
||||
"tailored_model_influence": tailored_model_influence,
|
||||
"id_strength": id_strength,
|
||||
"seed": seed
|
||||
@@ -69,7 +70,7 @@ class TailoredPortraitNode():
|
||||
result_image = torch.from_numpy(result_image)[None,]
|
||||
return (result_image,)
|
||||
else:
|
||||
raise Exception(f"Error: API request failed with status code {response.status_code}")
|
||||
raise Exception(f"Error: API request failed with status code {response.status_code} {response.text}")
|
||||
|
||||
except Exception as e:
|
||||
raise Exception(f"{e}")
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
import requests
|
||||
|
||||
from .common import postprocess_image, preprocess_image, image_to_base64
|
||||
from .common import deserialize_and_get_comfy_key, postprocess_image, preprocess_image, image_to_base64
|
||||
|
||||
|
||||
class Text2ImageBaseNode():
|
||||
@@ -48,6 +48,7 @@ class Text2ImageBaseNode():
|
||||
image_prompt_mode=None, image_prompt_image=None, image_prompt_scale=None,
|
||||
content_moderation=0,
|
||||
):
|
||||
api_key = deserialize_and_get_comfy_key(api_key)
|
||||
payload = {
|
||||
"prompt": prompt,
|
||||
"num_results": 1,
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
import requests
|
||||
|
||||
from .common import postprocess_image, preprocess_image, image_to_base64
|
||||
from .common import deserialize_and_get_comfy_key, postprocess_image, preprocess_image, image_to_base64
|
||||
|
||||
|
||||
class Text2ImageFastNode():
|
||||
@@ -45,6 +45,7 @@ class Text2ImageFastNode():
|
||||
image_prompt_mode=None, image_prompt_image=None, image_prompt_scale=None,
|
||||
content_moderation=0,
|
||||
):
|
||||
api_key = deserialize_and_get_comfy_key(api_key)
|
||||
payload = {
|
||||
"prompt": prompt,
|
||||
"num_results": 1,
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
import requests
|
||||
|
||||
from .common import postprocess_image
|
||||
from .common import deserialize_and_get_comfy_key, postprocess_image
|
||||
|
||||
|
||||
class Text2ImageHDNode():
|
||||
@@ -29,12 +29,13 @@ class Text2ImageHDNode():
|
||||
FUNCTION = "execute"
|
||||
|
||||
def __init__(self):
|
||||
self.api_url = "https://engine.prod.bria-api.com/v1/text-to-image/hd/2.3" #"http://0.0.0.0:5000/v1/text-to-image/hd/2.3"
|
||||
self.api_url = "https://engine.prod.bria-api.com/v1/text-to-image/hd/2.2" #"http://0.0.0.0:5000/v1/text-to-image/hd/2.3"
|
||||
|
||||
def execute(
|
||||
self, api_key, prompt, aspect_ratio, seed, negative_prompt,
|
||||
steps_num, prompt_enhancement, text_guidance_scale, medium, content_moderation=0,
|
||||
):
|
||||
api_key = deserialize_and_get_comfy_key(api_key)
|
||||
payload = {
|
||||
"prompt": prompt,
|
||||
"num_results": 1,
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
import requests
|
||||
import torch
|
||||
from ..common import postprocess_image, preprocess_image, image_to_base64
|
||||
from ..common import deserialize_and_get_comfy_key, postprocess_image, preprocess_image, image_to_base64
|
||||
|
||||
shot_by_text_api_url = (
|
||||
"https://engine.prod.bria-api.com/v1/product/lifestyle_shot_by_text"
|
||||
@@ -68,6 +68,7 @@ def create_text_payload(
|
||||
|
||||
validate_api_key(api_key)
|
||||
|
||||
|
||||
# Process image
|
||||
if isinstance(image, torch.Tensor):
|
||||
image = preprocess_image(image)
|
||||
@@ -126,9 +127,11 @@ def create_image_payload(image, ref_image, api_key, placement_type, **kwargs):
|
||||
|
||||
def make_api_request(api_url, payload, api_key, Placement_type = None):
|
||||
"""Make API request and return processed image"""
|
||||
headers = {"Content-Type": "application/json", "api_token": f"{api_key}"}
|
||||
|
||||
|
||||
try:
|
||||
api_key = deserialize_and_get_comfy_key(api_key)
|
||||
headers = {"Content-Type": "application/json", "api_token": f"{api_key}"}
|
||||
response = requests.post(api_url, json=payload, headers=headers)
|
||||
|
||||
if response.status_code == 200:
|
||||
|
||||
Reference in New Issue
Block a user