feat(model_downloader): add model downloaders tool

This commit is contained in:
Vito Sansevero
2025-10-05 05:39:21 -07:00
parent af1bc6845a
commit 8f459c502a
7 changed files with 1217 additions and 0 deletions
@@ -0,0 +1,13 @@
"""Model Downloader Tool for ComfyUI-KikoTools
Downloads models from CivitAI, HuggingFace, and custom URLs.
"""
from .node import ModelDownloaderNode
__all__ = ["ModelDownloaderNode"]
# Node registration
NODE_CLASS_MAPPINGS = {"KikoModelDownloader": ModelDownloaderNode}
NODE_DISPLAY_NAME_MAPPINGS = {"KikoModelDownloader": "Model Downloader 🌐"}
+180
View File
@@ -0,0 +1,180 @@
"""Base downloader class with common functionality"""
from abc import ABC, abstractmethod
from pathlib import Path
from typing import Optional, Callable
from urllib.parse import urlparse, unquote
import os
class BaseDownloader(ABC):
"""Abstract base class for all downloaders"""
def __init__(self, token: Optional[str] = None):
"""Initialize downloader with optional API token
Args:
token: Optional API token for authentication
"""
self.token = token
self._progress_callback: Optional[Callable[[int, int, str], None]] = None
def set_progress_callback(self, callback: Callable[[int, int, str], None]) -> None:
"""Set callback function for progress updates
Args:
callback: Function(downloaded_bytes, total_bytes, message)
"""
self._progress_callback = callback
def report_progress(self, downloaded: int, total: int, message: str = "") -> None:
"""Report download progress to callback
Args:
downloaded: Bytes downloaded so far
total: Total bytes to download
message: Optional status message
"""
if self._progress_callback:
self._progress_callback(downloaded, total, message)
def extract_filename(self, url: str, default: str = "downloaded_file") -> str:
"""Extract filename from URL
Args:
url: URL to extract filename from
default: Default filename if extraction fails
Returns:
Extracted or default filename
"""
try:
parsed = urlparse(url)
path = unquote(parsed.path)
filename = os.path.basename(path)
# Remove query parameters from filename
if "?" in filename:
filename = filename.split("?")[0]
# Validate filename
if filename and len(filename) > 0 and "." in filename:
return filename
except Exception:
pass
return default
def extract_filename_from_header(self, content_disposition: str) -> Optional[str]:
"""Extract filename from Content-Disposition header
Args:
content_disposition: Content-Disposition header value
Returns:
Extracted filename or None
"""
try:
if "filename=" in content_disposition:
filename = content_disposition.split("filename=")[1]
# Remove quotes and whitespace
filename = filename.strip().strip('"').strip("'")
return filename
except Exception:
pass
return None
def validate_output_path(self, output_path: str) -> bool:
"""Validate and create output path if needed
Args:
output_path: Directory path to validate
Returns:
True if valid
Raises:
ValueError: If path exists but is not a directory
"""
path = Path(output_path)
if path.exists():
if not path.is_dir():
raise ValueError(
f"Output path {output_path} exists but is not a directory"
)
return True
# Create directory if it doesn't exist
path.mkdir(parents=True, exist_ok=True)
return True
def should_download(self, file_path: str, force: bool = False) -> bool:
"""Check if file should be downloaded
Args:
file_path: Full path to file
force: Force download even if file exists
Returns:
True if should download, False if file exists and force=False
"""
if force:
return True
return not Path(file_path).exists()
def format_size(self, size_bytes: int) -> str:
"""Format file size in human-readable format
Args:
size_bytes: Size in bytes
Returns:
Formatted size string (e.g., "5.00 MB")
"""
for unit in ["B", "KB", "MB", "GB"]:
if size_bytes < 1024.0:
return f"{size_bytes:.2f} {unit}"
size_bytes /= 1024.0
return f"{size_bytes:.2f} TB"
def calculate_speed(self, bytes_downloaded: int, elapsed_seconds: float) -> float:
"""Calculate download speed in MB/s
Args:
bytes_downloaded: Number of bytes downloaded
elapsed_seconds: Time elapsed in seconds
Returns:
Download speed in MB/s
"""
if elapsed_seconds <= 0:
return 0.0
mb_downloaded = bytes_downloaded / (1024 * 1024)
return mb_downloaded / elapsed_seconds
@abstractmethod
def download(
self,
url: str,
output_path: str,
filename: Optional[str] = None,
force: bool = False,
) -> str:
"""Download file from URL
Args:
url: URL to download from
output_path: Directory to save file
filename: Optional filename override
force: Force re-download if file exists
Returns:
Path to downloaded file
Raises:
NotImplementedError: Must be implemented by subclass
"""
raise NotImplementedError("Subclasses must implement download()")
+311
View File
@@ -0,0 +1,311 @@
"""CivitAI downloader implementation"""
import os
import sys
import json
import time
import urllib.request
import urllib.parse
import urllib.error
from typing import Optional, Dict, Any
from urllib.parse import urlparse, parse_qs, unquote
from .base import BaseDownloader
CHUNK_SIZE = 1638400
USER_AGENT = "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36"
API_BASE = "https://civitai.com/api/v1"
MAX_RETRIES = 3
RETRY_DELAY = 5
class CivitAIDownloader(BaseDownloader):
"""Downloader for CivitAI models"""
def __init__(self, token: Optional[str] = None):
"""Initialize CivitAI downloader
Args:
token: Optional CivitAI API token
"""
super().__init__(token)
def _make_request(
self, url: str, headers: Optional[Dict[str, str]] = None
) -> urllib.request.Request:
"""Create HTTP request with authentication
Args:
url: URL to request
headers: Optional additional headers
Returns:
urllib Request object
"""
if headers is None:
headers = {}
headers["User-Agent"] = USER_AGENT
if self.token:
headers["Authorization"] = f"Bearer {self.token}"
return urllib.request.Request(url, headers=headers)
def _parse_civitai_url(self, url: str) -> Dict[str, Optional[int]]:
"""Extract model and version IDs from CivitAI URL
Args:
url: CivitAI URL to parse
Returns:
Dict with 'model_id' and 'version_id' keys
"""
parsed = urlparse(url)
result = {"model_id": None, "version_id": None}
# Handle different URL patterns
# 1. Direct API download URL: /api/download/models/123456
if "/api/download/models/" in url:
match = url.split("/api/download/models/")[-1].split("?")[0]
if match.isdigit():
result["version_id"] = int(match)
return result
# 2. Model page URL: /models/123456 or /models/123456/model-name
if "/models/" in url:
parts = parsed.path.split("/")
if "models" in parts:
idx = parts.index("models")
if idx + 1 < len(parts) and parts[idx + 1].isdigit():
result["model_id"] = int(parts[idx + 1])
# 3. Version specific URL with ?modelVersionId=789012
query_params = parse_qs(parsed.query)
if "modelVersionId" in query_params:
version_id = query_params["modelVersionId"][0]
if version_id.isdigit():
result["version_id"] = int(version_id)
return result
def get_model_details(self, model_id: int) -> Dict[str, Any]:
"""Get model details from API
Args:
model_id: CivitAI model ID
Returns:
Model details dictionary
Raises:
Exception: If API request fails
"""
url = f"{API_BASE}/models/{model_id}"
request = self._make_request(url)
try:
with urllib.request.urlopen(request) as response:
return json.loads(response.read().decode())
except urllib.error.HTTPError as e:
if e.code == 404:
raise Exception(f"Model {model_id} not found")
raise Exception(f"API request failed: {e}")
def download(
self,
url: str,
output_path: str,
filename: Optional[str] = None,
force: bool = False,
) -> str:
"""Download file from CivitAI
Args:
url: CivitAI URL to download
output_path: Directory to save file
filename: Optional filename override
force: Force re-download if file exists
Returns:
Path to downloaded file
Raises:
Exception: If download fails
"""
# Validate output path
self.validate_output_path(output_path)
# Convert web URL to API URL if needed
if "civitai.com" in url and "/api/download/models/" not in url:
ids = self._parse_civitai_url(url)
# If we have a version ID, use it directly
if ids["version_id"]:
url = f"https://civitai.com/api/download/models/{ids['version_id']}"
# If we only have a model ID, get the latest version
elif ids["model_id"]:
try:
model_details = self.get_model_details(ids["model_id"])
if model_details.get("modelVersions"):
version_id = model_details["modelVersions"][0]["id"]
url = f"https://civitai.com/api/download/models/{version_id}"
else:
raise Exception(
f"No versions found for model {ids['model_id']}"
)
except Exception as e:
raise Exception(f"Failed to get model details: {e}")
else:
raise Exception("Could not parse model or version ID from URL")
headers = {"User-Agent": USER_AGENT}
if self.token:
headers["Authorization"] = f"Bearer {self.token}"
# Disable automatic redirect handling
class NoRedirection(urllib.request.HTTPErrorProcessor):
def http_response(self, request, response):
return response
https_response = http_response
request = urllib.request.Request(url, headers=headers)
opener = urllib.request.build_opener(NoRedirection)
try:
response = opener.open(request)
except urllib.error.HTTPError as e:
if e.code == 401:
raise Exception(
"Authentication required. Please provide a valid API token."
)
elif e.code == 403:
raise Exception(
"Access forbidden. The model might be restricted or require special permissions."
)
elif e.code == 404:
raise Exception(
"Model not found. The URL might be incorrect or the model was removed."
)
elif e.code == 429:
raise Exception(
"Rate limited. Please wait a moment before trying again."
)
else:
raise Exception(f"HTTP error {e.code}: {e.reason}")
# Handle redirects
if response.status in [301, 302, 303, 307, 308]:
redirect_url = response.getheader("Location")
# Handle relative redirects
if redirect_url.startswith("/"):
base_url = urlparse(url)
redirect_url = f"{base_url.scheme}://{base_url.netloc}{redirect_url}"
# Extract filename from redirect URL if not provided
if not filename:
parsed_url = urlparse(redirect_url)
query_params = parse_qs(parsed_url.query)
content_disposition = query_params.get(
"response-content-disposition", [None]
)[0]
if content_disposition and "filename=" in content_disposition:
filename = unquote(
content_disposition.split("filename=")[1].strip('"')
)
else:
# Fallback: extract filename from URL path
path = parsed_url.path
if path and "/" in path:
filename = path.split("/")[-1]
else:
filename = "downloaded_file.safetensors"
response = urllib.request.urlopen(redirect_url)
elif response.status == 404:
raise Exception("File not found")
elif response.status != 200:
raise Exception(f"Download failed with status {response.status}")
# Use provided filename or extracted filename
if not filename:
filename = self.extract_filename(url, default="model.safetensors")
output_file = os.path.join(output_path, filename)
# Check if should download
if not self.should_download(output_file, force):
print(f"File already exists: {output_file}")
return output_file
total_size = response.getheader("Content-Length")
if total_size is not None:
total_size = int(total_size)
print(f"Downloading: {filename}")
print(f"Destination: {output_file}")
if total_size:
print(f"Size: {self.format_size(total_size)}")
# Download with progress
with open(output_file, "wb") as f:
downloaded = 0
start_time = time.time()
while True:
chunk_start_time = time.time()
buffer = response.read(CHUNK_SIZE)
chunk_end_time = time.time()
if not buffer:
break
downloaded += len(buffer)
f.write(buffer)
chunk_time = chunk_end_time - chunk_start_time
# Calculate speed
speed = self.calculate_speed(len(buffer), chunk_time)
# Report progress
if total_size is not None:
progress = downloaded / total_size
sys.stdout.write(
f'\r[{"=" * int(progress * 50):<50}] {progress * 100:.2f}% - {speed:.2f} MB/s'
)
sys.stdout.flush()
self.report_progress(downloaded, total_size, f"{speed:.2f} MB/s")
else:
sys.stdout.write(
f"\rDownloaded: {self.format_size(downloaded)} - {speed:.2f} MB/s"
)
sys.stdout.flush()
self.report_progress(downloaded, 0, f"{speed:.2f} MB/s")
end_time = time.time()
time_taken = end_time - start_time
hours, remainder = divmod(time_taken, 3600)
minutes, seconds = divmod(remainder, 60)
if hours > 0:
time_str = f"{int(hours)}h {int(minutes)}m {int(seconds)}s"
elif minutes > 0:
time_str = f"{int(minutes)}m {int(seconds)}s"
else:
time_str = f"{int(seconds)}s"
sys.stdout.write("\n")
print(f"✓ Download completed in {time_str}")
print(f"✓ File saved as: {output_file}")
# Verify file size
actual_size = os.path.getsize(output_file)
if total_size and actual_size != total_size:
raise Exception(
f"Download incomplete. Expected {total_size} bytes, got {actual_size} bytes"
)
return output_file
+181
View File
@@ -0,0 +1,181 @@
"""Custom URL downloader - best effort for direct download links"""
import os
import sys
import time
import urllib.request
import urllib.error
from typing import Optional
from .base import BaseDownloader
CHUNK_SIZE = 1638400
USER_AGENT = "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36"
class CustomDownloader(BaseDownloader):
"""Best-effort downloader for custom/direct URLs"""
def __init__(self, token: Optional[str] = None):
"""Initialize custom downloader
Args:
token: Optional authentication token (will be sent as Bearer token)
"""
super().__init__(token)
def download(
self,
url: str,
output_path: str,
filename: Optional[str] = None,
force: bool = False,
) -> str:
"""Download file from custom URL
Args:
url: Direct download URL
output_path: Directory to save file
filename: Optional filename override
force: Force re-download if file exists
Returns:
Path to downloaded file
Raises:
Exception: If download fails
"""
# Validate output path
self.validate_output_path(output_path)
# Determine filename
if not filename:
filename = self.extract_filename(
url, default="downloaded_model.safetensors"
)
output_file = os.path.join(output_path, filename)
# Check if should download
if not self.should_download(output_file, force):
print(f"File already exists: {output_file}")
return output_file
# Prepare headers
headers = {"User-Agent": USER_AGENT}
# Add authentication if token provided
if self.token:
headers["Authorization"] = f"Bearer {self.token}"
# Create request
request = urllib.request.Request(url, headers=headers)
try:
# First request to check if file exists and get metadata
response = urllib.request.urlopen(request)
# Try to extract filename from Content-Disposition header if not provided
if not filename:
content_disposition = response.getheader("Content-Disposition")
if content_disposition:
extracted_filename = self.extract_filename_from_header(
content_disposition
)
if extracted_filename:
filename = extracted_filename
output_file = os.path.join(output_path, filename)
except urllib.error.HTTPError as e:
if e.code == 401:
raise Exception(
"Authentication required. Please provide a valid token if needed."
)
elif e.code == 403:
raise Exception(
"Access forbidden. The URL might require authentication or special permissions."
)
elif e.code == 404:
raise Exception("File not found. Please check the URL.")
elif e.code == 429:
raise Exception(
"Rate limited. Please wait a moment before trying again."
)
else:
raise Exception(f"HTTP error {e.code}: {e.reason}")
except urllib.error.URLError as e:
raise Exception(f"Network error: {e.reason}")
# Get file size
total_size = response.getheader("Content-Length")
if total_size is not None:
total_size = int(total_size)
print(f"Downloading: {filename}")
print(f"Destination: {output_file}")
if total_size:
print(f"Size: {self.format_size(total_size)}")
else:
print("Size: Unknown")
# Download with progress
with open(output_file, "wb") as f:
downloaded = 0
start_time = time.time()
while True:
chunk_start_time = time.time()
buffer = response.read(CHUNK_SIZE)
chunk_end_time = time.time()
if not buffer:
break
downloaded += len(buffer)
f.write(buffer)
chunk_time = chunk_end_time - chunk_start_time
# Calculate speed
speed = self.calculate_speed(len(buffer), chunk_time)
# Report progress
if total_size is not None:
progress = downloaded / total_size
sys.stdout.write(
f'\r[{"=" * int(progress * 50):<50}] {progress * 100:.2f}% - {speed:.2f} MB/s'
)
sys.stdout.flush()
self.report_progress(downloaded, total_size, f"{speed:.2f} MB/s")
else:
sys.stdout.write(
f"\rDownloaded: {self.format_size(downloaded)} - {speed:.2f} MB/s"
)
sys.stdout.flush()
self.report_progress(downloaded, 0, f"{speed:.2f} MB/s")
end_time = time.time()
time_taken = end_time - start_time
hours, remainder = divmod(time_taken, 3600)
minutes, seconds = divmod(remainder, 60)
if hours > 0:
time_str = f"{int(hours)}h {int(minutes)}m {int(seconds)}s"
elif minutes > 0:
time_str = f"{int(minutes)}m {int(seconds)}s"
else:
time_str = f"{int(seconds)}s"
sys.stdout.write("\n")
print(f"✓ Download completed in {time_str}")
print(f"✓ File saved as: {output_file}")
# Verify file size if known
actual_size = os.path.getsize(output_file)
if total_size and actual_size != total_size:
print(
f"⚠ Warning: Downloaded size ({actual_size} bytes) doesn't match expected size ({total_size} bytes)"
)
# Don't raise error for custom URLs as size mismatch might be acceptable
return output_file
@@ -0,0 +1,129 @@
"""URL detection and downloader selection logic"""
from __future__ import annotations
from enum import Enum
from typing import TYPE_CHECKING, Optional
from urllib.parse import urlparse
if TYPE_CHECKING:
from .base import BaseDownloader
class DownloaderType(Enum):
"""Types of supported downloaders"""
CIVITAI = "civitai"
HUGGINGFACE = "huggingface"
CUSTOM = "custom"
class URLDetector:
"""Detects URL type and returns appropriate downloader"""
def detect(self, url: Optional[str]) -> DownloaderType:
"""Detect which downloader to use based on URL
Args:
url: URL to analyze
Returns:
DownloaderType enum value
Raises:
ValueError: If URL is invalid or empty
"""
if not url:
raise ValueError("URL cannot be empty")
url = url.strip()
if not url:
raise ValueError("URL cannot be empty")
try:
parsed = urlparse(url)
if not parsed.scheme or not parsed.netloc:
raise ValueError("Invalid URL format")
except Exception:
raise ValueError("Invalid URL")
# Check for CivitAI
if self._is_civitai_url(url, parsed):
return DownloaderType.CIVITAI
# Check for HuggingFace
if self._is_huggingface_url(url, parsed):
return DownloaderType.HUGGINGFACE
# Default to custom downloader
return DownloaderType.CUSTOM
def _is_civitai_url(self, url: str, parsed) -> bool:
"""Check if URL is from CivitAI
Args:
url: Full URL string
parsed: Parsed URL object
Returns:
True if CivitAI URL
"""
if "civitai.com" not in parsed.netloc:
return False
# Check for API download endpoint
if "/api/download/models/" in url:
return True
# Check for model page
if "/models/" in url:
return True
return False
def _is_huggingface_url(self, url: str, parsed) -> bool:
"""Check if URL is from HuggingFace
Args:
url: Full URL string
parsed: Parsed URL object
Returns:
True if HuggingFace URL
"""
# Check main domain and CDN
if "huggingface.co" in parsed.netloc:
return True
return False
def get_downloader(
self, url: str, api_token: Optional[str] = None
) -> "BaseDownloader":
"""Get appropriate downloader instance for URL
Args:
url: URL to download from
api_token: Optional API token for authentication
Returns:
Appropriate downloader instance
Raises:
ValueError: If URL is invalid
"""
downloader_type = self.detect(url)
if downloader_type == DownloaderType.CIVITAI:
from .civitai import CivitAIDownloader
return CivitAIDownloader(token=api_token)
elif downloader_type == DownloaderType.HUGGINGFACE:
from .huggingface import HuggingFaceDownloader
return HuggingFaceDownloader(token=api_token)
else: # CUSTOM
from .custom import CustomDownloader
return CustomDownloader(token=api_token)
@@ -0,0 +1,248 @@
"""HuggingFace downloader implementation"""
import os
import sys
import time
import urllib.request
import urllib.error
from typing import Optional
from urllib.parse import urlparse, quote
from .base import BaseDownloader
CHUNK_SIZE = 1638400
USER_AGENT = "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36"
class HuggingFaceDownloader(BaseDownloader):
"""Downloader for HuggingFace models"""
def __init__(self, token: Optional[str] = None):
"""Initialize HuggingFace downloader
Args:
token: Optional HuggingFace API token
"""
super().__init__(token)
def _parse_huggingface_url(self, url: str) -> dict:
"""Parse HuggingFace URL to extract repo and file information
Args:
url: HuggingFace URL
Returns:
Dict with 'repo_id', 'filename', 'revision' keys
"""
parsed = urlparse(url)
parts = parsed.path.strip("/").split("/")
result = {"repo_id": None, "filename": None, "revision": "main"}
# Handle blob URLs (web UI format) - convert to resolve format
# /{username}/{repo}/blob/{revision}/{file_path}
if len(parts) >= 5 and "blob" in parts:
blob_idx = parts.index("blob")
if blob_idx >= 2:
# Extract repo_id (username/repo)
result["repo_id"] = "/".join(parts[:blob_idx])
# Extract revision
if blob_idx + 1 < len(parts):
result["revision"] = parts[blob_idx + 1]
# Extract filename (everything after revision)
if blob_idx + 2 < len(parts):
result["filename"] = "/".join(parts[blob_idx + 2 :])
# Standard HF URL format: /{username}/{repo}/resolve/{revision}/{file_path}
elif len(parts) >= 5 and "resolve" in parts:
resolve_idx = parts.index("resolve")
if resolve_idx >= 2:
# Extract repo_id (username/repo)
result["repo_id"] = "/".join(parts[:resolve_idx])
# Extract revision
if resolve_idx + 1 < len(parts):
result["revision"] = parts[resolve_idx + 1]
# Extract filename (everything after revision)
if resolve_idx + 2 < len(parts):
result["filename"] = "/".join(parts[resolve_idx + 2 :])
# Alternative CDN format: Extract what we can
elif "cdn" in parsed.netloc:
# CDN URLs might have different structure
# Try to extract filename from path
if len(parts) > 0:
result["filename"] = parts[-1]
return result
def _construct_download_url(
self, repo_id: str, filename: str, revision: str = "main"
) -> str:
"""Construct HuggingFace download URL
Args:
repo_id: Repository ID (username/repo)
filename: File path within repo
revision: Branch/tag/commit (default: main)
Returns:
Download URL
"""
# URL encode the filename to handle special characters
encoded_filename = quote(filename, safe="/")
return f"https://huggingface.co/{repo_id}/resolve/{revision}/{encoded_filename}"
def download(
self,
url: str,
output_path: str,
filename: Optional[str] = None,
force: bool = False,
) -> str:
"""Download file from HuggingFace
Args:
url: HuggingFace URL to download
output_path: Directory to save file
filename: Optional filename override
force: Force re-download if file exists
Returns:
Path to downloaded file
Raises:
Exception: If download fails
"""
# Validate output path
self.validate_output_path(output_path)
# Parse URL to get file information
url_info = self._parse_huggingface_url(url)
# Convert blob URL to resolve URL if needed
if url_info["repo_id"] and url_info["filename"]:
download_url = self._construct_download_url(
url_info["repo_id"], url_info["filename"], url_info["revision"]
)
print(f"[HuggingFace] Converted URL to: {download_url}")
else:
# Use original URL if parsing failed
download_url = url
# Determine filename
if not filename:
if url_info["filename"]:
# Use just the basename from the URL
filename = os.path.basename(url_info["filename"])
else:
filename = self.extract_filename(url, default="model.safetensors")
output_file = os.path.join(output_path, filename)
# Check if should download
if not self.should_download(output_file, force):
print(f"File already exists: {output_file}")
return output_file
# Prepare headers
headers = {"User-Agent": USER_AGENT}
if self.token:
headers["Authorization"] = f"Bearer {self.token}"
# Create request with converted download URL
request = urllib.request.Request(download_url, headers=headers)
try:
response = urllib.request.urlopen(request)
except urllib.error.HTTPError as e:
if e.code == 401:
raise Exception(
"Authentication required. Please provide a valid HuggingFace token."
)
elif e.code == 403:
raise Exception(
"Access forbidden. The model might be gated or require special permissions."
)
elif e.code == 404:
raise Exception(
"File not found. The URL might be incorrect or the file was removed."
)
elif e.code == 429:
raise Exception(
"Rate limited. Please wait a moment before trying again."
)
else:
raise Exception(f"HTTP error {e.code}: {e.reason}")
except urllib.error.URLError as e:
raise Exception(f"Network error: {e.reason}")
# Get file size
total_size = response.getheader("Content-Length")
if total_size is not None:
total_size = int(total_size)
print(f"Downloading: {filename}")
print(f"Destination: {output_file}")
if total_size:
print(f"Size: {self.format_size(total_size)}")
# Download with progress
with open(output_file, "wb") as f:
downloaded = 0
start_time = time.time()
while True:
chunk_start_time = time.time()
buffer = response.read(CHUNK_SIZE)
chunk_end_time = time.time()
if not buffer:
break
downloaded += len(buffer)
f.write(buffer)
chunk_time = chunk_end_time - chunk_start_time
# Calculate speed
speed = self.calculate_speed(len(buffer), chunk_time)
# Report progress
if total_size is not None:
progress = downloaded / total_size
sys.stdout.write(
f'\r[{"=" * int(progress * 50):<50}] {progress * 100:.2f}% - {speed:.2f} MB/s'
)
sys.stdout.flush()
self.report_progress(downloaded, total_size, f"{speed:.2f} MB/s")
else:
sys.stdout.write(
f"\rDownloaded: {self.format_size(downloaded)} - {speed:.2f} MB/s"
)
sys.stdout.flush()
self.report_progress(downloaded, 0, f"{speed:.2f} MB/s")
end_time = time.time()
time_taken = end_time - start_time
hours, remainder = divmod(time_taken, 3600)
minutes, seconds = divmod(remainder, 60)
if hours > 0:
time_str = f"{int(hours)}h {int(minutes)}m {int(seconds)}s"
elif minutes > 0:
time_str = f"{int(minutes)}m {int(seconds)}s"
else:
time_str = f"{int(seconds)}s"
sys.stdout.write("\n")
print(f"✓ Download completed in {time_str}")
print(f"✓ File saved as: {output_file}")
# Verify file size
actual_size = os.path.getsize(output_file)
if total_size and actual_size != total_size:
raise Exception(
f"Download incomplete. Expected {total_size} bytes, got {actual_size} bytes"
)
return output_file
+155
View File
@@ -0,0 +1,155 @@
"""ComfyUI Model Downloader Node"""
from ...base import ComfyAssetsBaseNode
from .detector import URLDetector
class ModelDownloaderNode(ComfyAssetsBaseNode):
"""ComfyUI node for downloading models from CivitAI, HuggingFace, and custom URLs"""
@classmethod
def INPUT_TYPES(cls):
"""Define input types for the node"""
return {
"required": {
"url": (
"STRING",
{
"default": "",
"multiline": False,
"placeholder": "https://civitai.com/... or https://huggingface.co/...",
},
),
"save_path": (
"STRING",
{
"default": "models/checkpoints",
"multiline": False,
"placeholder": "Path to save downloaded models",
},
),
},
"optional": {
"filename": (
"STRING",
{
"default": "",
"multiline": False,
"placeholder": "Leave empty for auto-detection",
},
),
"api_token": (
"STRING",
{
"default": "",
"multiline": False,
"placeholder": "API token (CivitAI or HuggingFace)",
},
),
"force_download": (
"BOOLEAN",
{
"default": False,
"label_on": "Force Redownload",
"label_off": "Skip if Exists",
},
),
},
}
RETURN_TYPES = ()
FUNCTION = "download_model"
CATEGORY = "🫶 ComfyAssets/🛠️ Utils"
OUTPUT_NODE = True
def download_model(
self,
url: str,
save_path: str,
filename: str = "",
api_token: str = "",
force_download: bool = False,
):
"""Download model from URL
Args:
url: URL to download from
save_path: Directory to save file
filename: Optional filename override
api_token: Optional API token
force_download: Force re-download if file exists
Returns:
Dictionary with 'ui' key for ComfyUI display
"""
# Validate inputs
if not url or not url.strip():
error_msg = "URL cannot be empty"
return {"ui": {"text": [error_msg]}}
if not save_path or not save_path.strip():
error_msg = "Save path cannot be empty"
return {"ui": {"text": [error_msg]}}
url = url.strip()
save_path = save_path.strip()
filename = filename.strip() if filename else None
api_token = api_token.strip() if api_token else None
try:
# Detect downloader type and get appropriate downloader
detector = URLDetector()
downloader_type = detector.detect(url)
print(
f"\n[Model Downloader] Detected downloader type: {downloader_type.value}"
)
print(f"[Model Downloader] URL: {url}")
print(f"[Model Downloader] Save path: {save_path}")
if filename:
print(f"[Model Downloader] Filename: {filename}")
if force_download:
print("[Model Downloader] Force download: enabled")
# Get downloader instance
downloader = detector.get_downloader(url, api_token=api_token)
# Download file
file_path = downloader.download(
url=url, output_path=save_path, filename=filename, force=force_download
)
message = f"Successfully downloaded to {file_path}"
print(f"[Model Downloader] {message}")
return {"ui": {"text": [message]}}
except ValueError as e:
error_msg = f"Invalid URL: {str(e)}"
print(f"[Model Downloader] Error: {error_msg}")
return {"ui": {"text": [error_msg]}}
except Exception as e:
error_msg = f"Download failed: {str(e)}"
print(f"[Model Downloader] Error: {error_msg}")
return {"ui": {"text": [error_msg]}}
@classmethod
def IS_CHANGED(
cls, url, save_path, filename="", api_token="", force_download=False
):
"""Force re-evaluation on every execution or when inputs change"""
# Include hash of inputs plus timestamp to force execution
# This ensures the node re-runs even if the download failed previously
import time
import hashlib
# Create a unique hash based on inputs and current time
input_str = (
f"{url}|{save_path}|{filename}|{api_token}|{force_download}|{time.time()}"
)
return hashlib.md5(input_str.encode()).hexdigest()
# Node display name
NODE_DISPLAY_NAME = "Model Downloader 🌐"