diff --git a/__init__.py b/__init__.py index cd511d8..b0c394b 100644 --- a/__init__.py +++ b/__init__.py @@ -15,8 +15,10 @@ from .nodes import ( TailoredPortraitNode, ReimagineNode, AttributionByImageNode, - FiboGenerateNode, - FiboRefineAndRegenerateNode, + GenerateImageNodeV2, + GenerateImageProNodeV2, + RefineImageNodeV2, + RefineImageProNodeV2, ) # Map the node class to a name used internally by ComfyUI @@ -37,8 +39,10 @@ NODE_CLASS_MAPPINGS = { "Text2ImageHDNode": Text2ImageHDNode, "ReimagineNode": ReimagineNode, "AttributionByImageNode": AttributionByImageNode, - "Text2ImageFiboGenerateNode": FiboGenerateNode, - "Text2ImageFiboRefineAndRegenerateNode": FiboRefineAndRegenerateNode, + "GenerateImageNodeV2": GenerateImageNodeV2, + "RefineImageNodeV2": RefineImageNodeV2, + "GenerateImageProNodeV2":GenerateImageProNodeV2, + "RefineImageProNodeV2":RefineImageProNodeV2 } # Map the node display name to the one shown in the ComfyUI node interface NODE_DISPLAY_NAME_MAPPINGS = { @@ -58,6 +62,8 @@ NODE_DISPLAY_NAME_MAPPINGS = { "Text2ImageHDNode": "Bria Text2Image HD", "ReimagineNode": "Bria Reimagine", "AttributionByImageNode": "Attribution By Image Node", - "Text2ImageFiboGenerateNode": "Bria FIBO - Generate", - "Text2ImageFiboRefineAndRegenerateNode": "Bria FIBO - Refine and Regenerate", + "GenerateImageNodeV2": "Generate Image", + "RefineImageNodeV2": "Refine and Regenerate Image", + "GenerateImageProNodeV2":"Generate Image - Pro", + "RefineImageProNodeV2":"Refine and Regenerate Image - Pro" } diff --git a/nodes/__init__.py b/nodes/__init__.py index 5c4e664..aaff270 100644 --- a/nodes/__init__.py +++ b/nodes/__init__.py @@ -14,5 +14,5 @@ from .text_2_image_fast_node import Text2ImageFastNode from .text_2_image_hd_node import Text2ImageHDNode from .reimagine_node import ReimagineNode from .attribution_by_image_node import AttributionByImageNode -from .fibo_generate_node import FiboGenerateNode -from .fibo_refine_node import FiboRefineAndRegenerateNode \ No newline at end of file +from .generate_image_node_v2 import GenerateImageNodeV2, GenerateImageProNodeV2 +from .refine_image_node_v2 import RefineImageNodeV2, RefineImageProNodeV2 \ No newline at end of file diff --git a/nodes/fibo_generate_node.py b/nodes/generate_image_node_v2.py similarity index 68% rename from nodes/fibo_generate_node.py rename to nodes/generate_image_node_v2.py index cfd4ce2..1baf513 100644 --- a/nodes/fibo_generate_node.py +++ b/nodes/generate_image_node_v2.py @@ -9,9 +9,13 @@ from .common import ( ) -class FiboGenerateNode: +class _BaseGenerateImageNodeV2: + """Base class for image generation nodes (standard & pro).""" + + api_url = None # Each subclass must define its API endpoint + @classmethod - def INPUT_TYPES(self): + def INPUT_TYPES(cls): return { "required": { "api_token": ("STRING", {"default": "BRIA_API_TOKEN"}), @@ -25,7 +29,7 @@ class FiboGenerateNode: ["1:1", "2:3", "3:2", "3:4", "4:3", "4:5", "5:4", "9:16", "16:9"], {"default": "1:1"}, ), - "steps_num": ("INT", {"default": 30, "min": 20, "max": 50}), + "steps_num": ("INT", {"default": 50, "min": 20, "max": 50}), "guidance_scale": ("INT", {"default": 5, "min": 3, "max": 5}), "seed": ("INT", {"default": 123456}), }, @@ -36,8 +40,37 @@ class FiboGenerateNode: CATEGORY = "API Nodes" FUNCTION = "execute" - def __init__(self): - self.api_url = "https://engine.prod.bria-api.com/v2/image/generate" + def _validate_token(self, api_token: str): + if api_token.strip() == "" or api_token.strip() == "BRIA_API_TOKEN": + raise Exception("Please insert a valid API token.") + + def _build_payload( + self, + prompt, + mode, + negative_prompt, + aspect_ratio, + steps_num, + guidance_scale, + seed, + image=None, + ): + payload = { + "prompt": prompt, + "mode": mode, + "negative_prompt": negative_prompt, + "aspect_ratio": aspect_ratio, + "steps_num": steps_num, + "guidance_scale": guidance_scale, + "seed": seed, + } + + if image is not None: + if isinstance(image, torch.Tensor): + image = preprocess_image(image) + payload["images"] = [image_to_base64(image)] + + return payload def execute( self, @@ -51,36 +84,26 @@ class FiboGenerateNode: seed, image=None, ): - - if api_token.strip() == "" or api_token.strip() == "BRIA_API_TOKEN": - raise Exception("Please insert a valid API token.") - - payload = { - "prompt": prompt, - "mode": mode, - "negative_prompt": negative_prompt, - "aspect_ratio": aspect_ratio, - "steps_num": steps_num, - "guidance_scale": guidance_scale, - "seed": seed, - } - - # Add image if provided - if image is not None: - if isinstance(image, torch.Tensor): - image = preprocess_image(image) - image_base64 = image_to_base64(image) - payload["image"] = image_base64 + self._validate_token(api_token) + payload = self._build_payload( + prompt, + mode, + negative_prompt, + aspect_ratio, + steps_num, + guidance_scale, + seed, + image, + ) headers = {"Content-Type": "application/json", "api_token": api_token} try: - # Send initial request to get status URL response = requests.post(self.api_url, json=payload, headers=headers) - if response.status_code == 200 or response.status_code == 202: + if response.status_code in (200, 202): print( - "Initial FIBO generate request successful, polling for completion..." + f"Initial request successful to {self.api_url}, polling for completion..." ) response_dict = response.json() status_url = response_dict.get("status_url") @@ -91,25 +114,33 @@ class FiboGenerateNode: print(f"Request ID: {request_id}, Status URL: {status_url}") - # Poll for completion final_response = poll_status_until_completed(status_url, api_token) - # Extract results result = final_response.get("result", {}) - print(result) result_image_url = result.get("image_url") structured_prompt = result.get("structured_prompt", "") used_seed = result.get("seed") - # Download and process the result image image_response = requests.get(result_image_url) result_image = postprocess_image(image_response.content) return (result_image, structured_prompt, used_seed) - else: - raise Exception( - f"Error: API request failed with status code {response.status_code} {response.text}" - ) + + raise Exception( + f"Error: API request failed with status code {response.status_code} {response.text}" + ) except Exception as e: raise Exception(f"{e}") + + +class GenerateImageNodeV2(_BaseGenerateImageNodeV2): + """Standard Image Generation Node""" + def __init__(self): + self.api_url = "https://engine.prod.bria-api.com/v2/image/generate" + + +class GenerateImageProNodeV2(_BaseGenerateImageNodeV2): + """Pro Image Generation Node""" + def __init__(self): + self.api_url = "https://engine.prod.bria-api.com/v2/image/generate/pro" diff --git a/nodes/fibo_refine_node.py b/nodes/refine_image_node_v2.py similarity index 64% rename from nodes/fibo_refine_node.py rename to nodes/refine_image_node_v2.py index 4f3f215..ec30da3 100644 --- a/nodes/fibo_refine_node.py +++ b/nodes/refine_image_node_v2.py @@ -1,11 +1,14 @@ import requests - from .common import poll_status_until_completed -class FiboRefineAndRegenerateNode: +class _BaseRefineImageNodeV2: + """Base class for refine image nodes (standard & pro).""" + + api_url = None # Must be overridden by subclasses + @classmethod - def INPUT_TYPES(self): + def INPUT_TYPES(cls): return { "required": { "api_token": ("STRING", {"default": "BRIA_API_TOKEN"}), @@ -19,7 +22,7 @@ class FiboRefineAndRegenerateNode: ["1:1", "2:3", "3:2", "3:4", "4:3", "4:5", "5:4", "9:16", "16:9"], {"default": "1:1"}, ), - "steps_num": ("INT", {"default": 30, "min": 20, "max": 50}), + "steps_num": ("INT", {"default": 50, "min": 20, "max": 50}), "guidance_scale": ("INT", {"default": 5, "min": 3, "max": 5}), "seed": ("INT", {"default": 123456}), }, @@ -30,8 +33,31 @@ class FiboRefineAndRegenerateNode: CATEGORY = "API Nodes" FUNCTION = "execute" - def __init__(self): - self.api_url = "https://engine.prod.bria-api.com/v2/structured_prompt/generate" + def _validate_token(self, api_token: str): + if api_token.strip() == "" or api_token.strip() == "BRIA_API_TOKEN": + raise Exception("Please insert a valid API token.") + + def _build_payload( + self, + prompt, + structured_prompt, + mode, + negative_prompt, + aspect_ratio, + steps_num, + guidance_scale, + seed, + ): + return { + "prompt": prompt, + "mode": mode, + "negative_prompt": negative_prompt, + "aspect_ratio": aspect_ratio, + "steps_num": steps_num, + "guidance_scale": guidance_scale, + "seed": seed, + "structured_prompt": structured_prompt, + } def execute( self, @@ -45,31 +71,25 @@ class FiboRefineAndRegenerateNode: guidance_scale, seed, ): - - if api_token.strip() == "" or api_token.strip() == "BRIA_API_TOKEN": - raise Exception("Please insert a valid API token.") - - payload = { - "prompt": prompt, - "mode": mode, - "negative_prompt": negative_prompt, - "aspect_ratio": aspect_ratio, - "steps_num": steps_num, - "guidance_scale": guidance_scale, - "seed": seed, - "structured_prompt":structured_prompt - } + self._validate_token(api_token) + payload = self._build_payload( + prompt, + structured_prompt, + mode, + negative_prompt, + aspect_ratio, + steps_num, + guidance_scale, + seed, + ) headers = {"Content-Type": "application/json", "api_token": api_token} try: - # Send initial request to get status URL response = requests.post(self.api_url, json=payload, headers=headers) - if response.status_code == 200 or response.status_code == 202: - print( - "Initial FIBO Refine and Regenerate request successful, polling for completion..." - ) + if response.status_code in (200, 202): + print(f"Initial refine request successful to {self.api_url}, polling for completion...") response_dict = response.json() status_url = response_dict.get("status_url") request_id = response_dict.get("request_id") @@ -79,20 +99,29 @@ class FiboRefineAndRegenerateNode: print(f"Request ID: {request_id}, Status URL: {status_url}") - # Poll for completion final_response = poll_status_until_completed(status_url, api_token) - # Extract results result = final_response.get("result", {}) - print(result) structured_prompt = result.get("structured_prompt", "") used_seed = result.get("seed", seed) return (structured_prompt, used_seed) - else: - raise Exception( - f"Error: API request failed with status code {response.status_code} {response.text}" - ) + + raise Exception( + f"Error: API request failed with status code {response.status_code} {response.text}" + ) except Exception as e: raise Exception(f"{e}") + + +class RefineImageNodeV2(_BaseRefineImageNodeV2): + """Standard Refine Image Node""" + def __init__(self): + self.api_url = "https://engine.prod.bria-api.com/v2/structured_prompt/generate" + + +class RefineImageProNodeV2(_BaseRefineImageNodeV2): + """Pro Refine Image Node""" + def __init__(self): + self.api_url = "https://engine.prod.bria-api.com/v2/structured_prompt/generate/pro"