SPB Git forge

spb/llm-api

Public
0commits 0branches 0releases
0 Bsize
maindefault branch
—last push
18.6 KB · 344 lines python
Raw Blame History
1"""Hugging Face inspection + downloads (with progress), model manifests, verification."""23from __future__ import annotations45import asyncio6import json7import logging8import os9import re10import shutil11import time12from pathlib import Path13from typing import Any1415from .config import Settings16from .errors import APIError, DownloadError, InsufficientDisk17from .jobs import Job, JobRunner18from .models import compat, formats19from .models.estimator import kv_bytes_per_token20from .models.scanner import slugify2122log = logging.getLogger("llm_api.downloads")23GB = 1024**32425REPO_RE = re.compile(r"^[A-Za-z0-9][A-Za-z0-9._-]{0,95}/[A-Za-z0-9][A-Za-z0-9._-]{0,127}$")26TRUSTED_GGUF_AUTHORS = {"unsloth", "bartowski", "ggml-org", "lmstudio-community", "Qwen", "google", "mistralai",27                        "microsoft", "TheBloke", "mradermacher", "QuantFactory", "nomic-ai", "BAAI", "jinaai",28                        "mixedbread-ai", "second-state", "openai", "deepseek-ai", "zai-org", "nvidia", "meta-llama"}293031def parse_repo(text: str) -> str:32    t = text.strip()33    m = re.match(r"^https?://huggingface\.co/([^/\s]+/[^/\s?#]+)", t)34    if m:35        t = m.group(1)36    t = t.removeprefix("hf.co/").removeprefix("huggingface.co/")37    if not REPO_RE.match(t) or ".." in t:38        raise APIError("Invalid Hugging Face repository id. Expected 'organization/model-name'.", param="repository")39    return t404142def _quant_rank(label: str | None) -> float:43    if not label:44        return 045    _, bits = formats.parse_quant_from_name("x-" + label)46    return bits or 0474849def pick_gguf_files(files: list[dict], preferred: str | None = None) -> list[dict]:50    """Choose one quantization from a GGUF repo (prefer Q4_K_M/Q5_K_M/Q6_K/Q8_0, or preferred)."""51    ggufs = [f for f in files if f["path"].lower().endswith(".gguf")]52    if not ggufs:53        return []54    mmproj = [f for f in ggufs if "mmproj" in f["path"].lower()]55    weights = [f for f in ggufs if "mmproj" not in f["path"].lower()]5657    def label(f):58        q, _ = formats.parse_quant_from_name(Path(f["path"]).stem)59        return (q or "").upper()6061    groups: dict[str, list[dict]] = {}62    for f in weights:63        base = re.sub(r"-\d{5}-of-\d{5}$", "", Path(f["path"]).stem)64        groups.setdefault(base, []).append(f)65    if preferred:66        pref = preferred.upper()67        for base, fs in groups.items():68            if pref in base.upper():69                return sorted(fs, key=lambda f: f["path"]) + mmproj[:1]70    order = ["Q4_K_M", "Q4_K_XL", "UD-Q4_K_XL", "Q5_K_M", "Q4_K_S", "Q6_K", "Q8_0", "MXFP4", "Q5_K_S", "IQ4_XS", "Q4_0", "BF16", "F16"]71    for want in order:72        for base, fs in groups.items():73            if want in base.upper():74                return sorted(fs, key=lambda f: f["path"]) + mmproj[:1]75    if not groups:76        return []77    base = sorted(groups.items(), key=lambda kv: sum(f["size"] or 0 for f in kv[1]))[0]78    return sorted(base[1], key=lambda f: f["path"]) + mmproj[:1]798081class Downloader:82    def __init__(self, settings: Settings, db, registry, jobs: JobRunner):83        self.settings = settings84        self.db = db85        self.registry = registry86        self.jobs = jobs8788    def _api(self):89        from huggingface_hub import HfApi90        return HfApi(token=self.settings.hf_token or os.environ.get("HF_TOKEN") or None)9192    # ------------------------------------------------------------------ inspect93    async def inspect(self, repo: str, quant: str | None = None) -> dict:94        repo = parse_repo(repo)95        api = self._api()96        try:97            info = await asyncio.to_thread(api.model_info, repo, files_metadata=True)98        except Exception as e:99            msg = str(e)100            if "401" in msg or "gated" in msg.lower():101                raise DownloadError(f"Repository '{repo}' is gated or private. Set HF_TOKEN and accept the license on Hugging Face.")102            if "404" in msg:103                raise DownloadError(f"Repository '{repo}' was not found on Hugging Face.")104            if "429" in msg:105                raise DownloadError("Hugging Face rate limit reached. Try again in a minute.")106            raise DownloadError(f"Could not inspect '{repo}': {msg[:200]}")107        files = [{"path": s.rfilename, "size": s.size or 0} for s in (info.siblings or [])]108        tags = list(info.tags or [])109        library = getattr(info, "library_name", None)110        cfg = getattr(info, "config", None) or {}111        has_st = any(f["path"].endswith(".safetensors") for f in files)112        has_gguf = any(f["path"].lower().endswith(".gguf") for f in files)113        is_mlx = "mlx" in tags or library == "mlx"114        runtime = None115        selected: list[dict] = []116        if has_gguf:117            runtime = "llamacpp"118            selected = pick_gguf_files(files, quant)119            weights_bytes = sum(f["size"] for f in selected if "mmproj" not in f["path"].lower())120            download_bytes = sum(f["size"] for f in selected)121        elif has_st and is_mlx:122            runtime = "mlx"123            selected = [f for f in files if not f["path"].endswith((".gguf", ".bin", ".pt", ".onnx", ".h5", ".msgpack"))124                        and ".cache" not in f["path"]]125            weights_bytes = sum(f["size"] for f in selected if f["path"].endswith(".safetensors"))126            download_bytes = sum(f["size"] for f in selected)127        elif has_st:128            runtime = "mlx"  # mlx_lm can load standard HF safetensors (quantizes nothing; runs bf16)129            selected = [f for f in files if not f["path"].endswith((".gguf", ".bin", ".pt", ".onnx", ".h5", ".msgpack"))]130            weights_bytes = sum(f["size"] for f in selected if f["path"].endswith(".safetensors"))131            download_bytes = sum(f["size"] for f in selected)132        else:133            raise DownloadError(f"Repository '{repo}' contains neither safetensors nor GGUF weights.")134135        parsed = formats.parse_hf_config(cfg) if cfg else {}136        model_type = parsed.get("model_type")137        base_model = next((t.split(":", 1)[1] for t in tags if t.startswith("base_model:") and "quantized:" not in t138                           and "finetune:" not in t), None)139        bm_quant = next((t.split(":", 2)[2] for t in tags if t.startswith("base_model:quantized:")), None)140        quant_label = parsed.get("quantization")141        bits = parsed.get("quant_bits")142        if runtime == "llamacpp" and selected:143            quant_label, bits = formats.parse_quant_from_name(Path(selected[0]["path"]).stem)144        if not quant_label:145            quant_label, bits = formats.parse_quant_from_name(repo.split("/")[-1])146        if not quant_label and runtime == "mlx":147            quant_label, bits = (parsed.get("torch_dtype") or "bf16"), 16148        total_p, active_p = formats.parse_param_count_from_name(repo.split("/")[-1])149        if not total_p and weights_bytes and bits:150            total_p = int(weights_bytes * 8 / bits)151        kv = kv_bytes_per_token(parsed.get("n_layers"), parsed.get("n_kv_heads"), parsed.get("head_dim"), 16,152                                parsed.get("full_attention_layers"), parsed.get("sliding_window"))153        if not kv and total_p:154            # crude fallback: ~ 130 KB/token for 7-9B, scale with params^0.6155            kv = int(130_000 * (total_p / 8e9) ** 0.6)156        pipeline = getattr(info, "pipeline_tag", None)157        name_l = repo.lower()158        vision = bool(parsed.get("vision")) or pipeline == "image-text-to-text" or any("mmproj" in f["path"].lower() for f in selected)159        embedding = pipeline in ("feature-extraction", "sentence-similarity") or "embed" in name_l160        reranker = pipeline == "text-ranking" or "rerank" in name_l161        budget, absolute = await self.registry.budgets()162        from .models.scanner import llamacpp_available163        comp = compat.evaluate(runtime=runtime, weights_bytes=weights_bytes, kv_per_token=kv, max_context=parsed.get("max_context"),164                               model_type=model_type, architecture=(parsed.get("architectures") or [None])[0] if runtime == "mlx" else model_type,165                               vision=vision, embedding=embedding, reranker=reranker, budget_gb=budget, absolute_gb=absolute,166                               llamacpp_available=llamacpp_available(self.settings.llama_server_bin), quant_bits=bits,167                               weights_file=selected[0]["path"] if runtime == "llamacpp" and selected else None)168        if runtime == "llamacpp" and comp.status == compat.INCOMPATIBLE and "Architecture" in comp.reason:169            pass170        du = shutil.disk_usage(self.settings.models_dir)171        min_free = float(await self.db.get_setting("min_free_disk_gb", self.settings.min_free_disk_gb))172        free_after = (du.free - download_bytes) / GB173        target = self.target_dir(repo, runtime, vision, embedding, reranker)174        existing = await self.db.fetchone("SELECT id FROM models WHERE repository=? AND installed=1", (repo,))175        return {176            "repository": repo, "runtime": runtime, "format": "gguf" if runtime == "llamacpp" else "safetensors",177            "files": selected, "all_files": files, "download_bytes": download_bytes, "weights_bytes": weights_bytes,178            "quantization": quant_label, "quant_bits": bits, "parameter_count": total_p, "active_parameter_count": active_p,179            "model_type": model_type, "pipeline_tag": pipeline, "library": library, "tags": tags[:40],180            "base_model": base_model or bm_quant, "vision": vision, "embedding": embedding, "reranker": reranker,181            "max_context": parsed.get("max_context"), "kv_bytes_per_token": kv, "compatibility": comp.to_dict(),182            "size_class": formats.size_class(comp.estimated_ram_gb, budget),183            "disk": {"free_gb": round(du.free / GB, 1), "free_after_gb": round(free_after, 1), "min_free_gb": min_free,184                     "ok": free_after >= min_free},185            "target_dir": str(target), "already_installed": existing["id"] if existing else None,186            "downloads": getattr(info, "downloads", None), "likes": getattr(info, "likes", None),187            "last_modified": str(getattr(info, "last_modified", "") or ""), "gated": bool(getattr(info, "gated", False)),188        }189190    def target_dir(self, repo: str, runtime: str, vision: bool, embedding: bool, reranker: bool) -> Path:191        name = repo.split("/")[-1]192        root = self.settings.models_dir193        if embedding:194            sub = root / "embeddings"195        elif reranker:196            sub = root / "rerankers"197        elif vision:198            sub = root / "vision"199        else:200            sub = root / ("gguf" if runtime == "llamacpp" else "mlx")201        family = formats.guess_family(name)202        return sub / family / name203204    # ------------------------------------------------------------------ download205    async def start_download(self, repo: str, quant: str | None = None, *, force: bool = False,206                             actor: str | None = None) -> Job:207        if not await self.db.get_setting("allow_downloads", self.settings.allow_downloads):208            raise APIError("Downloads are disabled in settings.", code="DOWNLOADS_DISABLED", status_code=403)209        insp = await self.inspect(repo, quant)210        if insp["already_installed"] and not force:211            raise APIError(f"'{repo}' is already installed as '{insp['already_installed']}'.", code="ALREADY_INSTALLED", status_code=409)212        if not insp["disk"]["ok"]:213            raise InsufficientDisk(f"Downloading {insp['download_bytes'] / GB:.1f} GB would leave {insp['disk']['free_after_gb']} GB free, "214                                   f"below the {insp['disk']['min_free_gb']:.0f} GB reserve.")215        if insp["compatibility"]["status"] == compat.INCOMPATIBLE and not force:216            raise APIError(f"'{repo}' is incompatible: {insp['compatibility']['reason']}", code="MODEL_INCOMPATIBLE", status_code=422)217        target = Path(insp["target_dir"])218        title = f"Download {repo}"219        payload = {"repository": repo, "quant": quant, "target": str(target), "runtime": insp["runtime"],220                   "download_bytes": insp["download_bytes"], "files": [f["path"] for f in insp["files"]]}221        await self.db.audit("download.start", actor=actor, target=repo, detail={"bytes": insp["download_bytes"]})222223        async def run(job: Job):224            return await self._download_job(job, insp, target)225226        return self.jobs.submit("download", title, payload, run, exclusive_download=True)227228    async def _download_job(self, job: Job, insp: dict, target: Path) -> dict:229        from huggingface_hub import hf_hub_download230        repo = insp["repository"]231        files = insp["files"]232        total = max(1, insp["download_bytes"])233        target.mkdir(parents=True, exist_ok=True)234        token = self.settings.hf_token or os.environ.get("HF_TOKEN") or None235        t0 = time.time()236        done_bytes = 0237        # progress poller: sums sizes of files (+ partial .incomplete blobs) in target238        stop = asyncio.Event()239240        def measure() -> int:241            n = 0242            for p in target.rglob("*"):243                if p.is_file():244                    try:245                        n += p.stat().st_size246                    except OSError:247                        pass248            return n249250        async def poll():251            last = 0252            last_t = time.time()253            while not stop.is_set():254                cur = await asyncio.to_thread(measure)255                now = time.time()256                speed = (cur - last) / max(0.001, now - last_t)257                last, last_t = cur, now258                eta = (total - cur) / speed if speed > 0 else None259                self.jobs.update(job, progress=min(0.99, cur / total), downloaded=cur, total=total,260                                 speed_bps=round(speed), eta_seconds=round(eta) if eta else None,261                                 elapsed=round(now - t0))262                await asyncio.sleep(1.0)263264        poller = asyncio.create_task(poll())265        try:266            for i, f in enumerate(files):267                if job.cancelled:268                    break269                self.jobs.update(job, current_file=f["path"], file_index=i + 1, file_count=len(files))270                await asyncio.to_thread(hf_hub_download, repo, f["path"], local_dir=str(target), token=token,271                                        force_download=False)272                done_bytes += f["size"]273        finally:274            stop.set()275            poller.cancel()276        if job.cancelled:277            # remove partial download278            shutil.rmtree(target, ignore_errors=True)279            return {"cancelled": True}280        # cleanup HF metadata cache folder inside local_dir281        shutil.rmtree(target / ".cache", ignore_errors=True)282        # verify283        missing = [f["path"] for f in files if not (target / f["path"]).exists()]284        bad = [f["path"] for f in files if (target / f["path"]).exists() and f["size"] and (target / f["path"]).stat().st_size != f["size"]]285        if missing or bad:286            raise DownloadError(f"Download incomplete: missing {missing[:3]} size-mismatch {bad[:3]}")287        manifest = {288            "schema_version": 1, "model_id": slugify(repo.split("/")[-1]), "runtime": insp["runtime"],289            "quantization": insp["quantization"], "download_source": "huggingface", "repository": repo,290            "downloaded_at": time.strftime("%Y-%m-%dT%H:%M:%S%z"), "verified": True,291            "files": [{"path": f["path"], "size": f["size"]} for f in files], "provider": repo.split("/")[0],292            "base_model": insp.get("base_model"), "task": "embedding" if insp["embedding"] else "reranking" if insp["reranker"] else293            "image-text-to-text" if insp["vision"] else "text-generation",294        }295        (target / "llm-api.json").write_text(json.dumps(manifest, indent=2))296        # keep a copy in models/manifests297        mdir = self.settings.models_dir / "manifests"298        mdir.mkdir(parents=True, exist_ok=True)299        (mdir / f"{slugify(repo)}.json").write_text(json.dumps(manifest, indent=2))300        self.jobs.update(job, progress=0.99, stage="registering")301        summary = await self.registry.rescan()302        m = await self.db.fetchone("SELECT id FROM models WHERE path=?", (str(target),))303        mid = m["id"] if m else None304        if mid:305            await self.registry.update(mid, repository=repo, provider=repo.split("/")[0], verified=False)306            if insp.get("embedding") or insp.get("reranker") or insp.get("vision"):307                ov = {"task": manifest["task"], "vision": insp["vision"], "embedding": insp["embedding"], "reranker": insp["reranker"]}308                await self.registry.update(mid, overrides=ov, task=manifest["task"], vision=insp["vision"],309                                           embedding=insp["embedding"], reranker=insp["reranker"])310            await self.db.model_event(mid, "downloaded", {"repository": repo, "bytes": insp["download_bytes"],311                                                          "seconds": round(time.time() - t0)})312        return {"model_id": mid, "path": str(target), "bytes": insp["download_bytes"], "seconds": round(time.time() - t0),313                "scan": summary}314315    # ------------------------------------------------------------------ delete316    async def delete_model(self, model_id: str, *, actor: str | None = None, keep_benchmarks: bool = True) -> dict:317        m = await self.registry.get(model_id)318        if not m:319            from .errors import ModelNotFound320            raise ModelNotFound(f"Model '{model_id}' not found.")321        path = Path(m["path"]).resolve()322        root = self.settings.models_dir.resolve()323        if root not in path.parents:324            raise APIError("Refusing to delete a path outside the model root.", code="PATH_NOT_ALLOWED", status_code=403)325        size = m["disk_size_bytes"]326        if m["format"] == "gguf" and m["weights_file"]:327            # delete only this quantization's files (+ mmproj if no other model uses the dir)328            others = await self.db.fetchall("SELECT id FROM models WHERE path=? AND id<>?", (m["path"], model_id))329            wf = Path(m["weights_file"])330            base = re.sub(r"-\d{5}-of-\d{5}$", "", wf.stem)331            for f in path.glob("*.gguf"):332                if f.stem == wf.stem or f.stem.startswith(base + "-0"):333                    f.unlink(missing_ok=True)334            if not others:335                shutil.rmtree(path, ignore_errors=True)336        else:337            shutil.rmtree(path, ignore_errors=True)338        await self.registry.delete(model_id)339        if not keep_benchmarks:340            await self.db.execute("DELETE FROM model_benchmarks WHERE model_id=?", (model_id,))341        await self.db.audit("model.delete", actor=actor, target=model_id, detail={"bytes": size, "path": str(path)})342        await self.db.model_event(model_id, "deleted", {"bytes": size})343        return {"deleted": model_id, "bytes": size}344