converted API nodes to async

This commit is contained in:
bigcat88 2025-07-31 18:22:53 +03:00
parent 97b8a2c26a
commit 99363dc394
No known key found for this signature in database
GPG Key ID: 1F0BF0EC3CF22721
4 changed files with 435 additions and 675 deletions

View File

@ -1,4 +1,5 @@
from __future__ import annotations from __future__ import annotations
import aiohttp
import io import io
import logging import logging
import mimetypes import mimetypes
@ -30,7 +31,7 @@ from io import BytesIO
import av import av
def download_url_to_video_output(video_url: str, timeout: int = None) -> VideoFromFile: async def download_url_to_video_output(video_url: str, timeout: int = None) -> VideoFromFile:
"""Downloads a video from a URL and returns a `VIDEO` output. """Downloads a video from a URL and returns a `VIDEO` output.
Args: Args:
@ -39,7 +40,7 @@ def download_url_to_video_output(video_url: str, timeout: int = None) -> VideoFr
Returns: Returns:
A Comfy node `VIDEO` output. A Comfy node `VIDEO` output.
""" """
video_io = download_url_to_bytesio(video_url, timeout) video_io = await download_url_to_bytesio(video_url, timeout)
if video_io is None: if video_io is None:
error_msg = f"Failed to download video from {video_url}" error_msg = f"Failed to download video from {video_url}"
logging.error(error_msg) logging.error(error_msg)
@ -62,7 +63,7 @@ def downscale_image_tensor(image, total_pixels=1536 * 1024) -> torch.Tensor:
return s return s
def validate_and_cast_response( async def validate_and_cast_response(
response, timeout: int = None, node_id: Union[str, None] = None response, timeout: int = None, node_id: Union[str, None] = None
) -> torch.Tensor: ) -> torch.Tensor:
"""Validates and casts a response to a torch.Tensor. """Validates and casts a response to a torch.Tensor.
@ -86,35 +87,24 @@ def validate_and_cast_response(
image_tensors: list[torch.Tensor] = [] image_tensors: list[torch.Tensor] = []
# Process each image in the data array # Process each image in the data array
for image_data in data: async with aiohttp.ClientSession(timeout=aiohttp.ClientTimeout(total=timeout)) as session:
image_url = image_data.url for img_data in data:
b64_data = image_data.b64_json img_bytes: bytes
if img_data.b64_json:
img_bytes = base64.b64decode(img_data.b64_json)
elif img_data.url:
if node_id:
PromptServer.instance.send_progress_text(f"Result URL: {img_data.url}", node_id)
async with session.get(img_data.url) as resp:
if resp.status != 200:
raise ValueError("Failed to download generated image")
img_bytes = await resp.read()
else:
raise ValueError("Invalid image payload neither URL nor base64 data present.")
if not image_url and not b64_data: pil_img = Image.open(BytesIO(img_bytes)).convert("RGBA")
raise ValueError("No image was generated in the response") arr = np.asarray(pil_img).astype(np.float32) / 255.0
image_tensors.append(torch.from_numpy(arr))
if b64_data:
img_data = base64.b64decode(b64_data)
img = Image.open(io.BytesIO(img_data))
elif image_url:
if node_id:
PromptServer.instance.send_progress_text(
f"Result URL: {image_url}", node_id
)
img_response = requests.get(image_url, timeout=timeout)
if img_response.status_code != 200:
raise ValueError("Failed to download the image")
img = Image.open(io.BytesIO(img_response.content))
img = img.convert("RGBA")
# Convert to numpy array, normalize to float32 between 0 and 1
img_array = np.array(img).astype(np.float32) / 255.0
img_tensor = torch.from_numpy(img_array)
# Add to list of tensors
image_tensors.append(img_tensor)
return torch.stack(image_tensors, dim=0) return torch.stack(image_tensors, dim=0)
@ -175,7 +165,7 @@ def mimetype_to_extension(mime_type: str) -> str:
return mime_type.split("/")[-1].lower() return mime_type.split("/")[-1].lower()
def download_url_to_bytesio(url: str, timeout: int = None) -> BytesIO: async def download_url_to_bytesio(url: str, timeout: int = None) -> BytesIO:
"""Downloads content from a URL using requests and returns it as BytesIO. """Downloads content from a URL using requests and returns it as BytesIO.
Args: Args:
@ -185,9 +175,11 @@ def download_url_to_bytesio(url: str, timeout: int = None) -> BytesIO:
Returns: Returns:
BytesIO object containing the downloaded content. BytesIO object containing the downloaded content.
""" """
response = requests.get(url, stream=True, timeout=timeout) timeout_cfg = aiohttp.ClientTimeout(total=timeout) if timeout else None
response.raise_for_status() # Raises HTTPError for bad responses (4XX or 5XX) async with aiohttp.ClientSession(timeout=timeout_cfg) as session:
return BytesIO(response.content) async with session.get(url) as resp:
resp.raise_for_status() # Raises HTTPError for bad responses (4XX or 5XX)
return BytesIO(await resp.read())
def bytesio_to_image_tensor(image_bytesio: BytesIO, mode: str = "RGBA") -> torch.Tensor: def bytesio_to_image_tensor(image_bytesio: BytesIO, mode: str = "RGBA") -> torch.Tensor:
@ -210,9 +202,9 @@ def bytesio_to_image_tensor(image_bytesio: BytesIO, mode: str = "RGBA") -> torch
return torch.from_numpy(image_array).unsqueeze(0) return torch.from_numpy(image_array).unsqueeze(0)
def download_url_to_image_tensor(url: str, timeout: int = None) -> torch.Tensor: async def download_url_to_image_tensor(url: str, timeout: int = None) -> torch.Tensor:
"""Downloads an image from a URL and returns a [B, H, W, C] tensor.""" """Downloads an image from a URL and returns a [B, H, W, C] tensor."""
image_bytesio = download_url_to_bytesio(url, timeout) image_bytesio = await download_url_to_bytesio(url, timeout)
return bytesio_to_image_tensor(image_bytesio) return bytesio_to_image_tensor(image_bytesio)
@ -336,10 +328,10 @@ def text_filepath_to_data_uri(filepath: str) -> str:
return f"data:{mime_type};base64,{base64_string}" return f"data:{mime_type};base64,{base64_string}"
def upload_file_to_comfyapi( async def upload_file_to_comfyapi(
file_bytes_io: BytesIO, file_bytes_io: BytesIO,
filename: str, filename: str,
upload_mime_type: str, upload_mime_type: Optional[str],
auth_kwargs: Optional[dict[str, str]] = None, auth_kwargs: Optional[dict[str, str]] = None,
) -> str: ) -> str:
""" """
@ -354,7 +346,10 @@ def upload_file_to_comfyapi(
Returns: Returns:
The download URL for the uploaded file. The download URL for the uploaded file.
""" """
request_object = UploadRequest(file_name=filename, content_type=upload_mime_type) if upload_mime_type is None:
request_object = UploadRequest(file_name=filename)
else:
request_object = UploadRequest(file_name=filename, content_type=upload_mime_type)
operation = SynchronousOperation( operation = SynchronousOperation(
endpoint=ApiEndpoint( endpoint=ApiEndpoint(
path="/customers/storage", path="/customers/storage",
@ -366,12 +361,8 @@ def upload_file_to_comfyapi(
auth_kwargs=auth_kwargs, auth_kwargs=auth_kwargs,
) )
response: UploadResponse = operation.execute() response: UploadResponse = await operation.execute()
upload_response = ApiClient.upload_file( await ApiClient.upload_file(response.upload_url, file_bytes_io, content_type=upload_mime_type)
response.upload_url, file_bytes_io, content_type=upload_mime_type
)
upload_response.raise_for_status()
return response.download_url return response.download_url
@ -399,7 +390,7 @@ def video_to_base64_string(
return base64.b64encode(video_bytes_io.getvalue()).decode("utf-8") return base64.b64encode(video_bytes_io.getvalue()).decode("utf-8")
def upload_video_to_comfyapi( async def upload_video_to_comfyapi(
video: VideoInput, video: VideoInput,
auth_kwargs: Optional[dict[str, str]] = None, auth_kwargs: Optional[dict[str, str]] = None,
container: VideoContainer = VideoContainer.MP4, container: VideoContainer = VideoContainer.MP4,
@ -439,9 +430,7 @@ def upload_video_to_comfyapi(
video.save_to(video_bytes_io, format=container, codec=codec) video.save_to(video_bytes_io, format=container, codec=codec)
video_bytes_io.seek(0) video_bytes_io.seek(0)
return upload_file_to_comfyapi( return await upload_file_to_comfyapi(video_bytes_io, filename, upload_mime_type, auth_kwargs)
video_bytes_io, filename, upload_mime_type, auth_kwargs
)
def audio_tensor_to_contiguous_ndarray(waveform: torch.Tensor) -> np.ndarray: def audio_tensor_to_contiguous_ndarray(waveform: torch.Tensor) -> np.ndarray:
@ -501,7 +490,7 @@ def audio_ndarray_to_bytesio(
return audio_bytes_io return audio_bytes_io
def upload_audio_to_comfyapi( async def upload_audio_to_comfyapi(
audio: AudioInput, audio: AudioInput,
auth_kwargs: Optional[dict[str, str]] = None, auth_kwargs: Optional[dict[str, str]] = None,
container_format: str = "mp4", container_format: str = "mp4",
@ -527,7 +516,7 @@ def upload_audio_to_comfyapi(
audio_data_np, sample_rate, container_format, codec_name audio_data_np, sample_rate, container_format, codec_name
) )
return upload_file_to_comfyapi(audio_bytes_io, filename, mime_type, auth_kwargs) return await upload_file_to_comfyapi(audio_bytes_io, filename, mime_type, auth_kwargs)
def audio_to_base64_string( def audio_to_base64_string(
@ -544,7 +533,7 @@ def audio_to_base64_string(
return base64.b64encode(audio_bytes).decode("utf-8") return base64.b64encode(audio_bytes).decode("utf-8")
def upload_images_to_comfyapi( async def upload_images_to_comfyapi(
image: torch.Tensor, image: torch.Tensor,
max_images=8, max_images=8,
auth_kwargs: Optional[dict[str, str]] = None, auth_kwargs: Optional[dict[str, str]] = None,
@ -561,55 +550,15 @@ def upload_images_to_comfyapi(
mime_type: Optional MIME type for the image. mime_type: Optional MIME type for the image.
""" """
# if batch, try to upload each file if max_images is greater than 0 # if batch, try to upload each file if max_images is greater than 0
idx_image = 0
download_urls: list[str] = [] download_urls: list[str] = []
is_batch = len(image.shape) > 3 is_batch = len(image.shape) > 3
batch_length = 1 batch_len = image.shape[0] if is_batch else 1
if is_batch:
batch_length = image.shape[0]
while True:
curr_image = image
if len(image.shape) > 3:
curr_image = image[idx_image]
# get BytesIO version of image
img_binary = tensor_to_bytesio(curr_image, mime_type=mime_type)
# first, request upload/download urls from comfy API
if not mime_type:
request_object = UploadRequest(file_name=img_binary.name)
else:
request_object = UploadRequest(
file_name=img_binary.name, content_type=mime_type
)
operation = SynchronousOperation(
endpoint=ApiEndpoint(
path="/customers/storage",
method=HttpMethod.POST,
request_model=UploadRequest,
response_model=UploadResponse,
),
request=request_object,
auth_kwargs=auth_kwargs,
)
response = operation.execute()
upload_response = ApiClient.upload_file( for idx in range(min(batch_len, max_images)):
response.upload_url, img_binary, content_type=mime_type tensor = image[idx] if is_batch else image
) img_io = tensor_to_bytesio(tensor, mime_type=mime_type)
# verify success url = await upload_file_to_comfyapi(img_io, img_io.name, mime_type, auth_kwargs)
try: download_urls.append(url)
upload_response.raise_for_status()
except requests.exceptions.HTTPError as e:
raise ValueError(f"Could not upload one or more images: {e}") from e
# add download_url to list
download_urls.append(response.download_url)
idx_image += 1
# stop uploading additional files if done
if is_batch and max_images > 0:
if idx_image >= max_images:
break
if idx_image >= batch_length:
break
return download_urls return download_urls

File diff suppressed because it is too large Load Diff

View File

@ -99,14 +99,14 @@ def validate_input_image(image: torch.Tensor) -> bool:
return image.shape[2] < 8000 and image.shape[1] < 8000 return image.shape[2] < 8000 and image.shape[1] < 8000
def poll_until_finished( async def poll_until_finished(
auth_kwargs: dict[str, str], auth_kwargs: dict[str, str],
api_endpoint: ApiEndpoint[Any, TaskStatusResponse], api_endpoint: ApiEndpoint[Any, TaskStatusResponse],
estimated_duration: Optional[int] = None, estimated_duration: Optional[int] = None,
node_id: Optional[str] = None, node_id: Optional[str] = None,
) -> TaskStatusResponse: ) -> TaskStatusResponse:
"""Polls the Runway API endpoint until the task reaches a terminal state, then returns the response.""" """Polls the Runway API endpoint until the task reaches a terminal state, then returns the response."""
return PollingOperation( return await PollingOperation(
poll_endpoint=api_endpoint, poll_endpoint=api_endpoint,
completed_statuses=[ completed_statuses=[
TaskStatus.SUCCEEDED.value, TaskStatus.SUCCEEDED.value,
@ -115,7 +115,7 @@ def poll_until_finished(
TaskStatus.FAILED.value, TaskStatus.FAILED.value,
TaskStatus.CANCELLED.value, TaskStatus.CANCELLED.value,
], ],
status_extractor=lambda response: (response.status.value), status_extractor=lambda response: response.status.value,
auth_kwargs=auth_kwargs, auth_kwargs=auth_kwargs,
result_url_extractor=get_video_url_from_task_status, result_url_extractor=get_video_url_from_task_status,
estimated_duration=estimated_duration, estimated_duration=estimated_duration,
@ -167,11 +167,11 @@ class RunwayVideoGenNode(ComfyNodeABC):
) )
return True return True
def get_response( async def get_response(
self, task_id: str, auth_kwargs: dict[str, str], node_id: Optional[str] = None self, task_id: str, auth_kwargs: dict[str, str], node_id: Optional[str] = None
) -> RunwayImageToVideoResponse: ) -> RunwayImageToVideoResponse:
"""Poll the task status until it is finished then get the response.""" """Poll the task status until it is finished then get the response."""
return poll_until_finished( return await poll_until_finished(
auth_kwargs, auth_kwargs,
ApiEndpoint( ApiEndpoint(
path=f"{PATH_GET_TASK_STATUS}/{task_id}", path=f"{PATH_GET_TASK_STATUS}/{task_id}",
@ -183,7 +183,7 @@ class RunwayVideoGenNode(ComfyNodeABC):
node_id=node_id, node_id=node_id,
) )
def generate_video( async def generate_video(
self, self,
request: RunwayImageToVideoRequest, request: RunwayImageToVideoRequest,
auth_kwargs: dict[str, str], auth_kwargs: dict[str, str],
@ -200,15 +200,15 @@ class RunwayVideoGenNode(ComfyNodeABC):
auth_kwargs=auth_kwargs, auth_kwargs=auth_kwargs,
) )
initial_response = initial_operation.execute() initial_response = await initial_operation.execute()
self.validate_task_created(initial_response) self.validate_task_created(initial_response)
task_id = initial_response.id task_id = initial_response.id
final_response = self.get_response(task_id, auth_kwargs, node_id) final_response = await self.get_response(task_id, auth_kwargs, node_id)
self.validate_response(final_response) self.validate_response(final_response)
video_url = get_video_url_from_task_status(final_response) video_url = get_video_url_from_task_status(final_response)
return (download_url_to_video_output(video_url),) return (await download_url_to_video_output(video_url),)
class RunwayImageToVideoNodeGen3a(RunwayVideoGenNode): class RunwayImageToVideoNodeGen3a(RunwayVideoGenNode):
@ -250,7 +250,7 @@ class RunwayImageToVideoNodeGen3a(RunwayVideoGenNode):
}, },
} }
def api_call( async def api_call(
self, self,
prompt: str, prompt: str,
start_frame: torch.Tensor, start_frame: torch.Tensor,
@ -265,7 +265,7 @@ class RunwayImageToVideoNodeGen3a(RunwayVideoGenNode):
validate_input_image(start_frame) validate_input_image(start_frame)
# Upload image # Upload image
download_urls = upload_images_to_comfyapi( download_urls = await upload_images_to_comfyapi(
start_frame, start_frame,
max_images=1, max_images=1,
mime_type="image/png", mime_type="image/png",
@ -274,7 +274,7 @@ class RunwayImageToVideoNodeGen3a(RunwayVideoGenNode):
if len(download_urls) != 1: if len(download_urls) != 1:
raise RunwayApiError("Failed to upload one or more images to comfy api.") raise RunwayApiError("Failed to upload one or more images to comfy api.")
return self.generate_video( return await self.generate_video(
RunwayImageToVideoRequest( RunwayImageToVideoRequest(
promptText=prompt, promptText=prompt,
seed=seed, seed=seed,
@ -333,7 +333,7 @@ class RunwayImageToVideoNodeGen4(RunwayVideoGenNode):
}, },
} }
def api_call( async def api_call(
self, self,
prompt: str, prompt: str,
start_frame: torch.Tensor, start_frame: torch.Tensor,
@ -348,7 +348,7 @@ class RunwayImageToVideoNodeGen4(RunwayVideoGenNode):
validate_input_image(start_frame) validate_input_image(start_frame)
# Upload image # Upload image
download_urls = upload_images_to_comfyapi( download_urls = await upload_images_to_comfyapi(
start_frame, start_frame,
max_images=1, max_images=1,
mime_type="image/png", mime_type="image/png",
@ -357,7 +357,7 @@ class RunwayImageToVideoNodeGen4(RunwayVideoGenNode):
if len(download_urls) != 1: if len(download_urls) != 1:
raise RunwayApiError("Failed to upload one or more images to comfy api.") raise RunwayApiError("Failed to upload one or more images to comfy api.")
return self.generate_video( return await self.generate_video(
RunwayImageToVideoRequest( RunwayImageToVideoRequest(
promptText=prompt, promptText=prompt,
seed=seed, seed=seed,
@ -382,10 +382,10 @@ class RunwayFirstLastFrameNode(RunwayVideoGenNode):
DESCRIPTION = "Upload first and last keyframes, draft a prompt, and generate a video. More complex transitions, such as cases where the Last frame is completely different from the First frame, may benefit from the longer 10s duration. This would give the generation more time to smoothly transition between the two inputs. Before diving in, review these best practices to ensure that your input selections will set your generation up for success: https://help.runwayml.com/hc/en-us/articles/34170748696595-Creating-with-Keyframes-on-Gen-3." DESCRIPTION = "Upload first and last keyframes, draft a prompt, and generate a video. More complex transitions, such as cases where the Last frame is completely different from the First frame, may benefit from the longer 10s duration. This would give the generation more time to smoothly transition between the two inputs. Before diving in, review these best practices to ensure that your input selections will set your generation up for success: https://help.runwayml.com/hc/en-us/articles/34170748696595-Creating-with-Keyframes-on-Gen-3."
def get_response( async def get_response(
self, task_id: str, auth_kwargs: dict[str, str], node_id: Optional[str] = None self, task_id: str, auth_kwargs: dict[str, str], node_id: Optional[str] = None
) -> RunwayImageToVideoResponse: ) -> RunwayImageToVideoResponse:
return poll_until_finished( return await poll_until_finished(
auth_kwargs, auth_kwargs,
ApiEndpoint( ApiEndpoint(
path=f"{PATH_GET_TASK_STATUS}/{task_id}", path=f"{PATH_GET_TASK_STATUS}/{task_id}",
@ -437,7 +437,7 @@ class RunwayFirstLastFrameNode(RunwayVideoGenNode):
}, },
} }
def api_call( async def api_call(
self, self,
prompt: str, prompt: str,
start_frame: torch.Tensor, start_frame: torch.Tensor,
@ -455,7 +455,7 @@ class RunwayFirstLastFrameNode(RunwayVideoGenNode):
# Upload images # Upload images
stacked_input_images = image_tensor_pair_to_batch(start_frame, end_frame) stacked_input_images = image_tensor_pair_to_batch(start_frame, end_frame)
download_urls = upload_images_to_comfyapi( download_urls = await upload_images_to_comfyapi(
stacked_input_images, stacked_input_images,
max_images=2, max_images=2,
mime_type="image/png", mime_type="image/png",
@ -464,7 +464,7 @@ class RunwayFirstLastFrameNode(RunwayVideoGenNode):
if len(download_urls) != 2: if len(download_urls) != 2:
raise RunwayApiError("Failed to upload one or more images to comfy api.") raise RunwayApiError("Failed to upload one or more images to comfy api.")
return self.generate_video( return await self.generate_video(
RunwayImageToVideoRequest( RunwayImageToVideoRequest(
promptText=prompt, promptText=prompt,
seed=seed, seed=seed,
@ -543,11 +543,11 @@ class RunwayTextToImageNode(ComfyNodeABC):
) )
return True return True
def get_response( async def get_response(
self, task_id: str, auth_kwargs: dict[str, str], node_id: Optional[str] = None self, task_id: str, auth_kwargs: dict[str, str], node_id: Optional[str] = None
) -> TaskStatusResponse: ) -> TaskStatusResponse:
"""Poll the task status until it is finished then get the response.""" """Poll the task status until it is finished then get the response."""
return poll_until_finished( return await poll_until_finished(
auth_kwargs, auth_kwargs,
ApiEndpoint( ApiEndpoint(
path=f"{PATH_GET_TASK_STATUS}/{task_id}", path=f"{PATH_GET_TASK_STATUS}/{task_id}",
@ -559,7 +559,7 @@ class RunwayTextToImageNode(ComfyNodeABC):
node_id=node_id, node_id=node_id,
) )
def api_call( async def api_call(
self, self,
prompt: str, prompt: str,
ratio: str, ratio: str,
@ -574,7 +574,7 @@ class RunwayTextToImageNode(ComfyNodeABC):
reference_images = None reference_images = None
if reference_image is not None: if reference_image is not None:
validate_input_image(reference_image) validate_input_image(reference_image)
download_urls = upload_images_to_comfyapi( download_urls = await upload_images_to_comfyapi(
reference_image, reference_image,
max_images=1, max_images=1,
mime_type="image/png", mime_type="image/png",
@ -605,19 +605,19 @@ class RunwayTextToImageNode(ComfyNodeABC):
auth_kwargs=kwargs, auth_kwargs=kwargs,
) )
initial_response = initial_operation.execute() initial_response = await initial_operation.execute()
self.validate_task_created(initial_response) self.validate_task_created(initial_response)
task_id = initial_response.id task_id = initial_response.id
# Poll for completion # Poll for completion
final_response = self.get_response( final_response = await self.get_response(
task_id, auth_kwargs=kwargs, node_id=unique_id task_id, auth_kwargs=kwargs, node_id=unique_id
) )
self.validate_response(final_response) self.validate_response(final_response)
# Download and return image # Download and return image
image_url = get_image_url_from_task_status(final_response) image_url = get_image_url_from_task_status(final_response)
return (download_url_to_image_tensor(image_url),) return (await download_url_to_image_tensor(image_url),)
NODE_CLASS_MAPPINGS = { NODE_CLASS_MAPPINGS = {

View File

@ -124,7 +124,7 @@ class StabilityStableImageUltraNode:
}, },
} }
def api_call(self, prompt: str, aspect_ratio: str, style_preset: str, seed: int, async def api_call(self, prompt: str, aspect_ratio: str, style_preset: str, seed: int,
negative_prompt: str=None, image: torch.Tensor = None, image_denoise: float=None, negative_prompt: str=None, image: torch.Tensor = None, image_denoise: float=None,
**kwargs): **kwargs):
validate_string(prompt, strip_whitespace=False) validate_string(prompt, strip_whitespace=False)
@ -163,7 +163,7 @@ class StabilityStableImageUltraNode:
content_type="multipart/form-data", content_type="multipart/form-data",
auth_kwargs=kwargs, auth_kwargs=kwargs,
) )
response_api = operation.execute() response_api = await operation.execute()
if response_api.finish_reason != "SUCCESS": if response_api.finish_reason != "SUCCESS":
raise Exception(f"Stable Image Ultra generation failed: {response_api.finish_reason}.") raise Exception(f"Stable Image Ultra generation failed: {response_api.finish_reason}.")
@ -257,7 +257,7 @@ class StabilityStableImageSD_3_5Node:
}, },
} }
def api_call(self, model: str, prompt: str, aspect_ratio: str, style_preset: str, seed: int, cfg_scale: float, async def api_call(self, model: str, prompt: str, aspect_ratio: str, style_preset: str, seed: int, cfg_scale: float,
negative_prompt: str=None, image: torch.Tensor = None, image_denoise: float=None, negative_prompt: str=None, image: torch.Tensor = None, image_denoise: float=None,
**kwargs): **kwargs):
validate_string(prompt, strip_whitespace=False) validate_string(prompt, strip_whitespace=False)
@ -302,7 +302,7 @@ class StabilityStableImageSD_3_5Node:
content_type="multipart/form-data", content_type="multipart/form-data",
auth_kwargs=kwargs, auth_kwargs=kwargs,
) )
response_api = operation.execute() response_api = await operation.execute()
if response_api.finish_reason != "SUCCESS": if response_api.finish_reason != "SUCCESS":
raise Exception(f"Stable Diffusion 3.5 Image generation failed: {response_api.finish_reason}.") raise Exception(f"Stable Diffusion 3.5 Image generation failed: {response_api.finish_reason}.")
@ -374,7 +374,7 @@ class StabilityUpscaleConservativeNode:
}, },
} }
def api_call(self, image: torch.Tensor, prompt: str, creativity: float, seed: int, negative_prompt: str=None, async def api_call(self, image: torch.Tensor, prompt: str, creativity: float, seed: int, negative_prompt: str=None,
**kwargs): **kwargs):
validate_string(prompt, strip_whitespace=False) validate_string(prompt, strip_whitespace=False)
image_binary = tensor_to_bytesio(image, total_pixels=1024*1024).read() image_binary = tensor_to_bytesio(image, total_pixels=1024*1024).read()
@ -403,7 +403,7 @@ class StabilityUpscaleConservativeNode:
content_type="multipart/form-data", content_type="multipart/form-data",
auth_kwargs=kwargs, auth_kwargs=kwargs,
) )
response_api = operation.execute() response_api = await operation.execute()
if response_api.finish_reason != "SUCCESS": if response_api.finish_reason != "SUCCESS":
raise Exception(f"Stability Upscale Conservative generation failed: {response_api.finish_reason}.") raise Exception(f"Stability Upscale Conservative generation failed: {response_api.finish_reason}.")
@ -480,7 +480,7 @@ class StabilityUpscaleCreativeNode:
}, },
} }
def api_call(self, image: torch.Tensor, prompt: str, creativity: float, style_preset: str, seed: int, negative_prompt: str=None, async def api_call(self, image: torch.Tensor, prompt: str, creativity: float, style_preset: str, seed: int, negative_prompt: str=None,
**kwargs): **kwargs):
validate_string(prompt, strip_whitespace=False) validate_string(prompt, strip_whitespace=False)
image_binary = tensor_to_bytesio(image, total_pixels=1024*1024).read() image_binary = tensor_to_bytesio(image, total_pixels=1024*1024).read()
@ -512,7 +512,7 @@ class StabilityUpscaleCreativeNode:
content_type="multipart/form-data", content_type="multipart/form-data",
auth_kwargs=kwargs, auth_kwargs=kwargs,
) )
response_api = operation.execute() response_api = await operation.execute()
operation = PollingOperation( operation = PollingOperation(
poll_endpoint=ApiEndpoint( poll_endpoint=ApiEndpoint(
@ -527,7 +527,7 @@ class StabilityUpscaleCreativeNode:
status_extractor=lambda x: get_async_dummy_status(x), status_extractor=lambda x: get_async_dummy_status(x),
auth_kwargs=kwargs, auth_kwargs=kwargs,
) )
response_poll: StabilityResultsGetResponse = operation.execute() response_poll: StabilityResultsGetResponse = await operation.execute()
if response_poll.finish_reason != "SUCCESS": if response_poll.finish_reason != "SUCCESS":
raise Exception(f"Stability Upscale Creative generation failed: {response_poll.finish_reason}.") raise Exception(f"Stability Upscale Creative generation failed: {response_poll.finish_reason}.")
@ -563,8 +563,7 @@ class StabilityUpscaleFastNode:
}, },
} }
def api_call(self, image: torch.Tensor, async def api_call(self, image: torch.Tensor, **kwargs):
**kwargs):
image_binary = tensor_to_bytesio(image, total_pixels=4096*4096).read() image_binary = tensor_to_bytesio(image, total_pixels=4096*4096).read()
files = { files = {
@ -583,7 +582,7 @@ class StabilityUpscaleFastNode:
content_type="multipart/form-data", content_type="multipart/form-data",
auth_kwargs=kwargs, auth_kwargs=kwargs,
) )
response_api = operation.execute() response_api = await operation.execute()
if response_api.finish_reason != "SUCCESS": if response_api.finish_reason != "SUCCESS":
raise Exception(f"Stability Upscale Fast failed: {response_api.finish_reason}.") raise Exception(f"Stability Upscale Fast failed: {response_api.finish_reason}.")