已合并
fix(om2): report visual JSON extraction errors #4403
ClarkXie创建于 19 天前
fix(om2): report visual JSON extraction errors #4403
已合并
ClarkXie创建于 19 天前
3 个文件变更+115-4
@@ -9,6 +9,7 @@
9 */9 */
10 10 
11#include "common/ge_common/string_util.h"11#include "common/ge_common/string_util.h"
12+#include "base/err_msg.h"
12#include "framework/common/helper/om2_package_helper.h"13#include "framework/common/helper/om2_package_helper.h"
13#include "framework/common/helper/model_save_helper_factory.h"14#include "framework/common/helper/model_save_helper_factory.h"
14#include "common/file_constant_utils/file_constant_utils.h"15#include "common/file_constant_utils/file_constant_utils.h"
@@ -550,12 +551,21 @@ Status Om2PackageHelper::RelocateExternalWeights(const std::string &output_file_
550}551}
551 552 
552Status Om2PackageHelper::ExtractVisualJson(const void *model_data, size_t model_len, std::string &json_out) {553Status Om2PackageHelper::ExtractVisualJson(const void *model_data, size_t model_len, std::string &json_out) {
554+ const auto report_extract_failed = [](const char *reason) {
555+ (void)REPORT_PREDEFINED_ERR_MSG("E10059", std::vector<const char *>({"stage", "reason"}),
556+ std::vector<const char *>({"ExtractVisualJson", reason}));
557+ GELOGE(FAILED, "[OM2] ExtractVisualJson failed. Reason: %s", reason);
558+ };
559+ 
553 GE_ASSERT_NOTNULL(model_data, "[OM2] model_data is nullptr");560 GE_ASSERT_NOTNULL(model_data, "[OM2] model_data is nullptr");
554 GE_ASSERT_TRUE(model_len > 0U, "[OM2] model_len is 0");561 GE_ASSERT_TRUE(model_len > 0U, "[OM2] model_len is 0");
555 562 
556 const auto *data = static_cast<const uint8_t *>(model_data);563 const auto *data = static_cast<const uint8_t *>(model_data);
557 SimpleZipArchiveReader reader(data, model_len);564 SimpleZipArchiveReader reader(data, model_len);
558- GE_ASSERT_TRUE(reader.IsGood(), "[OM2] Failed to open OM2 ZIP archive");565+ if (!reader.IsGood()) {
566+ report_extract_failed("Failed to open OM2 ZIP archive.");
567+ return FAILED;
568+ }
559 569 
560 const auto file_list = reader.ListFiles();570 const auto file_list = reader.ListFiles();
561 std::string entry_path;571 std::string entry_path;
@@ -566,12 +576,17 @@ Status Om2PackageHelper::ExtractVisualJson(const void *model_data, size_t model_
566 break;576 break;
567 }577 }
568 }578 }
569- GE_ASSERT_TRUE(!entry_path.empty(), "[OM2] visual JSON not found in OM2 archive");579+ if (entry_path.empty()) {
580+ report_extract_failed("visual JSON not found in OM2 archive.");
581+ return FAILED;
582+ }
570 583 
571 size_t json_size = 0U;584 size_t json_size = 0U;
572 auto json_buf = reader.ExtractToMem(entry_path, json_size);585 auto json_buf = reader.ExtractToMem(entry_path, json_size);
573- GE_ASSERT_NOTNULL(json_buf, "[OM2] Failed to extract %s from OM2 archive", entry_path.c_str());586+ if ((json_buf == nullptr) || (json_size == 0U)) {
574- GE_ASSERT_TRUE(json_size > 0U, "[OM2] Extracted visual JSON is empty");587+ report_extract_failed("Failed to extract visual JSON from OM2 archive.");
588+ return FAILED;
589+ }
575 590 
576 json_out.assign(reinterpret_cast<const char *>(json_buf.get()), json_size);591 json_out.assign(reinterpret_cast<const char *>(json_buf.get()), json_size);
577 GELOGI("[OM2] Extracted visual JSON, entry:%s, size:%zu", entry_path.c_str(), json_out.size());592 GELOGI("[OM2] Extracted visual JSON, entry:%s, size:%zu", entry_path.c_str(), json_out.size());
@@ -24,6 +24,7 @@
24#include <gtest/gtest.h>24#include <gtest/gtest.h>
25#include <algorithm>25#include <algorithm>
26#include <cerrno>26#include <cerrno>
27+#include <cstring>
27#include <cstdlib>28#include <cstdlib>
28#include <filesystem>29#include <filesystem>
29#include <fstream>30#include <fstream>
@@ -1667,6 +1668,65 @@ TEST_F(Om2St, Om2PackageHelper_Ok_ExtractVisualJsonFromMinimalOm2) {
1667 EXPECT_EQ(extracted.Raw().at("model").at("name"), JsonFile::json("visual_model"));1668 EXPECT_EQ(extracted.Raw().at("model").at("name"), JsonFile::json("visual_model"));
1668}1669}
1669 1670 
1671+TEST_F(Om2St, Om2PackageHelper_Fail_ExtractVisualJsonFromInvalidZip) {
1672+ const uint8_t garbage[] = {0xDE, 0xAD, 0xBE, 0xEF, 0x00, 0x01, 0x02, 0x03};
1673+ std::string json_out;
1674+ (void)ErrorManager::GetInstance().GetErrorMessage();
1675+ 
1676+ EXPECT_NE(Om2PackageHelper::ExtractVisualJson(garbage, sizeof(garbage), json_out), SUCCESS);
1677+ EXPECT_NE(ErrorManager::GetInstance().GetErrorMessage().find("E10059"), std::string::npos);
1678+}
1679+ 
1680+TEST_F(Om2St, Om2PackageHelper_Fail_ExtractVisualJsonWithoutVisualJson) {
1681+ const std::string output_file = PathUtils::Join({test_work_dir, "no_visual_json.om2"});
1682+ ModelBufferData model;
1683+ {
1684+ ZipArchiveWriter writer(output_file);
1685+ ASSERT_TRUE(writer.IsMemFileOpened());
1686+ const std::string manifest = R"({"om2_version":"0","model_num":1})";
1687+ ASSERT_TRUE(writer.WriteBytes("manifest.json", manifest.data(), manifest.size(), false));
1688+ ASSERT_TRUE(writer.SaveModelData(model, false));
1689+ }
1690+ ASSERT_NE(model.data, nullptr);
1691+ ASSERT_GT(model.length, 0U);
1692+ 
1693+ std::string json_out;
1694+ (void)ErrorManager::GetInstance().GetErrorMessage();
1695+ EXPECT_NE(Om2PackageHelper::ExtractVisualJson(model.data.get(), model.length, json_out), SUCCESS);
1696+ EXPECT_NE(ErrorManager::GetInstance().GetErrorMessage().find("E10059"), std::string::npos);
1697+}
1698+ 
1699+TEST_F(Om2St, Om2PackageHelper_Fail_ExtractCorruptedVisualJson) {
1700+ const std::string output_file = PathUtils::Join({test_work_dir, "corrupted_visual_json.om2"});
1701+ const std::string visual_json = "{}";
1702+ ModelBufferData model;
1703+ {
1704+ ZipArchiveWriter writer(output_file);
1705+ ASSERT_TRUE(writer.IsMemFileOpened());
1706+ ASSERT_TRUE(writer.WriteBytes("data/model_0/debug/ge_visual_00000000_graph_0.json", visual_json.data(),
1707+ visual_json.size(), true));
1708+ ASSERT_TRUE(writer.SaveModelData(model, false));
1709+ }
1710+ ASSERT_NE(model.data, nullptr);
1711+ 
1712+ constexpr uint8_t kLocalFileHeaderMagic[] = {0x50U, 0x4BU, 0x03U, 0x04U};
1713+ bool corrupted = false;
1714+ for (size_t i = 0U; i + sizeof(kLocalFileHeaderMagic) + 6U < model.length; ++i) {
1715+ if (std::memcmp(model.data.get() + i, kLocalFileHeaderMagic, sizeof(kLocalFileHeaderMagic)) == 0) {
1716+ model.data.get()[i + 8U] = 0xFFU;
1717+ model.data.get()[i + 9U] = 0U;
1718+ corrupted = true;
1719+ break;
1720+ }
1721+ }
1722+ ASSERT_TRUE(corrupted);
1723+ 
1724+ std::string json_out;
1725+ (void)ErrorManager::GetInstance().GetErrorMessage();
1726+ EXPECT_NE(Om2PackageHelper::ExtractVisualJson(model.data.get(), model.length, json_out), SUCCESS);
1727+ EXPECT_NE(ErrorManager::GetInstance().GetErrorMessage().find("E10059"), std::string::npos);
1728+}
1729+ 
1670TEST_F(Om2St, ConvertOm2Model_Ok_ConvertMinimalVisualOm2ToJson) {1730TEST_F(Om2St, ConvertOm2Model_Ok_ConvertMinimalVisualOm2ToJson) {
1671 const std::string output_file = PathUtils::Join({test_work_dir, "minimal_visual_json.om2"});1731 const std::string output_file = PathUtils::Join({test_work_dir, "minimal_visual_json.om2"});
1672 const std::string json_file = PathUtils::Join({test_work_dir, "minimal_visual_json.json"});1732 const std::string json_file = PathUtils::Join({test_work_dir, "minimal_visual_json.json"});
@@ -37,6 +37,7 @@
37#include "graph/utils/file_utils.h"37#include "graph/utils/file_utils.h"
38#include "graph/utils/graph_utils.h"38#include "graph/utils/graph_utils.h"
39#include <cstdio>39#include <cstdio>
40+#include <cstring>
40#include <sstream>41#include <sstream>
41#include <system_error>42#include <system_error>
42 43 
@@ -44,6 +45,7 @@
44#include "ge_runtime_stub/include/faker/ge_model_builder.h"45#include "ge_runtime_stub/include/faker/ge_model_builder.h"
45#include "ge_runtime_stub/include/faker/aicore_taskdef_faker.h"46#include "ge_runtime_stub/include/faker/aicore_taskdef_faker.h"
46#include "common/tbe_handle_store/tbe_kernel_store.h"47#include "common/tbe_handle_store/tbe_kernel_store.h"
48+#include "common/util/error_manager/error_manager.h"
47 49 
48namespace ge {50namespace ge {
49namespace {51namespace {
@@ -992,7 +994,9 @@ TEST_F(Om2PackageHelperUt, ExtractVisualJson_Fail_ZeroLen) {
992TEST_F(Om2PackageHelperUt, ExtractVisualJson_Fail_InvalidZip) {994TEST_F(Om2PackageHelperUt, ExtractVisualJson_Fail_InvalidZip) {
993 const uint8_t garbage[] = {0xDE, 0xAD, 0xBE, 0xEF, 0x00, 0x01, 0x02, 0x03};995 const uint8_t garbage[] = {0xDE, 0xAD, 0xBE, 0xEF, 0x00, 0x01, 0x02, 0x03};
994 std::string json_out;996 std::string json_out;
997+ (void)ErrorManager::GetInstance().GetErrorMessage();
995 EXPECT_NE(Om2PackageHelper::ExtractVisualJson(garbage, sizeof(garbage), json_out), SUCCESS);998 EXPECT_NE(Om2PackageHelper::ExtractVisualJson(garbage, sizeof(garbage), json_out), SUCCESS);
999+ EXPECT_NE(ErrorManager::GetInstance().GetErrorMessage().find("E10059"), std::string::npos);
996}1000}
997 1001 
998TEST_F(Om2PackageHelperUt, ExtractVisualJson_Fail_NoVisualJson) {1002TEST_F(Om2PackageHelperUt, ExtractVisualJson_Fail_NoVisualJson) {
@@ -1011,7 +1015,39 @@ TEST_F(Om2PackageHelperUt, ExtractVisualJson_Fail_NoVisualJson) {
1011 ASSERT_NE(file_buf, nullptr);1015 ASSERT_NE(file_buf, nullptr);
1012 1016 
1013 std::string json_out;1017 std::string json_out;
1018+ (void)ErrorManager::GetInstance().GetErrorMessage();
1014 EXPECT_NE(Om2PackageHelper::ExtractVisualJson(file_buf.get(), file_size, json_out), SUCCESS);1019 EXPECT_NE(Om2PackageHelper::ExtractVisualJson(file_buf.get(), file_size, json_out), SUCCESS);
1020+ EXPECT_NE(ErrorManager::GetInstance().GetErrorMessage().find("E10059"), std::string::npos);
1021+}
1022+ 
1023+TEST_F(Om2PackageHelperUt, ExtractVisualJson_Fail_CorruptedVisualJsonEntry) {
1024+ const std::string zip_path = PathUtils::Join({test_work_dir, "corrupted_visual_json.om2"});
1025+ ZipArchiveWriter writer(zip_path);
1026+ ASSERT_TRUE(writer.IsMemFileOpened());
1027+ const std::string visual_json = "{}";
1028+ ASSERT_TRUE(writer.WriteBytes("data/model_0/debug/ge_visual_00000000_graph_0.json", visual_json.data(),
1029+ visual_json.size(), true));
1030+ 
1031+ ModelBufferData model;
1032+ ASSERT_TRUE(writer.SaveModelData(model, false));
1033+ ASSERT_NE(model.data, nullptr);
1034+ 
1035+ constexpr uint8_t kLocalFileHeaderMagic[] = {0x50U, 0x4BU, 0x03U, 0x04U};
1036+ bool corrupted = false;
1037+ for (size_t i = 0U; i + sizeof(kLocalFileHeaderMagic) + 6U < model.length; ++i) {
1038+ if (std::memcmp(model.data.get() + i, kLocalFileHeaderMagic, sizeof(kLocalFileHeaderMagic)) == 0) {
1039+ model.data.get()[i + 8U] = 0xFFU;
1040+ model.data.get()[i + 9U] = 0U;
1041+ corrupted = true;
1042+ break;
1043+ }
1044+ }
1045+ ASSERT_TRUE(corrupted);
1046+ 
1047+ (void)ErrorManager::GetInstance().GetErrorMessage();
1048+ std::string json_out;
1049+ EXPECT_NE(Om2PackageHelper::ExtractVisualJson(model.data.get(), model.length, json_out), SUCCESS);
1050+ EXPECT_NE(ErrorManager::GetInstance().GetErrorMessage().find("E10059"), std::string::npos);
1015}1051}
1016 1052 
1017TEST_F(Om2PackageHelperUt, BuildModelMeta_SpecialInputSize) {1053TEST_F(Om2PackageHelperUt, BuildModelMeta_SpecialInputSize) {