from __future__ import annotations import json import re from dataclasses import asdict, dataclass from datetime import datetime, timezone from pathlib import Path from typing import Any from uuid import uuid4 def _now() -> str: return datetime.now(timezone.utc).isoformat() def _normal(value: str) -> str: return re.sub(r"[^a-z0-9]+", " ", value.casefold()).strip() @dataclass(slots=True) class Asset: id: str kind: str name: str path: str trainer: str = "" dataset_id: str = "" checkpoint: str = "" epochs: int = 0 created_at: str = "" @classmethod def from_dict(cls, payload: dict[str, Any]) -> "Asset": return cls( id=str(payload.get("id") or uuid4().hex[:12]), kind=str(payload.get("kind", "")), name=str(payload.get("name", "")), path=str(payload.get("path", "")), trainer=str(payload.get("trainer", "")), dataset_id=str(payload.get("dataset_id", "")), checkpoint=str(payload.get("checkpoint", "")), epochs=int(payload.get("epochs", 0) or 0), created_at=str(payload.get("created_at") or _now()), ) class AssetRegistry: """Persistent friendly-name index for datasets, models, and checkpoints.""" def __init__(self, root: Path) -> None: self.path = root.resolve() / "data" / "assets.json" self.assets: list[Asset] = [] self.load() def load(self) -> None: try: payload = json.loads(self.path.read_text(encoding="utf-8")) self.assets = [ Asset.from_dict(item) for item in payload.get("assets", []) if isinstance(item, dict) ] except (OSError, ValueError, TypeError, json.JSONDecodeError): self.assets = [] def save(self) -> None: self.path.parent.mkdir(parents=True, exist_ok=True) temporary = self.path.with_suffix(".tmp") temporary.write_text( json.dumps({"assets": [asdict(item) for item in self.assets]}, indent=2), encoding="utf-8", ) temporary.replace(self.path) def register( self, *, kind: str, name: str, path: str, trainer: str = "", dataset_id: str = "", checkpoint: str = "", epochs: int = 0, persist: bool = True, ) -> Asset: resolved = str(Path(path).expanduser().resolve()) existing = next( ( item for item in self.assets if item.kind == kind and Path(item.path) == Path(resolved) ), None, ) asset = existing or Asset(uuid4().hex[:12], kind, name, resolved) asset.name = name.strip() or Path(resolved).name asset.trainer = trainer asset.dataset_id = dataset_id asset.checkpoint = checkpoint asset.epochs = int(epochs) asset.created_at = asset.created_at or _now() if existing is None: self.assets.insert(0, asset) if persist: self.save() return asset def ingest_result(self, result: dict[str, Any]) -> None: entries = result.get("assets", []) if not isinstance(entries, list): return for item in entries: if not isinstance(item, dict): continue if item.get("kind") and item.get("path"): values = { key: item[key] for key in ( "kind", "name", "path", "trainer", "dataset_id", "checkpoint", "epochs", ) if key in item } dataset_path = str(item.get("dataset_path", "")) if item.get("kind") == "model" and dataset_path and Path(dataset_path).is_dir(): dataset = self.register( kind="dataset", name=Path(dataset_path).name, path=dataset_path, ) values["dataset_id"] = dataset.id values.setdefault("name", Path(str(item["path"])).name) self.register(**values) def find(self, kind: str, query: str, *, trainer: str = "") -> list[Asset]: wanted = _normal(query) matches = [] exact = [] for item in self.assets: if item.kind != kind or (trainer and item.trainer != trainer): continue haystacks = {_normal(item.name), _normal(Path(item.path).name)} if wanted in haystacks: exact.append(item) elif any(wanted and wanted in value for value in haystacks): matches.append(item) return exact or matches def discover(self, config: Any) -> None: folders = config.get("tool_folders", {}) if not isinstance(folders, dict): return app_root = self.path.parent.parent external_lora_root = app_root / "LoRAModelsHere" if external_lora_root.is_dir(): for path in external_lora_root.rglob("*.safetensors"): if path.is_file() and "_comfy" not in path.stem.casefold(): self.register( kind="model", name=path.stem.removesuffix("_cancelled"), path=str(path), trainer="lora", checkpoint=str(path), persist=False, ) base_model_root = app_root / "LoRA StableDiffusionModels Here" if base_model_root.is_dir(): for path in base_model_root.iterdir(): is_model_file = path.is_file() and path.suffix.casefold() in { ".safetensors", ".ckpt", ".pt", ".bin" } is_diffusers_folder = path.is_dir() and ( (path / "model_index.json").is_file() or (path / "unet" / "config.json").is_file() ) if is_model_file or is_diffusers_folder: self.register( kind="base_model", name=path.stem if path.is_file() else path.name, path=str(path), trainer="stable_diffusion", persist=False, ) flow_datasets = self._flow_dataset_paths() collector = Path(str(folders.get("dataset_collector", ""))) / "Datasets" if collector.is_dir(): for folder in collector.iterdir(): if folder.is_dir(): self.register( kind="dataset", name=folder.name, path=str(folder), persist=False ) for trainer, folder_name, output_name in ( ("ddpm", "ddpm_trainer", "output"), ("lora", "lora_trainer", "output"), ("flow", "flow_trainer", "output_flow_models"), ): root = Path(str(folders.get(folder_name, ""))) / output_name if not root.is_dir(): continue for folder in root.iterdir(): if not folder.is_dir(): continue name = folder.name dataset_path = "" if trainer == "ddpm": # DDPM writes a durable sidecar with the friendly model name and # source dataset. Prefer it over a filesystem-safe folder name. try: metadata = json.loads((folder / "model_info.json").read_text(encoding="utf-8")) name = str(metadata.get("model_name") or metadata.get("name") or name) dataset_path = str(metadata.get("dataset_dir") or "") except (OSError, ValueError, TypeError, json.JSONDecodeError): pass checkpoints = sorted( folder.glob("checkpoint-*"), key=lambda p: int(p.name.rsplit("-", 1)[-1]) if p.name.rsplit("-", 1)[-1].isdigit() else -1, ) elif trainer == "lora": checkpoints = sorted( ( path for path in folder.glob("*.safetensors") if "_comfy" not in path.stem.casefold() ), key=lambda p: p.stat().st_mtime, ) if checkpoints: name = checkpoints[-1].stem.removesuffix("_cancelled") else: checkpoints = [] try: metadata = json.loads( (folder / "flow_model_info.json").read_text(encoding="utf-8") ) if metadata.get("model_type") != "rectified_flow": continue if not (folder / "unet" / "config.json").is_file(): continue name = str(metadata.get("model_name") or metadata.get("name") or name) dataset_path = flow_datasets.get(str(folder.resolve()), "") except (OSError, ValueError, TypeError, json.JSONDecodeError): continue checkpoint = ( str(folder) if trainer == "flow" else str(checkpoints[-1]) if checkpoints else "" ) dataset_id = "" if dataset_path and Path(dataset_path).is_dir(): dataset_asset = self.register( kind="dataset", name=Path(dataset_path).name, path=dataset_path, persist=False, ) dataset_id = dataset_asset.id self.register( kind="model", name=name, path=str(folder), trainer=trainer, dataset_id=dataset_id, checkpoint=checkpoint, persist=False, ) self.save() def _flow_dataset_paths(self) -> dict[str, str]: """Recover source datasets for Flow models created by ADAM in older runs.""" jobs_path = self.path.parent / "jobs.json" try: payload = json.loads(jobs_path.read_text(encoding="utf-8")) jobs = payload.get("jobs", []) except (OSError, ValueError, TypeError, json.JSONDecodeError): return {} links: dict[str, str] = {} if not isinstance(jobs, list): return links for job in jobs: if not isinstance(job, dict): continue plan = job.get("plan", {}) steps = plan.get("steps", []) if isinstance(plan, dict) else [] if not isinstance(steps, list): continue for step in steps: if not isinstance(step, dict) or step.get("tool_id") != "flow_trainer": continue arguments = step.get("arguments", {}) if not isinstance(arguments, dict): continue output = str(arguments.get("output_dir", "")) dataset = str(arguments.get("dataset_dir", "")) if not output or not dataset or not Path(dataset).is_dir(): continue try: links[str(Path(output).expanduser().resolve())] = str(Path(dataset).expanduser().resolve()) except OSError: continue return links