From 080c17f4a4fc5447a38ca55ae5fd8c4823088834 Mon Sep 17 00:00:00 2001 From: Ubuntu Date: Tue, 16 Sep 2025 17:11:06 +0000 Subject: [PATCH] WAI-3976 --- nodes/common.py | 77 +++++++++++++++++++++++++----- nodes/eraser_node.py | 10 ++-- nodes/generative_fill_node.py | 41 +++++++++++----- nodes/image_expansion_node.py | 84 +++++++++++++++++++++++---------- nodes/remove_foreground_node.py | 40 ++++++++++++---- nodes/replace_bg_node.py | 82 +++++++++++++++++++------------- nodes/rmbg_node.py | 60 +++++++++++++++-------- 7 files changed, 282 insertions(+), 112 deletions(-) diff --git a/nodes/common.py b/nodes/common.py index b57a492..b7dbc61 100644 --- a/nodes/common.py +++ b/nodes/common.py @@ -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") diff --git a/nodes/eraser_node.py b/nodes/eraser_node.py index b6a3c69..638d9b8 100644 --- a/nodes/eraser_node.py +++ b/nodes/eraser_node.py @@ -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) \ No newline at end of file diff --git a/nodes/generative_fill_node.py b/nodes/generative_fill_node.py index b30f005..f1fa731 100644 --- a/nodes/generative_fill_node.py +++ b/nodes/generative_fill_node.py @@ -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 diff --git a/nodes/image_expansion_node.py b/nodes/image_expansion_node.py index 4af3e6b..562c0ce 100644 --- a/nodes/image_expansion_node.py +++ b/nodes/image_expansion_node.py @@ -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}") diff --git a/nodes/remove_foreground_node.py b/nodes/remove_foreground_node.py index 3a81b82..373c328 100644 --- a/nodes/remove_foreground_node.py +++ b/nodes/remove_foreground_node.py @@ -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,] diff --git a/nodes/replace_bg_node.py b/nodes/replace_bg_node.py index 2ad30bd..7e8fb7c 100644 --- a/nodes/replace_bg_node.py +++ b/nodes/replace_bg_node.py @@ -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}") diff --git a/nodes/rmbg_node.py b/nodes/rmbg_node.py index 9f0df9f..ce025ba 100644 --- a/nodes/rmbg_node.py +++ b/nodes/rmbg_node.py @@ -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}")