import asyncio import contextlib import logging import time import uuid from io import BytesIO from pathlib import Path from typing import IO, Optional, Union from urllib.parse import urlparse import aiohttp import torch from aiohttp.client_exceptions import ClientError, ContentTypeError from comfy_api.input_impl import VideoFromFile from comfy_api.latest import IO as ComfyIO from comfy_api_nodes.apis import request_logger from ._helpers import default_base_url, get_auth_header, is_processing_interrupted from .client import _diagnose_connectivity from .common_exceptions import ApiServerError, LocalNetworkError, ProcessingInterrupted from .conversions import bytesio_to_image_tensor _RETRY_STATUS = {408, 429, 500, 502, 503, 504} async def download_url_to_bytesio( url: str, dest: Optional[Union[BytesIO, IO[bytes], str, Path]], *, timeout: Optional[float] = None, max_retries: int = 3, retry_delay: float = 1.0, retry_backoff: float = 2.0, cls: type[ComfyIO.ComfyNode] = None, ) -> None: """Stream-download a URL to `dest`. `dest` must be one of: - a BytesIO (rewound to 0 after write), - a file-like object opened in binary write mode (must implement .write()), - a filesystem path (str | pathlib.Path), which will be opened with 'wb'. If `url` starts with `/proxy/`, `cls` must be provided so the URL can be expanded to an absolute URL and authentication headers can be applied. Raises: ProcessingInterrupted, LocalNetworkError, ApiServerError, Exception (HTTP and other errors) """ if not isinstance(dest, (str, Path)) and not hasattr(dest, "write"): raise ValueError("dest must be a path (str|Path) or a binary-writable object providing .write().") attempt = 0 delay = retry_delay headers = {} if url.startswith("/proxy/"): if cls is None: raise ValueError("For relative 'cloud' paths, the `cls` parameter is required.") url = default_base_url().rstrip("/") + url headers = get_auth_header(cls) while True: attempt += 1 op_id = _generate_operation_id("GET", url, attempt) timeout_cfg = aiohttp.ClientTimeout(total=timeout) is_path_sink = isinstance(dest, (str, Path)) fhandle = None try: with contextlib.suppress(Exception): request_logger.log_request_response(operation_id=op_id, request_method="GET", request_url=url) async with aiohttp.ClientSession(timeout=timeout_cfg) as session: async with session.get(url, headers=headers) as resp: if resp.status >= 400: with contextlib.suppress(Exception): try: body = await resp.json() except (ContentTypeError, ValueError): text = await resp.text() body = text if len(text) <= 4096 else f"[text {len(text)} bytes]" request_logger.log_request_response( operation_id=op_id, request_method="GET", request_url=url, response_status_code=resp.status, response_headers=dict(resp.headers), response_content=body, error_message=f"HTTP {resp.status}", ) if resp.status in _RETRY_STATUS and attempt <= max_retries: await _sleep_with_cancel(delay) delay *= retry_backoff continue raise Exception(f"Failed to download (HTTP {resp.status}).") if is_path_sink: p = Path(str(dest)) with contextlib.suppress(Exception): p.parent.mkdir(parents=True, exist_ok=True) fhandle = open(p, "wb") sink = fhandle else: sink = dest # BytesIO or file-like written = 0 async for chunk in resp.content.iter_chunked(1024 * 1024): sink.write(chunk) written += len(chunk) if is_processing_interrupted(): raise ProcessingInterrupted("Task cancelled") if isinstance(dest, BytesIO): with contextlib.suppress(Exception): dest.seek(0) with contextlib.suppress(Exception): request_logger.log_request_response( operation_id=op_id, request_method="GET", request_url=url, response_status_code=resp.status, response_headers=dict(resp.headers), response_content=f"[streamed {written} bytes to dest]", ) return except ProcessingInterrupted: logging.debug("Download was interrupted by user") raise except (ClientError, asyncio.TimeoutError) as e: if attempt <= max_retries: with contextlib.suppress(Exception): request_logger.log_request_response( operation_id=op_id, request_method="GET", request_url=url, error_message=f"{type(e).__name__}: {str(e)} (will retry)", ) await _sleep_with_cancel(delay) delay *= retry_backoff continue diag = await _diagnose_connectivity() if diag.get("is_local_issue"): raise LocalNetworkError( "Unable to connect to the network. Please check your internet connection and try again." ) from e raise ApiServerError("The remote service appears unreachable at this time.") from e finally: with contextlib.suppress(Exception): if fhandle: fhandle.flush() fhandle.close() async def download_url_to_image_tensor( url: str, *, timeout: float = None, cls: type[ComfyIO.ComfyNode] = None, ) -> torch.Tensor: """Downloads an image from a URL and returns a [B, H, W, C] tensor.""" result = BytesIO() await download_url_to_bytesio(url, result, timeout=timeout, cls=cls) return bytesio_to_image_tensor(result) async def download_url_to_video_output( video_url: str, *, timeout: float = None, cls: type[ComfyIO.ComfyNode] = None, ) -> VideoFromFile: """Downloads a video from a URL and returns a `VIDEO` output.""" result = BytesIO() await download_url_to_bytesio(video_url, result, timeout=timeout, cls=cls) return VideoFromFile(result) def _generate_operation_id(method: str, url: str, attempt: int) -> str: try: parsed = urlparse(url) slug = (parsed.path.rsplit("/", 1)[-1] or parsed.netloc or "download").strip("/").replace("/", "_") except Exception: slug = "download" return f"{method}_{slug}_try{attempt}_{uuid.uuid4().hex[:8]}" async def _sleep_with_cancel(seconds: float) -> None: """Sleep in 1s slices while checking for interruption.""" end = time.monotonic() + seconds while True: if is_processing_interrupted(): raise ProcessingInterrupted("Task cancelled") now = time.monotonic() if now >= end: return await asyncio.sleep(min(1.0, end - now))