WAI-4698:set User-Agent header
This commit is contained in:
@@ -1,6 +1,12 @@
|
||||
import requests
|
||||
|
||||
from .common import deserialize_and_get_comfy_key, image_to_base64, normalize_images_input, poll_status_until_completed, to_pil_safe
|
||||
from .common import (
|
||||
bria_json_headers,
|
||||
image_to_base64,
|
||||
normalize_images_input,
|
||||
poll_status_until_completed,
|
||||
to_pil_safe,
|
||||
)
|
||||
|
||||
class AttributionByImageNode():
|
||||
@classmethod
|
||||
@@ -24,8 +30,6 @@ class AttributionByImageNode():
|
||||
def execute(self, images, 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)
|
||||
|
||||
images = normalize_images_input(images)
|
||||
|
||||
batch_results = []
|
||||
@@ -39,10 +43,7 @@ class AttributionByImageNode():
|
||||
"model_version": model_version,
|
||||
}
|
||||
|
||||
headers = {
|
||||
"Content-Type": "application/json",
|
||||
"api_token": f"{api_key}"
|
||||
}
|
||||
headers = bria_json_headers(api_key)
|
||||
|
||||
response = requests.post(self.api_url, json=payload, headers=headers)
|
||||
|
||||
|
||||
+10
-28
@@ -6,16 +6,17 @@ import base64
|
||||
from torchvision.transforms import ToPILImage
|
||||
import requests
|
||||
import time
|
||||
import json
|
||||
|
||||
|
||||
COMFY_KEY_ERROR = (
|
||||
"Invalid Token Type\n\n"
|
||||
"The API token you’ve entered is not a ComfyUI token.\n"
|
||||
"Please use the valid token from your BRIA Account API Keys page:\n"
|
||||
"https://platform.bria.ai/console/account/api-keys"
|
||||
)
|
||||
BRIA_COMFYUI_USER_AGENT = "bria/ComfyUI"
|
||||
|
||||
def bria_json_headers(api_token: str) -> dict:
|
||||
"""Headers for JSON POST requests to Bria API."""
|
||||
return {
|
||||
"Content-Type": "application/json",
|
||||
"api_token": api_token,
|
||||
"User-Agent": BRIA_COMFYUI_USER_AGENT,
|
||||
}
|
||||
def postprocess_image(image):
|
||||
result_image = Image.open(io.BytesIO(image))
|
||||
result_image = result_image.convert("RGB")
|
||||
@@ -86,7 +87,6 @@ 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):
|
||||
@@ -106,10 +106,7 @@ def process_request(api_url, image, mask, api_key, visual_input_content_moderati
|
||||
"visual_output_content_moderation":visual_output_content_moderation
|
||||
}
|
||||
|
||||
headers = {
|
||||
"Content-Type": "application/json",
|
||||
"api_token": f"{api_key}"
|
||||
}
|
||||
headers = bria_json_headers(api_key)
|
||||
|
||||
try:
|
||||
response = requests.post(api_url, json=payload, headers=headers)
|
||||
@@ -161,7 +158,7 @@ def poll_status_until_completed(status_url, api_key, timeout=360, check_interval
|
||||
Exception: If timeout is reached or API request fails
|
||||
"""
|
||||
start_time = time.time()
|
||||
headers = {"api_token": api_key}
|
||||
headers = bria_json_headers(api_key)
|
||||
|
||||
while time.time() - start_time < timeout:
|
||||
try:
|
||||
@@ -185,21 +182,6 @@ def poll_status_until_completed(status_url, api_key, timeout=360, check_interval
|
||||
|
||||
raise Exception(f"Timeout reached after {timeout} seconds")
|
||||
|
||||
def deserialize_and_get_comfy_key(encoded: str) -> str:
|
||||
"""
|
||||
Decodes a base64-encoded JSON token and returns the ComfyUI API key.
|
||||
"""
|
||||
try:
|
||||
decoded = base64.b64decode(encoded).decode("utf-8")
|
||||
payload = json.loads(decoded)
|
||||
|
||||
if payload.get("type") != "comfy":
|
||||
raise Exception(COMFY_KEY_ERROR)
|
||||
|
||||
return payload.get("apiKey")
|
||||
|
||||
except Exception as e:
|
||||
raise Exception(COMFY_KEY_ERROR)
|
||||
|
||||
def normalize_images_input(images):
|
||||
"""
|
||||
|
||||
@@ -2,12 +2,12 @@ import requests
|
||||
import torch
|
||||
|
||||
from .common import (
|
||||
deserialize_and_get_comfy_key,
|
||||
postprocess_image,
|
||||
preprocess_image,
|
||||
bria_json_headers,
|
||||
image_to_base64,
|
||||
poll_status_until_completed,
|
||||
preprocess_image,
|
||||
preprocess_mask,
|
||||
postprocess_image,
|
||||
)
|
||||
|
||||
|
||||
@@ -124,9 +124,7 @@ class FIBOEditNode:
|
||||
guidance_scale,
|
||||
seed,
|
||||
)
|
||||
api_token = deserialize_and_get_comfy_key(api_token)
|
||||
|
||||
headers = {"Content-Type": "application/json", "api_token": api_token}
|
||||
headers = bria_json_headers(api_token)
|
||||
|
||||
try:
|
||||
response = requests.post(self.api_url, json=payload, headers=headers)
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
import requests
|
||||
from .common import (
|
||||
deserialize_and_get_comfy_key,
|
||||
bria_json_headers,
|
||||
image_to_base64,
|
||||
normalize_images_input,
|
||||
poll_status_until_completed,
|
||||
@@ -40,8 +40,6 @@ class FIBOEditStructuredInstructionNode:
|
||||
|
||||
def execute(self, api_token, images, instruction):
|
||||
self._validate_token(api_token)
|
||||
api_token = deserialize_and_get_comfy_key(api_token)
|
||||
|
||||
# Normalize input to list of PIL images
|
||||
images = normalize_images_input(images)
|
||||
|
||||
@@ -50,7 +48,7 @@ class FIBOEditStructuredInstructionNode:
|
||||
for idx, pil_image in enumerate(images):
|
||||
try:
|
||||
payload = self._build_payload(pil_image, instruction)
|
||||
headers = {"Content-Type": "application/json", "api_token": api_token}
|
||||
headers = bria_json_headers(api_token)
|
||||
|
||||
response = requests.post(self.api_url, json=payload, headers=headers)
|
||||
|
||||
|
||||
@@ -1,11 +1,11 @@
|
||||
import requests
|
||||
import torch
|
||||
from .common import (
|
||||
deserialize_and_get_comfy_key,
|
||||
normalize_images_input,
|
||||
postprocess_image,
|
||||
bria_json_headers,
|
||||
image_to_base64,
|
||||
normalize_images_input,
|
||||
poll_status_until_completed,
|
||||
postprocess_image,
|
||||
)
|
||||
class GenerateImageLiteNodeV2:
|
||||
"""Lite Image Generation Node (multi-image compatible)"""
|
||||
@@ -80,8 +80,6 @@ class GenerateImageLiteNodeV2:
|
||||
images=None,
|
||||
):
|
||||
self._validate_token(api_token)
|
||||
api_token = deserialize_and_get_comfy_key(api_token)
|
||||
|
||||
images_list = normalize_images_input(images) if images is not None else [None]
|
||||
|
||||
# Structured prompts per image
|
||||
@@ -120,7 +118,7 @@ class GenerateImageLiteNodeV2:
|
||||
ref_image,
|
||||
)
|
||||
|
||||
headers = {"Content-Type": "application/json", "api_token": api_token}
|
||||
headers = bria_json_headers(api_token)
|
||||
response = requests.post(self.api_url, json=payload, headers=headers)
|
||||
if response.status_code not in (200, 202):
|
||||
raise Exception(
|
||||
|
||||
@@ -2,11 +2,11 @@ import requests
|
||||
import torch
|
||||
|
||||
from .common import (
|
||||
deserialize_and_get_comfy_key,
|
||||
normalize_images_input,
|
||||
postprocess_image,
|
||||
bria_json_headers,
|
||||
image_to_base64,
|
||||
normalize_images_input,
|
||||
poll_status_until_completed,
|
||||
postprocess_image,
|
||||
)
|
||||
|
||||
|
||||
@@ -88,8 +88,6 @@ class GenerateImageNodeV2:
|
||||
images=None,
|
||||
):
|
||||
self._validate_token(api_token)
|
||||
api_token = deserialize_and_get_comfy_key(api_token)
|
||||
|
||||
images_list = normalize_images_input(images) if images is not None else [None]
|
||||
|
||||
# Structured prompts per image
|
||||
@@ -128,7 +126,7 @@ class GenerateImageNodeV2:
|
||||
ref_image,
|
||||
)
|
||||
|
||||
headers = {"Content-Type": "application/json", "api_token": api_token}
|
||||
headers = bria_json_headers(api_token)
|
||||
response = requests.post(self.api_url, json=payload, headers=headers)
|
||||
|
||||
if response.status_code not in (200, 202):
|
||||
|
||||
@@ -1,8 +1,8 @@
|
||||
import requests
|
||||
from .common import (
|
||||
deserialize_and_get_comfy_key,
|
||||
normalize_images_input,
|
||||
bria_json_headers,
|
||||
image_to_base64,
|
||||
normalize_images_input,
|
||||
poll_status_until_completed,
|
||||
)
|
||||
|
||||
@@ -45,8 +45,6 @@ class GenerateStructuredPromptLiteNodeV2:
|
||||
|
||||
def execute(self, api_token, prompt, seed, structured_prompt, images=None):
|
||||
self._validate_token(api_token)
|
||||
api_token = deserialize_and_get_comfy_key(api_token)
|
||||
|
||||
images_list = normalize_images_input(images) if images is not None else [None]
|
||||
|
||||
# Seeds per image
|
||||
@@ -80,7 +78,7 @@ class GenerateStructuredPromptLiteNodeV2:
|
||||
image,
|
||||
)
|
||||
|
||||
headers = {"Content-Type": "application/json", "api_token": api_token}
|
||||
headers = bria_json_headers(api_token)
|
||||
response = requests.post(self.api_url, json=payload, headers=headers)
|
||||
if response.status_code not in (200, 202):
|
||||
raise Exception(
|
||||
|
||||
@@ -1,9 +1,9 @@
|
||||
import requests
|
||||
|
||||
from .common import (
|
||||
deserialize_and_get_comfy_key,
|
||||
normalize_images_input,
|
||||
bria_json_headers,
|
||||
image_to_base64,
|
||||
normalize_images_input,
|
||||
poll_status_until_completed,
|
||||
)
|
||||
|
||||
@@ -46,8 +46,6 @@ class GenerateStructuredPromptNodeV2:
|
||||
|
||||
def execute(self, api_token, prompt, seed, structured_prompt, images=None):
|
||||
self._validate_token(api_token)
|
||||
api_token = deserialize_and_get_comfy_key(api_token)
|
||||
|
||||
images_list = normalize_images_input(images) if images is not None else [None]
|
||||
|
||||
# Seeds per image
|
||||
@@ -81,7 +79,7 @@ class GenerateStructuredPromptNodeV2:
|
||||
image
|
||||
)
|
||||
|
||||
headers = {"Content-Type": "application/json", "api_token": api_token}
|
||||
headers = bria_json_headers(api_token)
|
||||
response = requests.post(self.api_url, json=payload, headers=headers)
|
||||
if response.status_code not in (200, 202):
|
||||
raise Exception(
|
||||
|
||||
@@ -4,7 +4,13 @@ from PIL import Image
|
||||
import io
|
||||
import torch
|
||||
|
||||
from .common import deserialize_and_get_comfy_key, preprocess_image, preprocess_mask, image_to_base64, poll_status_until_completed
|
||||
from .common import (
|
||||
bria_json_headers,
|
||||
image_to_base64,
|
||||
poll_status_until_completed,
|
||||
preprocess_image,
|
||||
preprocess_mask,
|
||||
)
|
||||
|
||||
|
||||
class GenFillNode():
|
||||
@@ -39,8 +45,6 @@ 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):
|
||||
image = preprocess_image(image)
|
||||
@@ -64,10 +68,7 @@ class GenFillNode():
|
||||
"version": 2
|
||||
}
|
||||
|
||||
headers = {
|
||||
"Content-Type": "application/json",
|
||||
"api_token": f"{api_key}"
|
||||
}
|
||||
headers = bria_json_headers(api_key)
|
||||
|
||||
try:
|
||||
# Send initial request to get status URL
|
||||
|
||||
@@ -5,10 +5,10 @@ from PIL import Image
|
||||
import torch
|
||||
|
||||
from .common import (
|
||||
deserialize_and_get_comfy_key,
|
||||
bria_json_headers,
|
||||
image_to_base64,
|
||||
normalize_images_input,
|
||||
poll_status_until_completed
|
||||
poll_status_until_completed,
|
||||
)
|
||||
|
||||
class ImageEnhanceNode():
|
||||
@@ -51,8 +51,6 @@ class ImageEnhanceNode():
|
||||
# Validate 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)
|
||||
|
||||
# Normalize input to list of PIL images
|
||||
images = normalize_images_input(images)
|
||||
|
||||
@@ -73,10 +71,7 @@ class ImageEnhanceNode():
|
||||
"preserve_alpha": preserve_alpha
|
||||
}
|
||||
|
||||
headers = {
|
||||
"Content-Type": "application/json",
|
||||
"api_token": f"{api_key}"
|
||||
}
|
||||
headers = bria_json_headers(api_key)
|
||||
|
||||
# Send request
|
||||
response = requests.post(self.api_url, json=payload, headers=headers)
|
||||
|
||||
@@ -5,7 +5,7 @@ from PIL import Image
|
||||
import torch
|
||||
|
||||
from .common import (
|
||||
deserialize_and_get_comfy_key,
|
||||
bria_json_headers,
|
||||
image_to_base64,
|
||||
normalize_images_input,
|
||||
poll_status_until_completed,
|
||||
@@ -59,8 +59,6 @@ class ImageExpansionNode():
|
||||
):
|
||||
if api_key.strip() in ("", "BRIA_API_TOKEN"):
|
||||
raise Exception("Please insert a valid API key.")
|
||||
api_key = deserialize_and_get_comfy_key(api_key)
|
||||
|
||||
images = normalize_images_input(images)
|
||||
canvas_size = [int(x.strip()) for x in canvas_size.split(",")] if canvas_size else ()
|
||||
original_image_size = [int(x.strip()) for x in original_image_size.split(",")] if original_image_size else ()
|
||||
@@ -107,7 +105,7 @@ class ImageExpansionNode():
|
||||
"visual_output_content_moderation": visual_output_content_moderation
|
||||
}
|
||||
|
||||
headers = {"Content-Type": "application/json", "api_token": api_key}
|
||||
headers = bria_json_headers(api_key)
|
||||
response = requests.post(self.api_url, json=payload, headers=headers)
|
||||
if response.status_code not in (200, 202):
|
||||
raise Exception(f"API request failed with status {response.status_code}: {response.text}")
|
||||
|
||||
@@ -2,11 +2,11 @@ import requests
|
||||
import torch
|
||||
|
||||
from .common import (
|
||||
deserialize_and_get_comfy_key,
|
||||
postprocess_image,
|
||||
preprocess_image,
|
||||
bria_json_headers,
|
||||
image_to_base64,
|
||||
poll_status_until_completed,
|
||||
postprocess_image,
|
||||
preprocess_image,
|
||||
)
|
||||
|
||||
|
||||
@@ -81,8 +81,6 @@ class ProductIntegrateNode:
|
||||
seed,
|
||||
):
|
||||
self._validate_token(api_token)
|
||||
api_token = deserialize_and_get_comfy_key(api_token)
|
||||
|
||||
# Process single scene image
|
||||
if isinstance(scene, torch.Tensor):
|
||||
processed_scene = preprocess_image(scene)
|
||||
@@ -99,7 +97,7 @@ class ProductIntegrateNode:
|
||||
seed,
|
||||
)
|
||||
|
||||
headers = {"Content-Type": "application/json", "api_token": api_token}
|
||||
headers = bria_json_headers(api_token)
|
||||
|
||||
try:
|
||||
response = requests.post(self.api_url, json=payload, headers=headers)
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
import requests
|
||||
from .common import deserialize_and_get_comfy_key, poll_status_until_completed, postprocess_image
|
||||
|
||||
from .common import bria_json_headers, poll_status_until_completed, postprocess_image
|
||||
|
||||
|
||||
|
||||
@@ -96,8 +97,7 @@ class RefineImageLiteNodeV2:
|
||||
guidance_scale,
|
||||
seed,
|
||||
)
|
||||
api_token = deserialize_and_get_comfy_key(api_token)
|
||||
headers = {"Content-Type": "application/json", "api_token": api_token}
|
||||
headers = bria_json_headers(api_token)
|
||||
|
||||
try:
|
||||
response = requests.post(self.api_url, json=payload, headers=headers)
|
||||
@@ -130,7 +130,7 @@ class RefineImageLiteNodeV2:
|
||||
"seed": used_seed,
|
||||
}
|
||||
|
||||
headers = {"Content-Type": "application/json", "api_token": api_token}
|
||||
headers = bria_json_headers(api_token)
|
||||
|
||||
response = requests.post(self.generate_api_url, json=payloadForImageGenetrate, headers=headers)
|
||||
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
import requests
|
||||
from .common import deserialize_and_get_comfy_key, poll_status_until_completed, postprocess_image
|
||||
|
||||
from .common import bria_json_headers, poll_status_until_completed, postprocess_image
|
||||
|
||||
|
||||
class RefineImageNodeV2:
|
||||
@@ -97,8 +98,7 @@ class RefineImageNodeV2:
|
||||
guidance_scale,
|
||||
seed,
|
||||
)
|
||||
api_token = deserialize_and_get_comfy_key(api_token)
|
||||
headers = {"Content-Type": "application/json", "api_token": api_token}
|
||||
headers = bria_json_headers(api_token)
|
||||
|
||||
try:
|
||||
response = requests.post(self.api_url, json=payload, headers=headers)
|
||||
@@ -132,7 +132,7 @@ class RefineImageNodeV2:
|
||||
"negative_prompt":negative_prompt
|
||||
}
|
||||
|
||||
headers = {"Content-Type": "application/json", "api_token": api_token}
|
||||
headers = bria_json_headers(api_token)
|
||||
|
||||
response = requests.post(self.generate_api_url, json=payloadForImageGenetrate, headers=headers)
|
||||
|
||||
|
||||
@@ -1,6 +1,11 @@
|
||||
import requests
|
||||
|
||||
from .common import deserialize_and_get_comfy_key, postprocess_image, preprocess_image, image_to_base64
|
||||
from .common import (
|
||||
bria_json_headers,
|
||||
image_to_base64,
|
||||
postprocess_image,
|
||||
preprocess_image,
|
||||
)
|
||||
|
||||
|
||||
class ReimagineNode():
|
||||
@@ -38,7 +43,6 @@ class ReimagineNode():
|
||||
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,
|
||||
@@ -59,7 +63,7 @@ class ReimagineNode():
|
||||
response = requests.post(
|
||||
self.api_url,
|
||||
json=payload,
|
||||
headers={"api_token": api_key}
|
||||
headers=bria_json_headers(api_key),
|
||||
)
|
||||
if response.status_code == 200:
|
||||
response_dict = response.json()
|
||||
|
||||
@@ -5,10 +5,10 @@ from PIL import Image
|
||||
import torch
|
||||
|
||||
from .common import (
|
||||
deserialize_and_get_comfy_key,
|
||||
bria_json_headers,
|
||||
image_to_base64,
|
||||
normalize_images_input,
|
||||
poll_status_until_completed
|
||||
poll_status_until_completed,
|
||||
)
|
||||
|
||||
class RemoveForegroundNode():
|
||||
@@ -44,8 +44,6 @@ class RemoveForegroundNode():
|
||||
):
|
||||
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)
|
||||
|
||||
images = normalize_images_input(images)
|
||||
batch_results = []
|
||||
|
||||
@@ -60,10 +58,7 @@ class RemoveForegroundNode():
|
||||
"preserve_alpha": preserve_alpha
|
||||
}
|
||||
|
||||
headers = {
|
||||
"Content-Type": "application/json",
|
||||
"api_token": api_key
|
||||
}
|
||||
headers = bria_json_headers(api_key)
|
||||
|
||||
response = requests.post(self.api_url, json=payload, headers=headers)
|
||||
if response.status_code not in (200, 202):
|
||||
|
||||
@@ -5,7 +5,7 @@ from PIL import Image
|
||||
import torch
|
||||
|
||||
from .common import (
|
||||
deserialize_and_get_comfy_key,
|
||||
bria_json_headers,
|
||||
image_to_base64,
|
||||
normalize_images_input,
|
||||
poll_status_until_completed,
|
||||
@@ -60,8 +60,6 @@ class ReplaceBgNode():
|
||||
):
|
||||
if api_key.strip() in ("", "BRIA_API_TOKEN"):
|
||||
raise Exception("Please insert a valid API key.")
|
||||
api_key = deserialize_and_get_comfy_key(api_key)
|
||||
|
||||
images = normalize_images_input(images)
|
||||
|
||||
# Normalize reference images
|
||||
@@ -96,7 +94,7 @@ class ReplaceBgNode():
|
||||
"force_background_detection": force_background_detection
|
||||
}
|
||||
|
||||
headers = {"Content-Type": "application/json", "api_token": api_key}
|
||||
headers = bria_json_headers(api_key)
|
||||
response = requests.post(self.api_url, json=payload, headers=headers)
|
||||
if response.status_code not in (200, 202):
|
||||
raise Exception(f"API request failed with status {response.status_code}: {response.text}")
|
||||
|
||||
+8
-7
@@ -4,7 +4,13 @@ import numpy as np
|
||||
from PIL import Image
|
||||
import torch
|
||||
|
||||
from .common import deserialize_and_get_comfy_key, image_to_base64, normalize_images_input, poll_status_until_completed, to_pil_safe
|
||||
from .common import (
|
||||
bria_json_headers,
|
||||
image_to_base64,
|
||||
normalize_images_input,
|
||||
poll_status_until_completed,
|
||||
to_pil_safe,
|
||||
)
|
||||
|
||||
class RmbgNode():
|
||||
@classmethod
|
||||
@@ -32,8 +38,6 @@ class RmbgNode():
|
||||
def execute(self, images, 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)
|
||||
|
||||
# Normalize input to list of PIL images
|
||||
images = normalize_images_input(images)
|
||||
|
||||
@@ -50,10 +54,7 @@ class RmbgNode():
|
||||
"preserve_alpha": preserve_alpha
|
||||
}
|
||||
|
||||
headers = {
|
||||
"Content-Type": "application/json",
|
||||
"api_token": f"{api_key}"
|
||||
}
|
||||
headers = bria_json_headers(api_key)
|
||||
|
||||
response = requests.post(self.api_url, json=payload, headers=headers)
|
||||
if response.status_code not in (200, 202):
|
||||
|
||||
@@ -1,6 +1,11 @@
|
||||
import requests
|
||||
|
||||
from .common import deserialize_and_get_comfy_key, postprocess_image, preprocess_image, image_to_base64
|
||||
from .common import (
|
||||
bria_json_headers,
|
||||
image_to_base64,
|
||||
postprocess_image,
|
||||
preprocess_image,
|
||||
)
|
||||
|
||||
|
||||
class TailoredGenNode():
|
||||
@@ -45,7 +50,6 @@ 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,
|
||||
@@ -74,7 +78,7 @@ class TailoredGenNode():
|
||||
response = requests.post(
|
||||
self.api_url + model_id,
|
||||
json=payload,
|
||||
headers={"api_token": api_key}
|
||||
headers=bria_json_headers(api_key),
|
||||
)
|
||||
if response.status_code == 200:
|
||||
response_dict = response.json()
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
import requests
|
||||
from .common import deserialize_and_get_comfy_key
|
||||
from .common import bria_json_headers
|
||||
|
||||
class TailoredModelInfoNode():
|
||||
@classmethod
|
||||
@@ -21,10 +21,9 @@ 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}
|
||||
headers=bria_json_headers(api_key),
|
||||
)
|
||||
if response.status_code == 200:
|
||||
generation_prefix = response.json()["generation_prefix"]
|
||||
|
||||
@@ -4,7 +4,12 @@ import numpy as np
|
||||
from PIL import Image
|
||||
import torch
|
||||
|
||||
from .common import deserialize_and_get_comfy_key, image_to_base64, normalize_images_input, to_pil_safe
|
||||
from .common import (
|
||||
bria_json_headers,
|
||||
image_to_base64,
|
||||
normalize_images_input,
|
||||
to_pil_safe,
|
||||
)
|
||||
|
||||
class TailoredPortraitNode():
|
||||
@classmethod
|
||||
@@ -41,8 +46,6 @@ class TailoredPortraitNode():
|
||||
):
|
||||
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)
|
||||
|
||||
# Normalize images to list of PIL images
|
||||
images = normalize_images_input(images)
|
||||
|
||||
@@ -60,10 +63,7 @@ class TailoredPortraitNode():
|
||||
"seed": seed
|
||||
}
|
||||
|
||||
headers = {
|
||||
"Content-Type": "application/json",
|
||||
"api_token": api_key
|
||||
}
|
||||
headers = bria_json_headers(api_key)
|
||||
|
||||
response = requests.post(self.api_url, json=payload, headers=headers)
|
||||
if response.status_code != 200:
|
||||
|
||||
@@ -1,6 +1,11 @@
|
||||
import requests
|
||||
|
||||
from .common import deserialize_and_get_comfy_key, postprocess_image, preprocess_image, image_to_base64
|
||||
from .common import (
|
||||
bria_json_headers,
|
||||
image_to_base64,
|
||||
postprocess_image,
|
||||
preprocess_image,
|
||||
)
|
||||
|
||||
|
||||
class Text2ImageBaseNode():
|
||||
@@ -48,7 +53,6 @@ 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,
|
||||
@@ -84,7 +88,7 @@ class Text2ImageBaseNode():
|
||||
response = requests.post(
|
||||
self.api_url,
|
||||
json=payload,
|
||||
headers={"api_token": api_key}
|
||||
headers=bria_json_headers(api_key),
|
||||
)
|
||||
if response.status_code == 200:
|
||||
response_dict = response.json()
|
||||
|
||||
@@ -1,6 +1,11 @@
|
||||
import requests
|
||||
|
||||
from .common import deserialize_and_get_comfy_key, postprocess_image, preprocess_image, image_to_base64
|
||||
from .common import (
|
||||
bria_json_headers,
|
||||
image_to_base64,
|
||||
postprocess_image,
|
||||
preprocess_image,
|
||||
)
|
||||
|
||||
|
||||
class Text2ImageFastNode():
|
||||
@@ -45,7 +50,6 @@ 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,
|
||||
@@ -77,7 +81,7 @@ class Text2ImageFastNode():
|
||||
response = requests.post(
|
||||
self.api_url,
|
||||
json=payload,
|
||||
headers={"api_token": api_key}
|
||||
headers=bria_json_headers(api_key),
|
||||
)
|
||||
if response.status_code == 200:
|
||||
response_dict = response.json()
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
import requests
|
||||
|
||||
from .common import deserialize_and_get_comfy_key, postprocess_image
|
||||
from .common import bria_json_headers, postprocess_image
|
||||
|
||||
|
||||
class Text2ImageHDNode():
|
||||
@@ -35,7 +35,6 @@ class Text2ImageHDNode():
|
||||
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,
|
||||
@@ -53,7 +52,7 @@ class Text2ImageHDNode():
|
||||
response = requests.post(
|
||||
self.api_url,
|
||||
json=payload,
|
||||
headers={"api_token": api_key}
|
||||
headers=bria_json_headers(api_key),
|
||||
)
|
||||
if response.status_code == 200:
|
||||
response_dict = response.json()
|
||||
|
||||
@@ -1,6 +1,11 @@
|
||||
import requests
|
||||
import torch
|
||||
from ..common import deserialize_and_get_comfy_key, postprocess_image, preprocess_image, image_to_base64
|
||||
from ..common import (
|
||||
bria_json_headers,
|
||||
image_to_base64,
|
||||
postprocess_image,
|
||||
preprocess_image,
|
||||
)
|
||||
|
||||
shot_by_text_api_url = (
|
||||
"https://engine.prod.bria-api.com/v1/product/lifestyle_shot_by_text"
|
||||
@@ -130,8 +135,7 @@ def make_api_request(api_url, payload, api_key, Placement_type = None):
|
||||
|
||||
|
||||
try:
|
||||
api_key = deserialize_and_get_comfy_key(api_key)
|
||||
headers = {"Content-Type": "application/json", "api_token": f"{api_key}"}
|
||||
headers = bria_json_headers(api_key)
|
||||
response = requests.post(api_url, json=payload, headers=headers)
|
||||
|
||||
if response.status_code == 200:
|
||||
|
||||
@@ -2,7 +2,7 @@ import os
|
||||
import uuid
|
||||
import requests
|
||||
import folder_paths
|
||||
from ..common import deserialize_and_get_comfy_key, poll_status_until_completed
|
||||
from ..common import bria_json_headers, poll_status_until_completed
|
||||
from .video_utils import upload_video_to_s3
|
||||
|
||||
class RemoveVideoBackgroundNode():
|
||||
@@ -55,7 +55,6 @@ class RemoveVideoBackgroundNode():
|
||||
def execute(self, api_key, video_url, preserve_audio=True, output_container_and_codec="webm_vp9",):
|
||||
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)
|
||||
video_path = None
|
||||
|
||||
input_video_url = ""
|
||||
@@ -77,10 +76,7 @@ class RemoveVideoBackgroundNode():
|
||||
"output_container_and_codec": output_container_and_codec
|
||||
}
|
||||
|
||||
headers = {
|
||||
"Content-Type": "application/json",
|
||||
"api_token": f"{api_key}"
|
||||
}
|
||||
headers = bria_json_headers(api_key)
|
||||
|
||||
response = requests.post(self.api_url, json=payload, headers=headers)
|
||||
|
||||
|
||||
@@ -2,7 +2,7 @@ import os
|
||||
import uuid
|
||||
import requests
|
||||
import folder_paths
|
||||
from ..common import deserialize_and_get_comfy_key, poll_status_until_completed
|
||||
from ..common import bria_json_headers, poll_status_until_completed
|
||||
from .video_utils import upload_video_to_s3
|
||||
|
||||
class VideoEraseElementsNode():
|
||||
@@ -60,7 +60,6 @@ class VideoEraseElementsNode():
|
||||
def execute(self, api_key, video_url, mask_url="", output_container_and_codec="mp4_h264", preserve_audio=True):
|
||||
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)
|
||||
video_path = None
|
||||
|
||||
if video_url and video_url.strip() != "":
|
||||
@@ -86,10 +85,7 @@ class VideoEraseElementsNode():
|
||||
"preserve_audio": preserve_audio
|
||||
}
|
||||
|
||||
headers = {
|
||||
"Content-Type": "application/json",
|
||||
"api_token": f"{api_key}"
|
||||
}
|
||||
headers = bria_json_headers(api_key)
|
||||
|
||||
response = requests.post(self.api_url, json=payload, headers=headers)
|
||||
|
||||
|
||||
@@ -2,7 +2,7 @@ import os
|
||||
import uuid
|
||||
import requests
|
||||
import folder_paths
|
||||
from ..common import deserialize_and_get_comfy_key, poll_status_until_completed
|
||||
from ..common import bria_json_headers, poll_status_until_completed
|
||||
from .video_utils import upload_video_to_s3
|
||||
|
||||
class VideoIncreaseResolutionNode():
|
||||
@@ -57,7 +57,6 @@ class VideoIncreaseResolutionNode():
|
||||
def execute(self, api_key, video_url, desired_increase='2', output_container_and_codec="mp4_h264", preserve_audio=True):
|
||||
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)
|
||||
video_path = None
|
||||
|
||||
if video_url and video_url.strip() != "":
|
||||
@@ -83,10 +82,7 @@ class VideoIncreaseResolutionNode():
|
||||
"preserve_audio": preserve_audio
|
||||
}
|
||||
|
||||
headers = {
|
||||
"Content-Type": "application/json",
|
||||
"api_token": f"{api_key}"
|
||||
}
|
||||
headers = bria_json_headers(api_key)
|
||||
|
||||
response = requests.post(self.api_url, json=payload, headers=headers)
|
||||
|
||||
|
||||
@@ -2,7 +2,7 @@ import os
|
||||
import uuid
|
||||
import requests
|
||||
import folder_paths
|
||||
from ..common import deserialize_and_get_comfy_key, poll_status_until_completed
|
||||
from ..common import bria_json_headers, poll_status_until_completed
|
||||
from .video_utils import upload_video_to_s3
|
||||
import json
|
||||
|
||||
@@ -58,8 +58,6 @@ class VideoMaskByKeyPointsNode():
|
||||
def execute(self, key_points, api_key, video_url, output_container_and_codec="mp4_h264", preserve_audio=True):
|
||||
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)
|
||||
|
||||
try:
|
||||
key_points_array = json.loads(key_points)
|
||||
except json.JSONDecodeError as e:
|
||||
@@ -91,10 +89,7 @@ class VideoMaskByKeyPointsNode():
|
||||
"preserve_audio": preserve_audio
|
||||
}
|
||||
|
||||
headers = {
|
||||
"Content-Type": "application/json",
|
||||
"api_token": f"{api_key}"
|
||||
}
|
||||
headers = bria_json_headers(api_key)
|
||||
|
||||
response = requests.post(self.api_url, json=payload, headers=headers)
|
||||
|
||||
|
||||
@@ -2,7 +2,7 @@ import os
|
||||
import uuid
|
||||
import requests
|
||||
import folder_paths
|
||||
from ..common import deserialize_and_get_comfy_key, poll_status_until_completed
|
||||
from ..common import bria_json_headers, poll_status_until_completed
|
||||
from .video_utils import upload_video_to_s3
|
||||
|
||||
class VideoMaskByPromptNode():
|
||||
@@ -57,8 +57,6 @@ class VideoMaskByPromptNode():
|
||||
def execute(self, prompt, api_key, video_url, output_container_and_codec="mp4_h264", preserve_audio=True):
|
||||
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)
|
||||
|
||||
video_path = None
|
||||
|
||||
if video_url and video_url.strip() != "":
|
||||
@@ -85,10 +83,7 @@ class VideoMaskByPromptNode():
|
||||
"preserve_audio": preserve_audio
|
||||
}
|
||||
|
||||
headers = {
|
||||
"Content-Type": "application/json",
|
||||
"api_token": f"{api_key}"
|
||||
}
|
||||
headers = bria_json_headers(api_key)
|
||||
|
||||
response = requests.post(self.api_url, json=payload, headers=headers)
|
||||
|
||||
|
||||
@@ -2,7 +2,7 @@ import os
|
||||
import uuid
|
||||
import requests
|
||||
import folder_paths
|
||||
from ..common import deserialize_and_get_comfy_key, poll_status_until_completed
|
||||
from ..common import bria_json_headers, poll_status_until_completed
|
||||
from .video_utils import upload_video_to_s3
|
||||
|
||||
class VideoSolidColorBackgroundNode():
|
||||
@@ -69,7 +69,6 @@ class VideoSolidColorBackgroundNode():
|
||||
def execute(self, api_key, video_url, background_color="Transparent", output_container_and_codec="webm_vp9", preserve_audio=True):
|
||||
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)
|
||||
video_path = None
|
||||
|
||||
if video_url and video_url.strip() != "":
|
||||
@@ -96,10 +95,7 @@ class VideoSolidColorBackgroundNode():
|
||||
"preserve_audio": preserve_audio
|
||||
}
|
||||
|
||||
headers = {
|
||||
"Content-Type": "application/json",
|
||||
"api_token": f"{api_key}"
|
||||
}
|
||||
headers = bria_json_headers(api_key)
|
||||
|
||||
response = requests.post(self.api_url, json=payload, headers=headers)
|
||||
|
||||
|
||||
@@ -1,11 +1,14 @@
|
||||
import os
|
||||
import requests
|
||||
|
||||
from ..common import BRIA_COMFYUI_USER_AGENT
|
||||
|
||||
|
||||
def upload_video_to_s3(video_path, filename, api_token):
|
||||
api_url = "https://platform.prod.bria-api.com/upload-video/anonymous/presigned-url"
|
||||
headers = {
|
||||
"Content-Type": "application/json"
|
||||
"Content-Type": "application/json",
|
||||
"User-Agent": BRIA_COMFYUI_USER_AGENT,
|
||||
}
|
||||
extension = os.path.splitext(filename)[1].lower()
|
||||
content_type_map = {
|
||||
|
||||
+1
-1
@@ -1,7 +1,7 @@
|
||||
[project]
|
||||
name = "comfyui-bria-api"
|
||||
description = "Custom nodes for ComfyUI using BRIA's API."
|
||||
version = "2.1.15"
|
||||
version = "2.1.16"
|
||||
license = {file = "LICENSE"}
|
||||
|
||||
[project.urls]
|
||||
|
||||
Reference in New Issue
Block a user