diff --git a/__init__.py b/__init__.py index dcffdc3..62a369b 100644 --- a/__init__.py +++ b/__init__.py @@ -3,8 +3,8 @@ from aiohttp import web import server -from .utils import collections_path, browser_path, sources_path -from .routes import sources, collections, config, files +from .utils import collections_path, browser_path, sources_path, download_logs_path +from .routes import sources, collections, config, files, downloads browser_app = web.Application() browser_app.add_routes([ @@ -26,18 +26,15 @@ browser_app.add_routes([ web.get("/config", config.api_get_browser_config), web.put("/config", config.api_update_browser_config), + web.post("/downloads", downloads.api_create_new_download), + web.static("/web", path.join(browser_path, 'web/build')), ]) server.PromptServer.instance.app.add_subapp("/browser/", browser_app) -def init_path(): - if not path.exists(collections_path): - mkdir(collections_path) - - if not path.exists(sources_path): - mkdir(sources_path) - -init_path() +for dir in [collections_path, sources_path, download_logs_path]: + if not path.exists(dir): + mkdir(dir) WEB_DIRECTORY = "web" NODE_CLASS_MAPPINGS = {} diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 0000000..78620c4 --- /dev/null +++ b/requirements.txt @@ -0,0 +1 @@ +tqdm diff --git a/routes/downloads.py b/routes/downloads.py new file mode 100644 index 0000000..438a3e3 --- /dev/null +++ b/routes/downloads.py @@ -0,0 +1,107 @@ +from os import path +import requests +import time +import asyncio +import json +from aiohttp import web +from tqdm import tqdm + +import folder_paths + +from ..utils import download_logs_path + +def parse_options_header(content_disposition): + param, options = '', {} + split_header = content_disposition.split(';') + + # Extract the first parameter + if len(split_header) > 0: + param = split_header[0].strip() + + # Extract the options + for option in split_header[1:]: + option_split = option.split('=') + if len(option_split) == 2: + key = option_split[0].strip() + value = option_split[1].strip().strip('"') + options[key] = value + + return param, options + + +# credit: https://gist.github.com/phineas-pta/d73f9a035b05f8e923af8c01df057175 +async def download_by_requests(uuid:str, download_url:str, save_in:str, filename:str="", overwrite:bool=False, chunk_size:int=1): + log_file_path = path.join(download_logs_path, uuid + '.json') + base_info = { + 'uuid': uuid, + 'download_url': download_url, + 'save_in': save_in, + 'filename': filename, + 'overwrite': overwrite, + 'result': '', + 'total_size': 0, + 'downloaded_size': 0, + } + with open(log_file_path, 'w') as log_file: + json.dump(base_info, log_file) + + HEADERS = {"User-Agent": "Mozilla/5.0 (Macintosh; Intel Mac OS X 10_15_7) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/119.0.0.0 Safari/537.36"} + + with requests.get(download_url, headers=HEADERS, stream=True) as resp: + MISSING_FILENAME = f"unkwown_{uuid}" + # get file name + if filename == "": + if content_disposition := resp.headers.get("Content-Disposition"): + param, options = parse_options_header(content_disposition) + if param == "attachment": + filename = options.get("filename", MISSING_FILENAME) + else: + fileext = path.splitext(filename)[-1] + if fileext != "": + filename = path.basename(download_url) + if filename == "": + filename = MISSING_FILENAME + + base_info['filename'] = filename + with open(log_file_path, 'w') as log_file: + json.dump(base_info, log_file) + + target_path = path.join(folder_paths.models_dir, save_in, filename) + if not overwrite and path.exists(target_path): + base_info['result'] = f'{target_path} already exists' + return + + # download file + TOTAL_SIZE = int(resp.headers.get("Content-Length", 0)) + CHUNK_SIZE = chunk_size * 10**6 + with ( + open(target_path, mode="wb") as file, + open(log_file_path, 'w') as log_file, + tqdm(total=TOTAL_SIZE, desc=f"download {filename}", unit="B", unit_scale=True) as bar, + ): + base_info['total_size'] = TOTAL_SIZE + for data in resp.iter_content(chunk_size=CHUNK_SIZE): + size = file.write(data) + bar.update(size) + base_info['downloaded_size'] += size + log_file.seek(0) + json.dump(base_info, log_file) + log_file.write('\n' * 2) + +# download_url, filename, save_in, overwrite +async def api_create_new_download(request): + json_data = await request.json() + download_url = json_data.get('download_url', None) + save_in = json_data.get('save_in', None) + filename = json_data.get('filename', '') + overwrite = json_data.get('overwrite', False) + + if not (download_url and save_in): + return web.Response(status=400) + + if '..' in save_in: + return web.Response(status=400) + + asyncio.create_task(download_by_requests(str(int(time.time())), download_url, save_in, filename, overwrite)) + + return web.json_response(status=201) diff --git a/utils.py b/utils.py index 284f854..faed074 100644 --- a/utils.py +++ b/utils.py @@ -10,6 +10,7 @@ browser_path = path.dirname(__file__) collections_path = path.join(browser_path, 'collections') config_path = path.join(browser_path, 'config.json') sources_path = path.join(browser_path, 'sources') +download_logs_path = path.join(browser_path, 'download_logs') image_extensions = ['.jpg', '.jpeg', '.png', '.gif', '.webp'] video_extensions = ['.mp4', '.mov', '.avi', '.webm', '.mkv']