feat(model_downloader): add model downloaders tool
This commit is contained in:
@@ -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 🌐"}
|
||||
@@ -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()")
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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 🌐"
|
||||
Reference in New Issue
Block a user