From 9b888443ac7da8d4f9201b48286ee3e05be0bc7d Mon Sep 17 00:00:00 2001 From: Vito Sansevero Date: Sat, 18 Oct 2025 12:45:37 -0700 Subject: [PATCH] feat(model_downloader): add interrupt handling --- kikotools/tools/model_downloader/custom.py | 121 ++++++++++++--------- 1 file changed, 72 insertions(+), 49 deletions(-) diff --git a/kikotools/tools/model_downloader/custom.py b/kikotools/tools/model_downloader/custom.py index 0888e7e..8318d0c 100644 --- a/kikotools/tools/model_downloader/custom.py +++ b/kikotools/tools/model_downloader/custom.py @@ -9,6 +9,18 @@ from typing import Optional from .base import BaseDownloader +try: + import comfy.model_management + + COMFY_AVAILABLE = True + InterruptProcessingException = comfy.model_management.InterruptProcessingException +except ImportError: + COMFY_AVAILABLE = False + # Fallback exception type that will never be raised + InterruptProcessingException = type( + "InterruptProcessingException", (Exception,), {} + ) + CHUNK_SIZE = 1638400 USER_AGENT = "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36" @@ -120,62 +132,73 @@ class CustomDownloader(BaseDownloader): print("Size: Unknown") # Download with progress - with open(output_file, "wb") as f: - downloaded = 0 - start_time = time.time() + try: + 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() + while True: + chunk_start_time = time.time() + buffer = response.read(CHUNK_SIZE) + chunk_end_time = time.time() - if not buffer: - break + if not buffer: + break - downloaded += len(buffer) - f.write(buffer) - chunk_time = chunk_end_time - chunk_start_time + downloaded += len(buffer) + f.write(buffer) + chunk_time = chunk_end_time - chunk_start_time - # Calculate speed - speed = self.calculate_speed(len(buffer), chunk_time) + # Check for user cancellation + self.check_interrupt() - # 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") + # Calculate speed + speed = self.calculate_speed(len(buffer), chunk_time) - end_time = time.time() - time_taken = end_time - start_time - hours, remainder = divmod(time_taken, 3600) - minutes, seconds = divmod(remainder, 60) + # 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") - 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" + end_time = time.time() + time_taken = end_time - start_time + hours, remainder = divmod(time_taken, 3600) + minutes, seconds = divmod(remainder, 60) - sys.stdout.write("\n") - print(f"✓ Download completed in {time_str}") - print(f"✓ File saved as: {output_file}") + 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" - # 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 + sys.stdout.write("\n") + print(f"✓ Download completed in {time_str}") + print(f"✓ File saved as: {output_file}") - return 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 + except InterruptProcessingException: + # Clean up partial download on interrupt + if os.path.exists(output_file): + os.remove(output_file) + raise