feat(civitai): verify downloads in flight and cache under loader keys

The downloader hashes bytes as it writes them, seeding from the partial
on resume, and checks the declared SHA256 against that digest instead
of rereading the file. The declared hash, or the computed one when none
is given, is cached under the key the owning loader reads; loader_kind
and hash_cache_title derive that key for the downloader and the sweeps.
This commit is contained in:
CalamitousFelicitousness
2026-09-13 05:37:44 +01:00
parent cd86ea6321
commit b06b538104
3 changed files with 103 additions and 44 deletions
+48 -35
View File
@@ -28,7 +28,7 @@ class DownloadItem:
token: str | None = None
model_id: int = 0
version_id: int = 0
status: str = "queued" # queued | downloading | verifying | completed | failed | cancelled
status: str = "queued" # queued | downloading | completed | failed | cancelled
progress: float = 0.0
bytes_downloaded: int = 0
bytes_total: int = 0
@@ -195,6 +195,7 @@ class DownloadManager:
item.status = "downloading"
item.bytes_downloaded = starting_pos
digest = hashlib.sha256()
try:
r = shared.req(item.url, headers=headers if headers else None, stream=True)
@@ -222,6 +223,10 @@ class DownloadManager:
starting_pos = 0
item.bytes_downloaded = 0
os.truncate(temp_file, 0)
if starting_pos > 0: # a resumed download's digest must include the partial already on disk
with open(temp_file, 'rb') as partial:
for block in iter(lambda: partial.read(1024 * 1024), b''):
digest.update(block)
total_size = int(r.headers.get('content-length', 0))
item.bytes_total = starting_pos + total_size
@@ -255,6 +260,7 @@ class DownloadManager:
return
f.write(chunk)
digest.update(chunk)
written += len(chunk)
item.bytes_downloaded = written
if item.bytes_total > 0:
@@ -290,27 +296,19 @@ class DownloadManager:
log.error(f'CivitAI download error: id={item.id} {e}')
return
# Hash verification
if item.expected_hash:
item.status = "verifying"
try:
from modules import hashes
computed = hashes.calculate_sha256(temp_file, quiet=True)
if computed.upper() != item.expected_hash.upper():
discard = getattr(shared.opts, 'civitai_discard_hash_mismatch', True)
if discard:
try:
os.remove(temp_file)
except OSError:
pass
item.status = "failed"
item.error = f'hash mismatch: expected={item.expected_hash[:16]}... got={computed[:16]}...'
item.completed_at = datetime.now()
log.error(f'CivitAI download hash mismatch: id={item.id} expected={item.expected_hash[:16]} got={computed[:16]}')
return
log.warning(f'CivitAI download hash mismatch (kept): id={item.id} expected={item.expected_hash[:16]} got={computed[:16]}')
except Exception as e:
log.warning(f'CivitAI download hash check failed: id={item.id} {e}')
computed = digest.hexdigest()
if item.expected_hash and computed != item.expected_hash.lower():
if getattr(shared.opts, 'civitai_discard_hash_mismatch', True):
try:
os.remove(temp_file)
except OSError:
pass
item.status = "failed"
item.error = f'hash mismatch: expected={item.expected_hash[:16]}... got={computed[:16]}...'
item.completed_at = datetime.now()
log.error(f'CivitAI download hash mismatch: id={item.id} expected={item.expected_hash[:16]} got={computed[:16]}')
return
log.warning(f'CivitAI download hash mismatch (kept): id={item.id} expected={item.expected_hash[:16]} got={computed[:16]}')
# Move temp to final
try:
@@ -326,18 +324,16 @@ class DownloadManager:
item.completed_at = datetime.now()
log.info(f'CivitAI download complete: id={item.id} file="{final_file}" size={item.bytes_downloaded}')
# Write verified hash to cache so check-local finds it immediately
if item.expected_hash:
try:
from modules import hashes
model_type_map = {'Checkpoint': 'checkpoint', 'LORA': 'lora', 'TextualInversion': 'embedding', 'VAE': 'vae'}
prefix = model_type_map.get(item.model_type, item.model_type.lower())
name = os.path.splitext(item.filename)[0]
title = f"{prefix}/{name}"
hashes.cache().add_hash(title, os.path.getmtime(final_file), item.expected_hash.lower())
# the declared hash is cached even on a kept mismatch: it is what CivitAI knows the file by
try:
from modules import hashes
from modules.civitai.filemanage_civitai import loader_kind, hash_cache_title
title = hash_cache_title(loader_kind(final_file), final_file)
if title is not None:
hashes.cache().add_hash(title, os.path.getmtime(final_file), (item.expected_hash or computed).lower())
hashes.save_cache()
except Exception:
pass
except Exception as e:
log.warning(f'CivitAI download hash cache: id={item.id} {e}')
# Download metadata and preview
self._fetch_sidecar(item, final_file)
@@ -776,8 +772,23 @@ def download_civit_preview(model_path: str, preview_url: str, meta: dict | None
return 200, str(total_size), ''
def declared_sha256(version_id: int, url: str, filename: str, token: str | None = None) -> str:
"""SHA256 CivitAI declares for the version file behind url, matched by download URL, then by file name."""
if not version_id:
return ''
from modules.civitai.client_civitai import client
version = client.get_version(version_id, token=token)
if version is None:
return ''
for matches in (lambda f: f.download_url == url, lambda f: f.name == filename):
for f in version.files:
if f.hashes.sha256 and matches(f):
return f.hashes.sha256.lower()
return ''
def download_civit_model(model_url: str, model_name: str = '', model_path: str = '', model_type: str = '', token: str | None = None,
base_model: str = '', model_id: int = 0, version_id: int = 0):
base_model: str = '', model_id: int = 0, version_id: int = 0, expected_hash: str = ''):
"""Legacy function — delegates to DownloadManager for non-blocking downloads."""
if not model_url:
log.error('Model download: no url provided')
@@ -797,17 +808,19 @@ def download_civit_model(model_url: str, model_name: str = '', model_path: str =
folder = model_path
else:
folder = os.path.join(paths.models_path, model_path)
expected_hash = expected_hash or declared_sha256(version_id, model_url, model_name, token=token)
item = download_manager.enqueue(
url=model_url,
folder=folder,
filename=model_name or "Unknown",
model_type=model_type,
expected_hash=expected_hash,
token=token,
model_id=model_id,
version_id=version_id,
)
# Wait for completion (legacy blocking behavior)
while item.status in ("queued", "downloading", "verifying"):
while item.status in ("queued", "downloading"):
time.sleep(0.5)
if item.status == "completed" and not item.error:
from modules.sd_models import list_models
+49
View File
@@ -93,6 +93,55 @@ def iter_type_roots() -> set[Path]:
return {r for r in roots if r.is_dir()}
def path_under(filename: str, root: str | None) -> bool:
if not root:
return False
root = os.path.normcase(os.path.abspath(root)).rstrip(os.sep) + os.sep
return os.path.normcase(os.path.abspath(filename)).startswith(root)
def loader_kind(filename: str) -> str | None:
"""Model loader that lists this file, judged by the folder it is in."""
from modules import shared, paths
ckpt_roots = (getattr(shared.opts, 'ckpt_dir', ''), os.path.join(paths.models_path, 'Stable-diffusion'))
if path_under(filename, getattr(shared.opts, 'vae_dir', '')) or path_under(filename, os.path.join(paths.models_path, 'VAE')):
return 'vae'
if filename.endswith('.vae.safetensors') and any(path_under(filename, root) for root in ckpt_roots):
return 'vae'
if path_under(filename, getattr(shared.opts, 'unet_dir', '')):
return 'unet'
if path_under(filename, getattr(shared.cmd_opts, 'lora_dir', '')):
return 'lora'
if any(path_under(filename, root) for root in ckpt_roots):
return 'checkpoint'
return None
def hash_cache_title(kind: str | None, filename: str, name: str | None = None) -> str | None:
"""Hash cache key the loader of kind reads for filename, or None when it keeps none."""
from modules import shared, paths
basename = os.path.basename(filename)
stem = os.path.splitext(basename)[0]
if kind == 'lora': # lora_load registers the basename with dots replaced
return 'lora/' + stem.replace('.', '_')
if kind == 'unet': # sd_unet keeps the extension on anything but safetensors
return f"unet/{name or (stem if '.safetensors' in basename else basename)}"
if kind == 'vae':
return f'vae/{os.path.abspath(filename)}'
if kind == 'checkpoint':
if name is None: # CheckpointInfo matches the folder by string prefix, then drops the extension
relname = filename
ckpt_dir = getattr(shared.opts, 'ckpt_dir', '') or ''
model_path = os.path.abspath(os.path.join(paths.models_path, 'Stable-diffusion'))
if ckpt_dir and relname.startswith(ckpt_dir):
relname = os.path.relpath(filename, ckpt_dir)
elif relname.startswith(model_path):
relname = os.path.relpath(filename, model_path)
name = os.path.splitext(relname)[0]
return f'checkpoint/{name}'
return None
def resolve_save_path(model_type: str, model_name: str = "", base_model: str = "",
nsfw: bool = False, creator: str = "", model_id: int = 0,
version_id: int = 0, version_name: str = "") -> Path:
+6 -9
View File
@@ -5,6 +5,7 @@ import threading
import concurrent.futures
from modules.shared import log, opts, max_workers, state, cmd_opts
from modules.civitai.client_civitai import client
from modules.civitai.filemanage_civitai import hash_cache_title
GIB = 1024 ** 3
@@ -34,17 +35,13 @@ class CivitModel:
self.status = 'Not found'
PAGE_KINDS = {'model': 'checkpoint', 'lora': 'lora', 'unet/dit': 'unet', 'vae': 'vae'}
def cache_title(page: str, item: dict) -> str | None:
"""Hash cache key read by the page's own loader, or None when it keeps no cache entry."""
if page == 'model':
return f"checkpoint/{item.get('name')}"
if page == 'lora': # lora_load registers the basename with dots replaced
return 'lora/' + os.path.splitext(os.path.basename(item.get('filename', '')))[0].replace('.', '_')
if page == 'unet/dit':
return f"unet/{item.get('name')}"
if page == 'vae':
return f"vae/{item.get('filename')}"
return None
kind = PAGE_KINDS.get(page)
return hash_cache_title(kind, item.get('filename') or '', name=item.get('name') if kind in ('checkpoint', 'unet') else None)
def resolve_sha256(entries: list[tuple[str, dict]], size_limit: int | None = None) -> tuple[dict[str, str], dict[str, str]]: