Add logic to handle ANNOUNCE messages. Also let the outgoing ANNOUNCE callbacks access the error code. PiperOrigin-RevId: 608760663
diff --git a/quiche/quic/moqt/moqt_integration_test.cc b/quiche/quic/moqt/moqt_integration_test.cc index a690d73..013fcc1 100644 --- a/quiche/quic/moqt/moqt_integration_test.cc +++ b/quiche/quic/moqt/moqt_integration_test.cc
@@ -31,6 +31,7 @@ using ::quic::simulator::Simulator; using ::testing::_; using ::testing::Assign; +using ::testing::Return; class ClientEndpoint : public quic::simulator::QuicEndpointWithConnection { public: @@ -62,6 +63,7 @@ testing::MockFunction<void(absl::string_view)>& terminated_callback() { return callbacks_.session_terminated_callback; } + MockSessionCallbacks& callbacks() { return callbacks_; } private: MockSessionCallbacks callbacks_; @@ -104,6 +106,7 @@ testing::MockFunction<void(absl::string_view)>& terminated_callback() { return callbacks_.session_terminated_callback; } + MockSessionCallbacks& callbacks() { return callbacks_; } private: MockSessionCallbacks callbacks_; @@ -189,19 +192,44 @@ EXPECT_TRUE(success); } -TEST_F(MoqtIntegrationTest, AnnounceExchange) { +TEST_F(MoqtIntegrationTest, AnnounceSuccess) { EstablishSession(); - testing::MockFunction<void(absl::string_view track_namespace, - std::optional<absl::string_view> error_message)> + EXPECT_CALL(server_->callbacks().incoming_announce_callback, Call("foo")) + .WillOnce(Return(std::nullopt)); + testing::MockFunction<void( + absl::string_view track_namespace, + std::optional<MoqtAnnounceErrorReason> error_message)> announce_callback; client_->session()->Announce("foo", announce_callback.AsStdFunction()); bool matches = false; EXPECT_CALL(announce_callback, Call(_, _)) .WillOnce([&](absl::string_view track_namespace, - std::optional<absl::string_view> error_message) { + std::optional<MoqtAnnounceErrorReason> error) { matches = true; EXPECT_EQ(track_namespace, "foo"); - EXPECT_FALSE(error_message.has_value()); + EXPECT_FALSE(error.has_value()); + }); + bool success = + test_harness_.RunUntilWithDefaultTimeout([&]() { return matches; }); + EXPECT_TRUE(success); +} + +TEST_F(MoqtIntegrationTest, AnnounceFailure) { + EstablishSession(); + testing::MockFunction<void( + absl::string_view track_namespace, + std::optional<MoqtAnnounceErrorReason> error_message)> + announce_callback; + client_->session()->Announce("foo", announce_callback.AsStdFunction()); + bool matches = false; + EXPECT_CALL(announce_callback, Call(_, _)) + .WillOnce([&](absl::string_view track_namespace, + std::optional<MoqtAnnounceErrorReason> error) { + matches = true; + EXPECT_EQ(track_namespace, "foo"); + ASSERT_TRUE(error.has_value()); + EXPECT_EQ(error->error_code, + MoqtAnnounceErrorCode::kAnnounceNotSupported); }); bool success = test_harness_.RunUntilWithDefaultTimeout([&]() { return matches; });
diff --git a/quiche/quic/moqt/moqt_messages.h b/quiche/quic/moqt/moqt_messages.h index dcd23f9..081abcd 100644 --- a/quiche/quic/moqt/moqt_messages.h +++ b/quiche/quic/moqt/moqt_messages.h
@@ -90,6 +90,18 @@ kAuthorizationInfo = 0x2, }; +// TODO: those are non-standard; add standard error codes once those exist, see +// <https://github.com/moq-wg/moq-transport/issues/393>. +enum class MoqtAnnounceErrorCode : uint64_t { + kInternalError = 0, + kAnnounceNotSupported = 1, +}; + +struct MoqtAnnounceErrorReason { + MoqtAnnounceErrorCode error_code; + std::string reason_phrase; +}; + struct FullTrackName { std::string track_namespace; std::string track_name; @@ -276,7 +288,7 @@ struct QUICHE_EXPORT MoqtAnnounceError { absl::string_view track_namespace; - uint64_t error_code; + MoqtAnnounceErrorCode error_code; absl::string_view reason_phrase; };
diff --git a/quiche/quic/moqt/moqt_parser.cc b/quiche/quic/moqt/moqt_parser.cc index 96c8a92..3c8a31c 100644 --- a/quiche/quic/moqt/moqt_parser.cc +++ b/quiche/quic/moqt/moqt_parser.cc
@@ -518,9 +518,11 @@ if (!reader.ReadStringPieceVarInt62(&announce_error.track_namespace)) { return 0; } - if (!reader.ReadVarInt62(&announce_error.error_code)) { + uint64_t error_code; + if (!reader.ReadVarInt62(&error_code)) { return 0; } + announce_error.error_code = static_cast<MoqtAnnounceErrorCode>(error_code); if (!reader.ReadStringPieceVarInt62(&announce_error.reason_phrase)) { return 0; }
diff --git a/quiche/quic/moqt/moqt_session.cc b/quiche/quic/moqt/moqt_session.cc index 443dd50..1c7d00f 100644 --- a/quiche/quic/moqt/moqt_session.cc +++ b/quiche/quic/moqt/moqt_session.cc
@@ -133,10 +133,13 @@ // TODO: Create state that allows ANNOUNCE_OK/ERROR on spurious namespaces to // trigger session errors. void MoqtSession::Announce(absl::string_view track_namespace, - MoqtAnnounceCallback announce_callback) { + MoqtOutgoingAnnounceCallback announce_callback) { if (pending_outgoing_announces_.contains(track_namespace)) { std::move(announce_callback)( - track_namespace, "ANNOUNCE message already outstanding for namespace"); + track_namespace, + MoqtAnnounceErrorReason{ + MoqtAnnounceErrorCode::kInternalError, + "ANNOUNCE message already outstanding for namespace"}); return; } MoqtAnnounce message; @@ -648,6 +651,16 @@ if (!CheckIfIsControlStream()) { return; } + std::optional<MoqtAnnounceErrorReason> error = + session_->incoming_announce_callback_(message.track_namespace); + if (error.has_value()) { + MoqtAnnounceError reply; + reply.track_namespace = message.track_namespace; + reply.error_code = error->error_code; + reply.reason_phrase = error->reason_phrase; + SendOrBufferMessage(session_->framer_.SerializeAnnounceError(reply)); + return; + } MoqtAnnounceOk ok; ok.track_namespace = message.track_namespace; SendOrBufferMessage(session_->framer_.SerializeAnnounceOk(ok)); @@ -678,7 +691,10 @@ "Received ANNOUNCE_ERROR for nonexistent announce"); return; } - std::move(it->second)(message.track_namespace, message.reason_phrase); + std::move(it->second)( + message.track_namespace, + MoqtAnnounceErrorReason{message.error_code, + std::string(message.reason_phrase)}); session_->pending_outgoing_announces_.erase(it); }
diff --git a/quiche/quic/moqt/moqt_session.h b/quiche/quic/moqt/moqt_session.h index cb9163d..60e6dd1 100644 --- a/quiche/quic/moqt/moqt_session.h +++ b/quiche/quic/moqt/moqt_session.h
@@ -35,9 +35,19 @@ quiche::SingleUseCallback<void(absl::string_view error_message)>; using MoqtSessionDeletedCallback = quiche::SingleUseCallback<void()>; // If |error_message| is nullopt, the ANNOUNCE was successful. -using MoqtAnnounceCallback = quiche::SingleUseCallback<void( +using MoqtOutgoingAnnounceCallback = quiche::SingleUseCallback<void( absl::string_view track_namespace, - std::optional<absl::string_view> error_message)>; + std::optional<MoqtAnnounceErrorReason> error)>; +using MoqtIncomingAnnounceCallback = + quiche::MultiUseCallback<std::optional<MoqtAnnounceErrorReason>( + absl::string_view track_namespace)>; + +inline std::optional<MoqtAnnounceErrorReason> DefaultIncomingAnnounceCallback( + absl::string_view /*track_namespace*/) { + return std::optional(MoqtAnnounceErrorReason{ + MoqtAnnounceErrorCode::kAnnounceNotSupported, + "This endpoint does not accept incoming ANNOUNCE messages"}); +}; // Callbacks for session-level events. struct MoqtSessionCallbacks { @@ -45,6 +55,9 @@ MoqtSessionTerminatedCallback session_terminated_callback = +[](absl::string_view) {}; MoqtSessionDeletedCallback session_deleted_callback = +[] {}; + + MoqtIncomingAnnounceCallback incoming_announce_callback = + DefaultIncomingAnnounceCallback; }; class QUICHE_EXPORT MoqtSession : public webtransport::SessionVisitor { @@ -59,6 +72,8 @@ std::move(callbacks.session_terminated_callback)), session_deleted_callback_( std::move(callbacks.session_deleted_callback)), + incoming_announce_callback_( + std::move(callbacks.incoming_announce_callback)), framer_(quiche::SimpleBufferAllocator::Get(), parameters.using_webtrans) {} ~MoqtSession() { std::move(session_deleted_callback_)(); } @@ -87,7 +102,7 @@ // |announce_callback| when the response arrives. Will fail immediately if // there is already an unresolved ANNOUNCE for that namespace. void Announce(absl::string_view track_namespace, - MoqtAnnounceCallback announce_callback); + MoqtOutgoingAnnounceCallback announce_callback); bool HasSubscribers(const FullTrackName& full_track_name) const; // Returns true if SUBSCRIBE was sent. If there is already a subscription to @@ -217,6 +232,7 @@ MoqtSessionEstablishedCallback session_established_callback_; MoqtSessionTerminatedCallback session_terminated_callback_; MoqtSessionDeletedCallback session_deleted_callback_; + MoqtIncomingAnnounceCallback incoming_announce_callback_; MoqtFramer framer_; std::optional<webtransport::StreamId> control_stream_; @@ -246,7 +262,7 @@ uint64_t next_subscribe_id_ = 0; // Indexed by track namespace. - absl::flat_hash_map<std::string, MoqtAnnounceCallback> + absl::flat_hash_map<std::string, MoqtOutgoingAnnounceCallback> pending_outgoing_announces_; };
diff --git a/quiche/quic/moqt/moqt_session_test.cc b/quiche/quic/moqt/moqt_session_test.cc index 9401a24..a9a57c2 100644 --- a/quiche/quic/moqt/moqt_session_test.cc +++ b/quiche/quic/moqt/moqt_session_test.cc
@@ -313,8 +313,9 @@ } TEST_F(MoqtSessionTest, AnnounceWithOk) { - testing::MockFunction<void(absl::string_view track_namespace, - std::optional<absl::string_view> error_message)> + testing::MockFunction<void( + absl::string_view track_namespace, + std::optional<MoqtAnnounceErrorReason> error_message)> announce_resolved_callback; StrictMock<webtransport::test::MockStream> mock_stream; std::unique_ptr<MoqtParserVisitor> stream_input = @@ -337,18 +338,19 @@ correct_message = false; EXPECT_CALL(announce_resolved_callback, Call(_, _)) .WillOnce([&](absl::string_view track_namespace, - std::optional<absl::string_view> error_message) { + std::optional<MoqtAnnounceErrorReason> error) { correct_message = true; EXPECT_EQ(track_namespace, "foo"); - EXPECT_FALSE(error_message.has_value()); + EXPECT_FALSE(error.has_value()); }); stream_input->OnAnnounceOkMessage(ok); EXPECT_TRUE(correct_message); } TEST_F(MoqtSessionTest, AnnounceWithError) { - testing::MockFunction<void(absl::string_view track_namespace, - std::optional<absl::string_view> error_message)> + testing::MockFunction<void( + absl::string_view track_namespace, + std::optional<MoqtAnnounceErrorReason> error_message)> announce_resolved_callback; StrictMock<webtransport::test::MockStream> mock_stream; std::unique_ptr<MoqtParserVisitor> stream_input = @@ -367,14 +369,18 @@ MoqtAnnounceError error = { /*track_namespace=*/"foo", + /*error_code=*/MoqtAnnounceErrorCode::kInternalError, + /*reason_phrase=*/"Test error", }; correct_message = false; EXPECT_CALL(announce_resolved_callback, Call(_, _)) .WillOnce([&](absl::string_view track_namespace, - std::optional<absl::string_view> error_message) { + std::optional<MoqtAnnounceErrorReason> error) { correct_message = true; EXPECT_EQ(track_namespace, "foo"); - EXPECT_TRUE(error_message.has_value()); + ASSERT_TRUE(error.has_value()); + EXPECT_EQ(error->error_code, MoqtAnnounceErrorCode::kInternalError); + EXPECT_EQ(error->reason_phrase, "Test error"); }); stream_input->OnAnnounceErrorMessage(error); EXPECT_TRUE(correct_message); @@ -526,6 +532,8 @@ /*track_namespace=*/"foo", }; bool correct_message = false; + EXPECT_CALL(session_callbacks_.incoming_announce_callback, Call("foo")) + .WillOnce(Return(std::nullopt)); EXPECT_CALL(mock_stream, Writev(_, _)) .WillOnce([&](absl::Span<const absl::string_view> data, const quiche::StreamWriteOptions& options) {
diff --git a/quiche/quic/moqt/test_tools/moqt_test_message.h b/quiche/quic/moqt/test_tools/moqt_test_message.h index cd3d094..6a0891d 100644 --- a/quiche/quic/moqt/test_tools/moqt_test_message.h +++ b/quiche/quic/moqt/test_tools/moqt_test_message.h
@@ -818,7 +818,7 @@ MoqtAnnounceError announce_error_ = { /*track_namespace=*/"foo", - /*error_code=*/1, + /*error_code=*/MoqtAnnounceErrorCode::kAnnounceNotSupported, /*reason_phrase=*/"bar", }; };
diff --git a/quiche/quic/moqt/tools/chat_client_bin.cc b/quiche/quic/moqt/tools/chat_client_bin.cc index 8e639d6..9bc100c 100644 --- a/quiche/quic/moqt/tools/chat_client_bin.cc +++ b/quiche/quic/moqt/tools/chat_client_bin.cc
@@ -176,11 +176,11 @@ // By not sending a visitor, the application will not fulfill subscriptions // to previous objects. session_->AddLocalTrack(my_track_name_, nullptr); - moqt::MoqtAnnounceCallback announce_callback = + moqt::MoqtOutgoingAnnounceCallback announce_callback = [&](absl::string_view track_namespace, - std::optional<absl::string_view> message) { - if (message.has_value()) { - std::cout << "ANNOUNCE rejected, " << *message << "\n"; + std::optional<moqt::MoqtAnnounceErrorReason> reason) { + if (reason.has_value()) { + std::cout << "ANNOUNCE rejected, " << reason->reason_phrase << "\n"; session_->Error(moqt::MoqtError::kGenericError, "Local ANNOUNCE rejected"); return;
diff --git a/quiche/quic/moqt/tools/moqt_mock_visitor.h b/quiche/quic/moqt/tools/moqt_mock_visitor.h index 246e841..7d174a1 100644 --- a/quiche/quic/moqt/tools/moqt_mock_visitor.h +++ b/quiche/quic/moqt/tools/moqt_mock_visitor.h
@@ -21,11 +21,20 @@ testing::MockFunction<void()> session_established_callback; testing::MockFunction<void(absl::string_view)> session_terminated_callback; testing::MockFunction<void()> session_deleted_callback; + testing::MockFunction<std::optional<MoqtAnnounceErrorReason>( + absl::string_view)> + incoming_announce_callback; + + MockSessionCallbacks() { + ON_CALL(incoming_announce_callback, Call(testing::_)) + .WillByDefault(DefaultIncomingAnnounceCallback); + } MoqtSessionCallbacks AsSessionCallbacks() { return MoqtSessionCallbacks{session_established_callback.AsStdFunction(), session_terminated_callback.AsStdFunction(), - session_deleted_callback.AsStdFunction()}; + session_deleted_callback.AsStdFunction(), + incoming_announce_callback.AsStdFunction()}; } };