Skip to content

Commit 5704143

Browse files
nfelttensorflower-gardener
authored andcommitted
Check syntax when parsing ApiDef textprotos
This change ensures that syntax errors in ApiDef textprotos (e.g. mistyping a field name, which I did) fail the overall API generation rule instead of having parsing fail silently and produce an empty result for that single file. PiperOrigin-RevId: 228961785
1 parent 3f594e1 commit 5704143

9 files changed

Lines changed: 237 additions & 16 deletions

File tree

‎tensorflow/core/BUILD‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1017,6 +1017,7 @@ cc_library(
10171017
":lib",
10181018
":lib_internal",
10191019
":protos_all_cc",
1020+
"//tensorflow/core/util/proto:proto_utils",
10201021
],
10211022
)
10221023

‎tensorflow/core/framework/op_gen_lib.cc‎

Lines changed: 10 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -23,6 +23,7 @@ limitations under the License.
2323
#include "tensorflow/core/lib/strings/str_util.h"
2424
#include "tensorflow/core/lib/strings/strcat.h"
2525
#include "tensorflow/core/platform/protobuf.h"
26+
#include "tensorflow/core/util/proto/proto_utils.h"
2627

2728
namespace tensorflow {
2829

@@ -488,14 +489,21 @@ Status ApiDefMap::LoadFile(Env* env, const string& filename) {
488489
if (filename.empty()) return Status::OK();
489490
string contents;
490491
TF_RETURN_IF_ERROR(ReadFileToString(env, filename, &contents));
491-
TF_RETURN_IF_ERROR(LoadApiDef(contents));
492+
Status status = LoadApiDef(contents);
493+
if (!status.ok()) {
494+
// Return failed status annotated with filename to aid in debugging.
495+
return Status(status.code(),
496+
strings::StrCat("Error parsing ApiDef file ", filename, ": ",
497+
status.error_message()));
498+
}
492499
return Status::OK();
493500
}
494501

495502
Status ApiDefMap::LoadApiDef(const string& api_def_file_contents) {
496503
const string contents = PBTxtFromMultiline(api_def_file_contents);
497504
ApiDefs api_defs;
498-
protobuf::TextFormat::ParseFromString(contents, &api_defs);
505+
TF_RETURN_IF_ERROR(
506+
proto_utils::ParseTextFormatFromString(contents, &api_defs));
499507
for (const auto& api_def : api_defs.op()) {
500508
// Check if the op definition is loaded. If op definition is not
501509
// loaded, then we just skip this ApiDef.

‎tensorflow/core/framework/op_gen_lib_test.cc‎

Lines changed: 14 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -16,6 +16,7 @@ limitations under the License.
1616
#include "tensorflow/core/framework/op_gen_lib.h"
1717

1818
#include "tensorflow/core/framework/op_def.pb.h"
19+
#include "tensorflow/core/lib/core/error_codes.pb.h"
1920
#include "tensorflow/core/platform/test.h"
2021

2122
namespace tensorflow {
@@ -39,7 +40,7 @@ constexpr char kTestOpList[] = R"(op {
3940
version: 123
4041
explanation: "foo"
4142
}
42-
)";
43+
})";
4344

4445
constexpr char kTestApiDef[] = R"(op {
4546
graph_op_name: "testop"
@@ -455,6 +456,18 @@ op {
455456
ASSERT_EQ(tensorflow::error::FAILED_PRECONDITION, status.code());
456457
}
457458

459+
TEST(OpGenLibTest, ApiDefInvalidSyntax) {
460+
const string api_def = R"pb(
461+
op { bad_op_name: "testop" }
462+
)pb";
463+
464+
OpList op_list;
465+
ApiDefMap api_map(op_list);
466+
// Loading with invalid syntax (e.g. unrecognized field name) should fail.
467+
auto status = api_map.LoadApiDef(api_def);
468+
ASSERT_EQ(tensorflow::error::INVALID_ARGUMENT, status.code());
469+
}
470+
458471
TEST(OpGenLibTest, ApiDefUpdateDocs) {
459472
const string op_list1 = R"(op {
460473
name: "testop"

‎tensorflow/core/platform/default/protobuf.h‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -23,6 +23,7 @@ limitations under the License.
2323
#include "google/protobuf/descriptor.h"
2424
#include "google/protobuf/descriptor.pb.h"
2525
#include "google/protobuf/dynamic_message.h"
26+
#include "google/protobuf/io/tokenizer.h"
2627
#include "google/protobuf/text_format.h"
2728
#include "google/protobuf/util/json_util.h"
2829
#include "google/protobuf/util/type_resolver_util.h"

‎tensorflow/core/util/proto/BUILD‎

Lines changed: 15 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -68,5 +68,20 @@ cc_library(
6868
deps = [
6969
"//tensorflow/core:framework",
7070
"//tensorflow/core:lib",
71+
"//tensorflow/core:platform_base",
72+
"@com_google_absl//absl/strings",
73+
],
74+
)
75+
76+
tf_cc_test(
77+
name = "proto_utils_test",
78+
srcs = ["proto_utils_test.cc"],
79+
deps = [
80+
":proto_utils",
81+
"//tensorflow/core:lib",
82+
"//tensorflow/core:test",
83+
"//tensorflow/core:test_main",
84+
"//tensorflow/core:testlib",
85+
"@com_google_googletest//:gtest_main",
7186
],
7287
)

‎tensorflow/core/util/proto/proto_utils.cc‎

Lines changed: 49 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -13,11 +13,14 @@ See the License for the specific language governing permissions and
1313
limitations under the License.
1414
==============================================================================*/
1515

16+
#include "tensorflow/core/util/proto/proto_utils.h"
17+
18+
#include "absl/strings/string_view.h"
19+
#include "absl/strings/substitute.h"
1620
#include "tensorflow/core/framework/types.h"
21+
#include "tensorflow/core/platform/logging.h"
1722
#include "tensorflow/core/platform/protobuf.h"
1823

19-
#include "tensorflow/core/util/proto/proto_utils.h"
20-
2124
namespace tensorflow {
2225
namespace proto_utils {
2326

@@ -66,5 +69,49 @@ bool IsCompatibleType(FieldDescriptor::Type field_type, DataType dtype) {
6669
}
6770
}
6871

72+
Status ParseTextFormatFromString(absl::string_view input,
73+
protobuf::Message* output) {
74+
DCHECK(output != nullptr) << "output must be non NULL";
75+
// When checks are disabled, instead log the error and return an error status.
76+
if (output == nullptr) {
77+
LOG(ERROR) << "output must be non NULL";
78+
return Status(error::INVALID_ARGUMENT, "output must be non NULL");
79+
}
80+
string err;
81+
StringErrorCollector err_collector(&err, /*one-indexing=*/true);
82+
protobuf::TextFormat::Parser parser;
83+
parser.RecordErrorsTo(&err_collector);
84+
if (!parser.ParseFromString(string(input), output)) {
85+
return Status(error::INVALID_ARGUMENT, err);
86+
}
87+
return Status::OK();
88+
}
89+
90+
StringErrorCollector::StringErrorCollector(string* error_text)
91+
: StringErrorCollector(error_text, false) {}
92+
93+
StringErrorCollector::StringErrorCollector(string* error_text,
94+
bool one_indexing)
95+
: error_text_(error_text), index_offset_(one_indexing ? 1 : 0) {
96+
DCHECK(error_text_ != nullptr) << "error_text must be non NULL";
97+
// When checks are disabled, just log and then ignore added errors/warnings.
98+
if (error_text_ == nullptr) {
99+
LOG(ERROR) << "error_text must be non NULL";
100+
}
101+
}
102+
103+
void StringErrorCollector::AddError(int line, int column,
104+
const string& message) {
105+
if (error_text_ != nullptr) {
106+
absl::SubstituteAndAppend(error_text_, "$0($1): $2\n", line + index_offset_,
107+
column + index_offset_, message);
108+
}
109+
}
110+
111+
void StringErrorCollector::AddWarning(int line, int column,
112+
const string& message) {
113+
AddError(line, column, message);
114+
}
115+
69116
} // namespace proto_utils
70117
} // namespace tensorflow

‎tensorflow/core/util/proto/proto_utils.h‎

Lines changed: 31 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -16,7 +16,9 @@ limitations under the License.
1616
#ifndef TENSORFLOW_CORE_UTIL_PROTO_PROTO_UTILS_H_
1717
#define TENSORFLOW_CORE_UTIL_PROTO_PROTO_UTILS_H_
1818

19+
#include "absl/strings/string_view.h"
1920
#include "tensorflow/core/framework/types.h"
21+
#include "tensorflow/core/lib/core/status.h"
2022
#include "tensorflow/core/platform/protobuf.h"
2123

2224
namespace tensorflow {
@@ -27,6 +29,35 @@ using tensorflow::protobuf::FieldDescriptor;
2729
// Returns true if the proto field type can be converted to the tensor dtype.
2830
bool IsCompatibleType(FieldDescriptor::Type field_type, DataType dtype);
2931

32+
// Parses a text-formatted protobuf from a string into the given Message* output
33+
// and returns status OK if valid, or INVALID_ARGUMENT with an accompanying
34+
// parser error message if the text format is invalid.
35+
Status ParseTextFormatFromString(absl::string_view input,
36+
protobuf::Message* output);
37+
38+
class StringErrorCollector : public protobuf::io::ErrorCollector {
39+
public:
40+
// String error_text is unowned and must remain valid during the use of
41+
// StringErrorCollector.
42+
explicit StringErrorCollector(string* error_text);
43+
// If one_indexing is set to true, all line and column numbers will be
44+
// increased by one for cases when provided indices are 0-indexed and
45+
// 1-indexed error messages are desired
46+
StringErrorCollector(string* error_text, bool one_indexing);
47+
StringErrorCollector(const StringErrorCollector&) = delete;
48+
StringErrorCollector& operator=(const StringErrorCollector&) = delete;
49+
50+
// Implementation of protobuf::io::ErrorCollector::AddError.
51+
void AddError(int line, int column, const string& message) override;
52+
53+
// Implementation of protobuf::io::ErrorCollector::AddWarning.
54+
void AddWarning(int line, int column, const string& message) override;
55+
56+
private:
57+
string* const error_text_;
58+
const int index_offset_;
59+
};
60+
3061
} // namespace proto_utils
3162
} // namespace tensorflow
3263

Lines changed: 112 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,112 @@
1+
/* Copyright 2018 The TensorFlow Authors. All Rights Reserved.
2+
3+
Licensed under the Apache License, Version 2.0 (the "License");
4+
you may not use this file except in compliance with the License.
5+
You may obtain a copy of the License at
6+
7+
http://www.apache.org/licenses/LICENSE-2.0
8+
9+
Unless required by applicable law or agreed to in writing, software
10+
distributed under the License is distributed on an "AS IS" BASIS,
11+
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12+
See the License for the specific language governing permissions and
13+
limitations under the License.
14+
==============================================================================*/
15+
16+
#include "tensorflow/core/util/proto/proto_utils.h"
17+
18+
#include <gmock/gmock.h>
19+
#include "tensorflow/core/lib/core/status_test_util.h"
20+
#include "tensorflow/core/platform/protobuf.h"
21+
#include "tensorflow/core/platform/test.h"
22+
23+
namespace tensorflow {
24+
25+
using proto_utils::ParseTextFormatFromString;
26+
using proto_utils::StringErrorCollector;
27+
using ::testing::ContainsRegex;
28+
29+
TEST(ParseTextFormatFromStringTest, Success) {
30+
protobuf::DescriptorProto output;
31+
TF_ASSERT_OK(ParseTextFormatFromString("name: \"foo\"", &output));
32+
EXPECT_EQ(output.name(), "foo");
33+
}
34+
35+
TEST(ParseTextFormatFromStringTest, ErrorOnInvalidSyntax) {
36+
protobuf::DescriptorProto output;
37+
Status status = ParseTextFormatFromString("name: foo", &output);
38+
EXPECT_EQ(status.code(), error::INVALID_ARGUMENT);
39+
EXPECT_THAT(status.error_message(), ContainsRegex("foo"));
40+
EXPECT_FALSE(output.has_name());
41+
}
42+
43+
TEST(ParseTextFormatFromStringTest, ErrorOnUnknownFieldName) {
44+
protobuf::DescriptorProto output;
45+
Status status = ParseTextFormatFromString("badname: \"foo\"", &output);
46+
EXPECT_EQ(status.code(), error::INVALID_ARGUMENT);
47+
EXPECT_THAT(status.error_message(), ContainsRegex("badname"));
48+
EXPECT_FALSE(output.has_name());
49+
}
50+
51+
TEST(ParseTextFormatFromStringTest, DiesOnNullOutputPointer) {
52+
#ifndef NDEBUG
53+
ASSERT_DEATH(ParseTextFormatFromString("foo", nullptr).IgnoreError(),
54+
"output.*non NULL");
55+
#else
56+
// Under NDEBUG we don't die but should still return an error status.
57+
Status status = ParseTextFormatFromString("foo", nullptr);
58+
EXPECT_EQ(status.code(), error::INVALID_ARGUMENT);
59+
EXPECT_THAT(status.error_message(), ContainsRegex("output.*non NULL"));
60+
#endif
61+
}
62+
63+
TEST(StringErrorCollectorTest, AppendsError) {
64+
string err;
65+
StringErrorCollector collector(&err);
66+
collector.AddError(1, 2, "foo");
67+
EXPECT_EQ("1(2): foo\n", err);
68+
}
69+
70+
TEST(StringErrorCollectorTest, AppendsWarning) {
71+
string err;
72+
StringErrorCollector collector(&err);
73+
collector.AddWarning(1, 2, "foo");
74+
EXPECT_EQ("1(2): foo\n", err);
75+
}
76+
77+
TEST(StringErrorCollectorTest, AppendsMultipleError) {
78+
string err;
79+
StringErrorCollector collector(&err);
80+
collector.AddError(1, 2, "foo");
81+
collector.AddError(3, 4, "bar");
82+
EXPECT_EQ("1(2): foo\n3(4): bar\n", err);
83+
}
84+
85+
TEST(StringErrorCollectorTest, AppendsMultipleWarning) {
86+
string err;
87+
StringErrorCollector collector(&err);
88+
collector.AddWarning(1, 2, "foo");
89+
collector.AddWarning(3, 4, "bar");
90+
EXPECT_EQ("1(2): foo\n3(4): bar\n", err);
91+
}
92+
93+
TEST(StringErrorCollectorTest, OffsetWorks) {
94+
string err;
95+
StringErrorCollector collector(&err, true);
96+
collector.AddError(1, 2, "foo");
97+
collector.AddWarning(3, 4, "bar");
98+
EXPECT_EQ("2(3): foo\n4(5): bar\n", err);
99+
}
100+
101+
TEST(StringErrorCollectorTest, DiesOnNullErrorText) {
102+
#ifndef NDEBUG
103+
ASSERT_DEATH(StringErrorCollector(nullptr), "error_text.*non NULL");
104+
#else
105+
// Under NDEBUG we don't die and instead AddError/AddWarning just do nothing.
106+
StringErrorCollector collector(nullptr);
107+
collector.AddError(1, 2, "foo");
108+
collector.AddWarning(3, 4, "bar");
109+
#endif
110+
}
111+
112+
} // namespace tensorflow

‎tensorflow/js/ops/ts_op_gen_test.cc‎

Lines changed: 4 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -112,22 +112,15 @@ import {createTensorsTypeOpAttr, nodeBackend} from './op_utils';
112112
}
113113

114114
TEST(TsOpGenTest, InputSingleAndList) {
115-
const string api_def = R"(
116-
op {
117-
name: "Foo"
118-
input_arg {
119-
name: "images"
120-
type_attr: "T"
121-
number_attr: "N"
122-
}
123-
}
124-
)";
115+
const string api_def = R"pb(
116+
op { graph_op_name: "Foo" arg_order: "dim" arg_order: "images" }
117+
)pb";
125118

126119
string ts_file_text;
127120
GenerateTsOpFileText("", api_def, &ts_file_text);
128121

129122
const string expected = R"(
130-
export function Foo(images: tfc.Tensor[], dim: tfc.Tensor): tfc.Tensor {
123+
export function Foo(dim: tfc.Tensor, images: tfc.Tensor[]): tfc.Tensor {
131124
)";
132125
ExpectContainsStr(ts_file_text, expected);
133126
}

0 commit comments

Comments
 (0)