Project import generated by Copybara.

GitOrigin-RevId: d073f8e21be2fcc0e503cb97c6695078b6b75310
This commit is contained in:
MediaPipe Team
2021-02-27 03:30:05 -05:00
committed by chuoling
parent 39309bedba
commit 350fbb2100
755 changed files with 16391 additions and 11075 deletions
+1 -5
View File
@@ -3,15 +3,10 @@ package(
features = ["-use_header_modules"],
)
# API2 is in preview mode. Internal clients are welcome and encouraged to try
# it out, but be aware that there may be more changes before release. Please
# add your package to this list and reach out to the MediaPipe team (use
# camillol@ as the CL reviewer).
package_group(
name = "preview_users",
packages = [
"//mediapipe/...",
"//video/content_analysis/...",
],
)
@@ -134,6 +129,7 @@ cc_library(
":tuple",
"//mediapipe/framework:packet",
"//mediapipe/framework/port:logging",
"@com_google_absl//absl/meta:type_traits",
],
)
+6 -7
View File
@@ -481,8 +481,7 @@ class Graph {
std::string TaggedName(const TagIndexLocation& loc, const std::string& name) {
if (loc.tag.empty()) {
// ParseTagIndexName does not allow using explicit indices without tags,
// while ParseTagIndex does. There is no explanation for this discrepancy
// in the CLs that introduced them (cl/143209019, cl/156499931).
// while ParseTagIndex does.
// TODO: decide whether we should just allow it.
return name;
} else {
@@ -494,8 +493,8 @@ class Graph {
}
}
mediapipe::Status UpdateNodeConfig(const NodeBase& node,
CalculatorGraphConfig::Node* config) {
absl::Status UpdateNodeConfig(const NodeBase& node,
CalculatorGraphConfig::Node* config) {
config->set_calculator(node.type_);
node.in_streams_.Visit(
[&](const TagIndexLocation& loc, const DestinationBase& endpoint) {
@@ -521,8 +520,8 @@ class Graph {
return {};
}
mediapipe::Status UpdateNodeConfig(const PacketGenerator& node,
PacketGeneratorConfig* config) {
absl::Status UpdateNodeConfig(const PacketGenerator& node,
PacketGeneratorConfig* config) {
config->set_packet_generator(node.type_);
node.in_sides_.Visit([&](const TagIndexLocation& loc,
const DestinationBase& endpoint) {
@@ -540,7 +539,7 @@ class Graph {
}
// For special boundary node.
mediapipe::Status UpdateBoundaryConfig(CalculatorGraphConfig* config) {
absl::Status UpdateBoundaryConfig(CalculatorGraphConfig* config) {
graph_boundary_.in_streams_.Visit(
[&](const TagIndexLocation& loc, const DestinationBase& endpoint) {
CHECK(endpoint.source != nullptr);
+22 -27
View File
@@ -26,7 +26,7 @@ class StreamHandler {
const const_str& name() { return name_; }
mediapipe::Status AddToContract(CalculatorContract* cc) const {
absl::Status AddToContract(CalculatorContract* cc) const {
cc->SetInputStreamHandler(name_.data());
return {};
}
@@ -47,7 +47,7 @@ class TimestampChange {
return TimestampChange(kUnset);
}
mediapipe::Status AddToContract(CalculatorContract* cc) const {
absl::Status AddToContract(CalculatorContract* cc) const {
if (offset_ != kUnset) cc->SetTimestampOffset(offset_);
return {};
}
@@ -71,10 +71,9 @@ struct HasProcessMethod : std::false_type {};
template <class T>
struct HasProcessMethod<
T, std::void_t<decltype(mediapipe::Status(
std::declval<std::decay_t<T>>().Process(
std::declval<mediapipe::CalculatorContext*>())))>>
: std::true_type {};
T,
std::void_t<decltype(absl::Status(std::declval<std::decay_t<T>>().Process(
std::declval<mediapipe::CalculatorContext*>())))>> : std::true_type {};
template <class T, class = void>
struct HasNestedItems : std::false_type {};
@@ -142,9 +141,9 @@ class Contract {
constexpr Contract(T&&... args)
: Contract(std::tuple<T...>{std::move(args)...}) {}
mediapipe::Status GetContract(mediapipe::CalculatorContract* cc) const {
std::vector<mediapipe::Status> statuses;
auto store_status = [&statuses](mediapipe::Status status) {
absl::Status GetContract(mediapipe::CalculatorContract* cc) const {
std::vector<absl::Status> statuses;
auto store_status = [&statuses](absl::Status status) {
if (!status.ok()) statuses.push_back(std::move(status));
};
internal::tuple_for_each(
@@ -209,7 +208,7 @@ class TaggedContract {
public:
constexpr TaggedContract() = default;
static mediapipe::Status GetContract(mediapipe::CalculatorContract* cc) {
static absl::Status GetContract(mediapipe::CalculatorContract* cc) {
return c2.GetContract(cc);
}
@@ -272,34 +271,32 @@ class OutputSender {
OutputSender(std::tuple<P...>&& args) : outputs_(args) {}
template <class R, std::enable_if_t<sizeof...(P) == 1, int> = 0>
mediapipe::Status operator()(CalculatorContext* cc,
mediapipe::StatusOr<R>&& result) {
absl::Status operator()(CalculatorContext* cc, absl::StatusOr<R>&& result) {
if (result.ok()) {
return this(cc, result.ValueOrDie());
return this(cc, result.value());
} else {
return result.status();
}
}
template <class R, std::enable_if_t<sizeof...(P) == 1, int> = 0>
mediapipe::Status operator()(CalculatorContext* cc, R&& result) {
absl::Status operator()(CalculatorContext* cc, R&& result) {
std::get<0>(outputs_)(cc).Send(std::forward<R>(result));
return {};
}
template <class... R>
mediapipe::Status operator()(CalculatorContext* cc,
mediapipe::StatusOr<std::tuple<R...>>&& result) {
absl::Status operator()(CalculatorContext* cc,
absl::StatusOr<std::tuple<R...>>&& result) {
if (result.ok()) {
return this(cc, result.ValueOrDie());
return this(cc, result.value());
} else {
return result.status();
}
}
template <class... R>
mediapipe::Status operator()(CalculatorContext* cc,
std::tuple<R...>&& result) {
absl::Status operator()(CalculatorContext* cc, std::tuple<R...>&& result) {
static_assert(sizeof...(P) == sizeof...(R), "");
internal::tuple_for_each(
[cc, &result](const auto& port, auto i_const) {
@@ -345,9 +342,9 @@ class FunCaller {
auto inputs() const { return internal::filter_tuple<IsInputPort>(args_); }
auto outputs() const { return internal::filter_tuple<IsOutputPort>(args_); }
mediapipe::Status AddToContract(CalculatorContract* cc) const { return {}; }
absl::Status AddToContract(CalculatorContract* cc) const { return {}; }
mediapipe::Status Process(CalculatorContext* cc) const { return (*this)(cc); }
absl::Status Process(CalculatorContext* cc) const { return (*this)(cc); }
constexpr std::tuple<P...> nested_items() const { return args_; }
@@ -359,16 +356,14 @@ class FunCaller {
// TODO: implement multiple callers for syncsets.
template <class... T>
mediapipe::Status ProcessFnCallers(CalculatorContext* cc,
std::tuple<T...> callers);
absl::Status ProcessFnCallers(CalculatorContext* cc, std::tuple<T...> callers);
inline mediapipe::Status ProcessFnCallers(CalculatorContext* cc, std::tuple<>) {
return mediapipe::InternalError("Process unimplemented");
inline absl::Status ProcessFnCallers(CalculatorContext* cc, std::tuple<>) {
return absl::InternalError("Process unimplemented");
}
template <class T>
mediapipe::Status ProcessFnCallers(CalculatorContext* cc,
std::tuple<T> callers) {
absl::Status ProcessFnCallers(CalculatorContext* cc, std::tuple<T> callers) {
return std::get<0>(callers).Process(cc);
}
+1 -1
View File
@@ -10,7 +10,7 @@ namespace api2 {
namespace {
struct ProcessItem {
mediapipe::Status Process(CalculatorContext* cc) { return {}; }
absl::Status Process(CalculatorContext* cc) { return {}; }
};
struct ItemWithNested {
+3 -3
View File
@@ -34,7 +34,7 @@ class CalculatorBaseFactoryFor<
typename std::enable_if<std::is_base_of<mediapipe::api2::Node, T>{}>::type>
: public CalculatorBaseFactory {
public:
mediapipe::Status GetContract(CalculatorContract* cc) final {
absl::Status GetContract(CalculatorContract* cc) final {
auto status = T::Contract::GetContract(cc);
if (status.ok()) {
status = UpdateContract<T>(cc);
@@ -54,7 +54,7 @@ class CalculatorBaseFactoryFor<
return U::UpdateContract(cc);
}
template <typename U>
mediapipe::Status UpdateContract(...) {
absl::Status UpdateContract(...) {
return {};
}
};
@@ -142,7 +142,7 @@ class RegisteredNode<void> : public Node {};
template <class Impl>
struct FunctionNode : public RegisteredNode<Impl> {
mediapipe::Status Process(CalculatorContext* cc) override {
absl::Status Process(CalculatorContext* cc) override {
return internal::ProcessFnCallers(cc, Impl::kContract.process_items());
}
};
+55 -12
View File
@@ -12,6 +12,7 @@
#include "mediapipe/framework/port/gmock.h"
#include "mediapipe/framework/port/gtest.h"
#include "mediapipe/framework/port/parse_text_proto.h"
#include "mediapipe/framework/port/status_macros.h"
#include "mediapipe/framework/port/status_matchers.h"
namespace mediapipe {
@@ -32,7 +33,7 @@ std::vector<T> PacketValues(const std::vector<mediapipe::Packet>& packets) {
class FooImpl : public NodeImpl<Foo, FooImpl> {
public:
mediapipe::Status Process(CalculatorContext* cc) override {
absl::Status Process(CalculatorContext* cc) override {
float bias = kBias(cc).GetOr(0.0);
float scale = kScale(cc).GetOr(1.0);
kOut(cc).Send(*kBase(cc) * scale + bias);
@@ -80,7 +81,7 @@ class Foo5 : public FunctionNode<Foo5> {
class Foo2Impl : public NodeImpl<Foo2, Foo2Impl> {
public:
mediapipe::Status Process(CalculatorContext* cc) override {
absl::Status Process(CalculatorContext* cc) override {
float bias = SideIn(MPP_TAG("BIAS"), cc).GetOr(0.0);
float scale = In(MPP_TAG("SCALE"), cc).GetOr(1.0);
Out(MPP_TAG("OUT"), cc).Send(*In(MPP_TAG("BASE"), cc) * scale + bias);
@@ -90,7 +91,7 @@ class Foo2Impl : public NodeImpl<Foo2, Foo2Impl> {
class BarImpl : public NodeImpl<Bar, BarImpl> {
public:
mediapipe::Status Process(CalculatorContext* cc) override {
absl::Status Process(CalculatorContext* cc) override {
Packet p = kIn(cc);
kOut(cc).Send(p);
return {};
@@ -99,9 +100,9 @@ class BarImpl : public NodeImpl<Bar, BarImpl> {
class BazImpl : public NodeImpl<Baz> {
public:
static mediapipe::Status UpdateContract(CalculatorContract* cc) { return {}; }
static absl::Status UpdateContract(CalculatorContract* cc) { return {}; }
mediapipe::Status Process(CalculatorContext* cc) override {
absl::Status Process(CalculatorContext* cc) override {
for (int i = 0; i < kData(cc).Count(); ++i) {
kDataOut(cc)[i].Send(kData(cc)[i]);
}
@@ -112,7 +113,7 @@ MEDIAPIPE_NODE_IMPLEMENTATION(BazImpl);
class IntForwarderImpl : public NodeImpl<IntForwarder, IntForwarderImpl> {
public:
mediapipe::Status Process(CalculatorContext* cc) override {
absl::Status Process(CalculatorContext* cc) override {
kOut(cc).Send(*kIn(cc));
return {};
}
@@ -120,7 +121,7 @@ class IntForwarderImpl : public NodeImpl<IntForwarder, IntForwarderImpl> {
class ToFloatImpl : public NodeImpl<ToFloat, ToFloatImpl> {
public:
mediapipe::Status Process(CalculatorContext* cc) override {
absl::Status Process(CalculatorContext* cc) override {
kIn(cc).Visit([cc](auto x) { kOut(cc).Send(x); });
return {};
}
@@ -315,7 +316,7 @@ struct SideFallback : public Node {
MEDIAPIPE_NODE_CONTRACT(kIn, kFactor, kOut);
mediapipe::Status Process(CalculatorContext* cc) override {
absl::Status Process(CalculatorContext* cc) override {
kOut(cc).Send(kIn(cc).Get() * kFactor(cc).Get());
return {};
}
@@ -341,7 +342,7 @@ TEST(NodeTest, SideFallbackWithStream) {
MP_EXPECT_OK(
graph.ObserveOutputStream("out", [&outputs](const mediapipe::Packet& p) {
outputs.push_back(p.Get<int>());
return mediapipe::OkStatus();
return absl::OkStatus();
}));
MP_EXPECT_OK(graph.StartRun({}));
MP_EXPECT_OK(graph.AddPacketToInputStream(
@@ -372,7 +373,7 @@ TEST(NodeTest, SideFallbackWithSide) {
MP_EXPECT_OK(
graph.ObserveOutputStream("out", [&outputs](const mediapipe::Packet& p) {
outputs.push_back(p.Get<int>());
return mediapipe::OkStatus();
return absl::OkStatus();
}));
MP_EXPECT_OK(graph.StartRun({{"factor", mediapipe::MakePacket<int>(2)}}));
MP_EXPECT_OK(graph.AddPacketToInputStream(
@@ -451,7 +452,7 @@ struct DropEvenTimestamps : public Node {
MEDIAPIPE_NODE_CONTRACT(kIn, kOut);
mediapipe::Status Process(CalculatorContext* cc) override {
absl::Status Process(CalculatorContext* cc) override {
if (cc->InputTimestamp().Value() % 2) {
kOut(cc).Send(kIn(cc));
}
@@ -466,7 +467,7 @@ struct ListIntPackets : public Node {
MEDIAPIPE_NODE_CONTRACT(kIn, kOut);
mediapipe::Status Process(CalculatorContext* cc) override {
absl::Status Process(CalculatorContext* cc) override {
std::string result = absl::StrCat(cc->InputTimestamp().DebugString(), ":");
for (int i = 0; i < kIn(cc).Count(); ++i) {
if (kIn(cc)[i].IsEmpty()) {
@@ -522,6 +523,48 @@ TEST(NodeTest, DefaultTimestampChange0) {
MP_EXPECT_OK(graph.WaitUntilDone());
}
struct ConsumerNode : public Node {
static constexpr Input<int> kInt{"INT"};
static constexpr Input<AnyType> kGeneric{"ANY"};
static constexpr Input<OneOf<int, float>> kOneOf{"NUM"};
MEDIAPIPE_NODE_CONTRACT(kInt, kGeneric, kOneOf);
absl::Status Process(CalculatorContext* cc) override {
ASSIGN_OR_RETURN(auto maybe_int, kInt(cc).Consume());
ASSIGN_OR_RETURN(auto maybe_float, kGeneric(cc).Consume<float>());
ASSIGN_OR_RETURN(auto maybe_int2, kOneOf(cc).Consume<int>());
return {};
}
};
MEDIAPIPE_REGISTER_NODE(ConsumerNode);
TEST(NodeTest, ConsumeInputs) {
CalculatorGraphConfig config =
mediapipe::ParseTextProtoOrDie<CalculatorGraphConfig>(R"(
input_stream: "int"
input_stream: "any"
input_stream: "num"
node {
calculator: "ConsumerNode"
input_stream: "INT:int"
input_stream: "ANY:any"
input_stream: "NUM:num"
}
)");
mediapipe::CalculatorGraph graph;
MP_EXPECT_OK(graph.Initialize(config, {}));
MP_EXPECT_OK(graph.StartRun({}));
MP_EXPECT_OK(graph.AddPacketToInputStream(
"int", mediapipe::MakePacket<int>(10).At(Timestamp(0))));
MP_EXPECT_OK(graph.AddPacketToInputStream(
"any", mediapipe::MakePacket<float>(10).At(Timestamp(0))));
MP_EXPECT_OK(graph.AddPacketToInputStream(
"num", mediapipe::MakePacket<int>(10).At(Timestamp(0))));
MP_EXPECT_OK(graph.CloseAllPacketSources());
MP_EXPECT_OK(graph.WaitUntilDone());
}
} // namespace test
} // namespace api2
} // namespace mediapipe
+10
View File
@@ -7,9 +7,19 @@ PacketBase FromOldPacket(const mediapipe::Packet& op) {
return PacketBase(packet_internal::GetHolderShared(op)).At(op.Timestamp());
}
PacketBase FromOldPacket(mediapipe::Packet&& op) {
Timestamp t = op.Timestamp();
return PacketBase(packet_internal::GetHolderShared(std::move(op))).At(t);
}
mediapipe::Packet ToOldPacket(const PacketBase& p) {
return mediapipe::packet_internal::Create(p.payload_, p.timestamp_);
}
mediapipe::Packet ToOldPacket(PacketBase&& p) {
return mediapipe::packet_internal::Create(std::move(p.payload_),
p.timestamp_);
}
} // namespace api2
} // namespace mediapipe
+115 -1
View File
@@ -13,6 +13,7 @@
#include <functional>
#include <type_traits>
#include "absl/meta/type_traits.h"
#include "mediapipe/framework/api2/tuple.h"
#include "mediapipe/framework/packet.h"
#include "mediapipe/framework/port/logging.h"
@@ -58,7 +59,22 @@ class PacketBase {
const T& Get() const;
// Conversion to old Packet type.
operator mediapipe::Packet() const { return ToOldPacket(*this); }
operator mediapipe::Packet() const& { return ToOldPacket(*this); }
operator mediapipe::Packet() && { return ToOldPacket(std::move(*this)); }
// Note: Consume is included for compatibility with the old Packet; however,
// it relies on shared_ptr.unique(), which is deprecated and is not guaranteed
// to give exact results.
template <typename T>
absl::StatusOr<std::unique_ptr<T>> Consume() {
// Using the implementation in the old Packet for now.
mediapipe::Packet old =
packet_internal::Create(std::move(payload_), timestamp_);
auto result = old.Consume<T>();
if (!result.ok())
payload_ = packet_internal::GetHolderShared(std::move(old));
return result;
}
protected:
explicit PacketBase(std::shared_ptr<HolderBase> payload)
@@ -70,11 +86,15 @@ class PacketBase {
template <typename T>
friend PacketBase PacketBaseAdopting(const T* ptr);
friend PacketBase FromOldPacket(const mediapipe::Packet& op);
friend PacketBase FromOldPacket(mediapipe::Packet&& op);
friend mediapipe::Packet ToOldPacket(const PacketBase& p);
friend mediapipe::Packet ToOldPacket(PacketBase&& p);
};
PacketBase FromOldPacket(const mediapipe::Packet& op);
PacketBase FromOldPacket(mediapipe::Packet&& op);
mediapipe::Packet ToOldPacket(const PacketBase& p);
mediapipe::Packet ToOldPacket(PacketBase&& p);
template <typename T>
inline const T& PacketBase::Get() const {
@@ -132,6 +152,16 @@ struct Generic {
Generic() = delete;
};
template <class V, class U>
struct IsCompatibleType : std::false_type {};
template <class V>
struct IsCompatibleType<V, V> : std::true_type {};
template <class V>
struct IsCompatibleType<V, internal::Generic> : std::true_type {};
template <class V, class... U>
struct IsCompatibleType<V, OneOf<U...>>
: std::integral_constant<bool, (std::is_same_v<V, U> || ...)> {};
}; // namespace internal
template <typename T>
@@ -191,6 +221,13 @@ class Packet : public Packet<internal::Generic> {
return IsEmpty() ? static_cast<T>(absl::forward<U>(v)) : **this;
}
// Note: Consume is included for compatibility with the old Packet; however,
// it relies on shared_ptr.unique(), which is deprecated and is not guaranteed
// to give exact results.
absl::StatusOr<std::unique_ptr<T>> Consume() {
return PacketBase::Consume<T>();
}
private:
explicit Packet(std::shared_ptr<HolderBase> payload)
: Packet<internal::Generic>(std::move(payload)) {}
@@ -216,6 +253,44 @@ template <class T, class... U>
struct First {
using type = T;
};
template <class T>
struct AddStatus {
using type = StatusOr<T>;
};
template <class T>
struct AddStatus<StatusOr<T>> {
using type = StatusOr<T>;
};
template <>
struct AddStatus<Status> {
using type = Status;
};
template <>
struct AddStatus<void> {
using type = Status;
};
template <class R, class F, class... A>
struct CallAndAddStatusImpl {
typename AddStatus<R>::type operator()(const F& f, A&&... a) {
return f(std::forward<A>(a)...);
}
};
template <class F, class... A>
struct CallAndAddStatusImpl<void, F, A...> {
Status operator()(const F& f, A&&... a) {
f(std::forward<A>(a)...);
return {};
}
};
template <class F, class... A>
auto CallAndAddStatus(const F& f, A&&... a) {
return CallAndAddStatusImpl<absl::result_of_t<F(A...)>, F, A...>()(
f, std::forward<A>(a)...);
}
} // namespace internal
template <class... T>
@@ -276,6 +351,30 @@ class Packet<OneOf<T...>> : public PacketBase {
return Invoke<decltype(f), T...>(f);
}
// Note: Consume is included for compatibility with the old Packet; however,
// it relies on shared_ptr.unique(), which is deprecated and is not guaranteed
// to give exact results.
template <class U, class = AllowedType<U>>
absl::StatusOr<std::unique_ptr<U>> Consume() {
return PacketBase::Consume<U>();
}
template <class... F>
auto ConsumeAndVisit(const F&... args) {
CHECK(payload_);
auto f = internal::Overload{args...};
using FirstT = typename internal::First<T...>::type;
using VisitorResultType =
absl::result_of_t<decltype(f)(std::unique_ptr<FirstT>)>;
static_assert(
(std::is_same_v<VisitorResultType,
absl::result_of_t<decltype(f)(std::unique_ptr<T>)>> &&
...),
"All visitor overloads must have the same return type");
using ResultType = typename internal::AddStatus<VisitorResultType>::type;
return InvokeConsuming<ResultType, decltype(f), T...>(f);
}
protected:
explicit Packet(std::shared_ptr<HolderBase> payload)
: PacketBase(std::move(payload)) {}
@@ -292,6 +391,21 @@ class Packet<OneOf<T...>> : public PacketBase {
auto Invoke(const F& f) const {
return Has<U>() ? f(Get<U>()) : Invoke<F, V, W...>(f);
}
template <class R, class F, class U>
auto InvokeConsuming(const F& f) -> R {
auto maybe_value = Consume<U>();
if (maybe_value.ok())
return internal::CallAndAddStatus(f, std::move(maybe_value).value());
else
return maybe_value.status();
}
template <class R, class F, class U, class V, class... W>
auto InvokeConsuming(const F& f) -> R {
return Has<U>() ? InvokeConsuming<R, F, U>(f)
: InvokeConsuming<R, F, V, W...>(f);
}
};
template <>
+52
View File
@@ -1,7 +1,9 @@
#include "mediapipe/framework/api2/packet.h"
#include "absl/strings/str_cat.h"
#include "mediapipe/framework/port/gmock.h"
#include "mediapipe/framework/port/gtest.h"
#include "mediapipe/framework/port/status_matchers.h"
namespace mediapipe {
namespace api2 {
@@ -168,12 +170,26 @@ TEST(PacketTest, FromOldPacket) {
mediapipe::Packet op = mediapipe::MakePacket<int>(7);
Packet<int> p = FromOldPacket(op).As<int>();
EXPECT_EQ(p.Get(), 7);
EXPECT_EQ(op.Get<int>(), 7);
}
TEST(PacketTest, FromOldPacketConsume) {
mediapipe::Packet op = mediapipe::MakePacket<int>(7);
Packet<int> p = FromOldPacket(std::move(op)).As<int>();
MP_EXPECT_OK(p.Consume());
}
TEST(PacketTest, ToOldPacket) {
auto p = MakePacket<int>(7);
mediapipe::Packet op = ToOldPacket(p);
EXPECT_EQ(op.Get<int>(), 7);
EXPECT_EQ(p.Get(), 7);
}
TEST(PacketTest, ToOldPacketConsume) {
auto p = MakePacket<int>(7);
mediapipe::Packet op = ToOldPacket(std::move(p));
MP_EXPECT_OK(op.Consume<int>());
}
TEST(PacketTest, OldRefCounting) {
@@ -190,6 +206,42 @@ TEST(PacketTest, OldRefCounting) {
EXPECT_FALSE(alive);
}
TEST(PacketTest, Consume) {
auto p = MakePacket<int>(7);
auto maybe_int = p.Consume();
EXPECT_TRUE(p.IsEmpty());
ASSERT_TRUE(maybe_int.ok());
EXPECT_EQ(*maybe_int.value(), 7);
p = MakePacket<int>(3);
auto p2 = p;
maybe_int = p.Consume();
EXPECT_FALSE(maybe_int.ok());
EXPECT_FALSE(p.IsEmpty());
EXPECT_FALSE(p2.IsEmpty());
}
TEST(PacketTest, OneOfConsume) {
Packet<OneOf<std::string, int>> p = MakePacket<std::string>("hi");
EXPECT_TRUE(p.Has<std::string>());
EXPECT_FALSE(p.Has<int>());
EXPECT_EQ(p.Get<std::string>(), "hi");
absl::StatusOr<std::string> out = p.ConsumeAndVisit(
[](std::unique_ptr<std::string> s) {
return absl::StrCat("string: ", *s);
},
[](std::unique_ptr<int> i) { return absl::StrCat("int: ", *i); });
MP_EXPECT_OK(out);
EXPECT_EQ(out.value(), "string: hi");
EXPECT_TRUE(p.IsEmpty());
p = MakePacket<int>(3);
absl::Status out2 = p.ConsumeAndVisit([](std::unique_ptr<std::string> s) {},
[](std::unique_ptr<int> i) {});
MP_EXPECT_OK(out2);
EXPECT_TRUE(p.IsEmpty());
}
} // namespace
} // namespace api2
} // namespace mediapipe
+93 -15
View File
@@ -179,7 +179,7 @@ inline void SetType(CalculatorContract* cc, PacketType& pt) {
template <typename ValueT>
InputShardAccess<ValueT> SinglePortAccess(mediapipe::CalculatorContext* cc,
const InputStreamShard* stream) {
InputStreamShard* stream) {
return InputShardAccess<ValueT>(*cc, stream);
}
@@ -203,7 +203,7 @@ OutputSidePacketAccess<ValueT> SinglePortAccess(
template <typename ValueT>
InputShardOrSideAccess<ValueT> SinglePortAccess(
mediapipe::CalculatorContext* cc, const InputStreamShard* stream,
mediapipe::CalculatorContext* cc, InputStreamShard* stream,
const mediapipe::Packet* packet) {
return InputShardOrSideAccess<ValueT>(*cc, stream, packet);
}
@@ -226,19 +226,50 @@ auto AccessPort(std::false_type, const PortT& port, CC* cc) {
template <typename ValueT, typename X, class CC>
class MultiplePortAccess {
public:
using AccessT = decltype(SinglePortAccess<ValueT>(std::declval<CC*>(),
std::declval<X*>()));
MultiplePortAccess(CC* cc, X* first, int count)
: cc_(cc), first_(first), count_(count) {}
// TODO: maybe this should be size(), like in a standard C++
// container?
int Count() { return count_; }
auto operator[](int pos) {
AccessT operator[](int pos) {
CHECK_GE(pos, 0);
CHECK_LT(pos, count_);
return SinglePortAccess<ValueT>(cc_, &first_[pos]);
}
// TODO: add begin/end.
class Iterator {
public:
using iterator_category = std::input_iterator_tag;
using value_type = AccessT;
using difference_type = std::ptrdiff_t;
using pointer = AccessT*;
using reference = AccessT; // allowed; see e.g. std::istreambuf_iterator
Iterator(CC* cc, X* p) : cc_(cc), p_(p) {}
Iterator& operator++() {
++p_;
return *this;
}
Iterator operator++(int) {
Iterator res = *this;
++(*this);
return res;
}
bool operator==(const Iterator& other) const { return p_ == other.p_; }
bool operator!=(const Iterator& other) const { return !(*this == other); }
AccessT operator*() const { return SinglePortAccess<ValueT>(cc_, p_); }
private:
CC* cc_;
X* p_;
};
Iterator begin() { return Iterator(cc_, first_); }
Iterator end() { return Iterator(cc_, first_ + count_); }
private:
CC* cc_;
@@ -307,7 +338,7 @@ class PortCommon : public Base {
}
private:
mediapipe::Status AddToContract(CalculatorContract* cc) const {
absl::Status AddToContract(CalculatorContract* cc) const {
if (kMultiple) {
AddMultiple(cc);
} else {
@@ -385,17 +416,17 @@ class SideFallbackT : public Base {
side_port(tag) {}
protected:
mediapipe::Status AddToContract(CalculatorContract* cc) const {
absl::Status AddToContract(CalculatorContract* cc) const {
stream_port.AddToContract(cc);
side_port.AddToContract(cc);
int connected_count =
stream_port(cc).IsConnected() + side_port(cc).IsConnected();
if (connected_count > 1)
return mediapipe::InvalidArgumentError(absl::StrCat(
return absl::InvalidArgumentError(absl::StrCat(
Tag(),
" can be connected as a stream or as a side packet, but not both"));
if (!IsOptionalV && connected_count == 0)
return mediapipe::InvalidArgumentError(
return absl::InvalidArgumentError(
absl::StrCat(Tag(), " must be connected"));
return {};
}
@@ -452,6 +483,14 @@ class OutputShardAccess : public OutputShardAccessBase {
void Send(const T& payload) { Send(payload, context_.InputTimestamp()); }
void Send(T&& payload, Timestamp time) {
Send(api2::MakePacket<T>(std::move(payload)).At(time));
}
void Send(T&& payload) {
Send(std::move(payload), context_.InputTimestamp());
}
void Send(std::unique_ptr<T> payload, Timestamp time) {
Send(api2::PacketAdopting(std::move(payload)).At(time));
}
@@ -501,6 +540,7 @@ class OutputSidePacketAccess {
}
void Set(const T& payload) { Set(MakePacket<T>(payload)); }
void Set(T&& payload) { Set(MakePacket<T>(std::move(payload))); }
private:
OutputSidePacketAccess(OutputSidePacket* output) : output_(output) {}
@@ -523,15 +563,54 @@ class InputShardAccess : public Packet<T> {
PacketBase Header() const { return FromOldPacket(stream_->Header()); }
// "Consume" requires exclusive ownership of the packet's payload. In the
// current interim implementation, InputShardAccess creates a new reference to
// the payload (as a Packet<T> instead of a type-erased Packet), which means
// the conditions for Consume would never be satisfied. This helper class
// defines wrappers for the Consume methods in Packet which temporarily erase
// the reference held by the underlying InputStreamShard.
// Note that we cannot simply take over the reference when InputShardAccess is
// created, because it is currently created as a temporary and we might create
// more than one instance for the same stream.
template <class U = T,
class = std::enable_if_t<std::is_same<U, T>{},
decltype(&Packet<U>::Consume)>>
absl::StatusOr<std::unique_ptr<U>> Consume() {
return WrapConsumeCall(&Packet<T>::Consume);
}
template <class V, class U = T,
std::enable_if_t<internal::IsCompatibleType<V, U>{}, int> = 0>
absl::StatusOr<std::unique_ptr<V>> Consume() {
return WrapConsumeCall(&Packet<T>::template Consume<V>);
}
template <class... F>
auto ConsumeAndVisit(F&&... args) {
auto f = &Packet<T>::template ConsumeAndVisit<F...>;
return WrapConsumeCall(f, std::forward<F>(args)...);
}
private:
InputShardAccess(const CalculatorContext&, const InputStreamShard* stream)
InputShardAccess(const CalculatorContext&, InputStreamShard* stream)
: Packet<T>(stream ? FromOldPacket(stream->Value()).template As<T>()
: Packet<T>()),
stream_(stream) {}
const InputStreamShard* stream_;
template <class F, class... A>
auto WrapConsumeCall(F f, A&&... args) {
stream_->Value() = {};
auto result = (this->*f)(std::forward<A>(args)...);
if (!result.ok()) {
stream_->Value() = ToOldPacket(*this);
}
return result;
}
InputStreamShard* stream_;
friend InputShardAccess<T> internal::SinglePortAccess<T>(
mediapipe::CalculatorContext*, const InputStreamShard*);
mediapipe::CalculatorContext*, InputStreamShard*);
};
template <typename T>
@@ -566,19 +645,18 @@ class InputShardOrSideAccess : public Packet<T> {
PacketBase Header() const { return FromOldPacket(stream_->Header()); }
private:
InputShardOrSideAccess(const CalculatorContext&,
const InputStreamShard* stream,
InputShardOrSideAccess(const CalculatorContext&, InputStreamShard* stream,
const mediapipe::Packet* packet)
: Packet<T>(stream ? FromOldPacket(stream->Value()).template As<T>()
: packet ? FromOldPacket(*packet).template As<T>()
: Packet<T>()),
stream_(stream),
connected_(stream_ != nullptr || packet != nullptr) {}
const InputStreamShard* stream_;
InputStreamShard* stream_;
bool connected_;
friend InputShardOrSideAccess<T> internal::SinglePortAccess<T>(
mediapipe::CalculatorContext*, const InputStreamShard*,
mediapipe::CalculatorContext*, InputStreamShard*,
const mediapipe::Packet*);
};
+4 -4
View File
@@ -17,7 +17,7 @@ namespace test {
class FooBarImpl1 : public SubgraphImpl<FooBar1, FooBarImpl1> {
public:
mediapipe::StatusOr<CalculatorGraphConfig> GetConfig(
absl::StatusOr<CalculatorGraphConfig> GetConfig(
const SubgraphOptions& /*options*/) {
builder::Graph graph;
auto& foo = graph.AddNode("Foo");
@@ -31,7 +31,7 @@ class FooBarImpl1 : public SubgraphImpl<FooBar1, FooBarImpl1> {
class FooBarImpl2 : public SubgraphImpl<FooBar2, FooBarImpl2> {
public:
mediapipe::StatusOr<CalculatorGraphConfig> GetConfig(
absl::StatusOr<CalculatorGraphConfig> GetConfig(
const SubgraphOptions& /*options*/) {
builder::Graph graph;
auto& foo = graph.AddNode<Foo>();
@@ -44,7 +44,7 @@ class FooBarImpl2 : public SubgraphImpl<FooBar2, FooBarImpl2> {
};
TEST(SubgraphTest, SubgraphConfig) {
CalculatorGraphConfig subgraph = FooBarImpl1().GetConfig({}).ValueOrDie();
CalculatorGraphConfig subgraph = FooBarImpl1().GetConfig({}).value();
const CalculatorGraphConfig expected_graph =
mediapipe::ParseTextProtoOrDie<CalculatorGraphConfig>(R"(
input_stream: "IN:__stream_0"
@@ -64,7 +64,7 @@ TEST(SubgraphTest, SubgraphConfig) {
}
TEST(SubgraphTest, TypedSubgraphConfig) {
CalculatorGraphConfig subgraph = FooBarImpl2().GetConfig({}).ValueOrDie();
CalculatorGraphConfig subgraph = FooBarImpl2().GetConfig({}).value();
const CalculatorGraphConfig expected_graph =
mediapipe::ParseTextProtoOrDie<CalculatorGraphConfig>(R"(
input_stream: "IN:__stream_0"