diff --git a/mlx/io/safetensors.cpp b/mlx/io/safetensors.cpp index b8f91ba9..44872b83 100644 --- a/mlx/io/safetensors.cpp +++ b/mlx/io/safetensors.cpp @@ -1,7 +1,8 @@ // Copyright © 2023 Apple Inc. -// + #include #include +#include #include #include "mlx/backend/cuda/cuda.h" @@ -97,8 +98,9 @@ Dtype dtype_from_safetensor_str(std::string_view str) { } else if (str == ST_F8_E4M3) { return uint8; } else { - throw std::runtime_error( - "[safetensor] unsupported dtype " + std::string(str)); + std::ostringstream msg; + msg << "[safetensor] unsupported dtype" << str; + throw std::runtime_error(msg.str()); } } @@ -109,8 +111,9 @@ SafetensorsLoad load_safetensors( //////////////////////////////////////////////////////// // Open and check file if (!in_stream->good() || !in_stream->is_open()) { - throw std::runtime_error( - "[load_safetensors] Failed to open " + in_stream->label()); + std::ostringstream msg; + msg << "[load_safetensors] Failed to open " << in_stream->label(); + throw std::runtime_error(msg.str()); } auto stream = cu::is_available() ? to_stream(s) : to_stream(s, Device::cpu); @@ -120,8 +123,10 @@ SafetensorsLoad load_safetensors( constexpr uint64_t kMaxJsonHeaderLength = 100000000; in_stream->read(reinterpret_cast(&jsonHeaderLength), 8); if (jsonHeaderLength <= 0 || jsonHeaderLength >= kMaxJsonHeaderLength) { - throw std::runtime_error( - "[load_safetensors] Invalid json header length " + in_stream->label()); + std::ostringstream msg; + msg << "[load_safetensors] Invalid json header length " + << in_stream->label(); + throw std::runtime_error(msg.str()); } // Load the json metadata auto rawJson = std::make_unique(jsonHeaderLength); @@ -129,8 +134,9 @@ SafetensorsLoad load_safetensors( auto metadata = json::parse(rawJson.get(), rawJson.get() + jsonHeaderLength); // Should always be an object on the top-level if (!metadata.is_object()) { - throw std::runtime_error( - "[load_safetensors] Invalid json metadata " + in_stream->label()); + std::ostringstream msg; + msg << "[load_safetensors] Invalid json metadata " << in_stream->label(); + throw std::runtime_error(msg.str()); } size_t offset = jsonHeaderLength + 8; // Load the arrays using metadata @@ -147,6 +153,28 @@ SafetensorsLoad load_safetensors( const Shape& shape = item.value().at("shape"); const std::vector& data_offsets = item.value().at("data_offsets"); Dtype type = dtype_from_safetensor_str(dtype); + if (data_offsets.size() != 2) { + std::ostringstream msg; + msg << "[load_safetensors] Tensor '" << item.key() + << "' data_offsets must have exactly 2 entries but has " + << data_offsets.size(); + throw std::runtime_error(msg.str()); + } + { + size_t expected_nbytes = type.size(); + for (auto dim : shape) { + expected_nbytes *= static_cast(dim); + } + if (data_offsets[1] < data_offsets[0] || + data_offsets[1] - data_offsets[0] != expected_nbytes) { + std::ostringstream msg; + msg << "[load_safetensors] Tensor '" << item.key() + << "' invalid data offsets (" << data_offsets[0] << ", " + << data_offsets[1] << "). Expecting " << expected_nbytes + << " bytes."; + throw std::runtime_error(msg.str()); + } + } res.insert( {item.key(), array( @@ -170,8 +198,9 @@ void save_safetensors( //////////////////////////////////////////////////////// // Check file if (!out_stream->good() || !out_stream->is_open()) { - throw std::runtime_error( - "[save_safetensors] Failed to open " + out_stream->label()); + std::ostringstream msg; + msg << "[save_safetensors] Failed to open " << out_stream->label(); + throw std::runtime_error(msg.str()); } //////////////////////////////////////////////////////// @@ -196,8 +225,10 @@ void save_safetensors( size_t offset = 0; for (auto& [key, arr] : a) { if (arr.nbytes() == 0) { - throw std::invalid_argument( - "[save_safetensors] cannot serialize an empty array key: " + key); + std::ostringstream msg; + msg << "[save_safetensors] Cannot serialize an empty array ('" << key + << "')"; + throw std::invalid_argument(msg.str()); } json child; diff --git a/tests/load_tests.cpp b/tests/load_tests.cpp index 1531ce06..574d26b9 100644 --- a/tests/load_tests.cpp +++ b/tests/load_tests.cpp @@ -1,6 +1,7 @@ // Copyright © 2023 Apple Inc. #include +#include #include #include @@ -40,6 +41,67 @@ TEST_CASE("test save_safetensors") { CHECK(array_equal(test2, ones({2, 2})).item()); } +TEST_CASE("test safetensors rejects mismatched data_offsets") { + // Build a minimal safetensors file where data_offsets claim 4 bytes + // but shape declares 1000x1000 float32 (4,000,000 bytes). + // Verifies that load_safetensors() catches the mismatch. + std::string file_path = get_temp_file("test_bad_offsets.safetensors"); + + std::string header = + R"({"t":{"dtype":"F32","shape":[1000,1000],"data_offsets":[0,4]}})"; + uint64_t header_len = header.size(); + + { + std::ofstream f(file_path, std::ios::binary); + f.write(reinterpret_cast(&header_len), 8); + f.write(header.c_str(), header_len); + // Write only 4 bytes of data (the offsets claim [0,4]) + float one = 1.0f; + f.write(reinterpret_cast(&one), sizeof(float)); + } + + CHECK_THROWS_AS(load_safetensors(file_path), std::runtime_error); +} + +TEST_CASE("test safetensors rejects bad data_offsets count") { + // data_offsets has 3 entries instead of the required 2. + std::string file_path = get_temp_file("test_bad_offsets_count.safetensors"); + + std::string header = + R"({"t":{"dtype":"F32","shape":[1],"data_offsets":[0,4,8]}})"; + uint64_t header_len = header.size(); + + { + std::ofstream f(file_path, std::ios::binary); + f.write(reinterpret_cast(&header_len), 8); + f.write(header.c_str(), header_len); + float one = 1.0f; + f.write(reinterpret_cast(&one), sizeof(float)); + } + + CHECK_THROWS_AS(load_safetensors(file_path), std::runtime_error); +} + +TEST_CASE("test safetensors rejects inverted data_offsets") { + // data_offsets[0] > data_offsets[1] + std::string file_path = + get_temp_file("test_bad_offsets_inverted.safetensors"); + + std::string header = + R"({"t":{"dtype":"F32","shape":[1],"data_offsets":[4,0]}})"; + uint64_t header_len = header.size(); + + { + std::ofstream f(file_path, std::ios::binary); + f.write(reinterpret_cast(&header_len), 8); + f.write(header.c_str(), header_len); + float one = 1.0f; + f.write(reinterpret_cast(&one), sizeof(float)); + } + + CHECK_THROWS_AS(load_safetensors(file_path), std::runtime_error); +} + TEST_CASE("test gguf") { std::string file_path = get_temp_file("test_arr.gguf"); using dict = std::unordered_map;