lhallee commited on
Commit
ece09aa
·
verified ·
1 Parent(s): 3a9f52f

Accept huggingface_hub's shared Xet blob store in the CCD loader (FastPLMs PR 51)

Browse files
fastplms/models/esmfold2/esmfold2_conformers.py CHANGED
@@ -9,6 +9,7 @@ from __future__ import annotations
9
 
10
  import os
11
  import pickle
 
12
  import stat
13
  import tempfile
14
  import numpy as np
@@ -28,6 +29,7 @@ from .esmfold2_constants import RES_TYPE_TO_CCD
28
 
29
  _CCD_ENVIRONMENT_VARIABLE = "ESMCFOLD_CCD_PATH"
30
  _CCD_ASSET_ID = "esmfold2_ccd"
 
31
 
32
 
33
  @dataclass(frozen=True)
@@ -174,19 +176,57 @@ def _resolve_trusted_hub_snapshot_link(
174
  ) from error
175
 
176
  resolved = asset_path.resolve(strict=True)
177
- blob_root = (repository_cache / "blobs").resolve(strict=True)
178
- try:
179
- blob_root.relative_to(root)
180
- resolved.relative_to(blob_root)
181
- except ValueError as error:
182
  raise ValueError(
183
- f"CCD Hub snapshot link escapes its repository blob cache: {asset_path}"
184
- ) from error
 
185
  if not resolved.is_file() or resolved.is_symlink():
186
  raise ValueError(f"CCD Hub snapshot target must be a regular file: {resolved}")
187
  return resolved
188
 
189
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
190
  class _ChemicalComponentStore:
191
  def __init__(self) -> None:
192
  self.molecules: dict[str, Any] | None = None
 
9
 
10
  import os
11
  import pickle
12
+ import re
13
  import stat
14
  import tempfile
15
  import numpy as np
 
29
 
30
  _CCD_ENVIRONMENT_VARIABLE = "ESMCFOLD_CCD_PATH"
31
  _CCD_ASSET_ID = "esmfold2_ccd"
32
+ _HEX_DIGEST = re.compile(r"[0-9a-f]{64}")
33
 
34
 
35
  @dataclass(frozen=True)
 
176
  ) from error
177
 
178
  resolved = asset_path.resolve(strict=True)
179
+ if not (
180
+ _is_repository_blob(resolved, repository_cache, root)
181
+ or _is_shared_store_blob(resolved, repository_cache, root, contract.sha256)
182
+ ):
 
183
  raise ValueError(
184
+ "CCD Hub snapshot link escapes its repository blob cache and the shared "
185
+ f"Hub blob store: {asset_path}"
186
+ )
187
  if not resolved.is_file() or resolved.is_symlink():
188
  raise ValueError(f"CCD Hub snapshot target must be a regular file: {resolved}")
189
  return resolved
190
 
191
 
192
+ def _is_repository_blob(target: Path, repository_cache: Path, root: Path) -> bool:
193
+ """Return whether ``target`` lies in the repository's own blob directory."""
194
+
195
+ blob_root = (repository_cache / "blobs").resolve(strict=True)
196
+ return blob_root.is_relative_to(root) and target.is_relative_to(blob_root)
197
+
198
+
199
+ def _is_shared_store_blob(
200
+ target: Path,
201
+ repository_cache: Path,
202
+ root: Path,
203
+ sha256: str,
204
+ ) -> bool:
205
+ """Return whether ``target`` is the shared Xet store entry of the pinned repository blob.
206
+
207
+ Since huggingface_hub 1.32, Xet downloads live once per cache at
208
+ ``<root>/blobs/<xet hash[:2]>/<xet hash>``, and ``models--<repo>/blobs/<sha256>``
209
+ becomes a relative symlink to that entry. The store name is a Xet hash, not the
210
+ SHA-256, so the repository blob named by the pinned digest binds the entry to
211
+ this asset. ``_open_verified_asset`` still hashes the bytes themselves.
212
+ """
213
+
214
+ if _HEX_DIGEST.fullmatch(sha256) is None:
215
+ return False
216
+ store_root = (root / "blobs").resolve()
217
+ if not store_root.is_relative_to(root) or not target.is_relative_to(store_root):
218
+ return False
219
+ store_parts = target.relative_to(store_root).parts
220
+ if (
221
+ len(store_parts) != 2
222
+ or _HEX_DIGEST.fullmatch(store_parts[1]) is None
223
+ or store_parts[0] != store_parts[1][:2]
224
+ ):
225
+ return False
226
+ repository_blob = repository_cache / "blobs" / sha256
227
+ return repository_blob.is_symlink() and repository_blob.resolve() == target
228
+
229
+
230
  class _ChemicalComponentStore:
231
  def __init__(self) -> None:
232
  self.molecules: dict[str, Any] | None = None
fastplms_bundle.py CHANGED
The diff for this file is too large to render. See raw diff
 
modeling_fastplms.py CHANGED
@@ -13,7 +13,7 @@ from zipfile import ZIP_DEFLATED, ZipFile
13
 
14
  from .fastplms_bundle import RUNTIME_DATA, RUNTIME_HASH
15
 
16
- if RUNTIME_HASH != "2ef7cccec8dfd129c6621d7bbb4b9f91a0ed1be1fd3f62d670259a4a9e199af0":
17
  raise RuntimeError("FastPLMs runtime identity differs from the bridge.")
18
 
19
  _RUNTIME_TEMPORARIES = []
 
13
 
14
  from .fastplms_bundle import RUNTIME_DATA, RUNTIME_HASH
15
 
16
+ if RUNTIME_HASH != "9eaca231522bbd565f9359301d8c51391e8e512d522ebca5f0d8119d666537fd":
17
  raise RuntimeError("FastPLMs runtime identity differs from the bridge.")
18
 
19
  _RUNTIME_TEMPORARIES = []