Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 0 additions & 1 deletion onnxruntime/core/flatbuffers/flatbuffers_utils.cc
Original file line number Diff line number Diff line change
Expand Up @@ -282,7 +282,6 @@ Status LoadValueInfoOrtFormat(const fbs::ValueInfo& fbs_value_info,
Status LoadOpsetImportOrtFormat(const flatbuffers::Vector<flatbuffers::Offset<fbs::OperatorSetId>>* fbs_op_set_ids,
std::unordered_map<std::string, int>& domain_to_version) {
ORT_RETURN_IF(nullptr == fbs_op_set_ids, "Model must have opset imports. Invalid ORT format model.");
ORT_RETURN_IF_ERROR(ValidateRequiredTableOffsets(fbs_op_set_ids, "opset import"));

domain_to_version.clear();
domain_to_version.reserve(fbs_op_set_ids->size());
Expand Down
18 changes: 0 additions & 18 deletions onnxruntime/core/flatbuffers/flatbuffers_utils.h
Original file line number Diff line number Diff line change
Expand Up @@ -50,24 +50,6 @@ onnxruntime::common::Status LoadOpsetImportOrtFormat(
const flatbuffers::Vector<flatbuffers::Offset<fbs::OperatorSetId>>* fbs_op_set_ids,
std::unordered_map<std::string, int>& domain_to_version);

template <typename T>
inline onnxruntime::common::Status ValidateRequiredTableOffsets(
const flatbuffers::Vector<flatbuffers::Offset<T>>* fbs_entries,
const char* entry_description) {
if (fbs_entries == nullptr) {
return onnxruntime::common::Status::OK();
}

const auto* raw_offsets = reinterpret_cast<const uint8_t*>(fbs_entries->Data());
for (flatbuffers::uoffset_t i = 0; i < fbs_entries->size(); ++i) {
const auto entry_offset =
flatbuffers::ReadScalar<flatbuffers::uoffset_t>(raw_offsets + i * sizeof(flatbuffers::uoffset_t));
ORT_RETURN_IF(entry_offset == 0, "Null ", entry_description, " entry. Invalid ORT format model.");
}

return onnxruntime::common::Status::OK();
}

// check if filename ends in .ort
bool IsOrtFormatModel(const PathString& filename);

Expand Down
5 changes: 0 additions & 5 deletions onnxruntime/core/graph/graph.cc
Original file line number Diff line number Diff line change
Expand Up @@ -6659,10 +6659,8 @@ common::Status Graph::LoadFromOrtFormat(const onnxruntime::fbs::Graph& fbs_graph

// Initializers
auto fbs_initializers = fbs_graph.initializers();
ORT_RETURN_IF_ERROR(fbs::utils::ValidateRequiredTableOffsets(fbs_initializers, "initializer"));
#if !defined(DISABLE_SPARSE_TENSORS)
auto fbs_sparse_initializers = fbs_graph.sparse_initializers();
ORT_RETURN_IF_ERROR(fbs::utils::ValidateRequiredTableOffsets(fbs_sparse_initializers, "sparse initializer"));
flatbuffers::uoffset_t map_size = (fbs_initializers != nullptr ? fbs_initializers->size() : 0U) +
(fbs_sparse_initializers != nullptr ? fbs_sparse_initializers->size() : 0U);
#else
Expand Down Expand Up @@ -6732,7 +6730,6 @@ common::Status Graph::LoadFromOrtFormat(const onnxruntime::fbs::Graph& fbs_graph
// NodeArgs
auto fbs_node_args = fbs_graph.node_args();
if (fbs_node_args) {
ORT_RETURN_IF_ERROR(fbs::utils::ValidateRequiredTableOffsets(fbs_node_args, "node arg"));
node_args_.reserve(fbs_node_args->size());
for (const auto* fbs_value_info : *fbs_node_args) {
ORT_RETURN_IF(nullptr == fbs_value_info, "NodeArg is missing. Invalid ORT format model.");
Expand All @@ -6751,9 +6748,7 @@ common::Status Graph::LoadFromOrtFormat(const onnxruntime::fbs::Graph& fbs_graph
// referenced indices. We compute the required slot count from actual node and edge data rather
// than trusting the serialized max_node_index field.
auto* fbs_nodes = fbs_graph.nodes();
ORT_RETURN_IF_ERROR(fbs::utils::ValidateRequiredTableOffsets(fbs_nodes, "node"));
auto* fbs_node_edges = fbs_graph.node_edges();
ORT_RETURN_IF_ERROR(fbs::utils::ValidateRequiredTableOffsets(fbs_node_edges, "node edge"));

uint32_t max_referenced_node_index = 0;
bool has_referenced_node_index = false;
Expand Down
1 change: 0 additions & 1 deletion onnxruntime/core/graph/model.cc
Original file line number Diff line number Diff line change
Expand Up @@ -998,7 +998,6 @@ common::Status Model::LoadFromOrtFormat(const fbs::Model& fbs_model,

// Load the model metadata
if (const auto* fbs_metadata_props = fbs_model.metadata_props()) {
ORT_RETURN_IF_ERROR(fbs::utils::ValidateRequiredTableOffsets(fbs_metadata_props, "metadata property"));
model->model_metadata_.reserve(fbs_metadata_props->size());
for (const auto* prop : *fbs_metadata_props) {
ORT_RETURN_IF(nullptr == prop, "Null entry in metadata_props. Invalid ORT format model.");
Expand Down
37 changes: 9 additions & 28 deletions onnxruntime/test/framework/ort_model_only_test.cc
Original file line number Diff line number Diff line change
@@ -1,8 +1,6 @@
// Copyright (c) Microsoft Corporation. All rights reserved.
// Licensed under the MIT License.

#include <algorithm>

#include "core/flatbuffers/ort_format_version.h"
#include "core/flatbuffers/schema/ort.fbs.h"
#include "core/framework/data_types.h"
Expand Down Expand Up @@ -76,9 +74,14 @@ std::vector<uint8_t> BuildOrtModelBuffer(
return std::vector<uint8_t>(builder.GetBufferPointer(), builder.GetBufferPointer() + builder.GetSize());
}

Status LoadOrtBuffer(const std::vector<uint8_t>& buffer) {
Status LoadOrtBuffer(const std::vector<uint8_t>& buffer, bool use_buffer_for_initializers = false) {
SessionOptions so;
ORT_RETURN_IF_ERROR(so.config_options.AddConfigEntry(kOrtSessionOptionsConfigLoadModelFormat, "ORT"));
if (use_buffer_for_initializers) {
ORT_RETURN_IF_ERROR(so.config_options.AddConfigEntry(kOrtSessionOptionsConfigUseORTModelBytesDirectly, "1"));
ORT_RETURN_IF_ERROR(so.config_options.AddConfigEntry(kOrtSessionOptionsConfigUseORTModelBytesForInitializers,
"1"));
}

InferenceSessionWrapper session_object{so, GetEnvironment()};
return session_object.Load(buffer.data(), static_cast<int>(buffer.size()));
Expand Down Expand Up @@ -135,40 +138,18 @@ static void RunOrtModel(const OrtModelTestInfo& test_info) {

TEST(OrtModelTest, RejectsInitializerRawDataSizeMismatch) {
const auto buffer = BuildOrtModelBuffer([](flatbuffers::FlatBufferBuilder& builder) {
std::vector<int64_t> dims{1};
std::vector<uint8_t> raw_data(sizeof(float) * 2, 0);
std::vector<int64_t> dims{32};
std::vector<uint8_t> raw_data(sizeof(float) * 33, 0);
std::vector<flatbuffers::Offset<fbs::Tensor>> initializers{
fbs::CreateTensorDirect(builder, "bad_initializer", "", &dims, fbs::TensorDataType::FLOAT, &raw_data)};
return fbs::CreateGraphDirect(builder, &initializers);
});

const auto status = LoadOrtBuffer(buffer);
const auto status = LoadOrtBuffer(buffer, true);
ASSERT_FALSE(status.IsOK());
EXPECT_THAT(status.ErrorMessage(), testing::HasSubstr("raw data size mismatch"));
}

TEST(OrtModelTest, RejectsNullNodeArgTableEntry) {
auto buffer = BuildOrtModelBuffer([](flatbuffers::FlatBufferBuilder& builder) {
std::vector<flatbuffers::Offset<fbs::ValueInfo>> node_args{
fbs::CreateValueInfoDirect(builder, "X", "", CreateFloatTensorTypeInfo(builder, 1))};
return fbs::CreateGraphDirect(builder, nullptr, &node_args);
});

const auto* fbs_session = fbs::GetInferenceSession(buffer.data());
ASSERT_NE(fbs_session, nullptr);
const auto* fbs_node_args = fbs_session->model()->graph()->node_args();
ASSERT_NE(fbs_node_args, nullptr);

auto* raw_offsets = const_cast<uint8_t*>(reinterpret_cast<const uint8_t*>(fbs_node_args->Data()));
std::fill_n(raw_offsets, sizeof(flatbuffers::uoffset_t), 0);

const auto status = LoadOrtBuffer(buffer);
ASSERT_FALSE(status.IsOK());
EXPECT_THAT(status.ErrorMessage(),
testing::AnyOf(testing::HasSubstr("Null node arg entry"),
testing::HasSubstr("verification failed")));
}

TEST(OrtModelTest, RejectsDanglingNodeEdge) {
const auto buffer = BuildOrtModelBuffer([](flatbuffers::FlatBufferBuilder& builder) {
std::vector<flatbuffers::Offset<fbs::NodeEdge>> node_edges{
Expand Down
Loading