285 lines
15 KiB
Python
285 lines
15 KiB
Python
from __future__ import annotations
|
|
|
|
from ..components.tree import TREE_API
|
|
from ..components.tree import PRIMERE_ROOT
|
|
import os
|
|
from ..components import utility
|
|
from ..components.API import api_helper
|
|
import folder_paths
|
|
|
|
import random
|
|
import argparse
|
|
import json
|
|
import copy
|
|
from pathlib import Path
|
|
from typing import Any
|
|
import requests
|
|
import sys
|
|
from PIL import Image
|
|
import numpy as np
|
|
|
|
from ..components.API import api_json_to_requestbody
|
|
from ..components.API import external_api_backend
|
|
from ..components.API import api_schema_registry
|
|
|
|
class PrimereApiProcessor:
|
|
CATEGORY = TREE_API
|
|
RETURN_TYPES = ("APICLIENT", "STRING", "TUPLE", "TUPLE", "TUPLE", "TUPLE", "IMAGE")
|
|
RETURN_NAMES = ("CLIENT", "PROVIDER", "SCHEMA", "RENDERED", "API_SCHEMAS", "API_RESULT", "RESULT_IMAGE")
|
|
FUNCTION = "process_uniapi"
|
|
|
|
API_RESULT = api_helper.get_api_config("apiconfig.json")
|
|
API_SCHEMAS_RAW = utility.json2tuple(os.path.join(PRIMERE_ROOT, 'front_end', 'api_schemas.json'))
|
|
API_SCHEMA_REGISTRY = api_schema_registry.normalize_registry(API_SCHEMAS_RAW)
|
|
|
|
@classmethod
|
|
def INPUT_TYPES(cls):
|
|
cls.required_inputs = {
|
|
"processor": ("BOOLEAN", {"default": True, "label_on": "ON", "label_off": "OFF"}),
|
|
"debug_mode": ("BOOLEAN", {"default": False, "label_on": "DEBUG ONLY", "label_off": "PRODUCTION"}),
|
|
"api_provider": (external_api_backend.provider_list(cls),),
|
|
"api_service": (external_api_backend.service_list(cls),),
|
|
"prompt": ("STRING", {"forceInput": True}),
|
|
"batch": ("INT", {"default": 1, "max": 10, "min": 1, "step": 1})
|
|
}
|
|
|
|
cls.optional_inputs = {
|
|
"negative_prompt": ("STRING", {"default": None, "forceInput": True}),
|
|
"reference_images": ("IMAGE", {"default": None, "forceInput": True}),
|
|
"first_image": ("IMAGE", {"default": None, "forceInput": True}),
|
|
"last_image": ("IMAGE", {"default": None, "forceInput": True}),
|
|
"width": ("INT", {"default": 1024, "max": 8192, "min": 64, "step": 64, "forceInput": True}),
|
|
"height": ("INT", {"default": 1024, "max": 8192, "min": 64, "step": 64, "forceInput": True}),
|
|
"aspect_ratio": ("STRING", {"forceInput": True, "default": "1:1"}),
|
|
"seed": ("INT", {"default": 1, "min": 0, "max": (2 ** 32) - 1, "forceInput": True})
|
|
}
|
|
|
|
hidden_inputs = {
|
|
"extra_pnginfo": "EXTRA_PNGINFO",
|
|
"prompt_extra": "PROMPT"
|
|
}
|
|
|
|
return {"required": cls.required_inputs, "optional": cls.optional_inputs, "hidden": hidden_inputs}
|
|
|
|
def process_uniapi(self, processor, api_provider, api_service, prompt, negative_prompt = None, batch = 1, reference_images = None, first_image = None, last_image = None, width = 1024, height = 1024, aspect_ratio = '1:1', seed = None, debug_mode = False, **kwargs):
|
|
API_SCHEMAS_PATH = os.path.join(PRIMERE_ROOT, 'front_end', 'api_schemas.json')
|
|
API_CONFIG_PATH = os.path.join(PRIMERE_ROOT, 'json', 'apiconfig.json')
|
|
API_SCHEMA_REGISTRY = api_schema_registry.load_and_validate_api_schema_registry(API_SCHEMAS_PATH, API_CONFIG_PATH)
|
|
|
|
img_binary_api = []
|
|
|
|
WORKFLOWDATA = kwargs['extra_pnginfo']['workflow']['nodes']
|
|
custom_values = utility.getInputsFromWorkflowByNode(WORKFLOWDATA, 'PrimereApiProcessor', kwargs['prompt_extra'])
|
|
|
|
custom_user_inputs = {k: v for k, v in custom_values.items() if k not in self.required_inputs}
|
|
custom_user_inputs = {k: v for k, v in custom_user_inputs.items() if k not in self.optional_inputs}
|
|
# return (None, api_provider, None, custom_user_inputs, None, None, None)
|
|
del kwargs['extra_pnginfo']
|
|
del kwargs['prompt_extra']
|
|
|
|
if not processor:
|
|
return (None, api_provider, None, None, None, None, reference_images)
|
|
|
|
config_json = self.API_RESULT
|
|
client, api_provider = api_helper.create_api_client(api_provider, config_json)
|
|
|
|
schema, selected_service = api_schema_registry.get_schema(API_SCHEMA_REGISTRY, api_provider, api_service)
|
|
if schema is None:
|
|
schema = {
|
|
"provider": api_provider,
|
|
"service": api_service,
|
|
"request": {
|
|
"method": "SDK",
|
|
"endpoint": "",
|
|
"sdk_call": {"args": [], "kwargs": {}},
|
|
},
|
|
}
|
|
selected_service = api_service
|
|
|
|
schema["provider"] = api_provider
|
|
schema["service"] = selected_service or api_service
|
|
schema_import_modules = schema.get("import_modules", []) if isinstance(schema, dict) else []
|
|
if not isinstance(schema_import_modules, list):
|
|
raise RuntimeError("Schema key 'import_modules' must be a list of import statements.")
|
|
schema["import_modules"] = schema_import_modules
|
|
|
|
if not debug_mode:
|
|
imported_context, imported_roots = external_api_backend.load_import_modules(schema_import_modules)
|
|
endpoint_value = (((schema.get("request") or {}).get("endpoint")) if isinstance(schema, dict) else "") or ""
|
|
endpoint_root = endpoint_value.split(".", 1)[0] if isinstance(endpoint_value, str) and "." in endpoint_value else "client"
|
|
loaded_client_for_upload = imported_context.get(endpoint_root, client)
|
|
|
|
if reference_images is not None:
|
|
if (type(reference_images).__name__ == "list" or type(reference_images).__name__ == "Tensor") and len(reference_images) > 0:
|
|
source_images = []
|
|
if type(reference_images).__name__ == "list":
|
|
source_images = reference_images
|
|
else:
|
|
source_images.append(reference_images)
|
|
|
|
for single_image in source_images:
|
|
r1 = random.randint(1000, 9999)
|
|
if single_image is not None and type(single_image).__name__ == "Tensor":
|
|
ref_image = (single_image[0].numpy() * 255).astype(np.uint8)
|
|
ref_file = Image.fromarray(ref_image)
|
|
TEMP_FILE_REF = os.path.join(folder_paths.temp_directory, f"{api_provider}_edit_{r1}.png")
|
|
ref_file.save(TEMP_FILE_REF, format="PNG")
|
|
if api_provider == "Gemini":
|
|
gemini_image_data = Image.open(TEMP_FILE_REF)
|
|
img_binary_api.append(gemini_image_data)
|
|
if api_provider == "BlackForest":
|
|
encoded_image = None
|
|
single_image = source_images[0]
|
|
width_original = single_image.shape[2]
|
|
height_original = single_image.shape[1]
|
|
image_np_bf = (single_image.numpy() * 255).astype(np.uint8)
|
|
img_bf = Image.fromarray(image_np_bf)
|
|
img_byte_arr_bf = io.BytesIO()
|
|
img_bf.save(img_byte_arr_bf, format="PNG")
|
|
img_byte_arr_bf.seek(0)
|
|
encoded_string = base64.b64encode(img_byte_arr_bf.read())
|
|
img_binary_api = encoded_string.decode('ascii')
|
|
break
|
|
|
|
elif hasattr(loaded_client_for_upload, "upload_file"):
|
|
uploaded_reference = loaded_client_for_upload.upload_file(TEMP_FILE_REF)
|
|
img_binary_api.append(uploaded_reference)
|
|
else:
|
|
img_binary_api.append(TEMP_FILE_REF)
|
|
else:
|
|
img_binary_api = 'Debug mode, reference images ignored.'
|
|
|
|
selected_parameters = {"prompt": prompt}
|
|
selected_parameters = {"width": width}
|
|
selected_parameters = {"height": height}
|
|
|
|
local_inputs = locals()
|
|
required_keys = set(self.required_inputs.keys()) if isinstance(getattr(self, "required_inputs", None), dict) else set()
|
|
optional_keys = set(self.optional_inputs.keys()) if isinstance(getattr(self, "optional_inputs", None), dict) else set()
|
|
reserved_keys = {"processor", "api_provider", "api_service"}
|
|
|
|
for key in (required_keys | optional_keys):
|
|
if key in reserved_keys:
|
|
continue
|
|
if key in local_inputs and local_inputs[key] not in (None, ""):
|
|
selected_parameters[key] = local_inputs[key]
|
|
|
|
if isinstance(custom_user_inputs, dict):
|
|
for key, value in custom_user_inputs.items():
|
|
if value not in (None, ""):
|
|
selected_parameters[key] = value
|
|
|
|
if aspect_ratio not in (None, ""):
|
|
selected_aspect_ratio = aspect_ratio
|
|
schema_aspect_ratios = external_api_backend.schema_possible_values(self, api_provider, (selected_service or api_service), "aspect_ratio",)
|
|
if len(schema_aspect_ratios) > 0:
|
|
selected_aspect_ratio = external_api_backend.closest_valid_ratio(aspect_ratio, schema_aspect_ratios)
|
|
selected_parameters["aspect_ratio"] = selected_aspect_ratio
|
|
|
|
if len(img_binary_api) > 0:
|
|
selected_parameters["reference_images"] = img_binary_api
|
|
|
|
possible_parameters = schema.get("possible_parameters", {}) if isinstance(schema, dict) else {}
|
|
if isinstance(possible_parameters, dict):
|
|
for key in possible_parameters.keys():
|
|
if key in kwargs and kwargs[key] not in (None, ""):
|
|
selected_parameters[key] = kwargs[key]
|
|
|
|
rendered, used_values = api_json_to_requestbody.render_from_schema(schema, selected_parameters)
|
|
# rendered_payload = rendered.__dict__
|
|
rendered_payload = copy.deepcopy(rendered.__dict__)
|
|
if len(img_binary_api) > 0:
|
|
rendered_payload = external_api_backend.redact_reference_images(rendered_payload)
|
|
sdk_call = rendered_payload.get("sdk_call")
|
|
if isinstance(sdk_call, dict):
|
|
sdk_kwargs = sdk_call.get("kwargs")
|
|
if isinstance(sdk_kwargs, dict):
|
|
contents = sdk_kwargs.get("contents")
|
|
if isinstance(contents, list):
|
|
sanitized_contents = []
|
|
for content_item in contents:
|
|
if isinstance(content_item, list) and len(content_item) > 0 and all(isinstance(item, (Image.Image, str)) for item in content_item):
|
|
sanitized_contents.append("[reference_images omitted]")
|
|
else:
|
|
sanitized_contents.append(content_item)
|
|
sdk_kwargs["contents"] = sanitized_contents
|
|
|
|
api_result = None
|
|
api_error = None
|
|
result_image = None
|
|
batch = max(1, int(batch))
|
|
sdk_context = {}
|
|
response_url = None
|
|
loaded_client = client
|
|
|
|
try:
|
|
if rendered.method.upper() == "SDK":
|
|
provider_config = config_json.get(api_provider, {}) if isinstance(config_json, dict) else {}
|
|
provider_api_key = provider_config.get("APIKEY") if isinstance(provider_config, dict) else None
|
|
context = {"client": client, "provider_api_key": provider_api_key}
|
|
allowed_roots = {"client"}
|
|
imported_context, imported_roots = external_api_backend.load_import_modules(schema_import_modules)
|
|
context.update(imported_context)
|
|
allowed_roots.update(imported_roots)
|
|
|
|
auto_context, auto_roots = external_api_backend.build_sdk_context(rendered, client)
|
|
for root_name in auto_roots:
|
|
if root_name not in context and root_name in auto_context:
|
|
context[root_name] = auto_context[root_name]
|
|
allowed_roots.add(root_name)
|
|
|
|
sdk_context = dict(context)
|
|
sdk_call_data = rendered.sdk_call if isinstance(rendered.sdk_call, dict) else {}
|
|
sdk_args = sdk_call_data.get("args", []) if isinstance(sdk_call_data, dict) else []
|
|
if isinstance(sdk_args, list) and len(sdk_args) > 0 and isinstance(sdk_args[0], str):
|
|
response_url = sdk_args[0]
|
|
|
|
endpoint_root = rendered.endpoint.split(".", 1)[0] if isinstance(rendered.endpoint, str) and "." in rendered.endpoint else "client"
|
|
loaded_client = context.get(endpoint_root, client)
|
|
|
|
if debug_mode:
|
|
return (loaded_client, api_provider, schema, rendered_payload, None, api_result, reference_images)
|
|
api_result = external_api_backend.execute_sdk_request(rendered, context, allowed_roots)
|
|
else:
|
|
import requests
|
|
|
|
response = requests.request(
|
|
method=rendered.method,
|
|
url=rendered.endpoint,
|
|
headers=rendered.headers,
|
|
params=rendered.query,
|
|
json=rendered.body,
|
|
timeout=60,
|
|
)
|
|
api_result = {
|
|
"status_code": response.status_code,
|
|
"ok": response.ok,
|
|
"text": response.text,
|
|
}
|
|
except Exception as e:
|
|
api_error = str(e)
|
|
|
|
selected_parameters_output = {k: v for k, v in selected_parameters.items() if k != "reference_images"}
|
|
|
|
api_schemas = (
|
|
{
|
|
"schema": schema,
|
|
"selected_parameters": selected_parameters_output,
|
|
"used_values": used_values,
|
|
"selected_service": selected_service,
|
|
# "rendered": rendered.__dict__,
|
|
"rendered": rendered_payload,
|
|
"api_result": api_result,
|
|
"api_error": api_error,
|
|
},
|
|
)
|
|
|
|
if api_error is not None:
|
|
raise RuntimeError(f"API call failed for {api_provider}/{selected_service}: {api_error}")
|
|
|
|
if api_error is None:
|
|
# result_image = external_api_backend.apply_response_handler(schema, api_result, provider=api_provider, service=(selected_service or api_service))
|
|
response_context = {"response_url": response_url, "call_url": response_url, "loaded_client": loaded_client, "client": client, "sdk_context": sdk_context}
|
|
result_image = external_api_backend.apply_response_handler(schema, api_result, provider=api_provider, service=(selected_service or api_service), response_context=response_context)
|
|
|
|
return (client, api_provider, schema, rendered_payload, api_schemas, api_result, result_image) |