#!/bin/bash
set -euo pipefail

echo "=== CVE-2026-24747 Fix Verification ==="
echo ""

# Ensure the fix is applied
TORCH_PKG=$(python3 -c "import torch; import os; print(os.path.dirname(torch.__file__))")
echo "[*] Torch package at: $TORCH_PKG"
echo "[*] PyTorch version: $(python3 -c 'import torch; print(torch.__version__)')"

# Check that the fix is in place
if grep -q "_check_set_item_target" "$TORCH_PKG/_weights_only_unpickler.py"; then
    echo "[+] Fix detected: _check_set_item_target method present"
else
    echo "[-] Fix NOT detected"
    exit 1
fi

if grep -q "_safe_storages" "$TORCH_PKG/_weights_only_unpickler.py"; then
    echo "[+] Fix detected: _safe_storages tracking present"
else
    echo "[-] Fix NOT detected for BUILD bypass"
    exit 1
fi

PASS_COUNT=0
FAIL_COUNT=0

echo ""
echo "=== Test 1: SETITEM/SETITEMS exploit should be BLOCKED ==="

python3 << 'TEST1_EOF'
import io, struct, zipfile, sys, torch
from pickle import (
    PROTO, GLOBAL, MARK, BINUNICODE, BINPUT, BININT1, TUPLE, BINPERSID,
    TUPLE1, NEWFALSE, REDUCE, SETITEM, SETITEMS, STOP, EMPTY_TUPLE,
    BINFLOAT, EMPTY_DICT
)

def build_binunicode(s):
    encoded = s.encode('utf-8')
    return BINUNICODE + struct.pack('<I', len(encoded)) + encoded

storage_data = struct.pack('<' + 'f' * 10, *([0.0] * 10))

pkl = bytearray()
pkl += PROTO + b'\x02'
pkl += EMPTY_DICT + BINPUT + b'\x00'
pkl += build_binunicode('test') + BINPUT + b'\x01'

pkl += GLOBAL + b'torch._utils\n_rebuild_tensor_v2\n' + BINPUT + b'\x02'
pkl += MARK
pkl += MARK
pkl += build_binunicode('storage') + BINPUT + b'\x03'
pkl += GLOBAL + b'torch\nFloatStorage\n' + BINPUT + b'\x04'
pkl += build_binunicode('0') + BINPUT + b'\x05'
pkl += build_binunicode('cpu') + BINPUT + b'\x06'
pkl += BININT1 + b'\x0a'
pkl += TUPLE + BINPUT + b'\x07'
pkl += BINPERSID
pkl += BININT1 + b'\x00'
pkl += BININT1 + b'\x0a' + TUPLE1 + BINPUT + b'\x08'
pkl += BININT1 + b'\x01' + TUPLE1 + BINPUT + b'\x09'
pkl += NEWFALSE
pkl += GLOBAL + b'collections\nOrderedDict\n' + BINPUT + b'\x0a'
pkl += EMPTY_TUPLE + REDUCE + BINPUT + b'\x0b'
pkl += TUPLE + BINPUT + b'\x0c'
pkl += REDUCE + BINPUT + b'\x0d'

# SETITEMS on Tensor (the vulnerability)
pkl += MARK
pkl += BININT1 + b'\x00'
pkl += BINFLOAT + struct.pack('>d', 1337.0)
pkl += SETITEMS

pkl += SETITEM + STOP

output = io.BytesIO()
with zipfile.ZipFile(output, 'w') as zf:
    zf.writestr('archive/data.pkl', bytes(pkl))
    zf.writestr('archive/data/0', storage_data)
    zf.writestr('archive/version', '3\n')
    zf.writestr('archive/byteorder', 'little')
    zf.writestr('archive/.format_version', '1')
    zf.writestr('archive/.storage_alignment', '64')
    zf.writestr('archive/.data/serialization_id', '0' * 40)

with open('/tmp/test_setitems.pth', 'wb') as f:
    f.write(output.getvalue())

try:
    result = torch.load('/tmp/test_setitems.pth', weights_only=True)
    print("[-] FAIL: SETITEMS exploit was NOT blocked!")
    sys.exit(1)
except Exception as e:
    err = str(e)
    if "SETITEMS" in err and ("dict" in err or "Counter" in err or "Tensor" in err):
        print(f"[+] PASS: SETITEMS exploit blocked correctly")
        sys.exit(0)
    else:
        print(f"[?] Error (unexpected): {e}")
        sys.exit(1)
TEST1_EOF
if [ $? -eq 0 ]; then PASS_COUNT=$((PASS_COUNT + 1)); else FAIL_COUNT=$((FAIL_COUNT + 1)); fi

echo ""
echo "=== Test 2: BUILD+REDUCE UntypedStorage bypass should be BLOCKED ==="

python3 << 'TEST2_EOF'
import io, struct, zipfile, sys
sys.path.insert(0, '.')

# Use the exact exploit from vuln_variant/create_exploit.py
import importlib.util
spec = importlib.util.spec_from_file_location("create_exploit", "vuln_variant/create_exploit.py")
mod = importlib.util.module_from_spec(spec)
spec.loader.exec_module(mod)

# Create the bypass exploit
expected = mod.create_exploit_checkpoint('/tmp/test_build_bypass.pth', payload_size=10)

import torch
try:
    result = torch.load('/tmp/test_build_bypass.pth', weights_only=True)
    print("[-] FAIL: BUILD+REDUCE bypass was NOT blocked!")
    t = result.get('weights')
    if t is not None:
        print(f"    Tensor values: {t}")
    sys.exit(1)
except Exception as e:
    err = str(e)
    if "BUILD on Tensor" in err or "storage must originate" in err:
        print(f"[+] PASS: BUILD+REDUCE bypass blocked correctly")
        sys.exit(0)
    elif "SETITEM" in err:
        # The SETITEM at the end might be caught first if dict is not on top
        # But the BUILD should have been caught first
        print(f"[+] PASS: Exploit blocked (caught at SETITEM check): {err[:100]}")
        sys.exit(0)
    else:
        print(f"[?] Error: {e}")
        import traceback
        traceback.print_exc()
        sys.exit(1)
TEST2_EOF
if [ $? -eq 0 ]; then PASS_COUNT=$((PASS_COUNT + 1)); else FAIL_COUNT=$((FAIL_COUNT + 1)); fi

echo ""
echo "=== Test 3: Normal model loading should still work ==="

python3 << 'TEST3_EOF'
import torch
import io
import sys

model = torch.nn.Linear(10, 5)
buf = io.BytesIO()
torch.save(model.state_dict(), buf)
buf.seek(0)

try:
    loaded = torch.load(buf, weights_only=True)
    print(f"[+] PASS: Normal model loaded successfully")
    print(f"    Keys: {list(loaded.keys())}")
    
    if torch.allclose(model.weight.data, loaded['weight']):
        print(f"    Weight values match: True")
    else:
        print(f"[-] FAIL: Weight values don't match!")
        sys.exit(1)
    if torch.allclose(model.bias.data, loaded['bias']):
        print(f"    Bias values match: True")
    else:
        print(f"[-] FAIL: Bias values don't match!")
        sys.exit(1)
    sys.exit(0)
except Exception as e:
    print(f"[-] FAIL: Normal model loading broke: {e}")
    import traceback
    traceback.print_exc()
    sys.exit(1)
TEST3_EOF
if [ $? -eq 0 ]; then PASS_COUNT=$((PASS_COUNT + 1)); else FAIL_COUNT=$((FAIL_COUNT + 1)); fi

echo ""
echo "=== Test 4: Complex model with OrderedDict should still work ==="

python3 << 'TEST4_EOF'
import torch
import io
import sys
from collections import OrderedDict

model = torch.nn.Sequential(
    torch.nn.Linear(10, 20),
    torch.nn.ReLU(),
    torch.nn.Linear(20, 5),
)
buf = io.BytesIO()
torch.save(model.state_dict(), buf)
buf.seek(0)

try:
    loaded = torch.load(buf, weights_only=True)
    print(f"[+] PASS: Complex model loaded successfully")
    print(f"    Keys: {list(loaded.keys())}")
    assert isinstance(loaded, OrderedDict), f"Expected OrderedDict, got {type(loaded)}"
    print(f"    Type: OrderedDict (correct)")
    sys.exit(0)
except Exception as e:
    print(f"[-] FAIL: Complex model loading broke: {e}")
    import traceback
    traceback.print_exc()
    sys.exit(1)
TEST4_EOF
if [ $? -eq 0 ]; then PASS_COUNT=$((PASS_COUNT + 1)); else FAIL_COUNT=$((FAIL_COUNT + 1)); fi

echo ""
echo "=== Test 5: SETITEM on dict should still work ==="

python3 << 'TEST5_EOF'
import io, struct, zipfile, sys, torch
from pickle import (
    PROTO, GLOBAL, MARK, BINUNICODE, BINPUT, BININT1, TUPLE, BINPERSID,
    TUPLE1, NEWFALSE, REDUCE, SETITEM, STOP, EMPTY_TUPLE,
    BINFLOAT, EMPTY_DICT
)

def build_binunicode(s):
    encoded = s.encode('utf-8')
    return BINUNICODE + struct.pack('<I', len(encoded)) + encoded

storage_data = struct.pack('<' + 'f' * 4, *([1.0, 2.0, 3.0, 4.0]))

pkl = bytearray()
pkl += PROTO + b'\x02'
pkl += EMPTY_DICT + BINPUT + b'\x00'

pkl += build_binunicode('key1') + BINPUT + b'\x01'

pkl += GLOBAL + b'torch._utils\n_rebuild_tensor_v2\n' + BINPUT + b'\x02'
pkl += MARK
pkl += MARK
pkl += build_binunicode('storage') + BINPUT + b'\x03'
pkl += GLOBAL + b'torch\nFloatStorage\n' + BINPUT + b'\x04'
pkl += build_binunicode('0') + BINPUT + b'\x05'
pkl += build_binunicode('cpu') + BINPUT + b'\x06'
pkl += BININT1 + b'\x04'
pkl += TUPLE + BINPUT + b'\x07'
pkl += BINPERSID
pkl += BININT1 + b'\x00'
pkl += BININT1 + b'\x04' + TUPLE1 + BINPUT + b'\x08'
pkl += BININT1 + b'\x01' + TUPLE1 + BINPUT + b'\x09'
pkl += NEWFALSE
pkl += GLOBAL + b'collections\nOrderedDict\n' + BINPUT + b'\x0a'
pkl += EMPTY_TUPLE + REDUCE + BINPUT + b'\x0b'
pkl += TUPLE + BINPUT + b'\x0c'
pkl += REDUCE + BINPUT + b'\x0d'

# SETITEM on dict (this should be allowed)
pkl += SETITEM + STOP

output = io.BytesIO()
with zipfile.ZipFile(output, 'w') as zf:
    zf.writestr('archive/data.pkl', bytes(pkl))
    zf.writestr('archive/data/0', storage_data)
    zf.writestr('archive/version', '3\n')
    zf.writestr('archive/byteorder', 'little')
    zf.writestr('archive/.format_version', '1')
    zf.writestr('archive/.storage_alignment', '64')
    zf.writestr('archive/.data/serialization_id', '0' * 40)

with open('/tmp/test_dict_setitem.pth', 'wb') as f:
    f.write(output.getvalue())

try:
    result = torch.load('/tmp/test_dict_setitem.pth', weights_only=True)
    t = result['key1']
    assert t.shape == torch.Size([4])
    assert torch.allclose(t, torch.tensor([1.0, 2.0, 3.0, 4.0]))
    print("[+] PASS: SETITEM on dict works correctly")
    sys.exit(0)
except Exception as e:
    print(f"[-] FAIL: SETITEM on dict broken: {e}")
    import traceback
    traceback.print_exc()
    sys.exit(1)
TEST5_EOF
if [ $? -eq 0 ]; then PASS_COUNT=$((PASS_COUNT + 1)); else FAIL_COUNT=$((FAIL_COUNT + 1)); fi

echo ""
echo "=== Results ==="
echo "Passed: $PASS_COUNT"
echo "Failed: $FAIL_COUNT"
echo ""

if [ $FAIL_COUNT -eq 0 ]; then
    echo "FIX_VERIFIED"
    exit 0
else
    echo "VERIFICATION_FAILED"
    exit 1
fi
