diff --git a/mediapipe/framework/api2/builder.h b/mediapipe/framework/api2/builder.h index 2a98c416..da09acc8 100644 --- a/mediapipe/framework/api2/builder.h +++ b/mediapipe/framework/api2/builder.h @@ -206,6 +206,16 @@ class SourceImpl { return ConnectTo(dest); } + template + bool operator==(const SourceImpl& other) { + return base_ == other.base_; + } + + template + bool operator!=(const SourceImpl& other) { + return !(*this == other); + } + Src& SetName(std::string name) { base_->name_ = std::move(name); return *this; @@ -218,6 +228,9 @@ class SourceImpl { } private: + template + friend class SourceImpl; + // Never null. SourceBase* base_; }; diff --git a/mediapipe/framework/api2/builder_test.cc b/mediapipe/framework/api2/builder_test.cc index 08f4f0ca..194f1b8f 100644 --- a/mediapipe/framework/api2/builder_test.cc +++ b/mediapipe/framework/api2/builder_test.cc @@ -494,5 +494,51 @@ TEST(BuilderTest, SinglePortAccessWorksThroughSlicing) { EXPECT_THAT(graph.GetConfig(), EqualsProto(expected)); } +TEST(BuilderTest, TestStreamEqualsNotEqualsOperators) { + Graph graph; + Stream input0 = graph.In(0); + EXPECT_TRUE(input0 == input0); + EXPECT_FALSE(input0 != input0); + + EXPECT_TRUE(input0 == input0.Cast()); + EXPECT_FALSE(input0.Cast() != input0); + + EXPECT_TRUE(input0.Cast() == input0.Cast()); + EXPECT_FALSE(input0.Cast() != input0.Cast()); + + Stream input1 = graph.In(1); + EXPECT_FALSE(input0 == input1); + EXPECT_TRUE(input0 != input1); + + input1 = input0; + EXPECT_TRUE(input0 == input1); + EXPECT_FALSE(input0 != input1); + EXPECT_TRUE(input0.Cast() == input1.Cast()); + EXPECT_FALSE(input0.Cast() != input1.Cast()); +} + +TEST(BuilderTest, TestSidePacketEqualsNotEqualsOperators) { + Graph graph; + SidePacket side_input0 = graph.SideIn(0); + EXPECT_TRUE(side_input0 == side_input0); + EXPECT_FALSE(side_input0 != side_input0); + + EXPECT_TRUE(side_input0 == side_input0.Cast()); + EXPECT_FALSE(side_input0.Cast() != side_input0); + + EXPECT_TRUE(side_input0.Cast() == side_input0.Cast()); + EXPECT_FALSE(side_input0.Cast() != side_input0.Cast()); + + SidePacket side_input1 = graph.SideIn(1); + EXPECT_FALSE(side_input0 == side_input1); + EXPECT_TRUE(side_input0 != side_input1); + + side_input1 = side_input0; + EXPECT_TRUE(side_input0 == side_input1); + EXPECT_FALSE(side_input0 != side_input1); + EXPECT_TRUE(side_input0.Cast() == side_input1.Cast()); + EXPECT_FALSE(side_input0.Cast() != side_input1.Cast()); +} + } // namespace } // namespace mediapipe::api2::builder