From 176b0b08c2b2541d24dadc2ce1e4c95e3f96a0c5 Mon Sep 17 00:00:00 2001 From: danielsongmicrosoft Date: Sun, 27 Sep 2026 22:25:20 -0700 Subject: [PATCH] Account for string payloads in initializer size limits --- .../core/framework/tensorprotoutils.cc | 35 ++++++++--- onnxruntime/core/framework/tensorprotoutils.h | 5 ++ .../test/framework/tensorutils_test.cc | 63 +++++++++++++++++++ 3 files changed, 94 insertions(+), 9 deletions(-) diff --git a/onnxruntime/core/framework/tensorprotoutils.cc b/onnxruntime/core/framework/tensorprotoutils.cc index 6a411f0d92efc..2e3aefaa3746b 100644 --- a/onnxruntime/core/framework/tensorprotoutils.cc +++ b/onnxruntime/core/framework/tensorprotoutils.cc @@ -1443,7 +1443,9 @@ common::Status GetSizeInBytesFromTensorTypeProto(const ONNX_NAMESPACE::TypeProto template Status GetSizeInBytesFromTensorTypeProto<0>(const ONNX_NAMESPACE::TypeProto_Tensor& tensor_proto, size_t* out); -common::Status ValidateEmbeddedTensorProtoDataSizeAndShape(const ONNX_NAMESPACE::TensorProto& tensor_proto) { +common::Status ValidateEmbeddedTensorProtoDataSizeAndShape( + const ONNX_NAMESPACE::TensorProto& tensor_proto, + size_t max_embedded_initializer_size_in_bytes) { ORT_RETURN_IF(HasExternalData(tensor_proto), "Expected to validate an embedded (non-external) TensorProto"); TensorShape tensor_shape = GetTensorShapeFromTensorProto(tensor_proto); @@ -1468,21 +1470,32 @@ common::Status ValidateEmbeddedTensorProtoDataSizeAndShape(const ONNX_NAMESPACE: ORT_RETURN_IF_ERROR(GetSizeInBytesFromTensorElemCountAndType<0>(num_elems_unsigned, tensor_proto.data_type(), &byte_size_from_shape)); } - ORT_RETURN_IF_NOT(byte_size_from_shape <= kMaxEmbeddedInitializerSizeInBytes, + ORT_RETURN_IF_NOT(byte_size_from_shape <= max_embedded_initializer_size_in_bytes, "Initializer '", tensor_proto.name(), "' declares a size of ", byte_size_from_shape, - " bytes which exceeds the ", kMaxEmbeddedInitializerSizeInBytes, + " bytes which exceeds the ", max_embedded_initializer_size_in_bytes, " byte limit for embedded initializer data. Use external data for large initializers."); - if (HasRawData(tensor_proto)) { - ORT_RETURN_IF_NOT(tensor_proto.raw_data().size() == byte_size_from_shape, - "Initializer '", tensor_proto.name(), "': raw_data size (", tensor_proto.raw_data().size(), - " bytes) does not match expected size from shape and data type (", - byte_size_from_shape, " bytes)"); - } else if (HasString(tensor_proto)) { + if (HasString(tensor_proto)) { + ORT_RETURN_IF(HasRawData(tensor_proto), + "Initializer '", tensor_proto.name(), "': string tensor can not have raw data"); ORT_RETURN_IF_NOT(tensor_proto.string_data_size() == num_elems_signed, "Initializer '", tensor_proto.name(), "': string_data count (", tensor_proto.string_data_size(), ") does not match expected count from shape (", num_elems_signed, ")"); + + size_t total_string_storage_size = byte_size_from_shape; + for (const auto& string_data : tensor_proto.string_data()) { + ORT_RETURN_IF(string_data.size() > max_embedded_initializer_size_in_bytes - total_string_storage_size, + "Initializer '", tensor_proto.name(), "': string_data shape bytes + payload exceeds the ", + max_embedded_initializer_size_in_bytes, + " byte limit for embedded initializer data. Use external data for large initializers."); + total_string_storage_size += string_data.size(); + } + } else if (HasRawData(tensor_proto)) { + ORT_RETURN_IF_NOT(tensor_proto.raw_data().size() == byte_size_from_shape, + "Initializer '", tensor_proto.name(), "': raw_data size (", tensor_proto.raw_data().size(), + " bytes) does not match expected size from shape and data type (", + byte_size_from_shape, " bytes)"); } else { // Typed data fields. Each data type maps to a specific repeated field in the proto. int64_t expected_count = 0; @@ -1564,6 +1577,10 @@ common::Status ValidateEmbeddedTensorProtoDataSizeAndShape(const ONNX_NAMESPACE: return Status::OK(); } +common::Status ValidateEmbeddedTensorProtoDataSizeAndShape(const ONNX_NAMESPACE::TensorProto& tensor_proto) { + return ValidateEmbeddedTensorProtoDataSizeAndShape(tensor_proto, kMaxEmbeddedInitializerSizeInBytes); +} + TensorShape GetTensorShapeFromTensorShapeProto(const ONNX_NAMESPACE::TensorShapeProto& tensor_shape_proto) { const auto& dims = tensor_shape_proto.dim(); TensorShapeVector tensor_shape_vec(static_cast(dims.size())); diff --git a/onnxruntime/core/framework/tensorprotoutils.h b/onnxruntime/core/framework/tensorprotoutils.h index 384136951ec7d..19aff139cfad1 100644 --- a/onnxruntime/core/framework/tensorprotoutils.h +++ b/onnxruntime/core/framework/tensorprotoutils.h @@ -281,6 +281,11 @@ Status GetSizeInBytesFromTensorTypeProto(const ONNX_NAMESPACE::TypeProto_Tensor& /// declared shape when the actual data is absent or much smaller. /// The caller must ensure that the TensorProto does not use external data; if it does, this function will /// return an error status. +/// max_embedded_initializer_size_in_bytes limits the total in-memory initializer size. For STRING tensors, +/// the limit includes both the std::string object storage implied by the shape and all string payload bytes. +common::Status ValidateEmbeddedTensorProtoDataSizeAndShape(const ONNX_NAMESPACE::TensorProto& tensor_proto, + size_t max_embedded_initializer_size_in_bytes); + common::Status ValidateEmbeddedTensorProtoDataSizeAndShape(const ONNX_NAMESPACE::TensorProto& tensor_proto); /** diff --git a/onnxruntime/test/framework/tensorutils_test.cc b/onnxruntime/test/framework/tensorutils_test.cc index d238627f8e6df..f73126dd3054b 100644 --- a/onnxruntime/test/framework/tensorutils_test.cc +++ b/onnxruntime/test/framework/tensorutils_test.cc @@ -363,6 +363,69 @@ TEST(TensorProtoUtilsTest, UnpackTensor) { EXPECT_FALSE(status.IsOK()); } +namespace { +TensorProto CreateStringTensorProto(std::initializer_list payload_sizes) { + TensorProto tensor_proto; + tensor_proto.set_name("string_initializer"); + tensor_proto.set_data_type(TensorProto_DataType_STRING); + tensor_proto.add_dims(static_cast(payload_sizes.size())); + + for (size_t payload_size : payload_sizes) { + tensor_proto.add_string_data(std::string(payload_size, 'a')); + } + + return tensor_proto; +} +} // namespace + +TEST(TensorProtoUtilsTest, ValidateEmbeddedStringTensorProtoRejectsOversizedPayload) { + constexpr size_t kPayloadBudgetBytes = 128; + constexpr size_t kTestBudgetBytes = sizeof(std::string) + kPayloadBudgetBytes; + const TensorProto tensor_proto = CreateStringTensorProto({kPayloadBudgetBytes + 1}); + + const Status status = ValidateEmbeddedTensorProtoDataSizeAndShape(tensor_proto, kTestBudgetBytes); + + ASSERT_STATUS_NOT_OK_AND_HAS_SUBSTR(status, "string_data shape bytes + payload exceeds"); +} + +TEST(TensorProtoUtilsTest, ValidateEmbeddedStringTensorProtoRejectsCombinedShapeAndPayloadOverflow) { + constexpr size_t kFirstPayloadBytes = 64; + constexpr size_t kTestBudgetBytes = 2 * sizeof(std::string) + kFirstPayloadBytes; + const TensorProto tensor_proto = CreateStringTensorProto({kFirstPayloadBytes, 1}); + + // STRING tensors account for the shape-declared bytes and the aggregate payload bytes. + // The first payload exactly fills the remaining allowance after object storage; the second + // payload proves that validation uses cumulative accounting. + const Status status = ValidateEmbeddedTensorProtoDataSizeAndShape(tensor_proto, kTestBudgetBytes); + + ASSERT_STATUS_NOT_OK_AND_HAS_SUBSTR(status, "string_data shape bytes + payload exceeds"); +} + +TEST(TensorProtoUtilsTest, ValidateEmbeddedStringTensorProtoAcceptsExactPayloadLimit) { + constexpr size_t kPayloadBytes = 17; + constexpr size_t kTestBudgetBytes = sizeof(std::string) + kPayloadBytes; + const TensorProto tensor_proto = CreateStringTensorProto({kPayloadBytes}); + + ASSERT_STATUS_OK(ValidateEmbeddedTensorProtoDataSizeAndShape(tensor_proto, kTestBudgetBytes)); +} + +TEST(TensorProtoUtilsTest, ValidateEmbeddedStringTensorProtoAcceptsNormalPayload) { + constexpr size_t kTestBudgetBytes = 3 * sizeof(std::string) + 32; + const TensorProto tensor_proto = CreateStringTensorProto({7, 11, 13}); + + ASSERT_STATUS_OK(ValidateEmbeddedTensorProtoDataSizeAndShape(tensor_proto, kTestBudgetBytes)); +} + +TEST(TensorProtoUtilsTest, ValidateEmbeddedStringTensorProtoRejectsRawData) { + TensorProto tensor_proto = CreateStringTensorProto({1}); + tensor_proto.set_raw_data(std::string(sizeof(std::string), '\0')); + + const Status status = ValidateEmbeddedTensorProtoDataSizeAndShape( + tensor_proto, sizeof(std::string) + 1); + + ASSERT_STATUS_NOT_OK_AND_HAS_SUBSTR(status, "string tensor can not have raw data"); +} + // A bool initializer supplied through raw_data is copied verbatim, so its bytes are not // restricted to {0, 1}. UnpackTensor must normalize them so downstream consumers (which assume // canonical bool values) all observe the same result regardless of how they read the byte.