Cupret commited on
Commit
d3e7294
·
verified ·
1 Parent(s): f5bf795

Upload civit_image_downloader2.py

Browse files
Files changed (1) hide show
  1. Misc/civit_image_downloader2.py +1023 -0
Misc/civit_image_downloader2.py ADDED
@@ -0,0 +1,1023 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import httpx
2
+ import os
3
+ import asyncio
4
+ import json
5
+ from tqdm import tqdm
6
+ import shutil
7
+ import re
8
+ from datetime import datetime
9
+ import logging
10
+ import csv
11
+ from threading import Lock
12
+ import argparse
13
+ import re
14
+
15
+ # Setup logging
16
+ script_dir = os.path.dirname(os.path.abspath(__file__))
17
+ log_file_path = os.path.join(script_dir, "civit_image_downloader_log_1.2.txt")
18
+ logger_cid = logging.getLogger('cid')
19
+ logger_cid.setLevel(logging.DEBUG)
20
+ file_handler_cid = logging.FileHandler(log_file_path, encoding='utf-8')
21
+ file_handler_cid.setLevel(logging.DEBUG)
22
+ formatter = logging.Formatter('%(asctime)s - %(levelname)s - %(message)s')
23
+ file_handler_cid.setFormatter(formatter)
24
+ logger_cid.addHandler(file_handler_cid)
25
+
26
+
27
+ ##########################################
28
+ # CivitAi API is fixed!#
29
+ # civit_image_downloader_1.2
30
+ ##########################################
31
+
32
+
33
+ # API endpoint for retrieving image URLs
34
+ base_url = "https://civitai.com/api/v1/images"
35
+
36
+ headers = {
37
+ "User-Agent": "Mozilla/5.0 (X11; Linux x86_64; rv:124.0) Gecko/20100101 Firefox/124.0",
38
+ "Content-Type": "application/json"
39
+ }
40
+
41
+ semaphore = asyncio.Semaphore(5)
42
+
43
+ # Directory for image downloads
44
+ output_dir = "image_downloads"
45
+ os.makedirs(output_dir, exist_ok=True)
46
+
47
+
48
+ def is_command_line_mode():
49
+ return any(vars(args).values())
50
+
51
+ def parse_arguments():
52
+ parser = argparse.ArgumentParser(description="CivitAI Image Downloader")
53
+ parser.add_argument("--timeout", type=int, help="Timeout value in seconds")
54
+ parser.add_argument("--quality", type=int, choices=[1, 2], help="Image quality (1 for SD, 2 for HD)")
55
+ parser.add_argument("--redownload", type=int, choices=[1, 2], help="Allow re-downloading of images (1 for Yes, 2 for No)")
56
+ parser.add_argument("--mode", type=int, choices=[1, 2, 3, 4, 5], help="Choose mode (1 for username, 2 for model ID, 3 for Model tag search, 4 for model version ID)")
57
+ parser.add_argument("--tags", help="Tags for Model tag search (comma-separated)")
58
+ parser.add_argument("--disable_prompt_check", choices=['y', 'n'], help="Disable prompt check (y/n)")
59
+ parser.add_argument("--username", help="Username for mode 1")
60
+ parser.add_argument("--model_id", help="Model ID for mode 2")
61
+ parser.add_argument("--model_version_id", help="Model Version ID for mode 4")
62
+ parser.add_argument("--datetime_min", help="Minimun date time image created at")
63
+ parser.add_argument("--datetime_max", help="Maximun date time image created at")
64
+ parser.add_argument("--reaction_min", type=int, help="Minimun date time image created at")
65
+ parser.add_argument("--condi", type=int, choices=[0, 1], help="Skip Filter (0 for OR, 1 for AND)")
66
+ parser.add_argument("--excludes", help="specified image id want to be excluded (multiple seperate by coma)")
67
+ return parser.parse_args()
68
+ args = parse_arguments()
69
+
70
+
71
+ def create_option_folder(option_name, base_dir):
72
+ option_dir = os.path.join(base_dir, option_name)
73
+ os.makedirs(option_dir, exist_ok=True)
74
+ return option_dir
75
+
76
+ allow_redownload = False
77
+
78
+
79
+ # Function to download an image from the provided URL
80
+ async def download_image(url, output_path, timeout_value, quality='SD'):
81
+ logger_cid.info(f"Attempting to download: {url}")
82
+ file_extension = ".png" if quality == 'HD' else ".jpeg"
83
+ output_path_with_extension = re.sub(r'\.jpeg|\.png', file_extension, output_path, flags=re.IGNORECASE)
84
+ if quality == 'HD':
85
+ url = re.sub(r"width=\d{3,4}", "original=true", url)
86
+
87
+ async with semaphore:
88
+ try:
89
+ async with httpx.AsyncClient() as client:
90
+ response = await client.get(url, timeout=timeout_value, headers=headers)
91
+ response.raise_for_status()
92
+ total_size = int(response.headers.get('content-length', 0))
93
+ progress_bar = tqdm(total=total_size, unit='B', unit_scale=True, desc=f"Downloading {output_path_with_extension}")
94
+
95
+ with open(output_path_with_extension, "wb") as file:
96
+ for chunk in response.iter_bytes():
97
+ progress_bar.update(len(chunk))
98
+ file.write(chunk)
99
+ progress_bar.close()
100
+ logger_cid.info(f"Successfully downloaded: {output_path_with_extension}")
101
+ return True, None
102
+ except Exception as e:
103
+ reason = str(e)
104
+ if isinstance(e, httpx.RequestError) or isinstance(e, httpx.ConnectError):
105
+ reason = "Network error while downloading the image. Please check your internet connection and try again"
106
+ elif isinstance(e, httpx.HTTPStatusError):
107
+ reason = f"Error downloading the image. Server response: {e.response.status_code} Please try again later."
108
+ elif isinstance(e, ConnectionResetError):
109
+ reason = "The connection to the server was closed unexpectedly. This could be a temporary network problem. Please try again later"
110
+ logger_cid.error(f"Error downloading {url}: {reason}")
111
+ return False, reason
112
+
113
+
114
+ # Async function to write meta data to a text file. If no meta data is available, the URL to the image is written to the txt file
115
+ async def write_meta_data(meta, output_path, image_id, username):
116
+ if not meta or all(value == '' for value in meta.values()):
117
+ output_path = output_path.replace(".txt", "_no_meta.txt")
118
+ url = f"https://civitai.com/images/{image_id}?period=AllTime&periodMode=published&sort=Newest&view=feed&username={username}&withTags=false"
119
+ with open(output_path, "w", encoding='utf-8') as file:
120
+ file.write(f"No metadata available for this image.\nURL: {url}\n")
121
+ else:
122
+ with open(output_path, "w", encoding='utf-8') as file:
123
+ for key, value in meta.items():
124
+ file.write(f"{key}: {value}\n")
125
+
126
+ async def write_prompt(meta, output_path):
127
+ if not meta or not meta['prompt'] :
128
+ print("No Metadata")
129
+ else:
130
+ #prompt = re.sub(r"(?<!\\)\(|\:\d.+?\)|\n", "", meta['prompt'])
131
+ prompt = re.sub(r"(?<!\\)\(|:\d.+?\)|(?<!\\)\)|\[|\]|<lora:.+?>|\n", "", meta['prompt'])
132
+ with open(output_path, "w", encoding='utf-8') as file:
133
+ file.write(f"{prompt}\n")
134
+
135
+
136
+ TRACKING_JSON_FILE = os.path.join(os.path.dirname(os.path.realpath(__file__)), "downloaded_images.json")
137
+
138
+ downloaded_images_lock = Lock()
139
+ tag_model_mapping_lock = Lock()
140
+
141
+
142
+ def load_downloaded_images():
143
+ try:
144
+ with open(TRACKING_JSON_FILE, "r") as file:
145
+ return json.load(file)
146
+ except (FileNotFoundError, json.JSONDecodeError):
147
+ return {}
148
+
149
+
150
+ def check_if_image_downloaded(image_id, image_path, quality='SD'):
151
+ with downloaded_images_lock:
152
+ image_key = f"{image_id}_{quality}"
153
+ if image_key in downloaded_images:
154
+ existing_path = downloaded_images[image_key].get("path")
155
+ if existing_path == image_path:
156
+ return True
157
+ return False
158
+
159
+
160
+ def mark_image_as_downloaded(image_id, image_path, quality='SD', tags=None, url=None):
161
+ with downloaded_images_lock:
162
+ image_key = f"{image_id}_{quality}"
163
+ current_date = datetime.now().strftime("%Y-%m-%d - %H:%M")
164
+
165
+ merged_tags = list(set(downloaded_images.get(image_key, {}).get("tags", [])) | set(tags or []))
166
+
167
+ downloaded_images[image_key] = {
168
+ "path": image_path,
169
+ "quality": quality,
170
+ "download_date": current_date,
171
+ "tags": merged_tags,
172
+ "url": url
173
+ }
174
+ json_data_string = json.dumps(downloaded_images[image_key], indent=4)
175
+ with open(TRACKING_JSON_FILE, "w") as file:
176
+ json.dump(downloaded_images, file, indent=4)
177
+
178
+ logger_cid.info(f"Marked as downloaded: {image_id} at {image_path}")
179
+
180
+
181
+ SOURCE_MISSING_MESSAGE_SHOWN = False
182
+ NEW_IMAGES_DOWNLOADED = False
183
+
184
+ # Customization of the manual_copy function
185
+ def manual_copy(src, dst):
186
+ global SOURCE_MISSING_MESSAGE_SHOWN, NEW_IMAGES_DOWNLOADED
187
+ # Check whether the source file exists
188
+ if os.path.exists(src):
189
+ try:
190
+ shutil.copy2(src, dst) # Copies the file and retains the metadata
191
+ NEW_IMAGES_DOWNLOADED = True
192
+ return True, dst # Returns True if the copy was successful
193
+ except Exception as e:
194
+ print(f"Error copying the file {src} to {dst}. Error: {e}")
195
+ return False # Returns False if an error has occurred
196
+ else:
197
+ if not SOURCE_MISSING_MESSAGE_SHOWN:
198
+ print(f"Source file {src} does not exist. Copy skipped.")
199
+ SOURCE_MISSING_MESSAGE_SHOWN = True
200
+ return False # Also returns False if the source file does not exist
201
+
202
+ def clean_and_shorten_path(path, max_total_length=260, max_component_length=80):
203
+ # Replace %20 with a space
204
+ path = path.replace("%20", " ")
205
+ path = path.replace("%2B", "+")
206
+ # Replace non-permitted characters
207
+ invalid_chars = '<>:"/\\|?*'
208
+ for char in invalid_chars:
209
+ path = path.replace(char, '_')
210
+
211
+ # Remove control characters
212
+ path = re.sub(r'[\x00-\x1f\x7f]', '', path)
213
+
214
+ # Remove trailing spaces or dots
215
+ path = path.rstrip('. ')
216
+
217
+ # Separate the path into directory and file name
218
+ dir_name, file_name = os.path.split(path)
219
+
220
+ # Shorten the file name if necessary
221
+ if len(file_name) > max_component_length:
222
+ file_name = file_name[:max_component_length - 3] + "_"
223
+
224
+ # Shorten the directory name if necessary
225
+ if len(dir_name) > max_total_length - len(file_name):
226
+ dir_name = dir_name[:max_total_length - len(file_name) - 3] + "_"
227
+
228
+ # Reassemble the path
229
+ shortened_path = os.path.join(dir_name, file_name)
230
+
231
+ # Ensure that the total path does not exceed the maximum length
232
+ if len(shortened_path) > max_total_length:
233
+ excess_length = len(shortened_path) - max_total_length
234
+ shortened_path = shortened_path[:-excess_length].rstrip('. ')
235
+
236
+ return shortened_path
237
+
238
+
239
+
240
+ def move_to_invalid_meta(src, model_dir):
241
+ invalid_meta_dir = os.path.join(model_dir, 'invalid_meta')
242
+ os.makedirs(invalid_meta_dir, exist_ok=True)
243
+ new_dst = os.path.join(invalid_meta_dir, os.path.basename(src))
244
+ try:
245
+ shutil.move(src, new_dst)
246
+ except Exception as e:
247
+ print(f"Error moving the file {src} to the 'invalid_meta' folder. Error: {e}")
248
+ return new_dst
249
+
250
+
251
+
252
+ def sort_images_by_model_name(model_dir):
253
+ global NEW_IMAGES_DOWNLOADED
254
+ no_meta_dir = os.path.join(model_dir, 'no_meta_data')
255
+ os.makedirs(no_meta_dir, exist_ok=True)
256
+
257
+ if os.path.exists(model_dir) and os.listdir(model_dir):
258
+ all_files = os.listdir(model_dir)
259
+ meta_files = [f for f in all_files if f.endswith('_meta.txt') or f.endswith('_meta_no_meta.txt')]
260
+
261
+ for meta_file in meta_files:
262
+ with open(os.path.join(model_dir, meta_file), 'r', encoding='utf-8') as file:
263
+ content = file.read()
264
+ base_name = meta_file.replace('_meta_no_meta.txt', '').replace('_meta.txt', '').strip()
265
+
266
+ if "No metadata available for this image." in content or meta_file.endswith('_meta_no_meta.txt'):
267
+ image_moved = False
268
+ for ext in ['.jpeg', '.png']:
269
+ image_path = os.path.join(model_dir, base_name + ext)
270
+ if os.path.exists(image_path):
271
+ shutil.move(image_path, os.path.join(no_meta_dir, os.path.basename(image_path)))
272
+ image_moved = True
273
+ break
274
+ if not image_moved:
275
+ logger_cid.info(f"No image found for metadata file {meta_file} (This is expected for images that didn't pass the prompt check)")
276
+ shutil.move(os.path.join(model_dir, meta_file), os.path.join(no_meta_dir, meta_file))
277
+ else:
278
+ model_name_found = False
279
+ for line in content.split('\n'):
280
+ if "Model:" in line:
281
+ model_name = line.split(":")[1].strip()
282
+ model_name_found = True
283
+ break
284
+
285
+ if model_name_found:
286
+ model_name = clean_and_shorten_path(model_name)
287
+ target_dir = os.path.join(model_dir, model_name)
288
+ os.makedirs(target_dir, exist_ok=True)
289
+ process_image_and_meta(model_dir, meta_file, target_dir, valid_meta=True)
290
+ else:
291
+ process_image_and_meta(model_dir, meta_file, model_dir, valid_meta=False)
292
+
293
+ # Only check for orphaned images among those that were actually downloaded
294
+ downloaded_images_in_dir = [f for f in all_files if f.endswith(('.jpeg', '.png')) and os.path.join(model_dir, f) in download_stats["downloaded"]]
295
+ for file in downloaded_images_in_dir:
296
+ base_name = file.rsplit('.', 1)[0]
297
+ if not any(meta for meta in meta_files if meta.startswith(base_name)):
298
+ logger_cid.warning(f"Orphaned image found: {os.path.join(model_dir, file)}")
299
+ os.remove(os.path.join(model_dir, file))
300
+
301
+ if meta_files:
302
+ NEW_IMAGES_DOWNLOADED = True
303
+
304
+
305
+ def process_image_and_meta(model_dir, meta_file, target_dir, valid_meta):
306
+ base_name = meta_file.replace('_meta.txt', '').replace('_meta_no_meta.txt', '')
307
+ image_moved = False
308
+ new_image_path = None
309
+
310
+ extensions = ['.jpeg', '.png']
311
+ for extension in extensions:
312
+ image_path = os.path.join(model_dir, base_name + extension)
313
+ if os.path.exists(image_path):
314
+ try:
315
+ if valid_meta:
316
+ new_image_path = os.path.join(target_dir, os.path.basename(image_path))
317
+ shutil.move(image_path, new_image_path)
318
+ else:
319
+ new_image_path = move_to_invalid_meta(image_path, model_dir)
320
+ image_moved = True
321
+ logger_cid.info(f"Moved image {image_path} to {new_image_path}")
322
+ except Exception as e:
323
+ logger_cid.error(f"Error moving image {image_path}: {e}")
324
+ break
325
+
326
+ if not image_moved:
327
+ logger_cid.info(f"No image found for metadata file {meta_file} (This is expected for images that didn't pass the prompt check)")
328
+ return None, None
329
+
330
+ meta_path = os.path.join(model_dir, meta_file)
331
+ if valid_meta:
332
+ new_meta_path = os.path.join(target_dir, meta_file)
333
+ shutil.move(meta_path, new_meta_path)
334
+ else:
335
+ new_meta_path = move_to_invalid_meta(meta_path, model_dir)
336
+
337
+ return new_image_path, new_meta_path
338
+
339
+
340
+ visited_pages = set()
341
+
342
+
343
+ async def search_models_by_tag(tag, failed_search_requests=[]):
344
+ base_url = f"https://civitai.com/api/v1/models?tag={tag}&nsfw=true"
345
+ model_id = set()
346
+ async with httpx.AsyncClient() as client:
347
+ async with semaphore:
348
+ while base_url:
349
+ try:
350
+ response = await client.get(base_url, headers=headers)
351
+ if response.status_code == 200:
352
+ data = response.json()
353
+ items = data.get('items', [])
354
+ if items: # If items list is not empty
355
+ for model in items:
356
+ model_id.add(model['id'])
357
+ else: # If items list is empty
358
+ print(f"No models found for the tag '{tag}'.")
359
+ return model_id # Return the empty set
360
+ metadata = data.get('metadata', {})
361
+ nextPage = metadata.get('nextPage', None)
362
+ base_url = nextPage if nextPage else None
363
+ else:
364
+ logger_cid.error(f"Server response: {response.status_code} for URL: {base_url}")
365
+ print(f"Unexpected server response: {response.status_code}. Please try again later.")
366
+ failed_search_requests.append(base_url)
367
+ break
368
+ except httpx.RequestError as e:
369
+ logger_cid.exception(f"Request error occurred while fetching models for tag '{tag}': {e}")
370
+ print("A connection error occurred. Please check your internet connection and try again.")
371
+ failed_search_requests.append(base_url)
372
+ break
373
+ return model_id
374
+
375
+
376
+ tag_model_mapping = {}
377
+
378
+ async def download_images_for_model_with_tag_check(model_ids, option_folder, timeout_value, quality='SD', tag_to_check=None, tag_dir_name=None, sanitized_tag_dir_name=None, disable_prompt_check=False, allow_redownload=2):
379
+ global NEW_IMAGES_DOWNLOADED, download_stats
380
+ failed_urls = []
381
+ images_without_meta = 0
382
+ tasks = []
383
+ total_api_items = 0
384
+ total_downloaded_items = 0
385
+
386
+ if tag_to_check is None:
387
+ tag_to_check = tag_dir_name
388
+
389
+ for model_id in model_ids:
390
+ url = f"{base_url}?modelId={str(model_id)}&nsfw=X"
391
+ visited_pages = set() # Reset the visited_pages set for each model
392
+
393
+ while url:
394
+ if url in visited_pages:
395
+ logger_cid.info(f"URL {url} already visited. Ending loop.")
396
+ break
397
+ visited_pages.add(url)
398
+ async with httpx.AsyncClient() as client:
399
+ async with semaphore:
400
+ try:
401
+ response = await client.get(url, timeout=timeout_value, headers=headers)
402
+ if response.status_code == 200:
403
+ data = response.json()
404
+ items = data.get('items', [])
405
+ total_api_items += len(items)
406
+
407
+ # Create directory for model ID
408
+ model_name = "unknown_model"
409
+ if items and isinstance(items[0], dict):
410
+ model_name = items[0]["meta"].get("Model", "unknown_model") if isinstance(items[0].get("meta"), dict) else "unknown_model"
411
+
412
+ tag_dir = os.path.join(option_folder, sanitized_tag_dir_name)
413
+ os.makedirs(tag_dir, exist_ok=True)
414
+
415
+ with tag_model_mapping_lock:
416
+ if tag_dir_name not in tag_model_mapping:
417
+ tag_model_mapping[tag_dir_name] = []
418
+ tag_model_mapping[tag_dir_name].append((model_id, model_name))
419
+
420
+ model_dir = os.path.join(tag_dir, f"model_{(model_id)}")
421
+ os.makedirs(model_dir, exist_ok=True)
422
+
423
+ for item in items:
424
+ item_meta = item.get("meta") if item else None
425
+ prompt = item_meta.get("prompt", "").replace(" ", "_") if isinstance(item_meta, dict) else ""
426
+ if disable_prompt_check or (tag_to_check and all(word in prompt.lower() for word in tag_to_check.lower().split("_"))):
427
+ image_id = item['id']
428
+ image_url = item['url']
429
+ image_extension = ".png" if quality == 'HD' else ".jpeg"
430
+ image_path = os.path.join(model_dir, f"{image_id}{image_extension}")
431
+
432
+ if allow_redownload == 2 and check_if_image_downloaded(str(image_id), image_path, quality):
433
+ continue
434
+
435
+ task = download_image(image_url, image_path, timeout_value, quality)
436
+ tasks.append(task)
437
+ else:
438
+ logger_cid.info(f"Skipping image {item['id']} as it doesn't pass the prompt check")
439
+
440
+ # This will run all download tasks asynchronously
441
+ download_results = await asyncio.gather(*tasks)
442
+ tasks = [] # Reset the tasks list after gathering the results
443
+
444
+ # Check the results and mark the downloaded images
445
+ for idx, (download_success, reason) in enumerate(download_results):
446
+ image_id = items[idx]['id']
447
+ if download_success:
448
+ NEW_IMAGES_DOWNLOADED = True
449
+ total_downloaded_items += 1
450
+ image_path = os.path.join(model_dir, f"{image_id}{'.png' if quality == 'HD' else '.jpeg'}")
451
+ tags = [tag_to_check] if tag_to_check else []
452
+ mark_image_as_downloaded(str(image_id), image_path, quality, tags=tags, url=items[idx]['url'])
453
+ download_stats["downloaded"].append(image_path)
454
+ else:
455
+ logger_cid.error(f"Failed to download image {image_id}: {reason}")
456
+ download_stats["skipped"].append((items[idx]['url'], reason))
457
+
458
+ meta_output_path = os.path.join(model_dir, f"{image_id}_meta.txt")
459
+ await write_meta_data(items[idx].get("meta"), meta_output_path, image_id, items[idx].get('username', 'unknown'))
460
+ if not items[idx].get("meta"):
461
+ images_without_meta += 1
462
+
463
+ sort_images_by_model_name(model_dir)
464
+ metadata = data['metadata']
465
+ next_page = metadata.get('nextPage')
466
+ if next_page:
467
+ url = next_page
468
+ await asyncio.sleep(3) # Add a delay between requests
469
+ else:
470
+ break
471
+ except Exception as e:
472
+ logger_cid.error(f"Error processing URL {url}: {str(e)}")
473
+ failed_urls.append(url)
474
+ continue
475
+
476
+ return tasks, failed_urls, images_without_meta, sanitized_tag_dir_name, total_api_items, total_downloaded_items
477
+
478
+
479
+ def sort_images_by_tag(option_folder, tag_model_mapping):
480
+ with tag_model_mapping_lock:
481
+ for tag, model_ids in tag_model_mapping.items():
482
+ sanitized_tag = tag.replace(" ", "_")
483
+ tag_dir = os.path.join(option_folder, sanitized_tag)
484
+ if not os.listdir(tag_dir):
485
+ print(f"No images found for the tag: {tag}")
486
+
487
+
488
+ def write_summary_to_csv(tag, downloaded_images, option_folder, tag_model_mapping):
489
+ with tag_model_mapping_lock:
490
+ tag_dir = os.path.join(option_folder, tag.replace(" ", "_"))
491
+ for model_info in tag_model_mapping.get(tag, []):
492
+ model_id, model_name = model_info
493
+ model_dir = os.path.join(tag_dir, f"model_{(model_id)}")
494
+
495
+ # Check if the model directory exists
496
+ if not os.path.exists(model_dir):
497
+ print(f"The {model_dir} directory does not exist. Skip the creation of the CSV file.")
498
+ continue
499
+ csv_file = os.path.join(model_dir, f"{tag.replace(' ', '_')}_summary_{datetime.now().strftime('%Y%m%d')}.csv")
500
+ with open(csv_file, "w", newline="") as file:
501
+ writer = csv.writer(file)
502
+ writer.writerow(["Current Tag", "Previously Downloaded Tag", "Image Path", "Download URL"])
503
+ for image_id, image_info in downloaded_images.items():
504
+
505
+ # Here it is assumed that the "tags" are stored in the "downloaded_images" dictionary
506
+ if tag in image_info.get("tags", []):
507
+ for prev_tag in image_info.get("tags", []):
508
+ if prev_tag != tag:
509
+ relative_path = os.path.relpath(image_info["path"], model_dir)
510
+ writer.writerow([tag, prev_tag, relative_path, image_info["url"]])
511
+
512
+ failed_identifiers = [] # List for saving failed usernames and model IDs
513
+
514
+ async def is_valid_username(username):
515
+ url = f"{base_url}?username={username.strip()}&nsfw=X"
516
+ async with httpx.AsyncClient() as client:
517
+ try:
518
+ response = await client.get(url, headers=headers)
519
+ if response.status_code == 500:
520
+ response_data = response.json()
521
+ if 'error' in response_data and response_data['error'] == "User not found":
522
+ return False, "Username not found"
523
+ return True, None
524
+ except httpx.RequestError as e:
525
+ return False, f"Network error: {e}"
526
+ except json.JSONDecodeError as e:
527
+ return False, f"Error decoding response: {e}"
528
+ except Exception as e:
529
+ return False, f"Unexpected error: {str(e)}"
530
+
531
+
532
+ async def is_valid_model_id(identifier):
533
+ url = f"{base_url}?modelId={str(identifier)}&nsfw=X"
534
+ async with httpx.AsyncClient() as client:
535
+ try:
536
+ response = await client.get(url, headers=headers)
537
+ if response.status_code == 500:
538
+ # If the modelId fails, try modelVersionId
539
+ url = f"{base_url}?modelVersionId={str(identifier)}&nsfw=X"
540
+ response = await client.get(url, headers=headers)
541
+ if response.status_code == 500:
542
+ return False, f"Invalid input syntax for model ID or model version ID: {identifier}"
543
+ elif response.status_code == 304:
544
+ response_data = response.json()
545
+ if not response_data['items']:
546
+ return False, f"No items found for model ID or model version ID: {identifier}"
547
+ return True, None
548
+ except httpx.RequestError as e:
549
+ return False, f"Network error: {e}"
550
+ except json.JSONDecodeError as e:
551
+ return False, f"Error decoding response: {e}"
552
+ except Exception as e:
553
+ return False, f"Unexpected error: {str(e)}"
554
+
555
+
556
+ async def is_valid_model_version_id(model_version_id):
557
+ url = f"{base_url}?modelVersionId={str(model_version_id)}&nsfw=X"
558
+ async with httpx.AsyncClient() as client:
559
+ try:
560
+ response = await client.get(url, headers=headers)
561
+ if response.status_code == 500:
562
+ return False, f"Invalid input syntax for model version ID: {model_version_id}"
563
+ elif response.status_code == 304:
564
+ response_data = response.json()
565
+ if not response_data['items']:
566
+ return False, f"No items found for model version ID: {model_version_id}"
567
+ return True, None
568
+ except httpx.RequestError as e:
569
+ return False, f"Network error: {e}"
570
+ except json.JSONDecodeError as e:
571
+ return False, f"Error decoding response: {e}"
572
+ except Exception as e:
573
+ return False, f"Unexpected error: {str(e)}"
574
+
575
+
576
+ def get_url_for_identifier(identifier, identifier_type):
577
+ base_url = "https://civitai.com/api/v1/images"
578
+ if identifier_type == 'model':
579
+ return f"{base_url}?modelId={str(identifier)}&nsfw=X"
580
+ elif identifier_type == 'modelVersion':
581
+ return f"{base_url}?modelVersionId={str(identifier)}&nsfw=X"
582
+ elif identifier_type == 'username':
583
+ return f"{base_url}?username={identifier.strip()}&nsfw=X&sort=Newest"
584
+ else:
585
+ raise ValueError("Invalid identifier_type. Should be 'model', 'modelVersion', or 'username'.")
586
+
587
+
588
+ async def download_images(identifier, option_folder, identifier_type, timeout_value, quality='SD', allow_redownload=2, datetime_min = None, datetime_max = None, reaction_min = 0, condi = "0", excludes = [], mode5 = False):
589
+ global NEW_IMAGES_DOWNLOADED, download_stats
590
+ valid, error_message = True, None
591
+ if identifier_type == 'username':
592
+ valid, error_message = await is_valid_username(identifier)
593
+ elif identifier_type == 'model':
594
+ valid, error_message = await is_valid_model_id(identifier)
595
+ elif identifier_type == 'modelVersion':
596
+ valid, error_message = await is_valid_model_version_id(identifier)
597
+ if not valid:
598
+ logger_cid.warning(f"Skipping: {error_message}")
599
+ failed_identifiers.append((identifier_type, identifier))
600
+ return [], 0, 0, 0
601
+
602
+ url = get_url_for_identifier(identifier, identifier_type)
603
+ failed_urls = []
604
+ images_without_meta = 0
605
+ total_items = 0
606
+ total_downloaded = 0
607
+
608
+ # Define dir_name here, outside the while loop
609
+ if identifier_type in ['model', 'modelVersion']:
610
+ dir_name = os.path.join(option_folder, f"{identifier_type}_{identifier}")
611
+ elif identifier_type == 'username':
612
+ dir_name = os.path.join(option_folder, identifier.strip())
613
+ else:
614
+ logger_cid.error(f"Invalid identifier_type: {identifier_type}")
615
+ return [], 0, 0, 0
616
+
617
+ os.makedirs(dir_name, exist_ok=True)
618
+
619
+ while url:
620
+ if url in visited_pages:
621
+ logger_cid.info(f"URL {url} already visited. Ending loop.")
622
+ break
623
+ visited_pages.add(url)
624
+ async with httpx.AsyncClient() as client:
625
+ async with semaphore:
626
+ try:
627
+ response = await client.get(url, timeout=timeout_value)
628
+ if response.status_code == 200:
629
+ data = response.json()
630
+ items = data.get('items', [])
631
+ logger_cid.info(f"Received {len(items)} items from API for {identifier_type} {identifier}")
632
+ total_items += len(items)
633
+
634
+ tasks = []
635
+ mode5_tmp = []
636
+ for item in items:
637
+ image_id = item['id']
638
+ image_url = item['url']
639
+ image_extension = ".png" if quality == 'HD' else ".jpeg"
640
+ image_path = os.path.join(dir_name, f"{image_id}{image_extension}")
641
+
642
+ if(mode5):
643
+ if str(image_id) in excludes:
644
+ print("exclude: "+ str(image_id))
645
+ continue
646
+
647
+ skip_val = 0
648
+ skip_limit = 0
649
+ skip_txt = "id:" + str(image_id)
650
+
651
+ if(datetime_min or datetime_max):
652
+ skip_limit += 1
653
+ image_stamp = item['createdAt'].split("T")[0]
654
+ datetime_img = datetime.strptime(image_stamp, '%Y-%m-%d')
655
+
656
+ if(datetime_min):
657
+ datetime_limit = datetime.strptime(datetime_min, '%Y-%m-%d')
658
+ if(datetime_limit > datetime_img):
659
+ skip_txt += " | wrong date: "+ image_stamp
660
+ skip_val += 1
661
+ #continue
662
+
663
+ if(datetime_max and skip_limit > skip_val):
664
+ datetime_limit = datetime.strptime(datetime_max, '%Y-%m-%d')
665
+ if(datetime_limit < datetime_img):
666
+ skip_txt += " | wrong date: "+ image_stamp
667
+ skip_val += 1
668
+ #continue
669
+
670
+ if(reaction_min > 0):
671
+ skip_limit += 1
672
+ reaction_total = item['stats']['likeCount'] + item['stats']['heartCount']
673
+ if(reaction_min > reaction_total):
674
+ skip_txt += " | insufficient reaction: "+ str(reaction_total)
675
+ skip_val += 1
676
+ #continue
677
+
678
+ if(skip_val > 0 and (condi == "0" and skip_limit == skip_val or condi == "1")):
679
+ print(skip_txt)
680
+ continue
681
+
682
+ if allow_redownload == 2 and check_if_image_downloaded(str(image_id), image_path, quality):
683
+ continue
684
+
685
+ task = download_image(image_url, image_path, timeout_value, quality)
686
+ tasks.append(task)
687
+ mode5_tmp.append(item)
688
+
689
+ download_results = await asyncio.gather(*tasks)
690
+ total_downloaded += sum(1 for result, _ in download_results if result)
691
+
692
+ if(mode5):
693
+ items = mode5_tmp
694
+
695
+ for idx, (download_success, reason) in enumerate(download_results):
696
+ image_id = items[idx]['id']
697
+ if download_success:
698
+ NEW_IMAGES_DOWNLOADED = True
699
+ image_path = os.path.join(dir_name, f"{image_id}{'.png' if quality == 'HD' else '.jpeg'}")
700
+ mark_image_as_downloaded(str(image_id), image_path, quality)
701
+ download_stats["downloaded"].append(image_path)
702
+ else:
703
+ download_stats["skipped"].append((items[idx]['url'], reason))
704
+
705
+ if(mode5):
706
+ meta_output_path = os.path.join(dir_name, f"{image_id}.txt")
707
+ await write_prompt(items[idx].get("meta"), meta_output_path)
708
+ else:
709
+ meta_output_path = os.path.join(dir_name, f"{image_id}_meta.txt")
710
+ await write_meta_data(items[idx].get("meta"), meta_output_path, image_id, items[idx].get('username', 'unknown'))
711
+
712
+ if not items[idx].get("meta"):
713
+ images_without_meta += 1
714
+
715
+ metadata = data['metadata']
716
+ next_page = metadata.get('nextPage')
717
+
718
+ print("loading next page")
719
+
720
+ if next_page:
721
+ url = next_page
722
+ await asyncio.sleep(3)
723
+ else:
724
+ break
725
+ except Exception as e:
726
+ logger_cid.error(f"Error processing URL {url}: {str(e)}")
727
+ failed_urls.append(url)
728
+ continue
729
+
730
+ print("complete?")
731
+ if identifier_type in ['model', 'modelVersion', 'username'] and not mode5:
732
+ sort_images_by_model_name(dir_name)
733
+
734
+ logger_cid.info(f"Total items from API: {total_items}, Total downloaded: {total_downloaded}")
735
+ return failed_urls, images_without_meta, total_items, total_downloaded
736
+
737
+
738
+ def print_download_statistics():
739
+ print("\n")
740
+ print(f"Number of downloaded images: {len(download_stats['downloaded'])}")
741
+ print(f"Number of skipped images: {len(download_stats['skipped'])}")
742
+
743
+ if download_stats['skipped']:
744
+ print("\nReasons for skipping images:")
745
+ reasons = {}
746
+ for _, reason in download_stats['skipped']:
747
+ if reason in reasons:
748
+ reasons[reason] += 1
749
+ else:
750
+ reasons[reason] = 1
751
+
752
+ for reason, count in reasons.items():
753
+ print(f"- {reason}: {count} times")
754
+
755
+ # Main async function
756
+ async def main():
757
+ global downloaded_images, download_stats
758
+ downloaded_images = load_downloaded_images()
759
+ download_stats = {
760
+ "downloaded": [],
761
+ "skipped": [],
762
+ }
763
+
764
+
765
+ # Check if we're in mixed mode
766
+ provided_args = [arg for arg in vars(args) if getattr(args, arg) is not None]
767
+ if provided_args and len(provided_args) < len(vars(args)):
768
+ print("Mixed mode detected. Some arguments provided via command-line, others will be prompted.")
769
+ logger_cid.info("Running in mixed mode. Some arguments provided via command-line, others will be prompted.")
770
+
771
+ # Check for mismatched arguments
772
+ check_mismatched_arguments()
773
+
774
+ timeout_value = get_timeout_value()
775
+ quality = get_quality()
776
+ allow_redownload = get_redownload_option()
777
+
778
+ failed_search_requests = []
779
+ failed_urls = []
780
+ images_without_meta = 0
781
+ total_api_items = 0
782
+ total_downloaded_items = 0
783
+
784
+ choice = get_mode_choice()
785
+ tasks = []
786
+
787
+ if choice == "1":
788
+ usernames = get_usernames()
789
+ option_folder = create_option_folder('Username_Search', output_dir)
790
+ tasks.extend([download_images(username.strip(), option_folder, 'username', timeout_value, quality, allow_redownload) for username in usernames])
791
+
792
+ elif choice == "2":
793
+ model_ids = get_model_ids()
794
+ option_folder = create_option_folder('Model_ID_Search', output_dir)
795
+ tasks.extend([download_images(model_id.strip(), option_folder, 'model', timeout_value, quality, allow_redownload) for model_id in model_ids])
796
+
797
+ elif choice == "3":
798
+ tags = get_tags()
799
+ disable_prompt_check = get_disable_prompt_check()
800
+ option_folder = create_option_folder('Model_Tag_Search', output_dir)
801
+
802
+ for tag in tags:
803
+ sanitized_tag_dir_name = tag.replace(" ", "_")
804
+ model_ids = await search_models_by_tag(tag.replace("_", "%20"), failed_search_requests)
805
+ tag_to_check = None if disable_prompt_check else tag
806
+ tasks_for_tag, failed_urls_for_tag, images_without_meta_for_tag, sanitized_tag_dir_name_for_tag, api_items, downloaded_items = await download_images_for_model_with_tag_check(model_ids, option_folder, timeout_value, quality, tag_to_check, tag, sanitized_tag_dir_name, disable_prompt_check, allow_redownload)
807
+ tasks.extend(tasks_for_tag)
808
+ failed_urls.extend(failed_urls_for_tag)
809
+ images_without_meta += images_without_meta_for_tag
810
+ total_api_items += api_items
811
+ total_downloaded_items += downloaded_items
812
+
813
+ # Sort images into tag-related folders
814
+ sort_images_by_tag(option_folder, tag_model_mapping)
815
+
816
+ for tag, model_ids in tag_model_mapping.items():
817
+ write_summary_to_csv(tag, downloaded_images, option_folder, tag_model_mapping)
818
+
819
+ elif choice == "4":
820
+ model_version_ids = get_model_version_ids()
821
+ option_folder = create_option_folder('Model_Version_ID_Search', output_dir)
822
+ tasks.extend([download_images(model_version_id.strip(), option_folder, 'modelVersion', timeout_value, quality, allow_redownload) for model_version_id in model_version_ids])
823
+
824
+ elif choice == "5":
825
+ usernames = get_usernames()
826
+ option_folder = create_option_folder('Username_Search', output_dir)
827
+
828
+ datetime_min = get_datetime_min()
829
+ datetime_max = get_datetime_max()
830
+ reaction_min = get_reaction_min()
831
+ condi = get_condi()
832
+ excludes = get_excludes()
833
+
834
+ tasks.extend([download_images(username.strip(), option_folder, 'username', timeout_value, quality, allow_redownload, datetime_min, datetime_max, reaction_min, condi, excludes, True) for username in usernames])
835
+
836
+ else:
837
+ logger_cid.error("Invalid choice!")
838
+ return
839
+
840
+ # Execute all collected tasks
841
+ results = await asyncio.gather(*tasks)
842
+
843
+ # Extract failed_urls, images_without_meta, and download statistics from the results
844
+ for result in results:
845
+ failed_urls.extend(result[0])
846
+ images_without_meta += result[1]
847
+ total_api_items += result[2]
848
+ total_downloaded_items += result[3]
849
+
850
+
851
+ for tag, model_ids in tag_model_mapping.items():
852
+ write_summary_to_csv(tag, downloaded_images, option_folder, tag_model_mapping)
853
+
854
+ if failed_urls:
855
+ logger_cid.info("Retrying failed URLs...")
856
+ for url in failed_urls:
857
+ await download_image(url, option_folder, timeout_value=timeout_value, quality=quality)
858
+
859
+ # Attempt to retry failed search requests
860
+ if failed_search_requests:
861
+ logger_cid.info("Retrying failed search requests...")
862
+ for url in failed_search_requests:
863
+ tag = url.split("tag=")[-1]
864
+ await search_models_by_tag(tag, [])
865
+
866
+ if images_without_meta > 0:
867
+ logger_cid.info(f"{images_without_meta} images have no meta data.")
868
+
869
+ logger_cid.info(f"Total API items: {total_api_items}, Total downloaded: {total_downloaded_items}")
870
+ print(f"Total API items: {total_api_items}, Total downloaded: {total_downloaded_items}")
871
+ print_download_statistics()
872
+
873
+ # Helper functions for main
874
+ def check_mismatched_arguments():
875
+ if args.mode:
876
+ if args.mode == 1 and (args.model_id or args.model_version_id or args.tags):
877
+ print("Warning: --model_id, --model_version_id, and --tags are not used in username mode. These arguments will be ignored.")
878
+ logger_cid.warning("Warning: --model_id, --model_version_id, and --tags are not used in username mode. These arguments will be ignored.")
879
+ elif args.mode == 2 and (args.username or args.model_version_id or args.tags):
880
+ print("Warning: --username, --model_version_id, and --tags are not used in model ID mode. These arguments will be ignored.")
881
+ logger_cid.warning("Warning: --username, --model_version_id, and --tags are not used in model ID mode. These arguments will be ignored.")
882
+ elif args.mode == 3 and (args.username or args.model_id or args.model_version_id):
883
+ print("Warning: --username, --model_id, and --model_version_id are not used in tag search mode. These arguments will be ignored.")
884
+ logger_cid.warning("Warning: --username, --model_id, and --model_version_id are not used in tag search mode. These arguments will be ignored.")
885
+ elif args.mode == 4 and (args.username or args.model_id or args.tags):
886
+ print("Warning: --username, --model_id, and --tags are not used in model version ID mode. These arguments will be ignored.")
887
+ logger_cid.warning("Warning: --username, --model_id, and --tags are not used in model version ID mode. These arguments will be ignored.")
888
+
889
+ def get_datetime_min():
890
+ if args.datetime_min:
891
+ return args.datetime_min
892
+ else:
893
+ return None
894
+
895
+ def get_datetime_max():
896
+ if args.datetime_max:
897
+ return args.datetime_max
898
+ else:
899
+ return None
900
+
901
+ def get_reaction_min():
902
+ if args.reaction_min:
903
+ return args.reaction_min
904
+ else:
905
+ return 0
906
+
907
+ def get_condi():
908
+ if args.condi:
909
+ return str(args.condi)
910
+ else:
911
+ return "0"
912
+
913
+ def get_excludes():
914
+ if args.excludes:
915
+ return args.excludes.split(",")
916
+ else:
917
+ return []
918
+
919
+ def get_timeout_value():
920
+ if args.timeout:
921
+ return args.timeout
922
+ else:
923
+ timeout_input = input("Enter timeout value (in seconds): ")
924
+ if timeout_input.isdigit() and int(timeout_input) > 0:
925
+ return int(timeout_input)
926
+ else:
927
+ logger_cid.warning("Invalid timeout value. Using default value of 60 seconds.")
928
+ return 60
929
+
930
+ def get_quality():
931
+ if args.quality:
932
+ return 'HD' if args.quality == 2 else 'SD'
933
+ else:
934
+ quality_choice = input("Choose image quality (1 for SD, 2 for HD): ")
935
+ if quality_choice == '2':
936
+ return 'HD'
937
+ elif quality_choice == '1':
938
+ return 'SD'
939
+ else:
940
+ logger_cid.warning("Invalid quality choice. Using default quality SD.")
941
+ return 'SD'
942
+
943
+ def get_redownload_option():
944
+ if args.redownload:
945
+ return args.redownload
946
+ else:
947
+ allow_redownload_choice = input("Allow re-downloading of images already tracked (1 for Yes, 2 for No) [default: 2]: ")
948
+ if allow_redownload_choice == '1':
949
+ return 1
950
+ elif allow_redownload_choice == '2' or allow_redownload_choice.strip() == '':
951
+ return 2
952
+ else:
953
+ logger_cid.warning("Invalid choice. Using default value (2 - No).")
954
+ return 2
955
+
956
+ def get_mode_choice():
957
+ if args.mode:
958
+ return str(args.mode)
959
+ else:
960
+ return input("Choose mode (1 for username, 2 for model ID, 3 for Model tag search, 4 for model version ID): ")
961
+
962
+ def get_usernames():
963
+ if args.username:
964
+ return [args.username]
965
+ else:
966
+ return input("Enter username: ").split(",")
967
+
968
+ def get_model_ids():
969
+ if args.model_id:
970
+ return [args.model_id]
971
+ else:
972
+ while True:
973
+ model_ids_input = input("Enter model ID: ")
974
+ model_ids = model_ids_input.split(",")
975
+ if all(model_id.strip().isdigit() for model_id in model_ids):
976
+ return model_ids
977
+ else:
978
+ logger_cid.warning("Invalid input. Please enter only numeric model IDs.")
979
+
980
+ def get_tags():
981
+ if args.tags:
982
+ return [tag.strip().replace(" ", "_") for tag in args.tags.split(',')]
983
+ else:
984
+ tags_input = input("Enter tags (comma-separated): ")
985
+ return [tag.strip().replace(" ", "_") for tag in tags_input.split(',')]
986
+
987
+ def get_disable_prompt_check():
988
+ if args.disable_prompt_check is not None:
989
+ return args.disable_prompt_check.lower() == 'y'
990
+ else:
991
+ return input("Disable prompt check? (y/n): ").lower() in ['y', 'yes']
992
+
993
+ def get_model_version_ids():
994
+ if args.model_version_id:
995
+ return [args.model_version_id]
996
+ else:
997
+ while True:
998
+ model_version_ids_input = input("Enter model version ID: ")
999
+ model_version_ids = model_version_ids_input.split(",")
1000
+ if all(model_version_id.strip().isdigit() for model_version_id in model_version_ids):
1001
+ return model_version_ids
1002
+ else:
1003
+ logger_cid.warning("Invalid input. Please enter only numeric model version IDs.")
1004
+
1005
+
1006
+ if __name__ == "__main__":
1007
+ args = parse_arguments()
1008
+
1009
+ if is_command_line_mode():
1010
+ print("Running in command-line mode.")
1011
+ logger_cid.info("Running in command-line mode")
1012
+ else:
1013
+ print("Running in interactive mode.")
1014
+ logger_cid.info("Running in interactive mode")
1015
+
1016
+ asyncio.run(main())
1017
+
1018
+ if failed_identifiers:
1019
+ logger_cid.warning("Failed identifiers:")
1020
+ for id_type, id_value in failed_identifiers:
1021
+ logger_cid.warning(f"{id_type}: {id_value}")
1022
+
1023
+ logger_cid.info("Image download completed.")