diff --git a/deepcell_types/utils/__init__.py b/deepcell_types/utils/__init__.py index dcda82b..f0925fe 100644 --- a/deepcell_types/utils/__init__.py +++ b/deepcell_types/utils/__init__.py @@ -1,13 +1,15 @@ """Utilities for model / data access. -The registries below pin checksums of the paper-release checkpoints; -uploads to ``users.deepcell.org`` use the asset paths constructed below -(``models/``). The hash algorithm is auto-detected from the -digest length (32 hex → md5, 64 hex → sha256), so entries can be migrated -to the stronger sha256 individually; new entries should pin sha256. Some -baselines ship more than one file (e.g. ``maps`` needs its ``_stats.npz`` -companion; ``xgboost`` needs its ``.remap.json`` label-remap), so each -baseline entry is a list of ``(filename, hash)`` tuples. +Model, baseline, and training-data downloads are delegated to the shared +``deepcell-auth`` client, whose bundled ``asset_manifest.yaml`` is the single +source of truth for asset keys and integrity hashes. The functions here are +thin adapters that keep this package's public API: the ``ValueError`` on an +unknown identifier, the pointer to Nimbus' upstream weights, and the +``Path`` / ``list[Path]`` return types. + +The version and baseline names below are mirrored statically from that +manifest so an unknown identifier is rejected without a manifest read; +``tests/test_download_delegation.py`` fails if they drift. """ __all__ = [ @@ -21,70 +23,33 @@ "resolve_supported_marker", ] -_latest = "2026-06-15" +_DEFAULT_MODEL_VERSION = "2026-06-15" -# Main model checkpoints. Values are ``(asset_filename, md5)``. -# Two released residual-MLP checkpoints share the current vocabulary (51 cell -# types, 278 markers) and load via stock ``predict.py`` (the resMLP head is +# Released residual-MLP checkpoints, sharing the current vocabulary (51 cell +# types, 278 markers) and loading via stock ``predict.py`` (the resMLP head is # auto-detected): # - ``2026-06-15`` (default): the from-scratch resMLP (paper Fig-3c headline). -# - ``2026-06-23``: the pretrain -> finetune (SSL) resMLP arm. -# -# Each md5 tracks the vocab-bundled asset (carries ``ct2idx`` + -# ``canonical_channels`` so ``validate_checkpoint_vocabulary`` can verify -# ordering). The original un-bundled assets predated the guard and raised -# ``ValueError: ... does not bundle a ct2idx`` on every ``predict()`` call; -# they were re-packaged with ``scripts/repackage_release_checkpoint.py``. Each -# bundled asset must be uploaded at its filename below before its md5 is served. -_model_registry = { - "2026-06-15": ( - "deepcell-types_2026-06-15_resmlp.pt", - "b819a7e0b177ad5330394eab3c6c7ad8", - ), - "2026-06-23": ( - "deepcell-types_2026-06-23_resmlp_ptft.pt", - "402e94c103c5e489a57433cd107009d3", - ), -} - -# Baseline-model checkpoints. Values are lists of ``(asset_filename, md5)``. -# Single-file baselines have a one-element list; ``maps`` and ``xgboost`` -# additionally ship a companion file required at inference. The md5s track the -# 2026-06-30 DCT-sampler retrains -- the release-of-record that produced the -# published baseline prediction CSVs and headline numbers. +# - ``2026-06-23-ptft``: the pretrain -> finetune (SSL) resMLP arm. +_MODEL_VERSIONS = ("2026-06-15", "2026-06-23-ptft") + +# Also in the manifest, but deliberately not offered here: the 2025 CLIP +# checkpoints predate the current architecture and cannot be loaded by this +# code. They stay served for reproducibility with the matching historical +# commit, so ``download_model`` rejects them rather than handing back a +# checkpoint that fails at ``load_state_dict``. +_LEGACY_MODEL_VERSIONS = ("2025-06-09", "2025-06-09_public-data-only") + +# All three baselines are served as a single ``.tar.gz`` bundle, one +# subdirectory per baseline, holding each one's weights plus any companion file +# required at inference (maps -> ``_stats.npz``; xgboost -> ``.remap.json``). +# Requesting any one baseline therefore downloads the whole bundle; the other +# two are then served from the cache without a second transfer. # # Nimbus is intentionally absent: it is inference-only with pretrained weights # produced and distributed upstream (angelolab/Nimbus-Inference), which this # project does not redistribute. ``download_baseline_checkpoint("nimbus")`` # raises with a pointer to the official source instead (see below). -_baseline_registry = { - "cellsighter": [ - ( - "deepcell-types_baseline-cellsighter.pth", - "c5a105fb044ad82c82817a32aabbae7c", - ), - ], - "maps": [ - ( - "deepcell-types_baseline-maps.pth", - "1ad5e30ceaccfdb91050b663099258fb", - ), - ( - "deepcell-types_baseline-maps_stats.npz", - "1f462d0c8bf531af73026d415d52728d", - ), - ], - "xgboost": [ - ( - "deepcell-types_baseline-xgboost.json", - "9e51cf1af8c6c43b00871a821fde0f57", - ), - ( - "deepcell-types_baseline-xgboost.remap.json", - "fba1b0e705e5f7747eb2f9cae30815ba", - ), - ], -} +_BASELINE_NAMES = ("cellsighter", "maps", "xgboost") # Nimbus pretrained weights are distributed upstream (via the # ``nimbus-inference`` library / Hugging Face Hub), not re-hosted here. @@ -100,38 +65,39 @@ def download_model(*, version=None): ---------- version : str, optional Which checkpoint version to download. Defaults to ``None``, - which resolves to the most-recently-released version - (``_latest`` in this module). + which resolves to the default version + (``_DEFAULT_MODEL_VERSION`` in this module). Returns ------- pathlib.Path Local path to the downloaded checkpoint. """ - from ._auth import fetch_data + from deepcell_auth import download_deepcell_types_model - version = version if version is not None else _latest - if version not in _model_registry: + version = version if version is not None else _DEFAULT_MODEL_VERSION + if version not in _MODEL_VERSIONS: raise ValueError( f"Unknown model version {version!r}. " - f"Known versions: {sorted(_model_registry)}." + f"Known versions: {sorted(_MODEL_VERSIONS)}." ) - filename, md5 = _model_registry[version] - return fetch_data(f"models/{filename}", cache_subdir="models", file_hash=md5) + return download_deepcell_types_model(version) def list_model_versions(): - """Return the available pre-trained model versions, newest first. + """Return the available pre-trained model versions, default first. Returns ------- list of str Version identifiers accepted by :func:`download_model`. The first - element is always the default (``_latest``) version that - ``download_model()`` resolves to with no argument. + element is always the default version that ``download_model()`` + resolves to with no argument. """ - others = sorted((v for v in _model_registry if v != _latest), reverse=True) - return [_latest, *others] + others = sorted( + (v for v in _MODEL_VERSIONS if v != _DEFAULT_MODEL_VERSION), reverse=True + ) + return [_DEFAULT_MODEL_VERSION, *others] def download_baseline_checkpoint(name): @@ -155,13 +121,15 @@ def download_baseline_checkpoint(name): Returns ------- list[pathlib.Path] - Local paths to every file downloaded for this baseline, in the - order declared in ``_baseline_registry``. Note the asymmetry with - :func:`download_model`, which returns a single ``Path``: baselines - return a *list* because some ship companion files. Call - :func:`list_baseline_names` for the accepted identifiers. + Local paths to every checkpoint file for this baseline, sorted by + filename. Note the asymmetry with :func:`download_model`, which + returns a single ``Path``: baselines return a *list* because some + ship companion files. Call :func:`list_baseline_names` for the + accepted identifiers. """ - from ._auth import fetch_data + from deepcell_auth import download_deepcell_types_baseline + + from ._archive import extract_archive if name == "nimbus": raise ValueError( @@ -171,14 +139,31 @@ def download_baseline_checkpoint(name): "Python 3.11), which downloads the weights automatically; see " f"{_NIMBUS_UPSTREAM_URL}." ) - if name not in _baseline_registry: + if name not in _BASELINE_NAMES: + raise ValueError( + f"Unknown baseline {name!r}. Known baselines: {sorted(_BASELINE_NAMES)}." + ) + + (archive,) = download_deepcell_types_baseline(name) + bundle_dir = archive.parent / archive.name.removesuffix(".tar.gz") + # Mirrors cellSAM: the bundle is unpacked once, and its presence is what + # skips re-extraction on later calls. ``fetch_data`` still re-checks the + # archive's pinned hash on every call, so a corrupt *download* is caught; + # a hand-edited extraction directory is not. + if not bundle_dir.is_dir(): + extract_archive(archive, archive.parent) + baseline_dir = bundle_dir / name + contents = ( + sorted(p for p in baseline_dir.iterdir() if p.is_file()) + if baseline_dir.is_dir() + else [] + ) + if not contents: raise ValueError( - f"Unknown baseline {name!r}. Known baselines: {sorted(_baseline_registry)}." + f"Baseline bundle {archive.name} did not unpack to a directory of " + f"checkpoint files at {baseline_dir}. Delete {bundle_dir} and retry." ) - return [ - fetch_data(f"models/{filename}", cache_subdir="models", file_hash=md5) - for filename, md5 in _baseline_registry[name] - ] + return contents def list_baseline_names(): @@ -190,7 +175,7 @@ def list_baseline_names(): Names accepted by :func:`download_baseline_checkpoint` (the ``list``-returning counterpart to :func:`list_model_versions`). """ - return sorted(_baseline_registry) + return sorted(_BASELINE_NAMES) def list_supported_markers(*, zarr_path=None): @@ -269,9 +254,6 @@ def list_supported_cell_types(*, zarr_path=None): return sorted(DCTConfig(zarr_path=zarr_path).ct2idx) -_training_data_asset_key = "data/deepcell-types/public_data_v1.1.zip" - - def download_training_data(*, extract=False): """Download the public training-data corpus for deepcell-types (v1.1). @@ -297,9 +279,11 @@ def download_training_data(*, extract=False): Local path to the downloaded ``.zip`` (or, when ``extract=True``, the directory it was extracted into). """ - from ._auth import extract_archive, fetch_data + from deepcell_auth import download_deepcell_types_data + + from ._archive import extract_archive - zip_path = fetch_data(_training_data_asset_key, cache_subdir="data") + zip_path = download_deepcell_types_data() if extract: return extract_archive(zip_path) return zip_path diff --git a/deepcell_types/utils/_archive.py b/deepcell_types/utils/_archive.py new file mode 100644 index 0000000..254541c --- /dev/null +++ b/deepcell_types/utils/_archive.py @@ -0,0 +1,120 @@ +"""Path-traversal-safe extraction of downloaded archives. + +Kept in this package rather than delegated to ``deepcell_auth``, whose +``extract_archive`` calls ``extractall`` without vetting members. +""" + +import logging +import tarfile +import zipfile +from pathlib import Path + +logger = logging.getLogger(__name__) + +# Deliberately generous limits accommodate the public multi-terabyte corpus +# while still bounding hostile or accidentally unbounded archives. +MAX_ARCHIVE_MEMBERS = 100_000 +MAX_ARCHIVE_MEMBER_BYTES = 2 * 2**40 +MAX_ARCHIVE_TOTAL_BYTES = 10 * 2**40 + + +def extract_archive( + archive_path, + dest=None, + *, + max_members=MAX_ARCHIVE_MEMBERS, + max_member_bytes=MAX_ARCHIVE_MEMBER_BYTES, + max_total_bytes=MAX_ARCHIVE_TOTAL_BYTES, +): + """Safely extract a ``.zip`` or ``.tar[.gz]`` archive. + + Rejects members whose resolved path would escape ``dest`` (zip-slip / + tar path-traversal) and rejects tar symlink/hardlink members, so it is + safe to call on archives downloaded from a remote source. + + Parameters + ---------- + archive_path : str or pathlib.Path + Path to a ``.zip`` or ``.tar``/``.tar.gz`` archive. + dest : str or pathlib.Path, optional + Destination directory. Defaults to the archive's parent directory. + + Returns + ------- + pathlib.Path + The destination directory the archive was extracted into. + """ + archive_path = Path(archive_path) + dest = Path(dest) if dest is not None else archive_path.parent + dest.mkdir(parents=True, exist_ok=True) + dest_resolved = dest.resolve() + + def _within(name): + try: + (dest / name).resolve().relative_to(dest_resolved) + return True + except ValueError: + return False + + if zipfile.is_zipfile(archive_path): + with zipfile.ZipFile(archive_path) as zf: + infos = zf.infolist() + if len(infos) > max_members: + raise ValueError( + f"Refusing to extract archive with {len(infos)} members; " + f"limit is {max_members}." + ) + total = 0 + for member in infos: + is_symlink = (member.external_attr >> 16) & 0o170000 == 0o120000 + if not _within(member.filename): + raise ValueError( + f"Refusing to extract {member.filename!r}: path escapes {dest}." + ) + if is_symlink: + raise ValueError( + f"Refusing to extract unsafe zip member {member.filename!r}." + ) + if member.file_size > max_member_bytes: + raise ValueError( + f"Refusing to extract {member.filename!r}: declared size " + f"exceeds {max_member_bytes} bytes." + ) + total += member.file_size + if total > max_total_bytes: + raise ValueError( + f"Refusing to extract archive larger than " + f"{max_total_bytes} bytes." + ) + zf.extractall(dest) + elif tarfile.is_tarfile(archive_path): + with tarfile.open(archive_path) as tf: + members = tf.getmembers() + if len(members) > max_members: + raise ValueError( + f"Refusing to extract archive with {len(members)} members; " + f"limit is {max_members}." + ) + total = 0 + for member in members: + if not (member.isfile() or member.isdir()) or not _within(member.name): + raise ValueError( + f"Refusing to extract unsafe tar member {member.name!r}." + ) + if member.size > max_member_bytes: + raise ValueError( + f"Refusing to extract {member.name!r}: declared size " + f"exceeds {max_member_bytes} bytes." + ) + total += member.size + if total > max_total_bytes: + raise ValueError( + f"Refusing to extract archive larger than " + f"{max_total_bytes} bytes." + ) + tf.extractall(dest, filter="data") + else: + raise ValueError(f"{archive_path} is not a recognized .zip or .tar archive.") + + logger.info(f"Extracted {archive_path} to {dest}") + return dest diff --git a/deepcell_types/utils/_auth.py b/deepcell_types/utils/_auth.py deleted file mode 100644 index f096956..0000000 --- a/deepcell_types/utils/_auth.py +++ /dev/null @@ -1,369 +0,0 @@ -"""User interface to authentication layer for data/models.""" - -import os -import tarfile -import zipfile -import requests -from pathlib import Path -from hashlib import md5, sha256 -from tqdm import tqdm -import logging -import tempfile - -logger = logging.getLogger(__name__) - - -_api_endpoint = "https://users.deepcell.org/api/getData/" -_asset_location = Path.home() / ".deepcell" - -# Hash algorithm is selected by digest length so existing md5-pinned assets keep -# working while new assets can pin the stronger sha256. SHA-256 is preferred for -# new entries (md5 is collision-weak). -_HASH_BY_HEXLEN = {32: ("md5", md5), 64: ("sha256", sha256)} - -# Deliberately generous limits accommodate the public multi-terabyte corpus -# while still bounding hostile or accidentally unbounded responses/archives. -MAX_DOWNLOAD_BYTES = 10 * 2**40 -MAX_ARCHIVE_MEMBERS = 100_000 -MAX_ARCHIVE_MEMBER_BYTES = 2 * 2**40 -MAX_ARCHIVE_TOTAL_BYTES = 10 * 2**40 - - -def _hash_file(fpath, file_hash): - """Return ``(algo_name, hexdigest)`` for ``fpath`` using the algorithm - implied by the length of ``file_hash`` (32 hex → md5, 64 hex → sha256). - - Hashes in chunks so multi-GB assets are not read fully into memory. - """ - try: - algo_name, algo = _HASH_BY_HEXLEN[len(file_hash)] - except KeyError: - raise ValueError( - f"Unrecognized file_hash length {len(file_hash)}; expected an md5 " - "(32 hex chars) or sha256 (64 hex chars) digest." - ) - hasher = algo() - with open(fpath, "rb") as fh: - for chunk in iter(lambda: fh.read(1 << 20), b""): - hasher.update(chunk) - return algo_name, hasher.hexdigest() - - -def fetch_data( - asset_key: str, - cache_subdir=None, - file_hash=None, - *, - max_download_bytes=MAX_DOWNLOAD_BYTES, -): - """Fetch assets through users.deepcell.org authentication system. - - Download assets from the deepcell suite of datasets and models which - require user-authentication. - - .. note:: - - You must have a Deepcell Access Token set as an environment variable - with the name ``DEEPCELL_ACCESS_TOKEN`` in order to access assets. - - Access tokens can be created at _ - - Args: - :param asset_key: Key of the file to download. - The list of available assets can be found on the users.deepcell.org - homepage. - - :param cache_subdir: `str` indicating directory relative to - `~/.deepcell` where downloaded data will be cached. The default is - `None`, which means cache the data in `~/.deepcell`. - - :param file_hash: `str` representing the md5 (32 hex chars) or sha256 - (64 hex chars) checksum of datafile; the algorithm is auto-detected from - the length. The checksum is used to perform data caching. If no checksum - is provided or the checksum differs from that found in the data cache, - the data will be (re)-downloaded. - - :param max_download_bytes: Maximum number of response bytes accepted. - """ - download_location = _asset_location - if cache_subdir is not None: - download_location /= cache_subdir - download_location.mkdir(exist_ok=True, parents=True) - - # Extract the filename from the asset_key, which can be a full path - fname = os.path.split(asset_key)[-1] - fpath = download_location / fname - - # Check for cached data - if file_hash is not None: - logger.info("Checking for cached data") - try: - logger.info(f"Checking {fname} against provided file_hash...") - _, digest = _hash_file(fpath, file_hash) - if digest == file_hash: - logger.info(f"{fname} with hash {file_hash} already available.") - return fpath - logger.info( - f"{fname} with hash {file_hash} not found in {download_location}" - ) - except FileNotFoundError: - pass - elif fpath.exists(): - # No integrity hash is available for this asset (e.g. the large public - # training corpus, for which no authoritative digest is published). - # Reuse an existing local copy rather than re-downloading it — but its - # contents are not verified, so say so loudly. - logger.warning( - f"Reusing cached {fname} at {fpath} WITHOUT an integrity check " - "(no file_hash is available for this asset). Delete the file to " - "force a fresh download." - ) - return fpath - - # Check for access token - access_token = os.environ.get("DEEPCELL_ACCESS_TOKEN") - if access_token is None: - raise ValueError( - "\nDEEPCELL_ACCESS_TOKEN not found.\n" - "Please set your access token to the DEEPCELL_ACCESS_TOKEN\n" - "environment variable.\n" - "For example:\n\n" - "\texport DEEPCELL_ACCESS_TOKEN=.\n\n" - "If you don't yet have a token, you can create one at\n" - "https://users.deepcell.org" - ) - - # Request download URL - headers = {"X-Api-Key": access_token} - logger.info("Making request to server") - resp = requests.post( - _api_endpoint, - headers=headers, - data={"s3_key": asset_key}, - timeout=(30, 300), - ) - - def _safe_json(r): - # Gateways/proxies often return HTML or empty error bodies; don't let a - # JSONDecodeError mask the real HTTP status. - try: - return r.json() - except ValueError: - return {} - - # Raise informative exception for the specific case when the asset_key is - # not found in the bucket - if resp.status_code == 404 and _safe_json(resp).get("error") == "Key not found": - raise ValueError(f"Object {asset_key} not found.") - # Raise informative exception for the specific case when an invalid - # API token is provided. - if resp.status_code == 403 and ( - _safe_json(resp).get("detail") - == "Authentication credentials were not provided." - ): - raise ValueError( - "\n\nThe provided DEEPCELL_ACCESS_TOKEN is not valid.\n" - "The token may be expired - if so, create a new one at\n" - "https://users.deepcell.org" - ) - # Handle all other non-http-200 status - resp.raise_for_status() - - # Parse response - response_data = _safe_json(resp) - if "url" not in response_data: - raise ValueError( - f"Unexpected response from {_api_endpoint} (status {resp.status_code}): " - f"missing download URL. Body starts: {resp.text[:200]!r}" - ) - download_url = response_data["url"] - file_size = response_data.get("size") - # The server-side ``size`` field comes back as a string like "12.3 MB"; parse - # it into bytes for the progress bar, but fall back to an unknown total - # (``None``) rather than aborting an otherwise-valid download if the format - # is unexpected. - suffix_mapping = {"B": 1, "KB": 2**10, "MB": 2**20, "GB": 2**30, "TB": 2**40} - try: - val, suff = str(file_size).split(" ") - file_size_numerical = int(float(val) * suffix_mapping[suff.upper()]) - except (ValueError, KeyError, AttributeError): - file_size_numerical = None - - if file_hash is None: - logger.warning( - f"Downloading {asset_key} WITHOUT an integrity check " - "(no file_hash is available for this asset)." - ) - logger.info(f"Downloading {asset_key} with size {file_size} to {download_location}") - data_req = requests.get( - download_url, - headers={"user-agent": "Wget/1.20 (linux-gnu)"}, - stream=True, - timeout=(30, 300), - ) - data_req.raise_for_status() - - content_length = data_req.headers.get("Content-Length") - if content_length is not None: - try: - declared_bytes = int(content_length) - except ValueError as exc: - raise ValueError( - f"Download for {asset_key} returned an invalid Content-Length: " - f"{content_length!r}." - ) from exc - if declared_bytes > max_download_bytes: - raise ValueError( - f"Download for {asset_key} declares {declared_bytes} bytes, " - f"exceeding the {max_download_bytes}-byte safety limit." - ) - - chunk_size = 4096 - tmp_path = None - try: - with tempfile.NamedTemporaryFile( - mode="wb", dir=download_location, prefix=f".{fname}.", delete=False - ) as raw_fh: - tmp_path = Path(raw_fh.name) - with tqdm.wrapattr( - raw_fh, "write", miniters=1, total=file_size_numerical - ) as fh: - downloaded = 0 - for chunk in data_req.iter_content(chunk_size=chunk_size): - if not chunk: - continue - downloaded += len(chunk) - if downloaded > max_download_bytes: - raise ValueError( - f"Download for {asset_key} exceeded the " - f"{max_download_bytes}-byte safety limit." - ) - fh.write(chunk) - fh.flush() - os.fsync(fh.fileno()) - - if file_hash is not None: - algo_name, actual = _hash_file(tmp_path, file_hash) - if actual != file_hash: - raise ValueError( - f"Integrity check failed for {fname}: " - f"expected {algo_name}={file_hash}, got {algo_name}={actual}. " - "The downloaded file has been removed; please retry." - ) - os.replace(tmp_path, fpath) - tmp_path = None - except BaseException: - # Preserve an existing cache entry; only the unique temporary download - # is removed when transfer or validation fails. - if tmp_path is not None: - try: - tmp_path.unlink(missing_ok=True) - except OSError: - pass - raise - - logger.info(f"Successfully downloaded {fname} to {fpath}") - - return fpath - - -def extract_archive( - archive_path, - dest=None, - *, - max_members=MAX_ARCHIVE_MEMBERS, - max_member_bytes=MAX_ARCHIVE_MEMBER_BYTES, - max_total_bytes=MAX_ARCHIVE_TOTAL_BYTES, -): - """Safely extract a ``.zip`` or ``.tar[.gz]`` archive. - - Rejects members whose resolved path would escape ``dest`` (zip-slip / - tar path-traversal) and rejects tar symlink/hardlink members, so it is - safe to call on archives downloaded from a remote source. - - Parameters - ---------- - archive_path : str or pathlib.Path - Path to a ``.zip`` or ``.tar``/``.tar.gz`` archive. - dest : str or pathlib.Path, optional - Destination directory. Defaults to the archive's parent directory. - - Returns - ------- - pathlib.Path - The destination directory the archive was extracted into. - """ - archive_path = Path(archive_path) - dest = Path(dest) if dest is not None else archive_path.parent - dest.mkdir(parents=True, exist_ok=True) - dest_resolved = dest.resolve() - - def _within(name): - try: - (dest / name).resolve().relative_to(dest_resolved) - return True - except ValueError: - return False - - if zipfile.is_zipfile(archive_path): - with zipfile.ZipFile(archive_path) as zf: - infos = zf.infolist() - if len(infos) > max_members: - raise ValueError( - f"Refusing to extract archive with {len(infos)} members; " - f"limit is {max_members}." - ) - total = 0 - for member in infos: - is_symlink = (member.external_attr >> 16) & 0o170000 == 0o120000 - if not _within(member.filename): - raise ValueError( - f"Refusing to extract {member.filename!r}: path escapes {dest}." - ) - if is_symlink: - raise ValueError( - f"Refusing to extract unsafe zip member {member.filename!r}." - ) - if member.file_size > max_member_bytes: - raise ValueError( - f"Refusing to extract {member.filename!r}: declared size " - f"exceeds {max_member_bytes} bytes." - ) - total += member.file_size - if total > max_total_bytes: - raise ValueError( - f"Refusing to extract archive larger than " - f"{max_total_bytes} bytes." - ) - zf.extractall(dest) - elif tarfile.is_tarfile(archive_path): - with tarfile.open(archive_path) as tf: - members = tf.getmembers() - if len(members) > max_members: - raise ValueError( - f"Refusing to extract archive with {len(members)} members; " - f"limit is {max_members}." - ) - total = 0 - for member in members: - if not (member.isfile() or member.isdir()) or not _within(member.name): - raise ValueError( - f"Refusing to extract unsafe tar member {member.name!r}." - ) - if member.size > max_member_bytes: - raise ValueError( - f"Refusing to extract {member.name!r}: declared size " - f"exceeds {max_member_bytes} bytes." - ) - total += member.size - if total > max_total_bytes: - raise ValueError( - f"Refusing to extract archive larger than " - f"{max_total_bytes} bytes." - ) - tf.extractall(dest, filter="data") - else: - raise ValueError(f"{archive_path} is not a recognized .zip or .tar archive.") - - logger.info(f"Extracted {archive_path} to {dest}") - return dest diff --git a/docs/site/API-key.md b/docs/site/API-key.md index 13968d6..7130bfa 100644 --- a/docs/site/API-key.md +++ b/docs/site/API-key.md @@ -51,6 +51,16 @@ from deepcell_types.utils import download_baseline_checkpoint download_baseline_checkpoint("maps") ``` +All three baselines ship in a single compressed bundle, which holds each one's +weights plus any companion file needed at inference (`maps` ships `_stats.npz`; +`xgboost` ships `.remap.json`). The bundle is unpacked automatically and the +returned list gives the local path of every file for the baseline you asked +for. + +Because the bundle is shared, the first call downloads all three baselines +(646 MB) regardless of which one you request; the other two are then served +from the local cache without a further download. + The Nimbus baseline is not served here: its pretrained weights are distributed upstream, so install it with `pip install nimbus-inference==0.0.5` on Python 3.11 (which fetches the weights diff --git a/pyproject.toml b/pyproject.toml index 68dcb88..8f1ad99 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -25,6 +25,7 @@ dependencies = [ "scipy>=1.9", "scikit-image>=0.20", "tqdm>=4.60", + "deepcell_auth@git+https://github.com/vanvalenlab/deepcell-auth.git@main", "torch>=2.0", ] classifiers = [ diff --git a/scripts/repackage_release_checkpoint.py b/scripts/repackage_release_checkpoint.py index 50f095e..6ea9024 100644 --- a/scripts/repackage_release_checkpoint.py +++ b/scripts/repackage_release_checkpoint.py @@ -14,8 +14,8 @@ archive. The checkpoint's ``config.n_celltypes`` must match the vocabulary size, which guards against pairing a checkpoint with a mismatched vocabulary (the count check that ``load_state_dict`` would otherwise only catch at load -time). Re-packaging changes the file's checksum, so the ``_model_registry`` -entry in ``deepcell_types/utils/__init__.py`` must be updated to the new md5 +time). Re-packaging changes the file's checksum, so the ``deepcell-auth`` +asset manifest (``asset_manifest.yaml``) entry must be updated to the new md5 and the new asset re-uploaded before ``download_model`` will serve it. Usage: diff --git a/tests/test_archive.py b/tests/test_archive.py new file mode 100644 index 0000000..f4bebed --- /dev/null +++ b/tests/test_archive.py @@ -0,0 +1,107 @@ +"""Unit tests for ``deepcell_types.utils._archive.extract_archive``. + +These guard the security-relevant, network-free extraction paths: zip-slip / +tar-traversal / tar-symlink rejection and the member-count and size bounds. +""" + +import io +import tarfile +import zipfile + +import pytest + +from deepcell_types.utils._archive import extract_archive + + +def test_extract_archive_accepts_benign_zip(tmp_path): + archive = tmp_path / "ok.zip" + with zipfile.ZipFile(archive, "w") as zf: + zf.writestr("inner/ok.txt", "hello") + dest = tmp_path / "out" + extract_archive(archive, dest) + assert (dest / "inner" / "ok.txt").read_text() == "hello" + + +def test_extract_archive_rejects_zip_slip(tmp_path): + archive = tmp_path / "evil.zip" + with zipfile.ZipFile(archive, "w") as zf: + zf.writestr("../escape.txt", "pwned") + dest = tmp_path / "out" + with pytest.raises(ValueError, match="escapes"): + extract_archive(archive, dest) + assert not (tmp_path / "escape.txt").exists() + + +def test_extract_archive_accepts_benign_tar(tmp_path): + archive = tmp_path / "ok.tar" + data = b"hi" + with tarfile.open(archive, "w") as tf: + info = tarfile.TarInfo("inner/ok.txt") + info.size = len(data) + tf.addfile(info, io.BytesIO(data)) + dest = tmp_path / "out" + extract_archive(archive, dest) + assert (dest / "inner" / "ok.txt").read_bytes() == data + + +def test_extract_archive_rejects_tar_symlink_member(tmp_path): + archive = tmp_path / "evil.tar" + with tarfile.open(archive, "w") as tf: + link = tarfile.TarInfo("link") + link.type = tarfile.SYMTYPE + link.linkname = "/etc/passwd" + tf.addfile(link) + dest = tmp_path / "out" + with pytest.raises(ValueError, match="unsafe tar member"): + extract_archive(archive, dest) + + +def test_extract_archive_rejects_tar_traversal_member(tmp_path): + archive = tmp_path / "evil2.tar" + data = b"x" + with tarfile.open(archive, "w") as tf: + info = tarfile.TarInfo("../escape.txt") + info.size = len(data) + tf.addfile(info, io.BytesIO(data)) + dest = tmp_path / "out" + with pytest.raises(ValueError, match="unsafe tar member"): + extract_archive(archive, dest) + + +def test_extract_archive_rejects_non_archive(tmp_path): + plain = tmp_path / "notes.txt" + plain.write_text("not an archive") + with pytest.raises(ValueError, match="not a recognized"): + extract_archive(plain, tmp_path / "out") + + +def test_extract_archive_enforces_member_and_size_limits(tmp_path): + archive = tmp_path / "bounded.zip" + with zipfile.ZipFile(archive, "w") as zf: + zf.writestr("one", b"123") + zf.writestr("two", b"456") + + with pytest.raises(ValueError, match="2 members"): + extract_archive(archive, tmp_path / "members", max_members=1) + with pytest.raises(ValueError, match="declared size"): + extract_archive(archive, tmp_path / "member-size", max_member_bytes=2) + with pytest.raises(ValueError, match="larger than"): + extract_archive(archive, tmp_path / "total-size", max_total_bytes=5) + + +def test_extract_archive_enforces_member_and_size_limits_tar(tmp_path): + # The tar branch duplicates the zip branch's bound checks; exercise it + # directly so a tar-only regression can't slip through. + archive = tmp_path / "bounded.tar" + with tarfile.open(archive, "w") as tf: + for name, data in (("one", b"123"), ("two", b"456")): + info = tarfile.TarInfo(name) + info.size = len(data) + tf.addfile(info, io.BytesIO(data)) + + with pytest.raises(ValueError, match="2 members"): + extract_archive(archive, tmp_path / "members", max_members=1) + with pytest.raises(ValueError, match="declared size"): + extract_archive(archive, tmp_path / "member-size", max_member_bytes=2) + with pytest.raises(ValueError, match="larger than"): + extract_archive(archive, tmp_path / "total-size", max_total_bytes=5) diff --git a/tests/test_auth.py b/tests/test_auth.py deleted file mode 100644 index 56ea0e5..0000000 --- a/tests/test_auth.py +++ /dev/null @@ -1,336 +0,0 @@ -"""Unit tests for the download / integrity / archive-extraction layer -(``deepcell_types.utils._auth`` + the model/baseline registries). These guard -network-free, security-relevant code paths that were previously untested: the -hash-algorithm dispatch, the zip-slip / tar-traversal / tar-symlink rejection in -``extract_archive``, the cache-hit and missing-token branches of ``fetch_data``, -and the registry digest shapes. -""" - -import hashlib -import io -import tarfile -import zipfile -import requests - -import pytest - -from deepcell_types.utils import _auth -from deepcell_types.utils._auth import _hash_file, extract_archive, fetch_data -from deepcell_types.utils import ( - _latest, - _model_registry, - list_model_versions, -) - - -# --- _hash_file: algorithm dispatch by digest length ------------------------ - - -def test_hash_file_dispatches_md5_and_sha256(tmp_path): - f = tmp_path / "blob.bin" - payload = b"deepcell-types integrity check" - f.write_bytes(payload) - - algo, digest = _hash_file(f, "0" * 32) # 32 hex -> md5 - assert algo == "md5" - assert digest == hashlib.md5(payload).hexdigest() - - algo, digest = _hash_file(f, "0" * 64) # 64 hex -> sha256 - assert algo == "sha256" - assert digest == hashlib.sha256(payload).hexdigest() - - -def test_hash_file_rejects_unknown_digest_length(tmp_path): - f = tmp_path / "blob.bin" - f.write_bytes(b"x") - with pytest.raises(ValueError, match="Unrecognized file_hash length"): - _hash_file(f, "abc123") # neither 32 nor 64 hex chars - - -# --- extract_archive: path-traversal / symlink rejection -------------------- - - -def test_extract_archive_accepts_benign_zip(tmp_path): - archive = tmp_path / "ok.zip" - with zipfile.ZipFile(archive, "w") as zf: - zf.writestr("inner/ok.txt", "hello") - dest = tmp_path / "out" - extract_archive(archive, dest) - assert (dest / "inner" / "ok.txt").read_text() == "hello" - - -def test_extract_archive_rejects_zip_slip(tmp_path): - archive = tmp_path / "evil.zip" - with zipfile.ZipFile(archive, "w") as zf: - zf.writestr("../escape.txt", "pwned") - dest = tmp_path / "out" - with pytest.raises(ValueError, match="escapes"): - extract_archive(archive, dest) - assert not (tmp_path / "escape.txt").exists() - - -def test_extract_archive_accepts_benign_tar(tmp_path): - archive = tmp_path / "ok.tar" - data = b"hi" - with tarfile.open(archive, "w") as tf: - info = tarfile.TarInfo("inner/ok.txt") - info.size = len(data) - tf.addfile(info, io.BytesIO(data)) - dest = tmp_path / "out" - extract_archive(archive, dest) - assert (dest / "inner" / "ok.txt").read_bytes() == data - - -def test_extract_archive_rejects_tar_symlink_member(tmp_path): - archive = tmp_path / "evil.tar" - with tarfile.open(archive, "w") as tf: - link = tarfile.TarInfo("link") - link.type = tarfile.SYMTYPE - link.linkname = "/etc/passwd" - tf.addfile(link) - dest = tmp_path / "out" - with pytest.raises(ValueError, match="unsafe tar member"): - extract_archive(archive, dest) - - -def test_extract_archive_rejects_tar_traversal_member(tmp_path): - archive = tmp_path / "evil2.tar" - data = b"x" - with tarfile.open(archive, "w") as tf: - info = tarfile.TarInfo("../escape.txt") - info.size = len(data) - tf.addfile(info, io.BytesIO(data)) - dest = tmp_path / "out" - with pytest.raises(ValueError, match="unsafe tar member"): - extract_archive(archive, dest) - - -def test_extract_archive_rejects_non_archive(tmp_path): - plain = tmp_path / "notes.txt" - plain.write_text("not an archive") - with pytest.raises(ValueError, match="not a recognized"): - extract_archive(plain, tmp_path / "out") - - -def test_extract_archive_enforces_member_and_size_limits(tmp_path): - archive = tmp_path / "bounded.zip" - with zipfile.ZipFile(archive, "w") as zf: - zf.writestr("one", b"123") - zf.writestr("two", b"456") - - with pytest.raises(ValueError, match="2 members"): - extract_archive(archive, tmp_path / "members", max_members=1) - with pytest.raises(ValueError, match="declared size"): - extract_archive(archive, tmp_path / "member-size", max_member_bytes=2) - with pytest.raises(ValueError, match="larger than"): - extract_archive(archive, tmp_path / "total-size", max_total_bytes=5) - - -def test_extract_archive_enforces_member_and_size_limits_tar(tmp_path): - # The tar branch duplicates the zip branch's bound checks; exercise it - # directly so a tar-only regression can't slip through. - archive = tmp_path / "bounded.tar" - with tarfile.open(archive, "w") as tf: - for name, data in (("one", b"123"), ("two", b"456")): - info = tarfile.TarInfo(name) - info.size = len(data) - tf.addfile(info, io.BytesIO(data)) - - with pytest.raises(ValueError, match="2 members"): - extract_archive(archive, tmp_path / "members", max_members=1) - with pytest.raises(ValueError, match="declared size"): - extract_archive(archive, tmp_path / "member-size", max_member_bytes=2) - with pytest.raises(ValueError, match="larger than"): - extract_archive(archive, tmp_path / "total-size", max_total_bytes=5) - - -# --- fetch_data: cache-hit and missing-token branches (no network) ---------- - - -def test_fetch_data_returns_cached_file_on_hash_match(tmp_path, monkeypatch): - monkeypatch.setattr(_auth, "_asset_location", tmp_path) - cache_dir = tmp_path / "models" - cache_dir.mkdir() - payload = b"cached checkpoint bytes" - (cache_dir / "model.pt").write_bytes(payload) - digest = hashlib.md5(payload).hexdigest() - - # Hash matches -> returns the cached path without ever needing a token. - monkeypatch.delenv("DEEPCELL_ACCESS_TOKEN", raising=False) - out = fetch_data("models/model.pt", cache_subdir="models", file_hash=digest) - assert out == cache_dir / "model.pt" - - -def test_fetch_data_requires_token_on_cache_miss(tmp_path, monkeypatch): - monkeypatch.setattr(_auth, "_asset_location", tmp_path) - monkeypatch.delenv("DEEPCELL_ACCESS_TOKEN", raising=False) - # No cached file -> falls through to the token check, which must raise - # (never hits the network in the test). - with pytest.raises(ValueError, match="DEEPCELL_ACCESS_TOKEN"): - fetch_data("models/missing.pt", cache_subdir="models", file_hash="0" * 32) - - -class _Response: - def __init__(self, *, status=200, json_data=None, text="", chunks=(), headers=None): - self.status_code = status - self._json_data = json_data - self.text = text - self._chunks = chunks - self.headers = headers or {} - - def json(self): - if isinstance(self._json_data, Exception): - raise self._json_data - return self._json_data - - def raise_for_status(self): - if self.status_code >= 400: - raise requests.HTTPError(f"HTTP {self.status_code}") - - def iter_content(self, chunk_size): - del chunk_size - yield from self._chunks - - -def _mock_download(monkeypatch, post_response, get_response): - monkeypatch.setenv("DEEPCELL_ACCESS_TOKEN", "test-token") - monkeypatch.setattr(_auth.requests, "post", lambda *args, **kwargs: post_response) - monkeypatch.setattr(_auth.requests, "get", lambda *args, **kwargs: get_response) - - -def test_fetch_data_reports_non_json_http_error(tmp_path, monkeypatch): - monkeypatch.setattr(_auth, "_asset_location", tmp_path) - _mock_download( - monkeypatch, - _Response(status=502, json_data=ValueError("not json"), text="gateway"), - _Response(), - ) - with pytest.raises(requests.HTTPError, match="502"): - fetch_data("models/model.pt") - - -def test_fetch_data_rejects_missing_download_url(tmp_path, monkeypatch): - monkeypatch.setattr(_auth, "_asset_location", tmp_path) - _mock_download(monkeypatch, _Response(json_data={}), _Response()) - with pytest.raises(ValueError, match="missing download URL"): - fetch_data("models/model.pt") - - -def test_fetch_data_preserves_cache_on_interrupted_refresh(tmp_path, monkeypatch): - monkeypatch.setattr(_auth, "_asset_location", tmp_path) - cached = tmp_path / "models" / "model.pt" - cached.parent.mkdir() - cached.write_bytes(b"previous") - - def interrupted(): - yield b"partial" - raise requests.ConnectionError("connection lost") - - _mock_download( - monkeypatch, - _Response(json_data={"url": "https://example.invalid/model"}), - _Response(chunks=interrupted()), - ) - with pytest.raises(requests.ConnectionError, match="connection lost"): - fetch_data("models/model.pt", cache_subdir="models", file_hash="0" * 32) - assert cached.read_bytes() == b"previous" - assert list(cached.parent.glob(".model.pt.*")) == [] - - -def test_fetch_data_rejects_oversized_stream_and_preserves_cache(tmp_path, monkeypatch): - monkeypatch.setattr(_auth, "_asset_location", tmp_path) - cached = tmp_path / "model.pt" - cached.write_bytes(b"previous") - _mock_download( - monkeypatch, - _Response(json_data={"url": "https://example.invalid/model"}), - _Response(chunks=[b"123", b"456"]), - ) - with pytest.raises(ValueError, match="safety limit"): - fetch_data("model.pt", file_hash="0" * 32, max_download_bytes=5) - assert cached.read_bytes() == b"previous" - - -def test_fetch_data_rejects_digest_mismatch_without_replacing_cache( - tmp_path, monkeypatch -): - monkeypatch.setattr(_auth, "_asset_location", tmp_path) - cached = tmp_path / "model.pt" - cached.write_bytes(b"previous") - _mock_download( - monkeypatch, - _Response(json_data={"url": "https://example.invalid/model"}), - _Response(chunks=[b"wrong"]), - ) - with pytest.raises(ValueError, match="Integrity check failed"): - fetch_data("model.pt", file_hash=hashlib.sha256(b"expected").hexdigest()) - assert cached.read_bytes() == b"previous" - - -def test_fetch_data_downloads_and_lands_atomically(tmp_path, monkeypatch): - # The happy path: a valid streamed download with a matching hash lands at - # download_location/fname with the exact bytes and leaves no temp file. - monkeypatch.setattr(_auth, "_asset_location", tmp_path) - payload = b"fresh model bytes" - digest = hashlib.md5(payload).hexdigest() - _mock_download( - monkeypatch, - _Response(json_data={"url": "https://example.invalid/model"}), - _Response(chunks=[payload[:5], payload[5:]]), - ) - out = fetch_data("models/model.pt", cache_subdir="models", file_hash=digest) - assert out == tmp_path / "models" / "model.pt" - assert out.read_bytes() == payload - assert list(out.parent.glob(".model.pt.*")) == [] - - -def test_fetch_data_reuses_unhashed_cache_without_network(tmp_path, monkeypatch): - # With no file_hash and an existing cached file, fetch_data must return the - # cached path directly — never reaching the token check or the network. - monkeypatch.setattr(_auth, "_asset_location", tmp_path) - monkeypatch.delenv("DEEPCELL_ACCESS_TOKEN", raising=False) - cache_dir = tmp_path / "data" - cache_dir.mkdir() - (cache_dir / "corpus.zip").write_bytes(b"corpus") - out = fetch_data("data/corpus.zip", cache_subdir="data") - assert out == cache_dir / "corpus.zip" - - -def test_fetch_data_rejects_oversized_content_length_before_writing( - tmp_path, monkeypatch -): - monkeypatch.setattr(_auth, "_asset_location", tmp_path) - _mock_download( - monkeypatch, - _Response(json_data={"url": "https://example.invalid/model"}), - _Response(headers={"Content-Length": "999"}, chunks=[b"x"]), - ) - with pytest.raises(ValueError, match="exceeding the"): - fetch_data("model.pt", max_download_bytes=5) - assert not (tmp_path / "model.pt").exists() - - -def test_fetch_data_rejects_non_numeric_content_length(tmp_path, monkeypatch): - monkeypatch.setattr(_auth, "_asset_location", tmp_path) - _mock_download( - monkeypatch, - _Response(json_data={"url": "https://example.invalid/model"}), - _Response(headers={"Content-Length": "not-a-number"}, chunks=[b"x"]), - ) - with pytest.raises(ValueError, match="invalid Content-Length"): - fetch_data("model.pt") - assert not (tmp_path / "model.pt").exists() - - -# --- model registry shape --------------------------------------------------- - - -def test_model_registry_entries_are_well_formed(): - assert list_model_versions()[0] == _latest - assert _latest in _model_registry - for version, entry in _model_registry.items(): - assert isinstance(version, str) - filename, file_hash = entry - assert isinstance(filename, str) and filename.endswith(".pt") - assert len(file_hash) in (32, 64) - int(file_hash, 16) # valid hex digest diff --git a/tests/test_download_delegation.py b/tests/test_download_delegation.py new file mode 100644 index 0000000..7a72deb --- /dev/null +++ b/tests/test_download_delegation.py @@ -0,0 +1,168 @@ +"""The static version / baseline lists in ``deepcell_types.utils`` mirror +deepcell-auth's bundled asset manifest. These read the packaged YAML via +``load_manifest`` (no network) and fail loudly if deepcell-auth adds or renames +an entry this repo hasn't mirrored -- the failure mode that would otherwise +surface as a ``KeyError`` from inside ``deepcell_auth`` at download time. +""" + +import io +import tarfile + +import pytest + +import deepcell_auth +from deepcell_auth._auth import load_manifest + +from deepcell_types.utils import ( + _DEFAULT_MODEL_VERSION, + _LEGACY_MODEL_VERSIONS, + _MODEL_VERSIONS, + download_baseline_checkpoint, + download_model, + list_baseline_names, + list_model_versions, +) + + +def test_model_versions_match_manifest(): + manifest_versions = set(load_manifest()["models"]["deepcell-types"]) + assert set(_MODEL_VERSIONS) | set(_LEGACY_MODEL_VERSIONS) == manifest_versions + + +def test_baseline_names_match_manifest(): + assert set(list_baseline_names()) == set( + load_manifest()["models"]["deepcell-types-baselines"] + ) + + +def test_default_version_present_in_manifest(): + assert _DEFAULT_MODEL_VERSION in load_manifest()["models"]["deepcell-types"] + assert list_model_versions()[0] == _DEFAULT_MODEL_VERSION + + +def test_manifest_records_carry_key_and_hash(): + models = load_manifest()["models"] + for version in _MODEL_VERSIONS: + record = models["deepcell-types"][version] + assert record["asset_key"].startswith("models/") + assert len(record["asset_hash"]) in (32, 64) + int(record["asset_hash"], 16) # valid hex digest + for records in models["deepcell-types-baselines"].values(): + # Each baseline resolves to exactly one .tar.gz bundle; + # download_baseline_checkpoint unpacks it and returns the files inside, + # so a manifest that split a baseline back into loose per-file records + # would break the unpacking. + assert len(records) == 1 + (record,) = records + assert record["asset_key"].startswith("models/") + assert record["asset_key"].endswith(".tar.gz") + assert len(record["asset_hash"]) in (32, 64) + int(record["asset_hash"], 16) + + +def test_all_baselines_share_one_bundle(): + # All three baselines ship in a single archive, so every name must resolve + # to the same asset -- otherwise requesting one baseline would unpack a + # bundle that does not contain the others' subdirectories. + baselines = load_manifest()["models"]["deepcell-types-baselines"] + assets = {records[0]["asset_key"] for records in baselines.values()} + hashes = {records[0]["asset_hash"] for records in baselines.values()} + assert len(assets) == 1 + assert len(hashes) == 1 + + +@pytest.mark.parametrize("version", _LEGACY_MODEL_VERSIONS) +def test_download_model_rejects_legacy_clip_versions(version): + # Served for reproducibility with the matching historical commit, but not + # loadable by this code -- reject before delegating (so: no network). + with pytest.raises(ValueError, match="Unknown model version"): + download_model(version=version) + + +def test_download_baseline_rejects_nimbus(): + with pytest.raises(ValueError, match="distributed upstream"): + download_baseline_checkpoint("nimbus") + + +_STEM = "deepcell-types_baselines_2026-06-30" +_TREE = { + "cellsighter": {"deepcell-types_baseline-cellsighter.pth": b"cs-weights"}, + "maps": { + "deepcell-types_baseline-maps.pth": b"maps-weights", + "deepcell-types_baseline-maps_stats.npz": b"maps-stats", + }, + "xgboost": { + "deepcell-types_baseline-xgboost.json": b"xgb-booster", + "deepcell-types_baseline-xgboost.remap.json": b"xgb-remap", + }, +} + + +def _make_bundle(models_dir, tree, stem=_STEM): + """Write a ``.tar.gz`` laid out the way the served bundle is: a single + top-level ``/`` directory with one subdirectory per baseline.""" + archive = models_dir / f"{stem}.tar.gz" + with tarfile.open(archive, "w:gz") as tf: + for baseline, members in tree.items(): + for filename, payload in members.items(): + info = tarfile.TarInfo(f"{stem}/{baseline}/{filename}") + info.size = len(payload) + tf.addfile(info, io.BytesIO(payload)) + return archive + + +@pytest.mark.parametrize("baseline", sorted(_TREE)) +def test_download_baseline_returns_only_that_baselines_files( + tmp_path, monkeypatch, baseline +): + archive = _make_bundle(tmp_path, _TREE) + monkeypatch.setattr( + deepcell_auth, "download_deepcell_types_baseline", lambda name: [archive] + ) + + paths = download_baseline_checkpoint(baseline) + + # One shared archive, but a request returns only the requested baseline -- + # and sorting puts weights ahead of the companion file, matching the order + # the loose per-file assets were declared in before bundling. + expected = _TREE[baseline] + assert [p.name for p in paths] == sorted(expected) + assert {p.name: p.read_bytes() for p in paths} == expected + assert all(p.parent.name == baseline for p in paths) + + +def test_download_baseline_reuses_unpacked_bundle(tmp_path, monkeypatch): + # Extraction is skipped once the bundle directory exists, so a later call + # succeeds even if the archive itself is gone from the cache. This is also + # what makes the other two baselines free after the first download. + archive = _make_bundle(tmp_path, _TREE) + monkeypatch.setattr( + deepcell_auth, "download_deepcell_types_baseline", lambda name: [archive] + ) + + first = download_baseline_checkpoint("cellsighter") + archive.unlink() + assert download_baseline_checkpoint("cellsighter") == first + assert [p.name for p in download_baseline_checkpoint("maps")] == sorted( + _TREE["maps"] + ) + + +@pytest.mark.parametrize("member_type", [tarfile.DIRTYPE, tarfile.REGTYPE]) +def test_download_baseline_rejects_bundle_without_checkpoints( + tmp_path, monkeypatch, member_type +): + # A bundle missing the requested baseline's subdirectory, or holding a + # plain file where that directory belongs, should surface a readable error + # rather than a NotADirectoryError from the unpack path. + archive = tmp_path / f"{_STEM}.tar.gz" + with tarfile.open(archive, "w:gz") as tf: + info = tarfile.TarInfo(f"{_STEM}/xgboost") + info.type = member_type + tf.addfile(info, io.BytesIO(b"") if member_type == tarfile.REGTYPE else None) + monkeypatch.setattr( + deepcell_auth, "download_deepcell_types_baseline", lambda name: [archive] + ) + + with pytest.raises(ValueError, match="did not unpack to a directory"): + download_baseline_checkpoint("xgboost")