This commit is contained in:
Ubuntu
2025-09-16 17:11:06 +00:00
parent 711080ffb0
commit 080c17f4a4
7 changed files with 282 additions and 112 deletions
+66 -11
View File
@@ -5,6 +5,7 @@ import torch
import base64
from torchvision.transforms import ToPILImage
import requests
import time
def postprocess_image(image):
result_image = Image.open(io.BytesIO(image))
@@ -44,7 +45,7 @@ def preprocess_mask(mask):
return mask
def process_request(api_url, image, mask, api_key):
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.")
@@ -58,10 +59,12 @@ def process_request(api_url, image, mask, api_key):
image_base64 = image_to_base64(image)
mask_base64 = image_to_base64(mask)
# Prepare the API request payload
# Prepare the API request payload for v2 API
payload = {
"file": f"{image_base64}",
"mask_file": f"{mask_base64}"
"image": image_base64,
"mask": mask_base64,
"visual_input_content_moderation":visual_input_content_moderation,
"visual_output_content_moderation":visual_output_content_moderation
}
headers = {
@@ -70,13 +73,23 @@ def process_request(api_url, image, mask, api_key):
}
try:
response = requests.post(api_url, json=payload, headers=headers)
# Check for successful response
if response.status_code == 200:
print('response is 200')
# Process the output image from API response
response = requests.post(api_url, json=payload, headers=headers)
if response.status_code == 200 or response.status_code == 202:
print('Initial request successful, polling for completion...')
response_dict = response.json()
image_response = requests.get(response_dict['result_url'])
status_url = response_dict.get('status_url')
request_id = response_dict.get('request_id')
if not status_url:
raise Exception("No status_url returned from API")
print(f"Request ID: {request_id}, Status URL: {status_url}")
final_response = poll_status_until_completed(status_url, api_key)
result_image_url = final_response['result']['image_url']
# Download and process the result image
image_response = requests.get(result_image_url)
result_image = Image.open(io.BytesIO(image_response.content))
result_image = result_image.convert("RGBA")
result_image = np.array(result_image).astype(np.float32) / 255.0
@@ -84,9 +97,51 @@ def process_request(api_url, image, mask, api_key):
# image_tensor = image_tensor = ToTensor()(output_image)
# image_tensor = image_tensor.permute(1, 2, 0) / 255.0 # Shape now becomes [1, 2200, 1548, 3]
# print(f"output tensor shape is: {image_tensor.shape}")
return (result_image,)
return (result_image,)
else:
raise Exception(f"Error: API request failed with status code {response.status_code}")
except Exception as e:
raise Exception(f"{e}")
def poll_status_until_completed(status_url, api_key, timeout=360, check_interval=2):
"""
Poll a status URL until the status is COMPLETED or timeout is reached.
Args:
status_url (str): The status URL to poll
api_key (str): API token for authentication
timeout (int): Maximum time to wait in seconds (default: 360)
check_interval (int): Time between checks in seconds (default: 2)
Returns:
dict: The final response containing the result
Raises:
Exception: If timeout is reached or API request fails
"""
start_time = time.time()
headers = {"api_token": api_key}
while time.time() - start_time < timeout:
try:
response = requests.get(status_url, headers=headers)
if response.status_code == 200 or response.status_code == 202:
response_dict = response.json()
status = response_dict.get("status", "").upper()
if status == "COMPLETED":
return response_dict
elif status == "FAILED":
raise Exception(f"Request failed: {response_dict}")
else:
print(f"Status: {status}, waiting...")
time.sleep(check_interval)
else:
raise Exception(f"Status check failed with status code {response.status_code}")
except requests.exceptions.RequestException as e:
raise Exception(f"Error checking status: {e}")
raise Exception(f"Timeout reached after {timeout} seconds")
+7 -3
View File
@@ -8,6 +8,10 @@ class EraserNode():
"image": ("IMAGE",), # Input image from another node
"mask": ("MASK",), # Binary mask input
"api_key": ("STRING", {"default": "BRIA_API_TOKEN"}) # API Key input with a default value
},
"optional": {
"visual_input_content_moderation": ("BOOLEAN", {"default": False}),
"visual_output_content_moderation": ("BOOLEAN", {"default": False}),
}
}
@@ -17,9 +21,9 @@ class EraserNode():
FUNCTION = "execute" # This is the method that will be executed
def __init__(self):
self.api_url = "https://engine.prod.bria-api.com/v1/eraser" # Eraser API URL
self.api_url = "https://engine.prod.bria-api.com/v2/image/edit/erase" # Eraser API URL
# Define the execute method as expected by ComfyUI
def execute(self, image, mask, api_key):
return process_request(self.api_url, image, mask, api_key)
def execute(self, image, mask, api_key, visual_input_content_moderation, visual_output_content_moderation):
return process_request(self.api_url, image, mask, api_key, visual_input_content_moderation, visual_output_content_moderation)
+29 -12
View File
@@ -4,7 +4,7 @@ from PIL import Image
import io
import torch
from .common import image_to_base64, preprocess_image, preprocess_mask
from .common import preprocess_image, preprocess_mask, image_to_base64, poll_status_until_completed
class GenFillNode():
@@ -18,7 +18,12 @@ class GenFillNode():
"api_key": ("STRING", {"default": "BRIA_API_TOKEN"}), # API Key input with a default value
},
"optional": {
"seed": ("INT", {"default": 123456})
"seed": ("INT", {"default": 123456}),
"prompt_content_moderation": ("BOOLEAN", {"default": True}),
"visual_input_content_moderation": ("BOOLEAN", {"default": False}),
"visual_output_content_moderation": ("BOOLEAN", {"default": False}),
}
}
@@ -28,10 +33,10 @@ class GenFillNode():
FUNCTION = "execute" # This is the method that will be executed
def __init__(self):
self.api_url = "https://engine.prod.bria-api.com/v1/gen_fill" # Eraser API URL
self.api_url = "https://engine.prod.bria-api.com/v2/image/edit/gen_fill"
# Define the execute method as expected by ComfyUI
def execute(self, image, mask, prompt, api_key, seed):
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.")
@@ -47,12 +52,14 @@ class GenFillNode():
# Prepare the API request payload
payload = {
"file": f"{image_base64}",
"mask_file": f"{mask_base64}",
"image": image_base64,
"mask": mask_base64,
"prompt": prompt,
"negative_prompt": "blurry",
"sync": True,
"seed": seed,
"prompt_content_moderation":prompt_content_moderation,
"visual_input_content_moderation":visual_input_content_moderation,
"visual_output_content_moderation":visual_output_content_moderation
}
headers = {
@@ -61,13 +68,23 @@ class GenFillNode():
}
try:
# Send initial request to get status URL
response = requests.post(self.api_url, json=payload, headers=headers)
# Check for successful response
if response.status_code == 200:
print('response is 200')
# Process the output image from API response
if response.status_code == 200 or response.status_code == 202:
print('Initial genfill request successful, polling for completion...')
response_dict = response.json()
image_response = requests.get(response_dict['urls'][0])
status_url = response_dict.get('status_url')
request_id = response_dict.get('request_id')
if not status_url:
raise Exception("No status_url returned from API")
print(f"Request ID: {request_id}, Status URL: {status_url}")
final_response = poll_status_until_completed(status_url, api_key)
result_image_url = final_response['result']['image_url']
image_response = requests.get(result_image_url)
result_image = Image.open(io.BytesIO(image_response.content))
result_image = result_image.convert("RGB")
result_image = np.array(result_image).astype(np.float32) / 255.0
+60 -24
View File
@@ -4,7 +4,7 @@ from PIL import Image
import io
import torch
from .common import image_to_base64, preprocess_image
from .common import image_to_base64, preprocess_image, poll_status_until_completed
class ImageExpansionNode():
@@ -13,17 +13,21 @@ class ImageExpansionNode():
return {
"required": {
"image": ("IMAGE",), # Input image from another node
"original_image_size": ("STRING",),
"original_image_location": ("STRING",),
"api_key": ("STRING", {"default": "BRIA_API_TOKEN"}), # API Key input with a default value
},
"optional": {
"original_image_size": ("STRING",),
"original_image_location": ("STRING",),
"canvas_size": ("STRING", {"default": "1000, 1000"}),
"aspect_ratio": (["1:1", "2:3", "3:2", "3:4", "4:3", "4:5", "5:4", "9:16", "16:9","None"], {"default": "None"}),
"prompt": ("STRING", {"default": ""}),
"seed": ("INT", {"default": 681794}),
"negative_prompt": ("STRING", {"default": "Ugly, mutated"}),
"content_moderation": ("BOOLEAN", {"default": False}),
"prompt_content_moderation": ("BOOLEAN", {"default": False}),
"preserve_alpha": ("BOOLEAN", {"default": True}),
"visual_input_content_moderation": ("BOOLEAN", {"default": False}),
"visual_output_content_moderation": ("BOOLEAN", {"default": False}),
}
}
@@ -33,24 +37,27 @@ class ImageExpansionNode():
FUNCTION = "execute" # This is the method that will be executed
def __init__(self):
self.api_url = "https://engine.prod.bria-api.com/v1/image_expansion" # Image Expansion API URL
self.api_url = "https://engine.prod.bria-api.com/v2/image/edit/expand" # Image Expansion API URL
# Define the execute method as expected by ComfyUI
def execute(self, image,
original_image_size,
original_image_location,
canvas_size,
aspect_ratio,
prompt,
seed,
negative_prompt,
content_moderation,
prompt_content_moderation,
preserve_alpha,
visual_input_content_moderation,
visual_output_content_moderation,
api_key):
if api_key.strip() == "" or api_key.strip() == "BRIA_API_TOKEN":
raise Exception("Please insert a valid API key.")
original_image_size = [int(x.strip()) for x in original_image_size.split(",")]
original_image_location = [int(x.strip()) for x in original_image_location.split(",")]
canvas_size = [int(x.strip()) for x in canvas_size.split(",")]
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 ()
if prompt == "":
prompt = None
@@ -63,18 +70,32 @@ class ImageExpansionNode():
# Convert the image directly to Base64 string
image_base64 = image_to_base64(image)
# Prepare the API request payload
payload = {
"file": f"{image_base64}",
"original_image_size": original_image_size,
"original_image_location": original_image_location,
"canvas_size": canvas_size,
if aspect_ratio and aspect_ratio != "None":
payload = {
"image": image_base64,
"aspect_ratio": aspect_ratio,
"prompt": prompt,
"negative_prompt": negative_prompt,
"seed": seed,
"content_moderation": content_moderation
"prompt_content_moderation": prompt_content_moderation,
"preserve_alpha": preserve_alpha,
"visual_input_content_moderation": visual_input_content_moderation,
"visual_output_content_moderation": visual_output_content_moderation
}
else:
payload = {
"image": image_base64,
"original_image_size": original_image_size,
"original_image_location": original_image_location,
"canvas_size": canvas_size,
"prompt": prompt,
"negative_prompt": negative_prompt,
"seed": seed,
"prompt_content_moderation": prompt_content_moderation,
"preserve_alpha": preserve_alpha,
"visual_input_content_moderation": visual_input_content_moderation,
"visual_output_content_moderation": visual_output_content_moderation
}
headers = {
"Content-Type": "application/json",
@@ -83,19 +104,34 @@ class ImageExpansionNode():
try:
response = requests.post(self.api_url, json=payload, headers=headers)
# Check for successful response
if response.status_code == 200:
print('response is 200')
# Process the output image from API response
if response.status_code == 200 or response.status_code == 202:
print('Initial image expansion request successful, polling for completion...')
response_dict = response.json()
image_response = requests.get(response_dict['result_url'])
status_url = response_dict.get('status_url')
request_id = response_dict.get('request_id')
if not status_url:
raise Exception("No status_url returned from API")
print(f"Request ID: {request_id}, Status URL: {status_url}")
# Poll status URL until completion
final_response = poll_status_until_completed(status_url, api_key)
# Get the result image URL
result_image_url = final_response['result']['image_url']
# Download and process the result image
image_response = requests.get(result_image_url)
result_image = Image.open(io.BytesIO(image_response.content))
result_image = result_image.convert("RGB")
result_image = np.array(result_image).astype(np.float32) / 255.0
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}")
+30 -10
View File
@@ -4,7 +4,7 @@ from PIL import Image
import io
import torch
from .common import preprocess_image, image_to_base64
from .common import preprocess_image, image_to_base64, poll_status_until_completed
class RemoveForegroundNode():
@@ -16,7 +16,9 @@ class RemoveForegroundNode():
"api_key": ("STRING", {"default": "BRIA_API_TOKEN"}), # API Key input with a default value
},
"optional": {
"content_moderation": ("BOOLEAN", {"default": False}),
"visual_input_content_moderation": ("BOOLEAN", {"default": False}),
"visual_output_content_moderation": ("BOOLEAN", {"default": False}),
"preserve_alpha": ("BOOLEAN", {"default": True}),
}
}
@@ -26,10 +28,10 @@ class RemoveForegroundNode():
FUNCTION = "execute" # This is the method that will be executed
def __init__(self):
self.api_url = "https://engine.prod.bria-api.com/v1/erase_foreground" # remove foreground API URL
self.api_url = "https://engine.prod.bria-api.com/v2/image/edit/erase_foreground" # remove foreground API URL
# Define the execute method as expected by ComfyUI
def execute(self, image, content_moderation, api_key):
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.")
@@ -44,7 +46,12 @@ class RemoveForegroundNode():
# files=[('file',('temp_img.jpeg', open(temp_img_path, 'rb'),'image/jpeg'))
# ]
payload = {"file": image_to_base64(image), "content_moderation": content_moderation}
payload = {
"image": image_to_base64(image),
"visual_input_content_moderation": visual_input_content_moderation,
"visual_output_content_moderation":visual_output_content_moderation,
"preserve_alpha": preserve_alpha
}
headers = {
"Content-Type": "application/json",
@@ -53,12 +60,25 @@ class RemoveForegroundNode():
try:
response = requests.post(self.api_url, json=payload, headers=headers)
# Check for successful response
if response.status_code == 200:
print('response is 200')
# Process the output image from API response
if response.status_code == 200 or response.status_code == 202:
print('Initial request successful, polling for completion...')
response_dict = response.json()
image_response = requests.get(response_dict['result_url'])
status_url = response_dict.get('status_url')
request_id = response_dict.get('request_id')
if not status_url:
raise Exception("No status_url returned from API")
print(f"Request ID: {request_id}, Status URL: {status_url}")
# Poll status URL until completion
final_response = poll_status_until_completed(status_url, api_key)
# Get the result image URL
result_image_url = final_response['result']['image_url']
image_response = requests.get(result_image_url)
result_image = Image.open(io.BytesIO(image_response.content))
result_image = np.array(result_image).astype(np.float32) / 255.0
result_image = torch.from_numpy(result_image)[None,]
+50 -32
View File
@@ -4,7 +4,7 @@ from PIL import Image
import io
import torch
from .common import image_to_base64, preprocess_image, preprocess_mask
from .common import image_to_base64, preprocess_image, preprocess_mask, poll_status_until_completed
class ReplaceBgNode():
@@ -16,16 +16,17 @@ class ReplaceBgNode():
"api_key": ("STRING", {"default": "BRIA_API_TOKEN"}), # API Key input with a default value
},
"optional": {
"mode": (["base", "fast", "high_control"], {"default": "high_control"}),
"bg_prompt": ("STRING",),
"ref_image": ("IMAGE",), # Input ref image from another node
"mode": (["base", "fast", "high_control"], {"default": "base"}),
"prompt": ("STRING",),
"ref_images": ("IMAGE",),
"refine_prompt": ("BOOLEAN", {"default": True}),
"enhance_ref_image": ("BOOLEAN", {"default": True}),
"enhance_ref_images": ("BOOLEAN", {"default": True}),
"original_quality": ("BOOLEAN", {"default": False}),
"force_rmbg": ("BOOLEAN", {"default": False}),
"negative_prompt": ("STRING", {"default": None}),
"seed": ("INT", {"default": 681794}),
"content_moderation": ("BOOLEAN", {"default": False}),
"visual_output_content_moderation": ("BOOLEAN", {"default": False}),
"prompt_content_moderation": ("BOOLEAN", {"default": False}),
"force_background_detection": ("BOOLEAN", {"default": False}),
}
}
@@ -35,20 +36,21 @@ class ReplaceBgNode():
FUNCTION = "execute" # This is the method that will be executed
def __init__(self):
self.api_url = "https://engine.prod.bria-api.com/v1/background/replace" # Replace BG API URL
self.api_url = "https://engine.prod.bria-api.com/v2/image/edit/replace_background" # Replace BG API URL
# Define the execute method as expected by ComfyUI
def execute(self, image, mode,
refine_prompt,
enhance_ref_image,
original_quality,
force_rmbg,
negative_prompt,
seed,
api_key,
content_moderation,
bg_prompt=None,
ref_image=None,):
visual_output_content_moderation,
prompt_content_moderation,
enhance_ref_images,
force_background_detection,
prompt=None,
ref_images=None,):
if api_key.strip() == "" or api_key.strip() == "BRIA_API_TOKEN":
raise Exception("Please insert a valid API key.")
@@ -56,28 +58,29 @@ class ReplaceBgNode():
if isinstance(image, torch.Tensor):
image = preprocess_image(image)
# Convert the image and mask directly to Base64 strings
# Convert the image to Base64 string
image_base64 = image_to_base64(image)
ref_image_file = None # initialization, will be updated if it is supplied
if ref_image is not None:
ref_image = preprocess_image(ref_image)
ref_image_file = image_to_base64(ref_image)
if ref_images is not None:
ref_images = preprocess_image(ref_images)
ref_images = [image_to_base64(ref_images)]
else:
ref_images=[]
# Prepare the API request payload
# Prepare the API request payload for v2 API
payload = {
"file": f"{image_base64}",
"image": image_base64,
"mode": mode,
"bg_prompt": bg_prompt,
"ref_image_file": ref_image_file,
"prompt": prompt,
"ref_images":ref_images,
"refine_prompt": refine_prompt,
"enhance_ref_image": enhance_ref_image,
"original_quality": original_quality,
"force_rmbg": force_rmbg,
"negative_prompt": negative_prompt,
"seed": seed,
"sync": True,
"num_results": 1,
"content_moderation": content_moderation
"prompt_content_moderation": prompt_content_moderation,
"visual_output_content_moderation":visual_output_content_moderation,
"enhance_ref_images":enhance_ref_images,
"force_background_detection": force_background_detection
}
headers = {
@@ -87,16 +90,31 @@ class ReplaceBgNode():
try:
response = requests.post(self.api_url, json=payload, headers=headers)
# Check for successful response
if response.status_code == 200:
print('response is 200')
# Process the output image from API response
if response.status_code == 200 or response.status_code == 202:
print('Initial replace background request successful, polling for completion...')
response_dict = response.json()
image_response = requests.get(response_dict['result'][0][0]) # first indexing for batched, second for url
status_url = response_dict.get('status_url')
request_id = response_dict.get('request_id')
if not status_url:
raise Exception("No status_url returned from API")
print(f"Request ID: {request_id}, Status URL: {status_url}")
# Poll status URL until completion
final_response = poll_status_until_completed(status_url, api_key)
# Get the result image URL
result_image_url = final_response['result']['image_url']
# Download and process the result image
image_response = requests.get(result_image_url)
result_image = Image.open(io.BytesIO(image_response.content))
result_image = result_image.convert("RGB")
result_image = np.array(result_image).astype(np.float32) / 255.0
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}")
+40 -20
View File
@@ -4,8 +4,7 @@ from PIL import Image
import io
import torch
from .common import preprocess_image
from io import BytesIO
from .common import preprocess_image, image_to_base64, poll_status_until_completed
class RmbgNode():
@classmethod
@@ -16,7 +15,10 @@ class RmbgNode():
"api_key": ("STRING", {"default": "BRIA_API_TOKEN"}), # API Key input with a default value
},
"optional": {
"content_moderation": ("BOOLEAN", {"default": False}),
"visual_input_content_moderation": ("BOOLEAN", {"default": False}),
"visual_output_content_moderation": ("BOOLEAN", {"default": False}),
"preserve_alpha": ("BOOLEAN", {"default": True}),
}
}
@@ -26,10 +28,10 @@ class RmbgNode():
FUNCTION = "execute" # This is the method that will be executed
def __init__(self):
self.api_url = "https://engine.prod.bria-api.com/v1/background/remove" # RMBG API URL
self.api_url = "https://engine.prod.bria-api.com/v2/image/edit/remove_background" # RMBG API URL
# Define the execute method as expected by ComfyUI
def execute(self, image, content_moderation, api_key):
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.")
@@ -37,28 +39,46 @@ class RmbgNode():
if isinstance(image, torch.Tensor):
image = preprocess_image(image)
# Prepare the API request payload
image_buffer = BytesIO()
image.save(image_buffer, format="JPEG")
# Convert image to base64 for the new API format
image_base64 = image_to_base64(image)
payload = {
"image": image_base64,
"visual_input_content_moderation": visual_input_content_moderation,
"visual_output_content_moderation":visual_output_content_moderation,
"preserve_alpha":preserve_alpha
}
# Get binary data from buffer
image_buffer.seek(0) # Move cursor to the start of the buffer
binary_data = image_buffer.read()
files=[('file',('temp_img.jpeg', BytesIO(binary_data),'image/jpeg'))]
payload = {"content_moderation": content_moderation}
headers = {
"Content-Type": "application/json",
"api_token": f"{api_key}"
}
try:
response = requests.post(self.api_url, data=payload, headers={"api_token": api_key}, files=files)
# Check for successful response
if response.status_code == 200:
print('response is 200')
# Process the output image from API response
response = requests.post(self.api_url, json=payload, headers=headers)
if response.status_code == 200 or response.status_code == 202:
print('Initial RMBG request successful, polling for completion...')
response_dict = response.json()
image_response = requests.get(response_dict['result_url'])
status_url = response_dict.get('status_url')
request_id = response_dict.get('request_id')
if not status_url:
raise Exception("No status_url returned from API")
print(f"Request ID: {request_id}, Status URL: {status_url}")
final_response = poll_status_until_completed(status_url, api_key)
# Get the result image URL
result_image_url = final_response['result']['image_url']
# Download and process the result image
image_response = requests.get(result_image_url)
result_image = Image.open(io.BytesIO(image_response.content))
result_image = np.array(result_image).astype(np.float32) / 255.0
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}")