Fix runway image upload and progress polling (#39)

This commit is contained in:
Christian Byrne 2025-04-29 18:05:00 -07:00 committed by Robin Huang
parent 0fcf869c9e
commit 22c7bacffc
4 changed files with 23 additions and 13 deletions

View File

@ -2685,7 +2685,7 @@ class RunwayPromptImageDetailedObject(BaseModel):
..., ...,
description="The position of the image in the output video. 'last' is currently supported for gen3a_turbo only.", description="The position of the image in the output video. 'last' is currently supported for gen3a_turbo only.",
) )
uri: AnyUrl = Field( uri: str = Field(
..., description='A HTTPS URL or data URI containing an encoded image.' ..., description='A HTTPS URL or data URI containing an encoded image.'
) )

View File

@ -738,7 +738,7 @@ class ApiClient:
def upload_file( def upload_file(
upload_url: str, upload_url: str,
file: io.BytesIO | str, file: io.BytesIO | str,
mime_type: str | None = None, content_type: str | None = None,
): ):
"""Upload a file to the API. Make sure the file has a filename equal to what the url expects. """Upload a file to the API. Make sure the file has a filename equal to what the url expects.
@ -748,8 +748,8 @@ class ApiClient:
mime_type: Optional mime type to set for the upload mime_type: Optional mime type to set for the upload
""" """
headers = {} headers = {}
if mime_type: if content_type:
headers["Content-Type"] = mime_type headers["Content-Type"] = content_type
if isinstance(file, io.BytesIO): if isinstance(file, io.BytesIO):
file.seek(0) # Ensure we're at the start of the file file.seek(0) # Ensure we're at the start of the file

View File

@ -327,13 +327,13 @@ def upload_images_to_comfyapi(
request_model=UploadRequest, request_model=UploadRequest,
response_model=UploadResponse, response_model=UploadResponse,
), ),
request=UploadRequest(filename=img_binary.name, mime_type=mime_type), request=UploadRequest(filename=img_binary.name, content_type=mime_type),
auth_token=auth_token, auth_token=auth_token,
) )
response = operation.execute() response = operation.execute()
upload_response = ApiClient.upload_file( upload_response = ApiClient.upload_file(
response.upload_url, img_binary, mime_type=mime_type response.upload_url, img_binary, content_type=mime_type
) )
# verify success # verify success
try: try:

View File

@ -30,6 +30,7 @@ from comfy_api.input_impl import VideoFromFile
from comfy_api_nodes.mapper_utils import model_field_to_node_input from comfy_api_nodes.mapper_utils import model_field_to_node_input
PATH_IMAGE_TO_VIDEO = "/proxy/runway/image-to-video" PATH_IMAGE_TO_VIDEO = "/proxy/runway/image-to-video"
PATH_GET_TASK_STATUS = "/proxy/runway/tasks"
class RunwayApiError(Exception): class RunwayApiError(Exception):
@ -38,6 +39,12 @@ class RunwayApiError(Exception):
pass pass
def extract_progress_from_task_status(response: TaskStatusResponse) -> float:
if hasattr(response, "progress") and response.progress is not None:
return response.progress * 100
return None
class RunwayImageToVideoNode(ComfyNodeABC): class RunwayImageToVideoNode(ComfyNodeABC):
""" """
Runway Image to Video Node. Runway Image to Video Node.
@ -86,7 +93,7 @@ class RunwayImageToVideoNode(ComfyNodeABC):
""" """
polling_operation = PollingOperation( polling_operation = PollingOperation(
poll_endpoint=ApiEndpoint( poll_endpoint=ApiEndpoint(
path=f"{PATH_IMAGE_TO_VIDEO}/{task_id}", path=f"{PATH_GET_TASK_STATUS}/{task_id}",
method=HttpMethod.GET, method=HttpMethod.GET,
request_model=EmptyRequest, request_model=EmptyRequest,
response_model=TaskStatusResponse, response_model=TaskStatusResponse,
@ -98,7 +105,7 @@ class RunwayImageToVideoNode(ComfyNodeABC):
TaskStatus.FAILED.value, TaskStatus.FAILED.value,
TaskStatus.CANCELLED.value, TaskStatus.CANCELLED.value,
], ],
progress_extractor=lambda response: (response.progress * 100), progress_extractor=extract_progress_from_task_status,
status_extractor=lambda response: (response.status.value), status_extractor=lambda response: (response.status.value),
auth_token=auth_token, auth_token=auth_token,
) )
@ -193,7 +200,10 @@ class RunwayImageToVideoNode(ComfyNodeABC):
prompt_images_tensor = torch.cat(prompt_images_tensors, dim=0) prompt_images_tensor = torch.cat(prompt_images_tensors, dim=0)
download_urls = upload_images_to_comfyapi( download_urls = upload_images_to_comfyapi(
prompt_images_tensor, max_images=2, auth_token=auth_token, mime_type="image/png" prompt_images_tensor,
max_images=2,
auth_token=auth_token,
mime_type="image/png",
) )
# Create a list of detailed image objects # Create a list of detailed image objects
@ -202,15 +212,15 @@ class RunwayImageToVideoNode(ComfyNodeABC):
] ]
if len(download_urls) > 1: if len(download_urls) > 1:
prompt_image_details.append( prompt_image_details.append(
RunwayPromptImageDetailedObject(uri=str(download_urls[1]), position="last") RunwayPromptImageDetailedObject(
uri=str(download_urls[1]), position="last"
)
) )
# Wrap the list in the main object if details exist # Wrap the list in the main object if details exist
prompt_image_object: Optional[RunwayPromptImageObject] = None prompt_image_object: Optional[RunwayPromptImageObject] = None
if prompt_image_details: if prompt_image_details:
prompt_image_object = RunwayPromptImageObject( prompt_image_object = RunwayPromptImageObject(root=prompt_image_details)
root=prompt_image_details
)
initial_operation = SynchronousOperation( initial_operation = SynchronousOperation(
endpoint=ApiEndpoint( endpoint=ApiEndpoint(