474 lines
20 KiB
Python
474 lines
20 KiB
Python
"""Reading a `.sientia` bundle: the seven gates, in order, in one function.
|
|
|
|
`BundleReader.open()` is the only entry point. Gate ordering is a property of that function rather
|
|
than of the caller, so it cannot be reordered by accident and can be unit-tested without Temporal:
|
|
|
|
1. object-level limits (`check_object_limits`, called by the download activity before the body is
|
|
fetched — the only gate that does not need the file)
|
|
2. header validation, including the chunk-stream digest, before any key derivation
|
|
3. AEAD decryption — the integrity check; a wrong password and a tampered file fail identically
|
|
4. archive inspection over the central directory, writing nothing to disk
|
|
5. extraction into a fresh, isolated directory
|
|
6. structure and schema of the extracted tree
|
|
7. content policy over the file set
|
|
|
|
Nothing here logs, notifies or emits a metric: rejections are raised as `BundleGateError` and the
|
|
activity layer owns the one failure channel.
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import json
|
|
import os
|
|
import shutil
|
|
import stat
|
|
import time
|
|
import unicodedata
|
|
import zipfile
|
|
from dataclasses import dataclass, field
|
|
from pathlib import Path
|
|
from typing import Any
|
|
|
|
import nacl.bindings as sodium
|
|
import nacl.exceptions
|
|
|
|
from laborious.utils.bundle import format as fmt
|
|
from laborious.utils.bundle.steps import GATE_TO_STEP, ImportStep
|
|
|
|
# Gate 6 — the layout `bundle.build_zip` writes and `model_repository.py` reads back.
|
|
METADATA_FILE_NAME = 'metadata.json'
|
|
ARTIFACTS_DIR_NAME = 'artifacts'
|
|
REQUIRED_TOP_LEVEL_KEYS = frozenset({'parameters', 'metadata'})
|
|
REQUIRED_METADATA_FIELDS = (
|
|
'model_name',
|
|
'experiment_name',
|
|
'run_id',
|
|
'model_version',
|
|
'model_project',
|
|
'export_timestamp',
|
|
)
|
|
REQUIRED_PARAMETERS = ('target_variable',)
|
|
REQUIRED_MODEL_DIRS = ('prediction_model', 'data_model')
|
|
REQUIRED_MODEL_FILES = ('MLmodel', 'model.pkl')
|
|
|
|
# The import record's `experiment_name` column is `VARCHAR(50)` with a length-3 floor, and the name
|
|
# is written onto the record before provisioning starts — so the column's limits are enforced here,
|
|
# not three activities later as a constraint violation.
|
|
EXPERIMENT_NAME_MIN_LENGTH = 3
|
|
EXPERIMENT_NAME_MAX_LENGTH = 50
|
|
|
|
# Gate 7 — an MLflow 2.x artifact tree for this platform: the model directories' known files, the
|
|
# model card, the CSVs the training template logs and the bundle's own metadata.
|
|
ALLOWED_FILE_NAMES = frozenset(
|
|
{
|
|
'MLmodel',
|
|
'metadata.json',
|
|
'model.pkl',
|
|
'conda.yaml',
|
|
'python_env.yaml',
|
|
'requirements.txt',
|
|
'model_card.json',
|
|
'model_card.svg',
|
|
}
|
|
)
|
|
ALLOWED_FILE_SUFFIXES = frozenset({'.csv', '.json', '.yaml', '.yml', '.txt', '.pkl', '.svg', '.md'})
|
|
|
|
_DECRYPT_FAILURE_MESSAGE = (
|
|
'could not open this bundle: wrong password, or the file is damaged or was altered'
|
|
)
|
|
|
|
|
|
class BundleGateError(Exception):
|
|
"""A gate rejected the bundle.
|
|
|
|
Carries the gate number and the import step that gate maps onto, so the workflow never has to
|
|
guess which step to record.
|
|
"""
|
|
|
|
def __init__(self, message: str, *, gate: int, step: ImportStep | None = None) -> None:
|
|
super().__init__(message)
|
|
self.gate = gate
|
|
self.step = step if step is not None else GATE_TO_STEP[gate]
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class BundleLimits:
|
|
"""Ceilings enforced by gates 1 and 4. Values come from `build_import_config()`."""
|
|
|
|
max_object_bytes: int = 1024 * 1024 * 1024
|
|
max_entries: int = 5000
|
|
max_uncompressed_bytes: int = 4 * 1024 * 1024 * 1024
|
|
max_compression_ratio: float = 200.0
|
|
object_prefix: str = 'imported_models/'
|
|
object_suffix: str = '.sientia'
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class OpenedBundle:
|
|
"""What a bundle yields once every gate has passed. Nothing secret is carried."""
|
|
|
|
extracted_dir: Path
|
|
metadata: dict[str, Any] = field(default_factory=dict)
|
|
parameters: dict[str, Any] = field(default_factory=dict)
|
|
digest: str = ''
|
|
header: fmt.BundleHeader | None = None
|
|
# How long Argon2id took, so the worker's most expensive step is measurable without the
|
|
# activity layer having to reach inside the reader.
|
|
kdf_seconds: float = 0.0
|
|
|
|
|
|
class BundleReader:
|
|
"""Opens `.sientia` bundles produced by the exporter in `sientia-projects-templates`."""
|
|
|
|
def __init__(self, limits: BundleLimits | None = None) -> None:
|
|
self.limits = limits or BundleLimits()
|
|
|
|
# ------------------------------------------------------------------ gate 1
|
|
|
|
def check_object_limits(self, *, size: int, key: str) -> None:
|
|
"""Gate 1, object level: reject before the object's body is fetched.
|
|
|
|
Args:
|
|
size: the object's size as reported by `stat_object`.
|
|
key: the object key.
|
|
|
|
Raises:
|
|
BundleGateError: oversized object, key outside the import prefix, or wrong suffix.
|
|
"""
|
|
if size > self.limits.max_object_bytes:
|
|
raise BundleGateError(
|
|
f'uploaded object is {size} bytes, above the '
|
|
f'{self.limits.max_object_bytes} byte limit',
|
|
gate=1,
|
|
)
|
|
if size <= 0:
|
|
raise BundleGateError('uploaded object is empty', gate=1)
|
|
if not key.startswith(self.limits.object_prefix):
|
|
raise BundleGateError('uploaded object is not under the import prefix', gate=1)
|
|
if not key.endswith(self.limits.object_suffix):
|
|
raise BundleGateError(
|
|
f'uploaded object does not end in {self.limits.object_suffix}', gate=1
|
|
)
|
|
|
|
# --------------------------------------------------------------------- open
|
|
|
|
def open(
|
|
self,
|
|
encrypted_path: Path | str,
|
|
password: str,
|
|
extract_root: Path | str,
|
|
work_dir: Path | str | None = None,
|
|
) -> OpenedBundle:
|
|
"""Run gates 2 to 7 in order and return the opened bundle.
|
|
|
|
Args:
|
|
encrypted_path: the downloaded `.sientia` file.
|
|
password: the plaintext bundle password. Never stored, never logged, never returned.
|
|
extract_root: the directory to extract into. It MUST NOT exist yet (gate 5).
|
|
work_dir: where the decrypted zip is written. Defaults to `extract_root`'s parent.
|
|
|
|
Returns:
|
|
OpenedBundle: extracted tree, `metadata`, `parameters`, digest and header.
|
|
|
|
Raises:
|
|
BundleGateError: with the gate that rejected the bundle and its import step.
|
|
"""
|
|
encrypted = Path(encrypted_path)
|
|
root = Path(extract_root)
|
|
staging = Path(work_dir) if work_dir is not None else root.parent
|
|
|
|
header = self._gate_2_header(encrypted)
|
|
zip_path, kdf_seconds = self._gate_3_decrypt(encrypted, password, header, staging)
|
|
try:
|
|
self._gate_4_inspect(zip_path)
|
|
self._gate_5_extract(zip_path, root)
|
|
metadata, parameters = self._gate_6_structure(root)
|
|
self._gate_7_content_policy(root)
|
|
except BaseException:
|
|
# A rejection after decryption leaves nothing extracted behind (gate 5's contract) and
|
|
# never leaves a decrypted zip on the worker for the next import to trip over.
|
|
shutil.rmtree(root, ignore_errors=True)
|
|
zip_path.unlink(missing_ok=True)
|
|
raise
|
|
zip_path.unlink(missing_ok=True)
|
|
|
|
return OpenedBundle(
|
|
extracted_dir=root,
|
|
metadata=metadata,
|
|
parameters=parameters,
|
|
digest=header.ciphertext_digest.hex(),
|
|
header=header,
|
|
kdf_seconds=kdf_seconds,
|
|
)
|
|
|
|
# ------------------------------------------------------------------ gate 2
|
|
|
|
def _gate_2_header(self, encrypted: Path) -> fmt.BundleHeader:
|
|
"""Parse the header and check the chunk-stream digest, before deriving anything."""
|
|
try:
|
|
header = fmt.read_header(encrypted)
|
|
fmt.verify_ciphertext_digest(encrypted, header)
|
|
except fmt.BundleFormatError as error:
|
|
raise BundleGateError(str(error), gate=2) from error
|
|
except OSError as error:
|
|
raise BundleGateError('downloaded bundle could not be read', gate=2) from error
|
|
return header
|
|
|
|
# ------------------------------------------------------------------ gate 3
|
|
|
|
def _gate_3_decrypt(
|
|
self, encrypted: Path, password: str, header: fmt.BundleHeader, staging: Path
|
|
) -> tuple[Path, float]:
|
|
"""Derive the key with the header's parameters and decrypt the stream to a file.
|
|
|
|
The AEAD *is* the integrity check: a wrong password and a tampered byte fail here with the
|
|
same message, which never claims to distinguish them.
|
|
"""
|
|
staging.mkdir(parents=True, exist_ok=True)
|
|
zip_path = staging / 'bundle.zip'
|
|
try:
|
|
started = time.monotonic()
|
|
key = self._derive_key(password, header)
|
|
kdf_seconds = time.monotonic() - started
|
|
self._pull_stream(encrypted, key, zip_path)
|
|
except BundleGateError:
|
|
zip_path.unlink(missing_ok=True)
|
|
raise
|
|
except (nacl.exceptions.CryptoError, ValueError, RuntimeError) as error:
|
|
zip_path.unlink(missing_ok=True)
|
|
raise BundleGateError(_DECRYPT_FAILURE_MESSAGE, gate=3) from error
|
|
except OSError as error:
|
|
zip_path.unlink(missing_ok=True)
|
|
raise BundleGateError('decrypted bundle could not be written', gate=3) from error
|
|
return zip_path, kdf_seconds
|
|
|
|
@staticmethod
|
|
def _derive_key(password: str, header: fmt.BundleHeader) -> bytes:
|
|
"""Argon2id over the NFC-normalised password, with the header's own cost parameters.
|
|
|
|
The password is normalised and **not** trimmed, matching the producer's `_derive_key`: the
|
|
same password typed on another OS must derive the same key, while a trailing space is part
|
|
of the secret.
|
|
"""
|
|
normalized = unicodedata.normalize('NFC', password)
|
|
return sodium.crypto_pwhash_alg(
|
|
fmt.KEY_SIZE,
|
|
normalized.encode('utf-8'),
|
|
header.salt,
|
|
header.kdf_opslimit,
|
|
header.kdf_memlimit_bytes,
|
|
sodium.crypto_pwhash_ALG_ARGON2ID13,
|
|
)
|
|
|
|
@staticmethod
|
|
def _pull_stream(encrypted: Path, key: bytes, zip_path: Path) -> None:
|
|
"""Decrypt the chunk stream into `zip_path`, requiring the FINAL tag."""
|
|
chunk_size = fmt.CHUNK_SIZE + fmt.ABYTES
|
|
state = sodium.crypto_secretstream_xchacha20poly1305_state()
|
|
with encrypted.open('rb') as source, zip_path.open('wb') as target:
|
|
source.seek(fmt.HEADER_SIZE)
|
|
stream_header = source.read(fmt.STREAM_HEADER_SIZE)
|
|
sodium.crypto_secretstream_xchacha20poly1305_init_pull(state, stream_header, key)
|
|
last_tag: int | None = None
|
|
while chunk := source.read(chunk_size):
|
|
plaintext, last_tag = sodium.crypto_secretstream_xchacha20poly1305_pull(
|
|
state, chunk
|
|
)
|
|
target.write(plaintext)
|
|
if last_tag != sodium.crypto_secretstream_xchacha20poly1305_TAG_FINAL:
|
|
# A stream that stops without its FINAL tag was cut short. The digest already catches
|
|
# truncation of a stored file; this catches a stream that was framed to look complete.
|
|
raise BundleGateError(_DECRYPT_FAILURE_MESSAGE, gate=3)
|
|
|
|
# ------------------------------------------------------------------ gate 4
|
|
|
|
def _gate_4_inspect(self, zip_path: Path) -> None:
|
|
"""Inspect the central directory only. Nothing is written to disk by this gate."""
|
|
try:
|
|
with zipfile.ZipFile(zip_path) as archive:
|
|
entries = archive.infolist()
|
|
except zipfile.BadZipFile as error:
|
|
raise BundleGateError('bundle content is not a readable archive', gate=4) from error
|
|
|
|
if len(entries) > self.limits.max_entries:
|
|
raise BundleGateError(
|
|
f'archive declares {len(entries)} entries, above the '
|
|
f'{self.limits.max_entries} limit',
|
|
gate=4,
|
|
)
|
|
if not entries:
|
|
raise BundleGateError('archive is empty', gate=4)
|
|
|
|
uncompressed = sum(entry.file_size for entry in entries)
|
|
compressed = sum(entry.compress_size for entry in entries)
|
|
if uncompressed > self.limits.max_uncompressed_bytes:
|
|
raise BundleGateError(
|
|
f'archive declares {uncompressed} uncompressed bytes, above the '
|
|
f'{self.limits.max_uncompressed_bytes} limit',
|
|
gate=4,
|
|
)
|
|
ratio = uncompressed / max(compressed, 1)
|
|
if ratio > self.limits.max_compression_ratio:
|
|
raise BundleGateError(
|
|
f'archive compression ratio {ratio:.1f} is above the '
|
|
f'{self.limits.max_compression_ratio} limit',
|
|
gate=4,
|
|
)
|
|
|
|
for entry in entries:
|
|
self._check_entry_name(entry.filename)
|
|
self._check_entry_mode(entry)
|
|
|
|
@staticmethod
|
|
def _check_entry_name(name: str) -> None:
|
|
"""Reject absolute names and traversal, from the declared name alone."""
|
|
normalised = name.replace('\\', '/')
|
|
if normalised.startswith('/') or (len(normalised) > 1 and normalised[1] == ':'):
|
|
raise BundleGateError('archive declares an absolute entry path', gate=4)
|
|
if any(part == '..' for part in normalised.split('/')):
|
|
raise BundleGateError('archive declares an entry escaping its own tree', gate=4)
|
|
|
|
@staticmethod
|
|
def _check_entry_mode(entry: zipfile.ZipInfo) -> None:
|
|
"""Reject symlinks and anything that is neither a regular file nor a directory."""
|
|
mode = entry.external_attr >> 16
|
|
if stat.S_IFMT(mode) == 0:
|
|
# No file-type bits stored. This is the normal case for two kinds of entry: a zip
|
|
# written by a tool that records no Unix mode at all (`mode == 0`), and
|
|
# `ZipFile.writestr`, which stores permissions only — the producer's `metadata.json`
|
|
# arrives as `0o600 << 16`. Nothing to check: the member is extracted as a plain file
|
|
# with the worker's own permissions either way.
|
|
return
|
|
if stat.S_ISLNK(mode):
|
|
raise BundleGateError('archive declares a symbolic link', gate=4)
|
|
if not (stat.S_ISREG(mode) or stat.S_ISDIR(mode)):
|
|
raise BundleGateError('archive declares an entry that is not a regular file', gate=4)
|
|
|
|
# ------------------------------------------------------------------ gate 5
|
|
|
|
def _gate_5_extract(self, zip_path: Path, root: Path) -> None:
|
|
"""Extract into a fresh directory, containing every member by resolved path.
|
|
|
|
A pre-existing extraction root is a failure rather than something to clear: it means
|
|
another import, or a previous attempt, owns that path.
|
|
"""
|
|
if root.exists():
|
|
raise BundleGateError('extraction directory already exists', gate=5)
|
|
root.mkdir(parents=True)
|
|
resolved_root = root.resolve()
|
|
try:
|
|
with zipfile.ZipFile(zip_path) as archive:
|
|
for entry in archive.infolist():
|
|
self._extract_member(archive, entry, root, resolved_root)
|
|
except BundleGateError:
|
|
shutil.rmtree(root, ignore_errors=True)
|
|
raise
|
|
except (OSError, zipfile.BadZipFile) as error:
|
|
shutil.rmtree(root, ignore_errors=True)
|
|
raise BundleGateError('bundle content could not be unpacked', gate=5) from error
|
|
|
|
@staticmethod
|
|
def _extract_member(
|
|
archive: zipfile.ZipFile, entry: zipfile.ZipInfo, root: Path, resolved_root: Path
|
|
) -> None:
|
|
"""Write one member, re-checking containment on the resolved path."""
|
|
name = entry.filename.replace('\\', '/')
|
|
target = root / name
|
|
resolved = Path(os.path.realpath(target))
|
|
if resolved != resolved_root and resolved_root not in resolved.parents:
|
|
raise BundleGateError('archive member resolves outside the extraction root', gate=5)
|
|
if entry.is_dir():
|
|
resolved.mkdir(parents=True, exist_ok=True)
|
|
return
|
|
if resolved.exists():
|
|
raise BundleGateError('archive member would overwrite an existing file', gate=5)
|
|
resolved.parent.mkdir(parents=True, exist_ok=True)
|
|
# Written as a plain file: the archive's stored permissions are never restored, so no
|
|
# member can arrive executable.
|
|
with archive.open(entry) as source, resolved.open('wb') as sink:
|
|
shutil.copyfileobj(source, sink)
|
|
|
|
# ------------------------------------------------------------------ gate 6
|
|
|
|
def _gate_6_structure(self, root: Path) -> tuple[dict[str, Any], dict[str, Any]]:
|
|
"""Validate the extracted layout and the metadata schema."""
|
|
metadata_path = root / METADATA_FILE_NAME
|
|
artifacts_dir = root / ARTIFACTS_DIR_NAME
|
|
if not metadata_path.is_file():
|
|
raise BundleGateError('bundle has no metadata document at its root', gate=6)
|
|
if not artifacts_dir.is_dir():
|
|
raise BundleGateError('bundle has no artifacts directory', gate=6)
|
|
|
|
document = self._load_metadata_document(metadata_path)
|
|
metadata = document['metadata']
|
|
parameters = document['parameters']
|
|
|
|
self._check_metadata_fields(metadata)
|
|
self._check_parameters(parameters)
|
|
self._check_model_directories(artifacts_dir)
|
|
return metadata, parameters
|
|
|
|
@staticmethod
|
|
def _load_metadata_document(metadata_path: Path) -> dict[str, Any]:
|
|
"""Read `metadata.json` and check its two top-level keys."""
|
|
try:
|
|
document = json.loads(metadata_path.read_text(encoding='utf-8'))
|
|
except (json.JSONDecodeError, UnicodeDecodeError, OSError) as error:
|
|
raise BundleGateError('bundle metadata is not readable JSON', gate=6) from error
|
|
if not isinstance(document, dict) or set(document) != REQUIRED_TOP_LEVEL_KEYS:
|
|
raise BundleGateError(
|
|
'bundle metadata does not carry exactly the parameters and metadata blocks', gate=6
|
|
)
|
|
if not isinstance(document['metadata'], dict) or not isinstance(
|
|
document['parameters'], dict
|
|
):
|
|
raise BundleGateError('bundle metadata blocks are not objects', gate=6)
|
|
return document
|
|
|
|
@staticmethod
|
|
def _check_metadata_fields(metadata: dict[str, Any]) -> None:
|
|
"""Every origin field the import records must be present and usable."""
|
|
for field_name in REQUIRED_METADATA_FIELDS:
|
|
value = metadata.get(field_name)
|
|
if not isinstance(value, str) or not value.strip():
|
|
raise BundleGateError(f'bundle metadata is missing the {field_name} field', gate=6)
|
|
experiment_name = metadata['experiment_name']
|
|
if not (EXPERIMENT_NAME_MIN_LENGTH <= len(experiment_name) <= EXPERIMENT_NAME_MAX_LENGTH):
|
|
raise BundleGateError(
|
|
"bundle experiment name does not fit the platform's limits", gate=6
|
|
)
|
|
|
|
@staticmethod
|
|
def _check_parameters(parameters: dict[str, Any]) -> None:
|
|
"""`target_variable` is what `model_config.target` is built from; nothing else supplies it."""
|
|
for name in REQUIRED_PARAMETERS:
|
|
value = parameters.get(name)
|
|
if value is None or not str(value).strip():
|
|
raise BundleGateError(f'bundle parameters are missing {name}', gate=6)
|
|
|
|
@staticmethod
|
|
def _check_model_directories(artifacts_dir: Path) -> None:
|
|
"""`prediction_model` and `data_model` must each be a loadable MLflow model directory."""
|
|
for directory in REQUIRED_MODEL_DIRS:
|
|
model_dir = artifacts_dir / directory
|
|
if not model_dir.is_dir():
|
|
raise BundleGateError(f'bundle has no {directory} artifacts', gate=6)
|
|
for required in REQUIRED_MODEL_FILES:
|
|
if not (model_dir / required).is_file():
|
|
raise BundleGateError(f'bundle {directory} artifacts are incomplete', gate=6)
|
|
|
|
# ------------------------------------------------------------------ gate 7
|
|
|
|
def _gate_7_content_policy(self, root: Path) -> None:
|
|
"""Only the file kinds an MLflow artifact tree calls for may be present."""
|
|
for path in sorted(root.rglob('*')):
|
|
if path.is_dir():
|
|
continue
|
|
if not path.is_file():
|
|
raise BundleGateError('bundle contains an entry that is not a file', gate=7)
|
|
if path.name in ALLOWED_FILE_NAMES:
|
|
continue
|
|
if path.suffix.lower() in ALLOWED_FILE_SUFFIXES:
|
|
continue
|
|
raise BundleGateError('bundle contains a file a model export should not carry', gate=7)
|