import argparse import json import os from pathlib import Path from modelscan.modelscan import ModelScan MARKER_FILE = "mlflow_compressed_pickle_marker.txt" def scan_with_modelscan(model_dir: Path) -> dict: scanner = ModelScan() return scanner.scan(model_dir) def remove_marker() -> None: marker = Path.cwd() / MARKER_FILE if marker.exists(): marker.unlink() def marker_state() -> tuple[bool, str]: marker = Path.cwd() / MARKER_FILE return marker.exists(), marker.read_text(encoding="utf-8") if marker.exists() else "" def load_with_mlflow(model_dir: Path, env_value: str | None) -> str: if env_value is None: os.environ.pop("MLFLOW_ALLOW_PICKLE_DESERIALIZATION", None) else: os.environ["MLFLOW_ALLOW_PICKLE_DESERIALIZATION"] = env_value import mlflow.pyfunc try: mlflow.pyfunc.load_model(str(model_dir)) return "loaded" except Exception as exc: return f"{type(exc).__name__}: {exc}" def main() -> None: parser = argparse.ArgumentParser() parser.add_argument("--model-dir", default="mlflow_compressed_pickle_poc") parser.add_argument("--json-out") args = parser.parse_args() model_dir = Path(args.model_dir).resolve() remove_marker() result = scan_with_modelscan(model_dir) print("ModelScan total issues:", result["summary"]["total_issues"]) print("ModelScan total errors:", len(result["errors"])) print("ModelScan total scanned:", result["summary"]["scanned"]["total_scanned"]) print("ModelScan total skipped:", result["summary"]["skipped"]["total_skipped"]) for skipped in result["summary"]["skipped"].get("skipped_files", []): print("Skipped:", skipped["source"], "-", skipped["description"]) blocked_load = load_with_mlflow(model_dir, env_value="false") exists_blocked, _ = marker_state() print("MLflow load with pickle gate disabled:", blocked_load) print("Marker after disabled-gate load:", exists_blocked) remove_marker() default_load = load_with_mlflow(model_dir, env_value=None) exists_default, marker_text = marker_state() print("Default MLflow load:", default_load) print("Marker after default load:", exists_default) print("Marker path:", Path.cwd() / MARKER_FILE) if marker_text: print("Marker contents:", marker_text) if args.json_out: Path(args.json_out).write_text(json.dumps(result, indent=2), encoding="utf-8") if __name__ == "__main__": main()