Annie Voigt
style: apply ruff lint --fix + ruff format across the tree
c3b49d6
Raw
History Blame Contribute Delete
4.6 kB
"""Shared, authenticated resolution of expression inputs to local paths.
Single source of truth for turning a path-or-URL into a readable local file, so
every loader (the bulk dataset tools, ``decoupler_differential_expression``, the
Mode-A integration builder, …) downloads private ``huggingface.co`` files with
the same ``HF_TOKEN`` auth instead of each rolling its own unauthenticated
``urllib`` call. Adding a new tool means calling this helper, not copying a
download snippet.
"""
from __future__ import annotations
import os
import tempfile
import urllib.request
from pathlib import Path
def resolve_to_local_path(path_or_url: str) -> tuple[str, bool]:
"""Resolve a path or URL to a readable local file path.
Returns ``(local_path, is_temp)``. ``is_temp`` is True only when a temporary
file was created that the caller is responsible for deleting; local paths and
HF-cache files return False.
Private ``huggingface.co`` ``/resolve/`` URLs are fetched with
``hf_hub_download`` (which handles the LFS redirect and ``HF_TOKEN`` auth)
when a token is available; everything else falls back to plain ``urllib`` so
public files and non-HF hosts are unaffected.
"""
s = str(path_or_url)
if not s.startswith(("http://", "https://", "ftp://")):
# Local paths are still integrity-verified when the manifest baselines
# this exact path (rare, but keeps the check uniform); no-op otherwise.
_verify_or_raise(s, s, is_temp=False)
return s, False
# Authenticated path for private HF repos.
if "huggingface.co/" in s and "/resolve/" in s:
# Env vars first (the Space sets HF_TOKEN as a secret); fall back to a
# cached `huggingface-cli login` token so local/dev runs authenticate too.
token = os.environ.get("HF_TOKEN") or os.environ.get("HUGGING_FACE_HUB_TOKEN")
if not token:
try:
from huggingface_hub import get_token as _hf_get_token
token = _hf_get_token()
except Exception:
token = None
if token:
local = None
try:
from huggingface_hub import hf_hub_download
after = s.split("huggingface.co/", 1)[1]
repo_type = "model"
if after.startswith("datasets/"):
repo_type, after = "dataset", after[len("datasets/") :]
elif after.startswith("spaces/"):
repo_type, after = "space", after[len("spaces/") :]
repo_id, file_part = after.split("/resolve/", 1)
revision, filename = file_part.split("/", 1)
local = hf_hub_download(
repo_id=repo_id,
filename=filename,
repo_type=repo_type,
revision=revision,
token=token,
)
except Exception:
local = None # fall through to unauthenticated urllib (e.g. public file)
# Verify OUTSIDE the try above so an IntegrityError is never swallowed
# into the urllib fallback (which fetches the same tampered content).
if local is not None:
_verify_or_raise(local, s, is_temp=False)
return local, False # HF cache file — caller must not delete it
suffix = Path(s.split("?", 1)[0]).suffix or ".h5ad"
with tempfile.NamedTemporaryFile(suffix=suffix, delete=False) as tmp:
with urllib.request.urlopen(s) as resp:
while True:
chunk = resp.read(4 * 1024 * 1024)
if not chunk:
break
tmp.write(chunk)
tmp_path = tmp.name
# Verify before returning so a tampered download is never handed to a loader
# or admitted to the cache; a bad temp file is deleted here, not leaked.
_verify_or_raise(tmp_path, s, is_temp=True)
return tmp_path, True
def _verify_or_raise(local_path: str, url: str, *, is_temp: bool) -> None:
"""Run the ADR-0010 integrity check; delete a bad temp file before re-raising.
No-op when the URL has no recorded baseline. On a hash mismatch the load is
refused (:class:`IntegrityError`); for a temp file we unlink first so the
tampered bytes are never left on disk or admitted to the cache.
"""
from src.core.integrity import IntegrityError, verify_file
try:
verify_file(local_path, url)
except IntegrityError:
if is_temp:
Path(local_path).unlink(missing_ok=True)
raise