mirror of
https://github.com/Comfy-Org/ComfyUI.git
synced 2026-08-12 12:53:50 +08:00
Add Cloud audio result download helper
Amp-Thread-ID: https://ampcode.com/threads/T-019fd01c-5e2c-77a8-b7c0-dd523218716f Co-authored-by: Amp <amp@ampcode.com>
This commit is contained in:
@@ -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",
|
||||
|
||||
@@ -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,
|
||||
*,
|
||||
|
||||
@@ -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}
|
||||
|
||||
Reference in New Issue
Block a user