diff --git a/comfy_api_nodes/util/__init__.py b/comfy_api_nodes/util/__init__.py index 1fb6b96cf..954bb789f 100644 --- a/comfy_api_nodes/util/__init__.py +++ b/comfy_api_nodes/util/__init__.py @@ -32,6 +32,7 @@ from .conversions import ( ) from .download_helpers import ( download_url_as_bytesio, + download_url_to_audio_input, download_url_to_bytesio, download_url_to_file_3d, download_url_to_image_tensor, @@ -76,6 +77,7 @@ __all__ = [ "upload_video_to_comfyapi", # Download helpers "download_url_as_bytesio", + "download_url_to_audio_input", "download_url_to_bytesio", "download_url_to_file_3d", "download_url_to_image_tensor", diff --git a/comfy_api_nodes/util/download_helpers.py b/comfy_api_nodes/util/download_helpers.py index 0ec3c6e66..3ed2c58c9 100644 --- a/comfy_api_nodes/util/download_helpers.py +++ b/comfy_api_nodes/util/download_helpers.py @@ -11,7 +11,7 @@ import torch from aiohttp.client_exceptions import ClientError, ContentTypeError from comfy_api.latest import IO as COMFY_IO -from comfy_api.latest import InputImpl, Types +from comfy_api.latest import Input, InputImpl, Types from folder_paths import get_output_directory from . import request_logger @@ -24,7 +24,7 @@ from ._helpers import ( ) from .client import _diagnose_connectivity from .common_exceptions import ApiServerError, LocalNetworkError, ProcessingInterrupted -from .conversions import bytesio_to_image_tensor +from .conversions import audio_bytes_to_audio_input, bytesio_to_image_tensor _RETRY_STATUS = {408, 429, 500, 502, 503, 504} @@ -241,6 +241,19 @@ async def download_url_to_video_output( return InputImpl.VideoFromFile(result) +async def download_url_to_audio_input( + audio_url: str, + *, + timeout: float = None, + max_retries: int = 5, + cls: type[COMFY_IO.ComfyNode] = None, +) -> Input.Audio: + """Downloads audio from a URL and decodes it into a Comfy AUDIO input.""" + result = BytesIO() + await download_url_to_bytesio(audio_url, result, timeout=timeout, max_retries=max_retries, cls=cls) + return audio_bytes_to_audio_input(result.getvalue()) + + async def download_url_as_bytesio( url: str, *, diff --git a/tests-unit/comfy_api_nodes_test/comfy_cloud_test.py b/tests-unit/comfy_api_nodes_test/comfy_cloud_test.py index 1b6650b42..91c37643b 100644 --- a/tests-unit/comfy_api_nodes_test/comfy_cloud_test.py +++ b/tests-unit/comfy_api_nodes_test/comfy_cloud_test.py @@ -1,5 +1,6 @@ import asyncio -from unittest.mock import AsyncMock +from io import BytesIO +from unittest.mock import AsyncMock, Mock import pytest import torch @@ -16,6 +17,7 @@ from comfy_api_nodes.apis.comfy_cloud import ( ComfyCloudWorkflowInputs, ) from comfy_api_nodes import nodes_comfy_cloud +from comfy_api_nodes.util import download_helpers @pytest.mark.parametrize( @@ -104,3 +106,67 @@ def test_contract_omits_optional_status_fields(): "inputs": {"prompt": "A lighthouse"}, } assert status.model_dump(exclude_none=True) == {"task_id": "task-1", "status": "queued"} + + +def test_download_cloud_audio_url_to_audio_input(monkeypatch): + node = nodes_comfy_cloud.ComfyCloudTextToImageNode + downloaded = b"encoded audio" + expected = {"waveform": torch.ones(1, 2, 3), "sample_rate": 48000} + download_call = Mock() + + async def download(url, dest, **kwargs): + download_call(url=url, dest=dest, **kwargs) + dest.write(downloaded) + dest.seek(0) + + audio_decode = Mock(return_value=expected) + monkeypatch.setattr(download_helpers, "download_url_to_bytesio", download) + monkeypatch.setattr(download_helpers, "audio_bytes_to_audio_input", audio_decode) + + output = asyncio.run( + download_helpers.download_url_to_audio_input( + "/proxy/comfy-cloud/results/task-1/audio.flac", + timeout=30, + max_retries=2, + cls=node, + ) + ) + + assert output is expected + download_call.assert_called_once() + assert download_call.call_args.kwargs["url"] == "/proxy/comfy-cloud/results/task-1/audio.flac" + assert isinstance(download_call.call_args.kwargs["dest"], BytesIO) + assert download_call.call_args.kwargs["timeout"] == 30 + assert download_call.call_args.kwargs["max_retries"] == 2 + assert download_call.call_args.kwargs["cls"] is node + audio_decode.assert_called_once_with(downloaded) + + +@pytest.mark.parametrize(("file_format", "expected_format"), [(".GLB", "glb"), ("SPZ", "spz")]) +def test_download_cloud_3d_url_to_file_3d(monkeypatch, file_format, expected_format): + node = nodes_comfy_cloud.ComfyCloudTextToImageNode + downloaded = b"3d result" + calls = [] + + async def download(url, dest, **kwargs): + calls.append((url, dest, kwargs)) + dest.write(downloaded) + dest.seek(0) + + monkeypatch.setattr(download_helpers, "download_url_to_bytesio", download) + + output = asyncio.run( + download_helpers.download_url_to_file_3d( + f"/proxy/comfy-cloud/results/task-1/model.{expected_format}", + file_format, + timeout=45, + max_retries=3, + cls=node, + ) + ) + + assert output.format == expected_format + assert output.get_bytes() == downloaded + assert calls[0][0] == f"/proxy/comfy-cloud/results/task-1/model.{expected_format}" + assert isinstance(calls[0][1], BytesIO) + assert calls[0][2] == {"timeout": 45, "max_retries": 3, "cls": node}