"""Explicitly trusted legacy checkpoint loading and guarded inference. Pickle executes code. Namespace compatibility and SHA256 comparison do not make an untrusted pickle safe; users must opt in before any unpickling takes place. """ from __future__ import annotations import hashlib import hmac import pickle import re from pathlib import Path from typing import Any, Mapping class WISPUnpickler(pickle.Unpickler): """Map historical method modules, without installing process-wide aliases. This is a compatibility loader, NOT a restricted or sandboxed unpickler. """ def find_class(self, module: str, name: str) -> Any: if module == "tsevolve.search.ts_evolve_v3": module = "wisp.search.controller" elif module == "tsevolve" or module.startswith("tsevolve."): module = "wisp" + module[len("tsevolve"):] return super().find_class(module, name) def sha256_file(path: str | Path) -> str: digest = hashlib.sha256() with Path(path).open("rb") as stream: for block in iter(lambda: stream.read(1024 * 1024), b""): digest.update(block) return digest.hexdigest() def load_checkpoint(path: str | Path, *, trust_pickle: bool = False, expected_sha256: str | None = None) -> Any: """Load a trusted fitted WISP estimator, optionally checking its SHA256. Supply ``expected_sha256`` from an independently trusted source. A digest supplied alongside a checkpoint proves integrity, not publisher identity. Historical checkpoint files are left unchanged on disk. """ if not trust_pickle: raise ValueError("Pickle can execute arbitrary code. Load only trusted WISP " "checkpoints and explicitly set trust_pickle=True.") path = Path(path) if expected_sha256 is not None: if not re.fullmatch(r"[0-9a-fA-F]{64}", expected_sha256): raise ValueError("expected_sha256 must be a 64-character hexadecimal digest") if not hmac.compare_digest(sha256_file(path), expected_sha256.lower()): raise ValueError(f"SHA256 mismatch for {path.name}; refusing to unpickle") with path.open("rb") as stream: model = WISPUnpickler(stream).load() if not callable(getattr(model, "predict", None)): raise TypeError("Checkpoint does not contain a fitted estimator with predict()") return model def _array_vector(value: Any, *, name: str, n: int) -> Any: import numpy as np array = np.asarray(value) if array.ndim != 1 or len(array) != n: raise ValueError(f"{name} must be a one-dimensional vector with {n} entries") return array def _check_X(X: Any, metadata: Mapping[str, Any] | None = None) -> Any: import numpy as np array = np.asarray(X, dtype=np.float64) if array.ndim != 3 or not all(array.shape): raise ValueError("X must have nonempty shape (n_windows, n_timepoints, n_channels)") if not np.isfinite(array).all(): raise ValueError("X contains NaN or infinite values; apply the dataset preprocessing first") if metadata and metadata.get("input_shape") is not None: expected = tuple(int(v) for v in metadata["input_shape"]) # A complete source array shape may include a variable number of windows. if len(expected) == 3: expected = expected[1:] if len(expected) != 2 or array.shape[1:] != expected: raise ValueError(f"Checkpoint expects time/channel shape {expected}; got {array.shape[1:]}") return array def _validate_class_axis(model: Any, metadata: Mapping[str, Any] | None) -> None: """Never redefine a fitted global class axis using a shorter input vocabulary.""" import numpy as np classes = getattr(model, "classes_", None) if classes is not None: classes = np.asarray(classes) if (classes.ndim != 1 or not len(classes) or not np.issubdtype(classes.dtype, np.integer) or not np.array_equal(classes, np.arange(len(classes)))): raise ValueError("Checkpoint classes_ must be the contiguous encoded global " "axis 0..K-1; refusing to silently remap a nonstandard class axis") names = None if metadata is None else metadata.get("label_names") if names is not None: if isinstance(names, (str, bytes)) or np.asarray(names).ndim != 1: raise ValueError("label_names must be a one-dimensional global class vocabulary") if classes is not None and len(names) != len(classes): raise ValueError(f"label_names has {len(names)} entries but checkpoint global " f"classes_ has {len(classes)}; preserve the complete fitted vocabulary") def require_sequence_metadata(model: Any, subject: Any, time_index: Any, *, n: int) -> tuple[Any, Any]: """Do not silently bypass the sequence decoder stored in a checkpoint.""" needs_sequence = getattr(model, "smoother_", None) is not None if needs_sequence and (subject is None or time_index is None): raise ValueError("This checkpoint uses the sequence decoder: both subject " "(sequence/session IDs) and time_index are required. " "Do not merge unrelated recording sessions into one sequence.") if subject is not None: subject = _array_vector(subject, name="subject", n=n) if time_index is not None: time_index = _array_vector(time_index, name="time_index", n=n) return subject, time_index def set_inference_jobs(model: Any, n_jobs: int) -> None: """Limit the known fitted estimator/ensemble members without refitting.""" if n_jobs < 1: raise ValueError("n_jobs must be positive") seen: set[int] = set() def visit(estimator: Any) -> None: if estimator is None or id(estimator) in seen: return seen.add(id(estimator)) if hasattr(estimator, "n_jobs"): estimator.n_jobs = n_jobs for attr in ("model_", "members_", "estimators_", "steps"): child = getattr(estimator, attr, None) if isinstance(child, (tuple, list)): for item in child: visit(item[-1] if isinstance(item, tuple) else item) elif child is not None: visit(child) visit(model) def predict_checkpoint(model: Any, X: Any, *, subject: Any = None, time_index: Any = None, metadata: Mapping[str, Any] | None = None, n_jobs: int = 1) -> Any: """Predict encoded class IDs using the fitted checkpoint's actual decoder. ``metadata`` may contain ``input_shape=[T,C]`` and ``label_names``; it never authorizes pickle loading, downloads data, or silently changes a decoder. """ import numpy as np _validate_class_axis(model, metadata) X = _check_X(X, metadata) subject, time_index = require_sequence_metadata(model, subject, time_index, n=len(X)) set_inference_jobs(model, n_jobs) prediction = np.asarray(model.predict(X, groups=subject, time_index=time_index)) if prediction.ndim != 1 or len(prediction) != len(X): raise ValueError("Model returned an invalid prediction shape") if not np.issubdtype(prediction.dtype, np.integer): raise ValueError("WISP checkpoints must return encoded integer class IDs") prediction = prediction.astype(np.int64, copy=False) names = None if metadata is None else metadata.get("label_names") if names is not None and len(prediction) and (prediction.min() < 0 or prediction.max() >= len(names)): raise ValueError("Predictions are incompatible with the supplied label_names") return prediction