File size: 3,608 Bytes
2407511
 
 
 
c306cf6
2407511
c306cf6
 
2407511
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
c306cf6
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
2407511
 
 
 
 
 
 
c306cf6
 
 
 
 
 
 
 
 
2407511
 
 
 
 
 
 
 
 
c306cf6
2407511
 
 
 
 
 
 
 
 
 
 
 
c306cf6
2407511
 
 
c306cf6
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
"""Deterministic cache and retrieval helpers for public upstream models."""

from __future__ import annotations

import hashlib
import os
import shutil
import uuid
from pathlib import Path
from typing import Any, Callable, Mapping


APPLICATION_ROOT = Path(__file__).resolve().parents[1]
DEFAULT_MODEL_CACHE = APPLICATION_ROOT / ".cache" / "satquery-models"


class PublicModelRetrievalError(RuntimeError):
    """A pinned public upstream artifact could not be retrieved safely."""


def model_cache_dir(environment: Mapping[str, str] | None = None) -> Path:
    """Return the application-controlled model cache, honoring its sole override."""

    values = os.environ if environment is None else environment
    configured = values.get("SATQUERY_MODEL_CACHE", "").strip()
    candidate = Path(configured).expanduser() if configured else DEFAULT_MODEL_CACHE
    if not candidate.is_absolute():
        candidate = APPLICATION_ROOT / candidate
    return candidate.resolve()


def _materialized_public_model_path(*, repo_id: str, filename: str, revision: str) -> Path:
    """Return a deterministic, suffix-preserving path for one immutable artifact."""

    if not filename or Path(filename).name != filename:
        raise PublicModelRetrievalError("Public model filename must be a single path component")
    identity = hashlib.sha256(f"{repo_id}\0{revision}\0{filename}".encode("utf-8")).hexdigest()
    return model_cache_dir() / "materialized" / identity / filename


def _materialize_cached_file(source: Path, target: Path) -> Path:
    """Atomically hard-link or copy a cached artifact to its stable named path."""

    if target.is_file():
        return target.absolute()

    target.parent.mkdir(parents=True, exist_ok=True)
    temporary = target.with_name(f".{target.name}.{os.getpid()}.{uuid.uuid4().hex}.tmp")
    try:
        try:
            os.link(source, temporary)
        except OSError:
            shutil.copyfile(source, temporary)
        os.replace(temporary, target)
    finally:
        temporary.unlink(missing_ok=True)
    return target.absolute()


def download_public_hf_file(
    *,
    repo_id: str,
    filename: str,
    revision: str,
    downloader: Callable[..., Any] | None = None,
) -> Path:
    """Resolve one immutable public file to a stable path retaining its filename."""

    target = _materialized_public_model_path(
        repo_id=repo_id,
        filename=filename,
        revision=revision,
    )
    if target.is_file():
        return target.absolute()

    if downloader is None:
        from huggingface_hub import hf_hub_download

        downloader = hf_hub_download

    cache_dir = model_cache_dir()
    try:
        cache_dir.mkdir(parents=True, exist_ok=True)
        source = Path(
            downloader(
                repo_id=repo_id,
                filename=filename,
                revision=revision,
                cache_dir=str(cache_dir),
                token=False,
            )
        ).resolve()
    except Exception as error:
        raise PublicModelRetrievalError(
            f"Pinned public model artifact could not be retrieved: {repo_id}@{revision}/{filename}"
        ) from error
    if not source.is_file():
        raise PublicModelRetrievalError(
            f"Pinned public model artifact was not materialized: {repo_id}@{revision}/{filename}"
        )
    try:
        return _materialize_cached_file(source, target)
    except OSError as error:
        raise PublicModelRetrievalError(
            f"Pinned public model artifact could not be named safely: {repo_id}@{revision}/{filename}"
        ) from error