File size: 5,843 Bytes
7ad3625
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
# Copyright 2026-present, the HuggingFace Inc. team.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
#     http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""On-disk cache for repository tree listings.

A tree listing is the set of files (with their download metadata) contained in a repo at a given commit. Because a
commit hash is immutable, its tree listing never changes and can be cached forever without any invalidation logic.

The listing is stored as a human-readable JSON file under `<tree_cache_folder>/trees/<commit_hash>.json`. The
folder depends on the download target: the per-repo `storage_folder` for `cache_dir` downloads, or
`local_dir/.cache/huggingface/` for `local_dir` downloads (see `tree_cache_folder_for_local_dir`).

```json
{
  "format_version": 1,
  "files": {
    "config.json": {"size": 519, "blob_id": "<git sha1>"},
    "model.safetensors": {"size": 1234, "blob_id": "...", "lfs_sha256": "<sha256>", "lfs_size": 1234, "xet_hash": "..."}
  }
}
```
"""

import json
import os
import tempfile
import threading
from dataclasses import dataclass

from .utils import logging
from .utils._xet import is_valid_xet_hash


logger = logging.get_logger(__name__)

TREE_CACHE_FORMAT_VERSION = 1

# In-memory cache of parsed tree listings, keyed by absolute file path.
_IN_MEMORY_TREE_CACHE: dict[str, "dict[str, TreeCacheEntry]"] = {}
_IN_MEMORY_TREE_CACHE_LOCK = threading.Lock()


@dataclass(frozen=True)
class TreeCacheEntry:
    """Raw metadata of a single file in a cached tree listing, mirroring the `/tree` endpoint fields."""

    size: int
    blob_id: str
    lfs_sha256: str | None = None
    lfs_size: int | None = None
    xet_hash: str | None = None

    def to_json(self) -> dict:
        info: dict = {"size": self.size, "blob_id": self.blob_id}
        if self.lfs_sha256 is not None:
            info["lfs_sha256"] = self.lfs_sha256
            info["lfs_size"] = self.lfs_size
        if self.xet_hash is not None:
            info["xet_hash"] = self.xet_hash
        return info

    @classmethod
    def from_json(cls, info: dict) -> "TreeCacheEntry":
        return cls(
            size=info["size"],
            blob_id=info["blob_id"],
            lfs_sha256=info.get("lfs_sha256"),
            lfs_size=info.get("lfs_size"),
            xet_hash=info.get("xet_hash"),
        )


def is_valid_tree_entries(entries: dict[str, TreeCacheEntry]) -> bool:
    """Return whether all Xet hashes in the tree listing are valid."""
    return all(entry.xet_hash is None or is_valid_xet_hash(entry.xet_hash) for entry in entries.values())


def _tree_cache_path(tree_cache_folder: str, commit_hash: str) -> str:
    return os.path.join(tree_cache_folder, "trees", f"{commit_hash}.json")


def tree_cache_folder_for_local_dir(local_dir: str) -> str:
    """Folder under which the `trees/` cache lives for a `local_dir` download."""
    return os.path.join(local_dir, ".cache", "huggingface")


def read_tree_cache(tree_cache_folder: str, commit_hash: str) -> dict[str, TreeCacheEntry] | None:
    """Return the cached tree listing for a commit hash, or `None` if not cached, invalid, or unreadable."""
    path = _tree_cache_path(tree_cache_folder, commit_hash)
    with _IN_MEMORY_TREE_CACHE_LOCK:
        if path in _IN_MEMORY_TREE_CACHE:
            cached_entries = _IN_MEMORY_TREE_CACHE[path]
            return cached_entries if is_valid_tree_entries(cached_entries) else None
    entries = _read_tree_cache_from_disk(path)
    if entries is None or not is_valid_tree_entries(entries):
        return None
    with _IN_MEMORY_TREE_CACHE_LOCK:
        _IN_MEMORY_TREE_CACHE[path] = entries
    return entries


def _read_tree_cache_from_disk(path: str) -> dict[str, TreeCacheEntry] | None:
    try:
        with open(path, encoding="utf-8") as f:
            data = json.load(f)
        if data.get("format_version") != TREE_CACHE_FORMAT_VERSION:
            # Unknown format (e.g. written by a newer version) => ignore and re-fetch.
            return None
        return {file_path: TreeCacheEntry.from_json(info) for file_path, info in data["files"].items()}
    except FileNotFoundError:
        return None
    except (OSError, ValueError, KeyError, TypeError) as e:
        logger.warning(f"Ignoring corrupted tree cache file {path}: {e}")
        return None


def write_tree_cache(tree_cache_folder: str, commit_hash: str, entries: dict[str, TreeCacheEntry]) -> None:
    """Write a valid tree listing to the cache (ignoring invalid entries and any failures)."""
    if not is_valid_tree_entries(entries):
        return

    path = _tree_cache_path(tree_cache_folder, commit_hash)
    data = {
        "format_version": TREE_CACHE_FORMAT_VERSION,
        "files": {file_path: entries[file_path].to_json() for file_path in sorted(entries)},
    }
    try:
        os.makedirs(os.path.dirname(path), exist_ok=True)
        tmp_fd, tmp_path = tempfile.mkstemp(dir=os.path.dirname(path), suffix=".tmp")
        with os.fdopen(tmp_fd, "w", encoding="utf-8") as f:
            json.dump(data, f, indent=1)
        os.replace(tmp_path, path)
    except OSError as e:
        logger.warning(f"Ignored error while writing tree cache file {path}: {e}")
        return

    # Seed the in-memory cache so later readers of this commit skip re-reading and re-parsing the file.
    with _IN_MEMORY_TREE_CACHE_LOCK:
        _IN_MEMORY_TREE_CACHE[path] = dict(entries)