Files
gpu-rent/src/gpu_rent/capture.py
T
Leonid Pershin 618e6e4806 Add Hugging Face support and enhance model resolution logic
- Introduced support for Hugging Face API integration, allowing fallback model resolution when Civitai fails.
- Updated configuration to include `HF_TOKEN` and `HF_TOKEN_PATH` for authentication.
- Enhanced model capture logic to differentiate between Civitai and Hugging Face sources.
- Improved error handling for model downloads, providing clearer messages for authentication issues.
- Updated documentation to reflect new environment variables and usage instructions for Hugging Face integration.
- Added tests to validate the new fallback mechanism and ensure robust model resolution.
2026-08-21 07:34:15 +03:00

690 lines
24 KiB
Python

"""Capture VM inventory into local models.yaml / extensions.yaml (links only)."""
from __future__ import annotations
import json
import shutil
from collections.abc import Callable
from dataclasses import dataclass, field
from importlib.resources import files
from pathlib import Path
from typing import Any
from urllib.parse import urlsplit, urlunsplit
import yaml
from gpu_rent.civitai import (
civitai_model_url,
fetch_model_version,
fetch_model_version_by_hash,
fetch_model_versions_by_hashes,
version_ids_from_payload,
)
from gpu_rent.config import Config
from gpu_rent.errors import CloudError, GpuRentError
from gpu_rent.huggingface import lookup_by_sha256
from gpu_rent.manifests import (
MODEL_TYPES,
extract_version_id,
parse_extensions,
parse_models,
)
from gpu_rent.ssh_ops import run_python, run_ssh
Log = Callable[[str], None]
INVENTORY_REMOTE = "/tmp/gpu-rent-inventory.json"
def _pkg_text(name: str) -> str:
return files("gpu_rent.remote").joinpath(name).read_text(encoding="utf-8")
def strip_git_auth(url: str) -> str:
parts = urlsplit(url.strip())
host = parts.hostname or ""
if parts.port:
host = f"{host}:{parts.port}"
return urlunsplit((parts.scheme, host, parts.path, parts.query, parts.fragment))
@dataclass
class ModelCaptureItem:
kind: str
url: str
title: str = ""
rel: str = ""
version_id: int | None = None
model_id: int | None = None
source: str = "civitai" # civitai | huggingface
@dataclass
class ExtCaptureItem:
kind: str
url: str
ref: str
directory: str
@dataclass
class ResolveOutcome:
"""ok → item; unknown → not on Civitai; api_error → network/HTTP (retry later)."""
item: ModelCaptureItem | None = None
status: str = "unknown" # ok | unknown | api_error
detail: str = ""
@dataclass
class CaptureReport:
models_new: list[ModelCaptureItem] = field(default_factory=list)
models_skip: list[str] = field(default_factory=list)
models_unknown: list[str] = field(default_factory=list)
models_api_errors: list[str] = field(default_factory=list)
ext_new: list[ExtCaptureItem] = field(default_factory=list)
ext_updated: list[ExtCaptureItem] = field(default_factory=list)
ext_skip: list[str] = field(default_factory=list)
ext_unknown: list[str] = field(default_factory=list)
models_path: Path | None = None
extensions_path: Path | None = None
wrote_models: bool = False
wrote_extensions: bool = False
def fetch_inventory(cfg: Config, host: str, log: Log) -> dict[str, Any]:
log("capture: сканирую Models / Extensions / DLNodes на VM…")
run_python(
cfg,
host,
_pkg_text("scan_inventory.py"),
remote_path="/tmp/gpu-rent-scan_inventory.py",
timeout=3600,
log=log,
)
raw = run_ssh(cfg, host, f"cat {INVENTORY_REMOTE}", timeout=60)
try:
data = json.loads(raw)
except json.JSONDecodeError as exc:
raise CloudError(f"не разобрать inventory JSON: {exc}") from exc
if not isinstance(data, dict):
raise CloudError("inventory: ожидался object")
return data
def resolve_model_item(
raw: dict,
*,
token: str,
api_host: str,
link_host: str,
hf_token: str | None = None,
) -> ResolveOutcome:
kind = str(raw.get("kind") or "")
if kind not in MODEL_TYPES:
return ResolveOutcome(status="unknown", detail="bad kind")
rel = str(raw.get("rel") or raw.get("name") or "?")
version_id = raw.get("version_id")
model_id = raw.get("model_id")
title = Path(str(raw.get("name") or rel)).stem
try:
vid = int(version_id) if version_id is not None else None
except (TypeError, ValueError):
vid = None
try:
mid = int(model_id) if model_id is not None else None
except (TypeError, ValueError):
mid = None
def _ok_civitai(v: int, m: int, name: str = title) -> ResolveOutcome:
return ResolveOutcome(
item=ModelCaptureItem(
kind=kind,
url=civitai_model_url(m, v, link_host),
title=name,
rel=rel,
version_id=v,
model_id=m,
source="civitai",
),
status="ok",
)
def _try_hf(sha: str, detail_prefix: str) -> ResolveOutcome | None:
"""Civitai miss → Hugging Face search by filename + LFS sha."""
try:
hit = lookup_by_sha256(
hf_token,
sha,
filename=str(raw.get("name") or rel),
)
except CloudError as exc:
return ResolveOutcome(
status="api_error",
detail=f"{detail_prefix} HF: {exc}",
)
if hit is None:
return None
return ResolveOutcome(
item=ModelCaptureItem(
kind=kind,
url=hit.url,
title=hit.title or title,
rel=rel,
version_id=None,
model_id=None,
source="huggingface",
),
status="ok",
)
if vid is not None and mid is not None:
return _ok_civitai(vid, mid)
# Partial sidecar: have version id → GET /model-versions/{id} for modelId.
if vid is not None and mid is None:
try:
_h, version = fetch_model_version(token or "", api_host, vid)
except CloudError as exc:
msg = str(exc)
if "404" in msg or "пустой files" in msg:
return ResolveOutcome(
status="unknown",
detail=f"{rel} version_id={vid} ({msg})",
)
return ResolveOutcome(
status="api_error",
detail=f"{rel} version_id={vid} {msg}",
)
vid2, mid2 = version_ids_from_payload(version)
if vid2 is not None and mid2 is not None:
name = str(version.get("name") or title)
return _ok_civitai(vid2, mid2, name)
return ResolveOutcome(
status="unknown",
detail=f"{rel} version_id={vid} (нет modelId в ответе)",
)
sha = raw.get("sha256")
if not sha:
return ResolveOutcome(
status="unknown",
detail=f"{rel} (нет sidecar и нет sha256)",
)
try:
_host, version = fetch_model_version_by_hash(token or None, api_host, str(sha))
except CloudError as exc:
msg = str(exc)
# 404 = genuinely not on Civitai → try HF
if "HTTP 404" in msg or msg.rstrip().endswith("404"):
hf_out = _try_hf(str(sha), f"{rel} sha={str(sha)[:12]}…")
if hf_out is not None:
return hf_out
return ResolveOutcome(
status="unknown",
detail=f"{rel} sha={str(sha)[:12]}… (нет на Civitai/HF)",
)
return ResolveOutcome(
status="api_error",
detail=f"{rel} sha={str(sha)[:12]}{msg}",
)
vid2, mid2 = version_ids_from_payload(version)
if vid2 is None or mid2 is None:
hf_out = _try_hf(str(sha), f"{rel} sha={str(sha)[:12]}…")
if hf_out is not None:
return hf_out
return ResolveOutcome(
status="unknown",
detail=f"{rel} sha={str(sha)[:12]}… (пустой payload)",
)
name = str(version.get("name") or title)
return _ok_civitai(vid2, mid2, name)
def _backup(path: Path) -> None:
if path.is_file() and path.stat().st_size > 0:
shutil.copy2(path, path.with_suffix(path.suffix + ".bak"))
def _keep_model_entry(it: dict) -> bool:
if not isinstance(it, dict):
return False
url = it.get("url")
if url:
return True
vid = it.get("version_id")
return vid not in (None, "", 0, "0")
def _model_dedupe_key(kind: str, *, version_id: int | None, url: str | None) -> tuple:
u = (url or "").rstrip("/").lower()
if u and ("huggingface.co" in u or "hf.co/" in u):
return ("hf", kind, u)
if version_id is not None:
return ("civitai", kind, int(version_id))
return ("url", kind, u or "?")
def merge_models_yaml(
path: Path,
new_items: list[ModelCaptureItem],
*,
dry_run: bool,
) -> tuple[list[ModelCaptureItem], list[str]]:
"""Return (actually_new, skip_msgs). Dedupe by Civitai version_id or HF url."""
existing = parse_models(path) if path.is_file() else []
have: set[tuple] = set()
for e in existing:
vid = e.version_id
if vid is None and e.url:
vid = extract_version_id(e.url)
have.add(_model_dedupe_key(e.kind, version_id=vid, url=e.url))
added: list[ModelCaptureItem] = []
skipped: list[str] = []
seen_new: set[tuple] = set()
for item in new_items:
key = _model_dedupe_key(item.kind, version_id=item.version_id, url=item.url)
if key in have or key in seen_new:
label = (
f"modelVersionId={item.version_id}"
if item.version_id is not None
else item.url
)
skipped.append(f"{item.kind} {item.title} {label}")
continue
seen_new.add(key)
added.append(item)
if dry_run or not added:
return added, skipped
data: dict[str, list[dict[str, str]]] = {k: [] for k in MODEL_TYPES}
if path.is_file():
raw = yaml.safe_load(path.read_text(encoding="utf-8")) or {}
if isinstance(raw, dict):
for kind in MODEL_TYPES:
items = raw.get(kind) or []
if isinstance(items, list):
for it in items:
if _keep_model_entry(it):
data[kind].append(dict(it))
for item in added:
data[item.kind].append({"url": item.url})
out = {k: v for k, v in data.items() if v}
_backup(path)
path.parent.mkdir(parents=True, exist_ok=True)
path.write_text(
yaml.safe_dump(out, allow_unicode=True, default_flow_style=False, sort_keys=False),
encoding="utf-8",
)
return added, skipped
def merge_extensions_yaml(
path: Path,
new_items: list[ExtCaptureItem],
*,
dry_run: bool,
) -> tuple[list[ExtCaptureItem], list[ExtCaptureItem], list[str]]:
"""Return (added, updated, skipped). Same dir + different URL → update url/ref."""
data: dict[str, list[dict[str, Any]]] = {"swarmui": [], "comfy": []}
if path.is_file():
raw = yaml.safe_load(path.read_text(encoding="utf-8")) or {}
if isinstance(raw, dict):
for kind in ("swarmui", "comfy"):
items = raw.get(kind) or []
if isinstance(items, list):
for it in items:
if isinstance(it, dict) and it.get("url"):
data[kind].append(dict(it))
def _dir_of(it: dict) -> str:
d = it.get("dir")
if d:
return str(d).lower()
url = strip_git_auth(str(it.get("url") or ""))
name = url.rstrip("/").rsplit("/", 1)[-1]
if name.endswith(".git"):
name = name[:-4]
return name.lower()
added: list[ExtCaptureItem] = []
updated: list[ExtCaptureItem] = []
skipped: list[str] = []
for item in new_items:
clean = strip_git_auth(item.url)
key_url = (item.kind, clean.rstrip("/").lower())
key_dir = (item.kind, item.directory.lower())
bucket = data.setdefault(item.kind, [])
matched_url = False
matched_dir_idx: int | None = None
for idx, it in enumerate(bucket):
it_url = strip_git_auth(str(it.get("url") or "")).rstrip("/").lower()
if (item.kind, it_url) == key_url:
matched_url = True
break
if (item.kind, _dir_of(it)) == key_dir:
matched_dir_idx = idx
if matched_url:
skipped.append(f"{item.kind} {item.directory} {clean}")
continue
if matched_dir_idx is not None:
old = bucket[matched_dir_idx]
old_url = strip_git_auth(str(old.get("url") or ""))
old_ref = str(old.get("ref") or "main")
if old_url.rstrip("/").lower() == clean.rstrip("/").lower() and old_ref == item.ref:
skipped.append(f"{item.kind} {item.directory} {clean}")
continue
old["url"] = clean
old["ref"] = item.ref
if not old.get("dir"):
old["dir"] = item.directory
updated.append(item)
continue
bucket.append({"url": clean, "ref": item.ref, "dir": item.directory})
added.append(item)
if dry_run or (not added and not updated):
return added, updated, skipped
out = {k: v for k, v in data.items() if v}
_backup(path)
path.parent.mkdir(parents=True, exist_ok=True)
path.write_text(
yaml.safe_dump(out, allow_unicode=True, default_flow_style=False, sort_keys=False),
encoding="utf-8",
)
return added, updated, skipped
def capture_models(
cfg: Config,
host: str,
inventory: dict[str, Any] | None,
*,
dry_run: bool,
kind_filter: str | None,
log: Log,
) -> CaptureReport:
report = CaptureReport(models_path=cfg.models_manifest)
inv = inventory if inventory is not None else fetch_inventory(cfg, host, log)
raw_models = inv.get("models") or []
if not isinstance(raw_models, list):
raise CloudError("inventory.models: ожидался list")
resolved: list[ModelCaptureItem] = []
need_hash: list[dict] = []
link_host = cfg.civitai_api_host or "civitai.red"
token = cfg.civitai_api_token
api_host = cfg.civitai_api_host
hf_token = cfg.hf_token or None
for raw in raw_models:
if not isinstance(raw, dict):
continue
if kind_filter and str(raw.get("kind")) != kind_filter:
continue
vid = raw.get("version_id")
mid = raw.get("model_id")
try:
vid_i = int(vid) if vid is not None else None
except (TypeError, ValueError):
vid_i = None
try:
mid_i = int(mid) if mid is not None else None
except (TypeError, ValueError):
mid_i = None
sha = raw.get("sha256")
# Full sidecar / ids → no API. Partial vid → GET version. Else batch by-hash.
if vid_i is not None and mid_i is not None:
outcome = resolve_model_item(
raw,
token=token,
api_host=api_host,
link_host=link_host,
hf_token=hf_token,
)
elif vid_i is not None:
outcome = resolve_model_item(
raw,
token=token,
api_host=api_host,
link_host=link_host,
hf_token=hf_token,
)
elif sha:
need_hash.append(raw)
continue
else:
outcome = resolve_model_item(
raw,
token=token,
api_host=api_host,
link_host=link_host,
hf_token=hf_token,
)
if outcome.status == "ok" and outcome.item is not None:
resolved.append(outcome.item)
elif outcome.status == "api_error":
report.models_api_errors.append(outcome.detail or "?")
else:
report.models_unknown.append(outcome.detail or str(raw.get("rel") or "?"))
if need_hash:
digests = [str(r["sha256"]).strip().lower() for r in need_hash if r.get("sha256")]
log(f"capture: by-hash batch {len(digests)} файл(ов)…")
try:
by_hash = fetch_model_versions_by_hashes(token or None, api_host, digests)
except CloudError as exc:
report.models_api_errors.append(f"by-hash batch: {exc}")
by_hash = {}
for raw in need_hash:
rel = str(raw.get("rel") or "?")
sha = str(raw.get("sha256") or "")[:12]
report.models_api_errors.append(f"{rel} sha={sha}… (batch failed)")
need_hash = []
hf_fallback = 0
for raw in need_hash:
sha = str(raw.get("sha256") or "").strip().lower()
rel = str(raw.get("rel") or raw.get("name") or "?")
version = by_hash.get(sha)
if version:
vid2, mid2 = version_ids_from_payload(version)
if vid2 is not None and mid2 is not None:
kind = str(raw.get("kind") or "")
if kind not in MODEL_TYPES:
continue
title = str(
version.get("name") or Path(str(raw.get("name") or rel)).stem
)
resolved.append(
ModelCaptureItem(
kind=kind,
url=civitai_model_url(mid2, vid2, link_host),
title=title,
rel=rel,
version_id=vid2,
model_id=mid2,
source="civitai",
)
)
continue
# Civitai miss → Hugging Face
try:
hit = lookup_by_sha256(
hf_token,
sha,
filename=str(raw.get("name") or rel),
)
except CloudError as exc:
report.models_api_errors.append(f"{rel} HF: {exc}")
report.models_unknown.append(f"{rel} sha={sha[:12]}…")
continue
if hit:
kind = str(raw.get("kind") or "")
if kind not in MODEL_TYPES:
continue
hf_fallback += 1
resolved.append(
ModelCaptureItem(
kind=kind,
url=hit.url,
title=hit.title or Path(str(raw.get("name") or rel)).stem,
rel=rel,
source="huggingface",
)
)
else:
report.models_unknown.append(f"{rel} sha={sha[:12]}… (нет на Civitai/HF)")
if hf_fallback:
log(f"capture: Hugging Face fallback — {hf_fallback} файл(ов)")
added, skipped = merge_models_yaml(
cfg.models_manifest, resolved, dry_run=dry_run
)
report.models_new = added
report.models_skip = skipped
report.wrote_models = bool(added) and not dry_run
return report
def capture_extensions(
cfg: Config,
host: str,
inventory: dict[str, Any] | None,
*,
dry_run: bool,
log: Log,
) -> CaptureReport:
report = CaptureReport(extensions_path=cfg.extensions_manifest)
inv = inventory if inventory is not None else fetch_inventory(cfg, host, log)
raw_ext = inv.get("extensions") or []
if not isinstance(raw_ext, list):
raise CloudError("inventory.extensions: ожидался list")
resolved: list[ExtCaptureItem] = []
for raw in raw_ext:
if not isinstance(raw, dict):
continue
if raw.get("unknown") or not raw.get("url"):
name = str(raw.get("dir") or "?")
reason = str(raw.get("reason") or "unknown")
report.ext_unknown.append(f"{raw.get('kind')}/{name} ({reason})")
continue
kind = str(raw.get("kind") or "")
if kind not in ("swarmui", "comfy"):
continue
resolved.append(
ExtCaptureItem(
kind=kind,
url=strip_git_auth(str(raw["url"])),
ref=str(raw.get("ref") or "main"),
directory=str(raw.get("dir") or "extension"),
)
)
added, updated, skipped = merge_extensions_yaml(
cfg.extensions_manifest, resolved, dry_run=dry_run
)
report.ext_new = added
report.ext_updated = updated
report.ext_skip = skipped
report.wrote_extensions = (bool(added) or bool(updated)) and not dry_run
return report
def merge_reports(a: CaptureReport, b: CaptureReport) -> CaptureReport:
return CaptureReport(
models_new=a.models_new + b.models_new,
models_skip=a.models_skip + b.models_skip,
models_unknown=a.models_unknown + b.models_unknown,
models_api_errors=a.models_api_errors + b.models_api_errors,
ext_new=a.ext_new + b.ext_new,
ext_updated=a.ext_updated + b.ext_updated,
ext_skip=a.ext_skip + b.ext_skip,
ext_unknown=a.ext_unknown + b.ext_unknown,
models_path=a.models_path or b.models_path,
extensions_path=a.extensions_path or b.extensions_path,
wrote_models=a.wrote_models or b.wrote_models,
wrote_extensions=a.wrote_extensions or b.wrote_extensions,
)
def capture_all(
cfg: Config,
host: str,
*,
dry_run: bool,
kind_filter: str | None = None,
log: Log,
) -> CaptureReport:
inv = fetch_inventory(cfg, host, log)
m = capture_models(
cfg, host, inv, dry_run=dry_run, kind_filter=kind_filter, log=log
)
e = capture_extensions(cfg, host, inv, dry_run=dry_run, log=log)
return merge_reports(m, e)
def print_report(
report: CaptureReport,
log: Log,
*,
dry_run: bool,
show_models: bool = True,
show_extensions: bool = True,
) -> None:
prefix = "[dry-run] " if dry_run else ""
if show_models:
log(
f"{prefix}models: +{len(report.models_new)} new, "
f"{len(report.models_skip)} already in yaml, "
f"{len(report.models_unknown)} unknown (no Civitai), "
f"{len(report.models_api_errors)} api errors"
)
for item in report.models_new:
log(f" + {item.kind} {item.title} modelVersionId={item.version_id}")
for line in report.models_unknown:
log(f" ? {line}")
for line in report.models_api_errors:
log(f" ! {line}")
if show_extensions:
log(
f"{prefix}extensions: +{len(report.ext_new)} new, "
f"~{len(report.ext_updated)} updated, "
f"{len(report.ext_skip)} already, "
f"{len(report.ext_unknown)} unknown"
)
for item in report.ext_new:
log(f" + {item.kind} {item.directory} {item.url} @{item.ref}")
for item in report.ext_updated:
log(f" ~ {item.kind} {item.directory} {item.url} @{item.ref}")
for line in report.ext_unknown:
log(f" ? {line}")
if dry_run:
log("dry-run — файлы не записаны")
return
if report.wrote_models and report.models_path:
log(f"wrote {report.models_path} (backup {report.models_path.name}.bak)")
if report.wrote_extensions and report.extensions_path:
log(
f"wrote {report.extensions_path} "
f"(backup {report.extensions_path.name}.bak)"
)
if not report.wrote_models and not report.wrote_extensions:
log("нечего добавлять")