373 lines
17 KiB
Python
373 lines
17 KiB
Python
"""Local workflow libraries and explicitly requested Civitai downloads."""
|
|
|
|
import asyncio
|
|
import hashlib
|
|
import json
|
|
import os
|
|
from pathlib import Path
|
|
import uuid
|
|
from urllib.parse import urlsplit
|
|
|
|
import aiohttp
|
|
from aiohttp import web
|
|
|
|
BASE = Path(__file__).resolve().parent
|
|
MODEL_TYPES = {
|
|
"Checkpoint": "checkpoints", "LORA": "loras", "LoCon": "loras",
|
|
"DoRA": "loras", "TextualInversion": "embeddings", "VAE": "vae",
|
|
"Controlnet": "controlnet", "Upscaler": "upscale_models",
|
|
"Hypernetwork": "hypernetworks", "MotionModule": "animatediff_models",
|
|
}
|
|
MAX_WORKFLOW_BYTES = 16 * 1024 * 1024
|
|
|
|
|
|
def legacy_settings():
|
|
path = BASE / "app" / "settings.json"
|
|
if not path.exists():
|
|
return {}
|
|
try:
|
|
data = json.loads(path.read_text(encoding="utf-8"))
|
|
return data if isinstance(data, dict) else {}
|
|
except (OSError, ValueError):
|
|
return {}
|
|
|
|
|
|
def workflow_roots():
|
|
configured = os.environ.get("NSIDEBAR_WORKFLOW_ROOTS", "")
|
|
configured += "\n" + str(legacy_settings().get("sb_wf_path", ""))
|
|
roots = []
|
|
for line in configured.splitlines():
|
|
if line.strip():
|
|
path = Path(line.strip()).expanduser().resolve()
|
|
if path.is_dir() and path not in roots:
|
|
roots.append(path)
|
|
return roots
|
|
|
|
|
|
def workflow_path(root, relative):
|
|
candidate = (root / relative).resolve()
|
|
if not candidate.is_relative_to(root) or candidate.suffix.lower() != ".json":
|
|
raise ValueError("Workflow must be a JSON file inside the selected library")
|
|
if candidate.stat().st_size > MAX_WORKFLOW_BYTES:
|
|
raise ValueError("Workflow exceeds the 16 MiB limit")
|
|
return candidate
|
|
|
|
|
|
def read_workflow(path):
|
|
data = json.loads(path.read_text(encoding="utf-8"))
|
|
if not isinstance(data, dict) or not isinstance(data.get("nodes"), list):
|
|
raise ValueError("This JSON is not a ComfyUI workflow")
|
|
return data
|
|
|
|
|
|
def list_workflows(roots):
|
|
files = []
|
|
for index, root in enumerate(roots):
|
|
for path in sorted(root.rglob("*.json")):
|
|
try:
|
|
relative = path.relative_to(root).as_posix()
|
|
read_workflow(workflow_path(root, relative))
|
|
except (OSError, ValueError):
|
|
continue
|
|
files.append({"root": index, "path": relative, "name": path.stem})
|
|
return files
|
|
|
|
|
|
def safe_filename(name):
|
|
if (not isinstance(name, str) or not name or name in {".", ".."}
|
|
or any(char in name for char in '/\\\x00:') or name != Path(name).name):
|
|
raise ValueError("Invalid model filename")
|
|
if Path(name).suffix.lower() not in {".safetensors", ".ckpt", ".pt", ".pth", ".bin", ".gguf"}:
|
|
raise ValueError("Unsupported model file type")
|
|
return name
|
|
|
|
|
|
def managed_models(folder_paths):
|
|
"""Only files downloaded by this extension have verifiable sidecar metadata."""
|
|
roots = set()
|
|
for folder in set(MODEL_TYPES.values()):
|
|
try:
|
|
roots.update(Path(root).resolve() for root in folder_paths.get_folder_paths(folder))
|
|
except KeyError:
|
|
continue
|
|
models = []
|
|
for root in sorted(roots):
|
|
for sidecar in root.rglob("*.n-sidebar.json"):
|
|
try:
|
|
if not sidecar.resolve().is_relative_to(root) or sidecar.stat().st_size > 65536:
|
|
continue
|
|
metadata = json.loads(sidecar.read_text(encoding="utf-8"))
|
|
model = sidecar.with_name(sidecar.name.removesuffix(".n-sidebar.json"))
|
|
if not model.resolve().is_relative_to(root):
|
|
continue
|
|
expected = metadata.get("size")
|
|
status = "missing" if not model.is_file() else "valid size" if model.stat().st_size == expected else "size mismatch"
|
|
models.append({"id": hashlib.sha256(str(sidecar).encode()).hexdigest()[:24],
|
|
"name": model.name, "path": str(model), "status": status,
|
|
"model": metadata.get("model", {}).get("name", ""),
|
|
"sha256": metadata.get("sha256", "")})
|
|
except (OSError, ValueError, TypeError, AttributeError):
|
|
continue
|
|
return models
|
|
|
|
|
|
def model_digest(path):
|
|
digest = hashlib.sha256()
|
|
with Path(path).open("rb") as model:
|
|
for chunk in iter(lambda: model.read(1024 * 1024), b""):
|
|
digest.update(chunk)
|
|
return digest.hexdigest()
|
|
|
|
|
|
class Downloads:
|
|
def __init__(self, folder_paths):
|
|
self.folder_paths = folder_paths
|
|
self.jobs = {}
|
|
self.tasks = {}
|
|
self.start_lock = asyncio.Lock()
|
|
|
|
def owner(self, request):
|
|
import server
|
|
return server.PromptServer.instance.user_manager.get_request_user_id(request)
|
|
|
|
def public(self, job):
|
|
return {key: value for key, value in job.items() if key not in {"owner", "gate", "token"}}
|
|
|
|
async def version(self, version_id, token=""):
|
|
if not isinstance(version_id, int) or isinstance(version_id, bool) or version_id <= 0:
|
|
raise ValueError("Use a positive Civitai model version ID")
|
|
headers = {"Authorization": f"Bearer {token}"} if token else {}
|
|
async with aiohttp.ClientSession(headers=headers, timeout=aiohttp.ClientTimeout(total=60)) as session:
|
|
async with session.get(f"https://civitai.com/api/v1/model-versions/{version_id}") as response:
|
|
response.raise_for_status()
|
|
data = await response.json()
|
|
files = [file for file in data.get("files", []) if file.get("type") == "Model"]
|
|
return {
|
|
"id": data["id"], "name": data.get("name", ""),
|
|
"model": data.get("model", {}), "baseModel": data.get("baseModel", ""),
|
|
"files": [{key: file.get(key) for key in ("id", "name", "sizeKB", "hashes", "primary", "downloadUrl")} for file in files],
|
|
}
|
|
|
|
def destination(self, model_type, name, selected=None):
|
|
folder = MODEL_TYPES.get(model_type)
|
|
if folder is None:
|
|
raise ValueError("Unsupported Civitai model type")
|
|
paths = self.folder_paths.get_folder_paths(folder)
|
|
if not paths:
|
|
raise ValueError(f"No ComfyUI model directory configured for {folder}")
|
|
local = Path(self.folder_paths.models_dir) / folder
|
|
if selected is not None:
|
|
if not isinstance(selected, int) or isinstance(selected, bool) or not 0 <= selected < len(paths):
|
|
raise ValueError("Unknown model destination")
|
|
root = Path(paths[selected])
|
|
else:
|
|
root = local if str(local) in paths else Path(paths[0])
|
|
root.mkdir(parents=True, exist_ok=True)
|
|
destination = root / safe_filename(name)
|
|
if destination.exists() or destination.is_symlink():
|
|
raise ValueError("A file with this name already exists")
|
|
return destination
|
|
|
|
async def start(self, payload, owner):
|
|
async with self.start_lock:
|
|
return await self._start(payload, owner)
|
|
|
|
async def _start(self, payload, owner):
|
|
if not isinstance(payload, dict):
|
|
raise ValueError("Invalid download request")
|
|
if any(job["status"] in {"queued", "running", "paused"} for job in self.jobs.values()):
|
|
raise ValueError("Finish or cancel the active download first")
|
|
token = os.environ.get("CIVITAI_TOKEN", "")
|
|
version = await self.version(payload.get("version"), token)
|
|
file = next((f for f in version["files"] if f["id"] == payload.get("file")), None)
|
|
if file is None:
|
|
raise ValueError("Choose a file from this model version")
|
|
url = urlsplit(file.get("downloadUrl") or "")
|
|
if url.scheme != "https" or url.netloc != "civitai.com" or url.path != f"/api/download/models/{version['id']}":
|
|
raise ValueError("Invalid Civitai file download URL")
|
|
destination = self.destination(version["model"].get("type"), file["name"], payload.get("destination"))
|
|
job_id = uuid.uuid4().hex
|
|
job = {
|
|
"id": job_id, "owner": owner, "name": file["name"], "status": "queued",
|
|
"received": 0, "total": 0, "error": "", "gate": asyncio.Event(),
|
|
}
|
|
job["gate"].set()
|
|
self.jobs[job_id] = job
|
|
for old_id in list(self.jobs)[:-20]:
|
|
if old_id not in self.tasks:
|
|
self.jobs.pop(old_id)
|
|
self.tasks[job_id] = asyncio.create_task(self.run(job, destination, version, file, token))
|
|
def finished(task):
|
|
if task.cancelled():
|
|
job["status"] = "cancelled"
|
|
self.tasks.pop(job_id, None)
|
|
self.tasks[job_id].add_done_callback(finished)
|
|
return self.public(job)
|
|
|
|
async def run(self, job, destination, version, file, token):
|
|
temporary = destination.with_name(destination.name + f".{job['id']}.part")
|
|
digest = hashlib.sha256()
|
|
headers = {"Authorization": f"Bearer {token}"} if token else {}
|
|
try:
|
|
if job["status"] == "queued":
|
|
job["status"] = "running"
|
|
timeout = aiohttp.ClientTimeout(total=None, connect=60, sock_read=120)
|
|
async with aiohttp.ClientSession(timeout=timeout) as session:
|
|
async with session.get(file["downloadUrl"], headers=headers) as response:
|
|
response.raise_for_status()
|
|
job["total"] = int(response.headers.get("Content-Length", 0))
|
|
with temporary.open("xb") as output:
|
|
async for chunk in response.content.iter_chunked(1024 * 1024):
|
|
await job["gate"].wait()
|
|
write = asyncio.create_task(asyncio.to_thread(output.write, chunk))
|
|
try:
|
|
await asyncio.shield(write)
|
|
except asyncio.CancelledError:
|
|
await write
|
|
raise
|
|
digest.update(chunk)
|
|
job["received"] += len(chunk)
|
|
if not job["received"]:
|
|
raise ValueError("Empty model download")
|
|
if job["total"] and job["received"] != job["total"]:
|
|
raise ValueError("Incomplete download")
|
|
expected = (file.get("hashes") or {}).get("SHA256")
|
|
if expected and digest.hexdigest().lower() != expected.lower():
|
|
raise ValueError("SHA256 verification failed")
|
|
os.link(temporary, destination)
|
|
temporary.unlink()
|
|
metadata = {"version": version["id"], "model": version["model"],
|
|
"baseModel": version["baseModel"], "size": job["received"],
|
|
"sha256": digest.hexdigest()}
|
|
destination.with_name(destination.name + ".n-sidebar.json").write_text(json.dumps(metadata, indent=2))
|
|
job["status"] = "completed"
|
|
except asyncio.CancelledError:
|
|
job["status"] = "cancelled"
|
|
except (aiohttp.ClientError, OSError, ValueError, asyncio.TimeoutError) as error:
|
|
job["status"] = "failed"
|
|
job["error"] = str(error)
|
|
finally:
|
|
temporary.unlink(missing_ok=True)
|
|
self.tasks.pop(job["id"], None)
|
|
|
|
|
|
def register_routes():
|
|
import folder_paths
|
|
import server
|
|
|
|
routes = server.PromptServer.instance.routes
|
|
downloads = Downloads(folder_paths)
|
|
|
|
@routes.get("/n-sidebar/legacy")
|
|
async def legacy(request):
|
|
data = legacy_settings()
|
|
keys = ("sb_pinnedItems", "sb_categoryNodeMap", "sb_ColorCustomCategories",
|
|
"sb_templateNodeMap", "sb_ColorCustomTemplates", "sb_workflowNodeMap", "sb_ColorCustomWorkflows")
|
|
return web.json_response({key: data[key] for key in keys if key in data})
|
|
|
|
@routes.get("/n-sidebar/workflows")
|
|
async def workflows(request):
|
|
roots = workflow_roots()
|
|
files = await asyncio.to_thread(list_workflows, roots)
|
|
return web.json_response({"roots": [str(root) for root in roots], "files": files})
|
|
|
|
@routes.get("/n-sidebar/workflow")
|
|
async def workflow(request):
|
|
try:
|
|
roots = workflow_roots()
|
|
index = int(request.query.get("root", "-1"))
|
|
if not 0 <= index < len(roots):
|
|
raise ValueError("Unknown workflow library")
|
|
path = workflow_path(roots[index], request.query.get("path", ""))
|
|
data = await asyncio.to_thread(read_workflow, path)
|
|
return web.json_response(data)
|
|
except (ValueError, OSError) as error:
|
|
return web.json_response({"error": str(error)}, status=400)
|
|
|
|
@routes.get("/n-sidebar/civitai/search")
|
|
async def search(request):
|
|
params = {"limit": "12", "nsfw": "false"}
|
|
for key in ("query", "types", "page"):
|
|
if request.query.get(key):
|
|
params[key] = request.query[key]
|
|
try:
|
|
token = os.environ.get("CIVITAI_TOKEN", "")
|
|
headers = {"Authorization": f"Bearer {token}"} if token else {}
|
|
async with aiohttp.ClientSession(headers=headers, timeout=aiohttp.ClientTimeout(total=60)) as session:
|
|
async with session.get("https://civitai.com/api/v1/models", params=params) as response:
|
|
response.raise_for_status()
|
|
data = await response.json()
|
|
return web.json_response(data)
|
|
except (aiohttp.ClientError, ValueError, asyncio.TimeoutError) as error:
|
|
return web.json_response({"error": str(error)}, status=502)
|
|
|
|
@routes.get("/n-sidebar/civitai/version/{version}")
|
|
async def version(request):
|
|
try:
|
|
data = await downloads.version(int(request.match_info["version"]), os.environ.get("CIVITAI_TOKEN", ""))
|
|
data["files"] = [{key: value for key, value in file.items() if key != "downloadUrl"} for file in data["files"]]
|
|
return web.json_response(data)
|
|
except (ValueError, aiohttp.ClientError, asyncio.TimeoutError) as error:
|
|
return web.json_response({"error": str(error)}, status=400)
|
|
|
|
@routes.get("/n-sidebar/model-paths")
|
|
async def model_paths(request):
|
|
folder = MODEL_TYPES.get(request.query.get("type"))
|
|
if folder is None:
|
|
return web.json_response({"error": "Unsupported model type"}, status=400)
|
|
try:
|
|
paths = folder_paths.get_folder_paths(folder)
|
|
except KeyError:
|
|
paths = []
|
|
local = str(Path(folder_paths.models_dir) / folder)
|
|
return web.json_response({"paths": paths, "default": paths.index(local) if local in paths else 0})
|
|
|
|
@routes.get("/n-sidebar/downloads")
|
|
async def progress(request):
|
|
owner = downloads.owner(request)
|
|
return web.json_response([downloads.public(job) for job in downloads.jobs.values() if job["owner"] == owner])
|
|
|
|
@routes.get("/n-sidebar/models")
|
|
async def models(request):
|
|
entries = await asyncio.to_thread(managed_models, folder_paths)
|
|
return web.json_response([{key: value for key, value in entry.items() if key != "path"} for entry in entries])
|
|
|
|
@routes.post("/n-sidebar/models/{id}/verify")
|
|
async def verify(request):
|
|
entries = await asyncio.to_thread(managed_models, folder_paths)
|
|
model = next((entry for entry in entries if entry["id"] == request.match_info["id"]), None)
|
|
if model is None:
|
|
raise web.HTTPNotFound()
|
|
if model["status"] == "missing":
|
|
return web.json_response({"status": "missing"})
|
|
try:
|
|
digest = await asyncio.to_thread(model_digest, model["path"])
|
|
return web.json_response({"status": "verified" if digest == model["sha256"] else "hash mismatch"})
|
|
except OSError as error:
|
|
return web.json_response({"error": str(error)}, status=400)
|
|
|
|
@routes.post("/n-sidebar/downloads")
|
|
async def download(request):
|
|
try:
|
|
return web.json_response(await downloads.start(await request.json(), downloads.owner(request)), status=202)
|
|
except (ValueError, OSError, aiohttp.ClientError, asyncio.TimeoutError) as error:
|
|
return web.json_response({"error": str(error)}, status=400)
|
|
|
|
@routes.post("/n-sidebar/downloads/{id}/{action}")
|
|
async def control(request):
|
|
job = downloads.jobs.get(request.match_info["id"])
|
|
if job is None or job["owner"] != downloads.owner(request):
|
|
raise web.HTTPNotFound()
|
|
action = request.match_info["action"]
|
|
if action not in {"pause", "resume", "cancel"} or job["status"] not in {"queued", "running", "paused"}:
|
|
return web.json_response({"error": "Invalid download action"}, status=400)
|
|
if action == "cancel":
|
|
downloads.tasks[job["id"]].cancel()
|
|
elif action == "pause":
|
|
job["gate"].clear()
|
|
job["status"] = "paused"
|
|
else:
|
|
job["gate"].set()
|
|
job["status"] = "running"
|
|
return web.json_response(downloads.public(job))
|