Move SUBSCRIBE to a dedicated bidirectional stream. This change refactors MoQT to handle SUBSCRIBE messages on their own bidirectional streams, rather than on the control stream. Key changes include: - New `MoqtSubscribeRequestStream` and `MoqtSubscribeResponseStream` classes. - `MoqtSession::Subscribe` now opens a bidi stream and sends the SUBSCRIBE message on it. - `MoqtSession::Unsubscribe` now resets the bidi stream associated with the subscription. - `MoqtSession::RequestUpdate` sends the update on the existing subscribe bidi stream. - The control stream no longer processes SUBSCRIBE, SUBSCRIBE_OK, or UNSUBSCRIBE messages. - `MoqtUnsubscribe` message type and serialization/parsing are removed. - MoqtSession no longer owns any subscriptions - All external state is now destroyed in Detach(). On teardown, streams are destroyed in arbitrary order, making clearing such state in the destructor dangerous. I wrote better helper functions in MoqtSessionTest to reduce the toil in sending/receiving messages, setting up bidi streams, and not hard-coding message formats. PiperOrigin-RevId: 949593252
diff --git a/build/source_list.bzl b/build/source_list.bzl index 30e8046..b1fa635 100644 --- a/build/source_list.bzl +++ b/build/source_list.bzl
@@ -1607,6 +1607,7 @@ "quic/moqt/moqt_session_callbacks.h", "quic/moqt/moqt_session_interface.h", "quic/moqt/moqt_stream_map.h", + "quic/moqt/moqt_subscribe_stream.h", "quic/moqt/moqt_subscription.h", "quic/moqt/moqt_trace_recorder.h", "quic/moqt/moqt_track.h", @@ -1643,6 +1644,7 @@ "quic/moqt/moqt_relay_track_publisher.cc", "quic/moqt/moqt_session.cc", "quic/moqt/moqt_stream_map.cc", + "quic/moqt/moqt_subscribe_stream.cc", "quic/moqt/moqt_subscription.cc", "quic/moqt/moqt_trace_recorder.cc", "quic/moqt/moqt_track.cc", @@ -1678,6 +1680,7 @@ "quic/moqt/moqt_relay_track_publisher_test.cc", "quic/moqt/moqt_session_test.cc", "quic/moqt/moqt_stream_map_test.cc", + "quic/moqt/moqt_subscribe_stream_test.cc", "quic/moqt/moqt_subscription_test.cc", "quic/moqt/moqt_track_test.cc", "quic/moqt/moqt_uni_stream_test.cc",
diff --git a/build/source_list.gni b/build/source_list.gni index 96933eb..43addc5 100644 --- a/build/source_list.gni +++ b/build/source_list.gni
@@ -1611,6 +1611,7 @@ "src/quiche/quic/moqt/moqt_session_callbacks.h", "src/quiche/quic/moqt/moqt_session_interface.h", "src/quiche/quic/moqt/moqt_stream_map.h", + "src/quiche/quic/moqt/moqt_subscribe_stream.h", "src/quiche/quic/moqt/moqt_subscription.h", "src/quiche/quic/moqt/moqt_trace_recorder.h", "src/quiche/quic/moqt/moqt_track.h", @@ -1647,6 +1648,7 @@ "src/quiche/quic/moqt/moqt_relay_track_publisher.cc", "src/quiche/quic/moqt/moqt_session.cc", "src/quiche/quic/moqt/moqt_stream_map.cc", + "src/quiche/quic/moqt/moqt_subscribe_stream.cc", "src/quiche/quic/moqt/moqt_subscription.cc", "src/quiche/quic/moqt/moqt_trace_recorder.cc", "src/quiche/quic/moqt/moqt_track.cc", @@ -1683,6 +1685,7 @@ "src/quiche/quic/moqt/moqt_relay_track_publisher_test.cc", "src/quiche/quic/moqt/moqt_session_test.cc", "src/quiche/quic/moqt/moqt_stream_map_test.cc", + "src/quiche/quic/moqt/moqt_subscribe_stream_test.cc", "src/quiche/quic/moqt/moqt_subscription_test.cc", "src/quiche/quic/moqt/moqt_track_test.cc", "src/quiche/quic/moqt/moqt_uni_stream_test.cc",
diff --git a/build/source_list.json b/build/source_list.json index f8008d3..ea8c649 100644 --- a/build/source_list.json +++ b/build/source_list.json
@@ -1610,6 +1610,7 @@ "quiche/quic/moqt/moqt_session_callbacks.h", "quiche/quic/moqt/moqt_session_interface.h", "quiche/quic/moqt/moqt_stream_map.h", + "quiche/quic/moqt/moqt_subscribe_stream.h", "quiche/quic/moqt/moqt_subscription.h", "quiche/quic/moqt/moqt_trace_recorder.h", "quiche/quic/moqt/moqt_track.h", @@ -1646,6 +1647,7 @@ "quiche/quic/moqt/moqt_relay_track_publisher.cc", "quiche/quic/moqt/moqt_session.cc", "quiche/quic/moqt/moqt_stream_map.cc", + "quiche/quic/moqt/moqt_subscribe_stream.cc", "quiche/quic/moqt/moqt_subscription.cc", "quiche/quic/moqt/moqt_trace_recorder.cc", "quiche/quic/moqt/moqt_track.cc", @@ -1682,6 +1684,7 @@ "quiche/quic/moqt/moqt_relay_track_publisher_test.cc", "quiche/quic/moqt/moqt_session_test.cc", "quiche/quic/moqt/moqt_stream_map_test.cc", + "quiche/quic/moqt/moqt_subscribe_stream_test.cc", "quiche/quic/moqt/moqt_subscription_test.cc", "quiche/quic/moqt/moqt_track_test.cc", "quiche/quic/moqt/moqt_uni_stream_test.cc",
diff --git a/quiche/quic/moqt/moqt_bidi_stream.cc b/quiche/quic/moqt/moqt_bidi_stream.cc index df12569..58bafea 100644 --- a/quiche/quic/moqt/moqt_bidi_stream.cc +++ b/quiche/quic/moqt/moqt_bidi_stream.cc
@@ -13,6 +13,7 @@ #include "absl/strings/string_view.h" #include "quiche/quic/core/quic_time.h" #include "quiche/quic/moqt/moqt_error.h" +#include "quiche/quic/moqt/moqt_fetch_task.h" #include "quiche/quic/moqt/moqt_key_value_pair.h" #include "quiche/quic/moqt/moqt_messages.h" #include "quiche/quic/moqt/moqt_parser.h" @@ -83,6 +84,42 @@ info.reason_phrase, fin); } +absl::Status MoqtBidiStreamBase::SendRequestUpdate( + uint64_t request_id, uint64_t existing_request_id, + const MessageParameters& parameters, MoqtResponseCallback callback) { + MoqtRequestUpdate request_update; + request_update.request_id = request_id; + request_update.existing_request_id = existing_request_id; + request_update.parameters = parameters; + pending_responses_.push_back(std::move(callback)); + pending_updates_.push_back(parameters); + return SendOrBufferMessage(framer_->SerializeRequestUpdate(request_update), + /*fin=*/false); +} + +absl::Status MoqtBidiStreamBase::OnControlMessage( + const MoqtRequestOk& message) { + if (pending_responses_.empty()) { + return absl::OkStatus(); + } + std::move(pending_responses_.front())(message.parameters); + pending_responses_.pop_front(); + return absl::OkStatus(); +} + +absl::Status MoqtBidiStreamBase::OnControlMessage( + const MoqtRequestError& message) { + if (pending_responses_.empty()) { + return absl::OkStatus(); + } + std::move(pending_responses_.front())(MoqtRequestErrorInfo{ + message.error_code, message.retry_interval, message.reason_phrase}); + pending_responses_.clear(); + pending_updates_.clear(); + Fin(); + return absl::OkStatus(); +} + void MoqtBidiStreamBase::OnFatalError(absl::Status status) { QUICHE_DCHECK(!status.ok()); if (session_error_callback_ == nullptr) {
diff --git a/quiche/quic/moqt/moqt_bidi_stream.h b/quiche/quic/moqt/moqt_bidi_stream.h index 0ebb74c..22af6b9 100644 --- a/quiche/quic/moqt/moqt_bidi_stream.h +++ b/quiche/quic/moqt/moqt_bidi_stream.h
@@ -13,11 +13,13 @@ #include "absl/base/nullability.h" #include "absl/status/status.h" +#include "absl/status/statusor.h" #include "absl/strings/str_cat.h" #include "absl/strings/string_view.h" #include "quiche/quic/core/quic_time.h" #include "quiche/quic/moqt/moqt_control_message_queue.h" #include "quiche/quic/moqt/moqt_error.h" +#include "quiche/quic/moqt/moqt_fetch_task.h" #include "quiche/quic/moqt/moqt_framer.h" #include "quiche/quic/moqt/moqt_key_value_pair.h" #include "quiche/quic/moqt/moqt_messages.h" @@ -25,6 +27,7 @@ #include "quiche/common/platform/api/quiche_logging.h" #include "quiche/common/quiche_buffer_allocator.h" #include "quiche/common/quiche_callbacks.h" +#include "quiche/common/quiche_circular_deque.h" #include "quiche/web_transport/web_transport.h" namespace moqt { @@ -102,7 +105,11 @@ absl::string_view reason_phrase, bool fin = false); absl::Status SendRequestError(uint64_t request_id, MoqtRequestErrorInfo info, bool fin = false); - + // Can be overridden for message-specific constraints. + virtual absl::Status SendRequestUpdate(uint64_t request_id, + uint64_t existing_request_id, + const MessageParameters& parameters, + MoqtResponseCallback callback); void Fin() { CheckStatus(outgoing_message_queue_.Fin()); Detach(); @@ -122,14 +129,13 @@ } } - // Removes any state in MoqtSession related to the stream. Overrides of this - // method must be robust to multiple invocations. - virtual void Detach() = 0; + MoqtFramer* framer() const { return framer_; } - // TODO(martinduke): Remove once SUBSCRIBE moves to a bidi stream. This is - // only needed to check whether or not to FIN the bidi stream. - bool is_control_stream() const { return control_stream_; } - void set_control_stream() { control_stream_ = true; } + // Removes any state in MoqtSession related to the stream. Overrides of this + // method must be robust to multiple invocations. Only called on sending a FIN + // or RESET. If otherwise destroyed, it's due to a larger cleanup where the + // state no longer matters. + virtual void Detach() = 0; protected: // Called when a WebTransport stream has been associated with the object. @@ -141,6 +147,18 @@ virtual absl::Status OnRawControlMessage( const MoqtRawControlMessage& message) = 0; + virtual absl::Status OnControlMessage(const MoqtRequestOk& message); + virtual absl::Status OnControlMessage(const MoqtRequestError& message); + + absl::StatusOr<MessageParameters> PopParameters() { + if (pending_updates_.empty()) { + return absl::NotFoundError("Too many REQUEST_OK received"); + } + MessageParameters parameters = pending_updates_.front(); + pending_updates_.pop_front(); + return parameters; + } + // Terminates the MoQT session due to a fatal error encountered. void OnFatalError(absl::Status status); @@ -148,7 +166,6 @@ const MoqtControlMessageParser& message_parser() const { return message_parser_; } - MoqtFramer* framer() const { return framer_; } webtransport::Stream* stream() const { return stream_parser_ != nullptr ? stream_parser_->stream() : nullptr; } @@ -156,12 +173,12 @@ private: friend class test::MoqtBidiStreamTestWrapper; - // TODO(martinduke): Remove once SUBSCRIBE moves to a bidi stream. - bool control_stream_ = false; MoqtFramer* absl_nonnull framer_; std::unique_ptr<MoqtControlStreamParser> absl_nullable stream_parser_; MoqtControlMessageParser message_parser_; MoqtControlMessageQueue outgoing_message_queue_; + quiche::QuicheCircularDeque<MoqtResponseCallback> pending_responses_; + quiche::QuicheCircularDeque<MessageParameters> pending_updates_; SessionErrorCallback session_error_callback_; };
diff --git a/quiche/quic/moqt/moqt_bidi_stream_test.cc b/quiche/quic/moqt/moqt_bidi_stream_test.cc index 8b5f011..9da7253 100644 --- a/quiche/quic/moqt/moqt_bidi_stream_test.cc +++ b/quiche/quic/moqt/moqt_bidi_stream_test.cc
@@ -6,15 +6,20 @@ #include <memory> #include <optional> +#include <utility> +#include <variant> #include "absl/status/status.h" #include "absl/strings/string_view.h" #include "quiche/quic/core/quic_types.h" #include "quiche/quic/moqt/moqt_error.h" +#include "quiche/quic/moqt/moqt_fetch_task.h" #include "quiche/quic/moqt/moqt_framer.h" +#include "quiche/quic/moqt/moqt_key_value_pair.h" #include "quiche/quic/moqt/moqt_messages.h" #include "quiche/quic/moqt/moqt_parser.h" #include "quiche/quic/moqt/moqt_session_interface.h" +#include "quiche/quic/moqt/test_tools/mock_moqt_session.h" #include "quiche/quic/moqt/test_tools/moqt_framer_utils.h" #include "quiche/common/platform/api/quiche_test.h" #include "quiche/common/test_tools/quiche_test_utils.h" @@ -27,11 +32,6 @@ public: using MoqtBidiStreamBase::MoqtBidiStreamBase; - absl::Status OnControlMessage(const MoqtRequestOk& message) { - ++ok_received_; - return absl::OkStatus(); - } - void OnStreamBound() override {} absl::Status OnRawControlMessage( const MoqtRawControlMessage& message) override { @@ -39,9 +39,14 @@ *this, message_parser(), message, "test"); } void Detach() override { detached_ = true; } - - int ok_received_ = 0; bool detached_ = false; + + absl::Status OnControlMessage(const MoqtRequestOk& message) { + return MoqtBidiStreamBase::OnControlMessage(message); + } + absl::Status OnControlMessage(const MoqtRequestError& message) { + return MoqtBidiStreamBase::OnControlMessage(message); + } }; class MoqtBidiStreamTest : public quiche::test::QuicheTest { @@ -118,7 +123,6 @@ MoqtFramer framer(/*using_webtrans=*/true, quic::Perspective::IS_SERVER); stream.Receive(framer.SerializeRequestOk(MoqtRequestOk()).AsStringView()); stream_->OnCanRead(); - EXPECT_EQ(stream_->ok_received_, 1u); stream.Receive(framer.SerializeGoAway(MoqtGoAway()).AsStringView()); EXPECT_CALL(error_callback_, Call) @@ -131,4 +135,98 @@ stream_->OnCanRead(); } +TEST_F(MoqtBidiStreamTest, SendRequestOk) { + stream_->BindStream(&mock_stream_); + EXPECT_CALL(mock_stream_, CanWrite).WillRepeatedly(testing::Return(true)); + EXPECT_CALL( + mock_stream_, + Writev(ControlMessageOfType(MoqtMessageType::kRequestOk), testing::_)); + MessageParameters parameters; + parameters.subscriber_priority = 20; + QUICHE_EXPECT_OK(stream_->SendRequestOk(1, parameters, /*fin=*/false)); + EXPECT_FALSE(stream_->detached_); +} + +TEST_F(MoqtBidiStreamTest, SendRequestOkFin) { + stream_->BindStream(&mock_stream_); + EXPECT_CALL(mock_stream_, CanWrite).WillRepeatedly(testing::Return(true)); + EXPECT_CALL( + mock_stream_, + Writev(ControlMessageOfType(MoqtMessageType::kRequestOk), testing::_)); + MessageParameters parameters; + QUICHE_EXPECT_OK(stream_->SendRequestOk(1, parameters, /*fin=*/true)); + EXPECT_TRUE(stream_->detached_); +} + +TEST_F(MoqtBidiStreamTest, SendRequestErrorOverload) { + stream_->BindStream(&mock_stream_); + EXPECT_CALL(mock_stream_, CanWrite).WillRepeatedly(testing::Return(true)); + EXPECT_CALL( + mock_stream_, + Writev(ControlMessageOfType(MoqtMessageType::kRequestError), testing::_)); + QUICHE_EXPECT_OK(stream_->SendRequestError(1, RequestErrorCode::kUnauthorized, + std::nullopt, "reason", + /*fin=*/true)); + EXPECT_TRUE(stream_->detached_); +} + +TEST_F(MoqtBidiStreamTest, SendRequestUpdateAndReceiveOk) { + stream_->BindStream(&mock_stream_); + EXPECT_CALL(mock_stream_, CanWrite).WillRepeatedly(testing::Return(true)); + EXPECT_CALL(mock_stream_, + Writev(ControlMessageOfType(MoqtMessageType::kRequestUpdate), + testing::_)); + MessageParameters parameters; + parameters.subscriber_priority = 20; + bool callback_called = false; + MoqtResponseCallback callback = + [&](std::variant<MessageParameters, MoqtRequestErrorInfo> res) { + callback_called = true; + ASSERT_TRUE(std::holds_alternative<MessageParameters>(res)); + EXPECT_EQ(std::get<MessageParameters>(res).subscriber_priority, 30); + }; + QUICHE_EXPECT_OK( + stream_->SendRequestUpdate(1, 0, parameters, std::move(callback))); + // Simulate receiving RequestOk + MoqtRequestOk request_ok; + request_ok.request_id = 1; + request_ok.parameters.subscriber_priority = 30; + QUICHE_EXPECT_OK(stream_->OnControlMessage(request_ok)); + EXPECT_TRUE(callback_called); + EXPECT_FALSE(stream_->detached_); +} + +TEST_F(MoqtBidiStreamTest, SendRequestUpdateAndReceiveError) { + stream_->BindStream(&mock_stream_); + EXPECT_CALL(mock_stream_, CanWrite).WillRepeatedly(testing::Return(true)); + EXPECT_CALL(mock_stream_, + Writev(ControlMessageOfType(MoqtMessageType::kRequestUpdate), + testing::_)); + MessageParameters parameters; + bool callback_called = false; + MoqtResponseCallback callback = + [&](std::variant<MessageParameters, MoqtRequestErrorInfo> res) { + callback_called = true; + ASSERT_TRUE(std::holds_alternative<MoqtRequestErrorInfo>(res)); + EXPECT_EQ(std::get<MoqtRequestErrorInfo>(res).error_code, + RequestErrorCode::kUnauthorized); + }; + QUICHE_EXPECT_OK( + stream_->SendRequestUpdate(1, 0, parameters, std::move(callback))); + // Simulate receiving RequestError + MoqtRequestError request_error; + request_error.request_id = 1; + request_error.error_code = RequestErrorCode::kUnauthorized; + request_error.reason_phrase = "unauthorized"; + ExpectFin(mock_stream_); + QUICHE_EXPECT_OK(stream_->OnControlMessage(request_error)); + EXPECT_TRUE(callback_called); + EXPECT_TRUE(stream_->detached_); +} + +TEST_F(MoqtBidiStreamTest, QueueIsFull) { + stream_->BindStream(&mock_stream_); + EXPECT_FALSE(stream_->QueueIsFull()); +} + } // namespace moqt::test
diff --git a/quiche/quic/moqt/moqt_framer.cc b/quiche/quic/moqt/moqt_framer.cc index 17fd0e0..d446967 100644 --- a/quiche/quic/moqt/moqt_framer.cc +++ b/quiche/quic/moqt/moqt_framer.cc
@@ -544,12 +544,6 @@ WireStringWithMoqVarIntLength(message.reason_phrase)); } -quiche::QuicheBuffer MoqtFramer::SerializeUnsubscribe( - const MoqtUnsubscribe& message) { - return SerializeControlMessage(MoqtMessageType::kUnsubscribe, - WireMoqVarInt(message.request_id)); -} - quiche::QuicheBuffer MoqtFramer::SerializePublishDone( const MoqtPublishDone& message) { return SerializeControlMessage( @@ -707,8 +701,8 @@ quiche::QuicheBuffer MoqtFramer::SerializeObjectAck( const MoqtObjectAck& message) { return SerializeControlMessage( - MoqtMessageType::kObjectAck, WireMoqVarInt(message.subscribe_id), - WireMoqVarInt(message.group_id), WireMoqVarInt(message.object_id), + MoqtMessageType::kObjectAck, WireMoqVarInt(message.group_id), + WireMoqVarInt(message.object_id), WireMoqVarInt(SignedVarintSerializedForm( message.delta_from_deadline.ToMicroseconds()))); }
diff --git a/quiche/quic/moqt/moqt_framer.h b/quiche/quic/moqt/moqt_framer.h index 96a7251..0d8df94 100644 --- a/quiche/quic/moqt/moqt_framer.h +++ b/quiche/quic/moqt/moqt_framer.h
@@ -54,7 +54,6 @@ quiche::QuicheBuffer SerializeSubscribeOk( const MoqtSubscribeOk& message, MoqtMessageType message_type = MoqtMessageType::kSubscribeOk); - quiche::QuicheBuffer SerializeUnsubscribe(const MoqtUnsubscribe& message); quiche::QuicheBuffer SerializePublishDone(const MoqtPublishDone& message); quiche::QuicheBuffer SerializeRequestUpdate(const MoqtRequestUpdate& message); quiche::QuicheBuffer SerializePublishNamespace(
diff --git a/quiche/quic/moqt/moqt_framer_test.cc b/quiche/quic/moqt/moqt_framer_test.cc index 482e5c4..2000e69 100644 --- a/quiche/quic/moqt/moqt_framer_test.cc +++ b/quiche/quic/moqt/moqt_framer_test.cc
@@ -46,7 +46,6 @@ MoqtMessageType::kRequestError, MoqtMessageType::kSubscribe, MoqtMessageType::kSubscribeOk, - MoqtMessageType::kUnsubscribe, MoqtMessageType::kPublishDone, MoqtMessageType::kPublishNamespace, MoqtMessageType::kPublishNamespaceDone, @@ -152,10 +151,6 @@ auto data = std::get<MoqtSubscribeOk>(structured_data); return framer_.SerializeSubscribeOk(data); } - case MoqtMessageType::kUnsubscribe: { - auto data = std::get<MoqtUnsubscribe>(structured_data); - return framer_.SerializeUnsubscribe(data); - } case MoqtMessageType::kPublishDone: { auto data = std::get<MoqtPublishDone>(structured_data); return framer_.SerializePublishDone(data);
diff --git a/quiche/quic/moqt/moqt_integration_test.cc b/quiche/quic/moqt/moqt_integration_test.cc index c00cd60..89a2d22 100644 --- a/quiche/quic/moqt/moqt_integration_test.cc +++ b/quiche/quic/moqt/moqt_integration_test.cc
@@ -849,6 +849,10 @@ test_harness_.RunUntilWithDefaultTimeout([&]() { return stream_reset; }); EXPECT_TRUE(success); EXPECT_EQ(bytes_received, 2000); + // On teardown, streams are destroyed in arbitrary order. If the uni stream + // is destroyed before the bidi stream, there will be a second stream reset + // notification to the visitor. If not, there won't be. + EXPECT_CALL(subscribe_visitor_, OnStreamReset).Times(testing::AnyNumber()); } TEST_F(MoqtIntegrationTest, BandwidthProbe) {
diff --git a/quiche/quic/moqt/moqt_messages.h b/quiche/quic/moqt/moqt_messages.h index edfe38b..b05a392 100644 --- a/quiche/quic/moqt/moqt_messages.h +++ b/quiche/quic/moqt/moqt_messages.h
@@ -386,10 +386,6 @@ TrackExtensions extensions; }; -struct QUICHE_EXPORT MoqtUnsubscribe { - uint64_t request_id; -}; - struct QUICHE_EXPORT MoqtPublishDone { uint64_t request_id; PublishDoneCode status_code; @@ -526,11 +522,10 @@ TrackExtensions extensions; }; -// All of the four values in this message are encoded as varints. +// All of the three values in this message are encoded as varints. // `delta_from_deadline` is encoded as an absolute value, with the lowest bit // indicating the sign (0 if positive). struct QUICHE_EXPORT MoqtObjectAck { - uint64_t subscribe_id; uint64_t group_id; uint64_t object_id; // Positive if the object has been received before the deadline.
diff --git a/quiche/quic/moqt/moqt_namespace_stream.h b/quiche/quic/moqt/moqt_namespace_stream.h index a80c4a1..34e64d1 100644 --- a/quiche/quic/moqt/moqt_namespace_stream.h +++ b/quiche/quic/moqt/moqt_namespace_stream.h
@@ -50,7 +50,7 @@ request_id_(request_id), remove_callback_(std::move(remove_callback)), response_callback_(std::move(response_callback)) {} - ~MoqtNamespaceSubscriberStream() override; + ~MoqtNamespaceSubscriberStream(); // MoqtBidiStreamBase overrides. void OnStreamBound() override; @@ -154,7 +154,7 @@ AddPrefixCallback add_callback, RemovePrefixCallback remove_callback, SessionErrorCallback session_error_callback, MoqtIncomingSubscribeNamespaceCallback& application); - ~MoqtNamespacePublisherStream() override { Detach(); } + ~MoqtNamespacePublisherStream() { Detach(); } void OnStreamBound() override { // TODO(martinduke): Set the priority for this stream.
diff --git a/quiche/quic/moqt/moqt_parser.cc b/quiche/quic/moqt/moqt_parser.cc index 13bfe26..5a297a6 100644 --- a/quiche/quic/moqt/moqt_parser.cc +++ b/quiche/quic/moqt/moqt_parser.cc
@@ -685,17 +685,6 @@ return request_error; } -absl::StatusOr<MoqtUnsubscribe> MoqtControlMessageParser::ProcessUnsubscribe( - absl::string_view data) const { - quic::QuicDataReader reader(data); - MoqtUnsubscribe unsubscribe; - if (!reader.ReadMoqVarInt(&unsubscribe.request_id)) { - return absl::InvalidArgumentError("Message missing fields"); - } - QUICHE_RETURN_IF_ERROR(CheckForTrailingData(reader)); - return unsubscribe; -} - absl::StatusOr<MoqtPublishDone> MoqtControlMessageParser::ProcessPublishDone( absl::string_view data) const { quic::QuicDataReader reader(data); @@ -1002,8 +991,7 @@ quic::QuicDataReader reader(data); MoqtObjectAck object_ack; uint64_t raw_delta; - if (!reader.ReadMoqVarInt(&object_ack.subscribe_id) || - !reader.ReadMoqVarInt(&object_ack.group_id) || + if (!reader.ReadMoqVarInt(&object_ack.group_id) || !reader.ReadMoqVarInt(&object_ack.object_id) || !reader.ReadMoqVarInt(&raw_delta)) { return absl::InvalidArgumentError("Message missing fields");
diff --git a/quiche/quic/moqt/moqt_parser.h b/quiche/quic/moqt/moqt_parser.h index f36415c..eca79f4 100644 --- a/quiche/quic/moqt/moqt_parser.h +++ b/quiche/quic/moqt/moqt_parser.h
@@ -127,8 +127,6 @@ absl::StatusOr<MoqtSubscribe> ProcessSubscribe(absl::string_view data) const; absl::StatusOr<MoqtSubscribeOk> ProcessSubscribeOk( absl::string_view data) const; - absl::StatusOr<MoqtUnsubscribe> ProcessUnsubscribe( - absl::string_view data) const; absl::StatusOr<MoqtPublishDone> ProcessPublishDone( absl::string_view data) const; absl::StatusOr<MoqtRequestUpdate> ProcessRequestUpdate( @@ -184,8 +182,6 @@ return parse(&MoqtControlMessageParser::ProcessSubscribe); case MoqtMessageType::kSubscribeOk: return parse(&MoqtControlMessageParser::ProcessSubscribeOk); - case MoqtMessageType::kUnsubscribe: - return parse(&MoqtControlMessageParser::ProcessUnsubscribe); case MoqtMessageType::kPublishDone: return parse(&MoqtControlMessageParser::ProcessPublishDone); case MoqtMessageType::kRequestUpdate:
diff --git a/quiche/quic/moqt/moqt_parser_test.cc b/quiche/quic/moqt/moqt_parser_test.cc index 7b9d69c..74fcd62 100644 --- a/quiche/quic/moqt/moqt_parser_test.cc +++ b/quiche/quic/moqt/moqt_parser_test.cc
@@ -52,7 +52,6 @@ MoqtMessageType::kSubscribe, MoqtMessageType::kSubscribeOk, MoqtMessageType::kRequestUpdate, - MoqtMessageType::kUnsubscribe, MoqtMessageType::kPublishDone, MoqtMessageType::kTrackStatus, MoqtMessageType::kPublishNamespace, @@ -1312,8 +1311,8 @@ TEST_F(MoqtMessageSpecificTest, ObjectAckNegativeDelta) { char object_ack[] = { - 0xb1, 0x84, 0x00, 0x05, // type - 0x01, 0x10, 0x20, // subscribe ID, group, object + 0xb1, 0x84, 0x00, 0x04, // type + 0x10, 0x20, // group, object 0x80, 0x81, // -0x40 time delta }; absl::StatusOr<std::vector<AnyMoqtControlMessage>> parsed = @@ -1322,7 +1321,6 @@ ASSERT_TRUE(parsed.ok()); ASSERT_EQ(parsed->size(), 1); MoqtObjectAck message = std::get<MoqtObjectAck>((*parsed)[0]); - EXPECT_EQ(message.subscribe_id, 0x01); EXPECT_EQ(message.group_id, 0x10); EXPECT_EQ(message.object_id, 0x20); EXPECT_EQ(message.delta_from_deadline,
diff --git a/quiche/quic/moqt/moqt_publish_stream.cc b/quiche/quic/moqt/moqt_publish_stream.cc index dd7c36e..1ffe56c 100644 --- a/quiche/quic/moqt/moqt_publish_stream.cc +++ b/quiche/quic/moqt/moqt_publish_stream.cc
@@ -38,7 +38,12 @@ response_callback_(std::move(response_callback)), stream_deleted_callback_(std::move(stream_deleted_callback)) {} -MoqtPublishPublisherStream::~MoqtPublishPublisherStream() { Detach(); } +MoqtPublishPublisherStream::~MoqtPublishPublisherStream() { + if (publisher_ != nullptr) { + publisher_->IgnoreResetAllStreams(); + } + Detach(); +} void MoqtPublishPublisherStream::OnStreamBound() { stream_parser()->set_allow_fin(true); @@ -165,6 +170,11 @@ }}, response); })); + } else { + // Since the application already called SUBSCRIBE, there will be no + // invocation of the request callback. Send REQUEST_OK immediately. + CheckStatus( + SendRequestOk(message.request_id, subscriber_->const_parameters())); } incoming_publish_callback_ = nullptr; if (subscriber_->visitor() == nullptr) { @@ -183,6 +193,10 @@ absl::Status MoqtPublishSubscriberStream::OnControlMessage( const MoqtRequestUpdate& message) { + if (subscriber_ == nullptr) { + // Stream is already closing. + return absl::OkStatus(); + } subscriber_->Update(message.parameters); CheckStatus(SendRequestOk(message.request_id, MessageParameters())); return absl::OkStatus();
diff --git a/quiche/quic/moqt/moqt_publish_stream.h b/quiche/quic/moqt/moqt_publish_stream.h index a3cef79..43e6dbe 100644 --- a/quiche/quic/moqt/moqt_publish_stream.h +++ b/quiche/quic/moqt/moqt_publish_stream.h
@@ -48,6 +48,10 @@ absl::Status OnControlMessage(const MoqtRequestOk& message); absl::Status OnControlMessage(const MoqtRequestError& message); absl::Status OnControlMessage(const MoqtRequestUpdate& message); + absl::Status OnControlMessage(const MoqtObjectAck& message) { + publisher_->ProcessObjectAck(message); + return absl::OkStatus(); + } void SetPublisher(std::unique_ptr<SubscriptionPublisher> publisher) { publisher_ = std::move(publisher); @@ -61,6 +65,8 @@ std::move(stream_deleted_callback_); stream_deleted_callback_ = nullptr; std::move(callback)(publisher_.get()); + publisher_->ResetAllStreams(); + publisher_ = nullptr; } private: @@ -96,6 +102,8 @@ absl::Status OnControlMessage(const MoqtRequestError& message); absl::Status OnControlMessage(const MoqtPublishDone& message); + SubscribeRemoteTrack* track() { return subscriber_.get(); } + void Detach() override { if (remove_callback_ != nullptr) { SubscribeRemoteTrack::RemoveCallback callback = @@ -103,6 +111,7 @@ remove_callback_ = nullptr; std::move(callback)(subscriber_.get()); } + subscriber_ = nullptr; } private:
diff --git a/quiche/quic/moqt/moqt_publish_stream_test.cc b/quiche/quic/moqt/moqt_publish_stream_test.cc index e59cb95..9de992f 100644 --- a/quiche/quic/moqt/moqt_publish_stream_test.cc +++ b/quiche/quic/moqt/moqt_publish_stream_test.cc
@@ -14,7 +14,6 @@ #include "absl/status/status.h" #include "absl/strings/string_view.h" #include "absl/types/span.h" -#include "quiche/quic/core/quic_alarm_factory.h" #include "quiche/quic/core/quic_time.h" #include "quiche/quic/core/quic_types.h" #include "quiche/quic/moqt/moqt_error.h" @@ -31,6 +30,7 @@ #include "quiche/quic/moqt/moqt_trace_recorder.h" #include "quiche/quic/moqt/moqt_track.h" #include "quiche/quic/moqt/moqt_types.h" +#include "quiche/quic/moqt/test_tools/mock_moqt_session.h" #include "quiche/quic/moqt/test_tools/moqt_framer_utils.h" #include "quiche/quic/moqt/test_tools/moqt_mock_visitor.h" #include "quiche/quic/test_tools/mock_clock.h" @@ -66,18 +66,6 @@ using ::testing::Return; using ::testing::StrictMock; -class MockSessionToPublisherInterface : public SessionToPublisherInterface { - public: - ~MockSessionToPublisherInterface() override = default; - MOCK_METHOD(bool, alternate_delivery_timeout, (), (const, override)); - MOCK_METHOD(void, UpdateTrackPriority, - (uint64_t, std::optional<MoqtTrackPriority>, MoqtTrackPriority), - (override)); - MOCK_METHOD(quic::QuicAlarmFactory*, alarm_factory, (), (override)); - MOCK_METHOD(void, PublishIsDone, (uint64_t), (override)); - MOCK_METHOD(webtransport::Session*, session, (), (override)); -}; - constexpr uint64_t kRequestId = 1; constexpr uint64_t kTrackAlias = 10; const FullTrackName kTrackName("foo", "bar"); @@ -105,8 +93,7 @@ EXPECT_CALL(visitor_, session).WillRepeatedly(Return(&webtrans_)); auto publisher = std::make_unique<SubscriptionPublisher>( framer_, track_publisher_, stream_.get(), kRequestId, kTrackAlias, - parameters_, &visitor_, /*monitoring_interface=*/nullptr, &mock_clock_, - trace_recorder_, /*is_publish=*/true); + parameters_, visitor_.weak_ptr_factory_.Create(), /*is_publish=*/true); publisher_ = publisher.get(); // Keep raw pointer for testing stream_->SetPublisher(std::move(publisher)); @@ -229,6 +216,30 @@ EXPECT_EQ(pub_params.subscription_filter->start(), Location(1, 3)); } +TEST_F(MoqtPublishPublisherStreamTest, ReceiveObjectAck) { + EXPECT_CALL(mock_stream_, + Writev(ControlMessageOfType(MoqtMessageType::kPublish), _)) + .WillOnce(Return(absl::OkStatus())); + stream_->BindStream(&mock_stream_); + + MoqtObjectAck ack; + ack.group_id = 1; + ack.object_id = 2; + ack.delta_from_deadline = quic::QuicTimeDelta::FromMilliseconds(100); + + EXPECT_CALL(visitor_, trace_recorder()) + .WillOnce(testing::ReturnRef(trace_recorder_)); + QUICHE_EXPECT_OK(stream_->OnControlMessage(ack)); +} + +TEST_F(MoqtPublishPublisherStreamTest, Detach) { + EXPECT_CALL(deleted_callback_, Call(publisher_)); + stream_->Detach(); + // Verifying second detach is a no-op + EXPECT_CALL(deleted_callback_, Call).Times(0); + stream_->Detach(); +} + class MoqtPublishSubscriberStreamTest : public quiche::test::QuicheTest { public: MoqtPublishSubscriberStreamTest() @@ -561,5 +572,34 @@ EXPECT_CALL(mock_subscribe_visitor_, OnPublishDone(kTrackName)); } +TEST_F(MoqtPublishSubscriberStreamTest, ReceiveRequestOkAndErrorTodo) { + MoqtRequestOk request_ok; + QUICHE_EXPECT_OK(stream_->OnControlMessage(request_ok)); + + MoqtRequestError request_error; + QUICHE_EXPECT_OK(stream_->OnControlMessage(request_error)); +} + +TEST_F(MoqtPublishSubscriberStreamTest, TrackAndDetach) { + MoqtPublish publish = DefaultPublish(); + EXPECT_CALL(incoming_publish_callback_mock_, Call(kTrackName, _, _, _)) + .WillOnce(Return(&mock_subscribe_visitor_)); + EXPECT_CALL(mock_add_callback_, Call(NotNull())).WillOnce(Return(true)); + EXPECT_CALL(mock_subscribe_visitor_, OnReply(kTrackName, _)) + .WillOnce( + [](const FullTrackName&, + const std::variant<SubscribeOkData, MoqtRequestErrorInfo>& reply) { + EXPECT_TRUE(std::holds_alternative<SubscribeOkData>(reply)); + }); + QUICHE_EXPECT_OK(stream_->OnControlMessage(publish)); + EXPECT_NE(stream_->track(), nullptr); + EXPECT_EQ(stream_->track()->track_alias(), kTrackAlias); + + EXPECT_CALL(mock_remove_callback_, Call(stream_->track())); + EXPECT_CALL(mock_subscribe_visitor_, OnPublishDone(kTrackName)); + stream_->Detach(); + EXPECT_EQ(stream_->track(), nullptr); +} + } // namespace } // namespace moqt::test
diff --git a/quiche/quic/moqt/moqt_session.cc b/quiche/quic/moqt/moqt_session.cc index f7f8333..9dd7d62 100644 --- a/quiche/quic/moqt/moqt_session.cc +++ b/quiche/quic/moqt/moqt_session.cc
@@ -4,13 +4,13 @@ #include "quiche/quic/moqt/moqt_session.h" +#include <array> #include <cstdint> #include <memory> #include <optional> #include <string> #include <utility> #include <variant> -#include <vector> #include "absl/base/casts.h" #include "absl/base/nullability.h" @@ -24,6 +24,7 @@ #include "absl/status/status.h" #include "absl/status/statusor.h" #include "absl/strings/string_view.h" +#include "absl/types/span.h" #include "quiche/quic/core/quic_alarm_factory.h" #include "quiche/quic/core/quic_time.h" #include "quiche/quic/core/quic_types.h" @@ -42,6 +43,7 @@ #include "quiche/quic/moqt/moqt_publisher.h" #include "quiche/quic/moqt/moqt_session_callbacks.h" #include "quiche/quic/moqt/moqt_session_interface.h" +#include "quiche/quic/moqt/moqt_subscribe_stream.h" #include "quiche/quic/moqt/moqt_subscription.h" #include "quiche/quic/moqt/moqt_track.h" #include "quiche/quic/moqt/moqt_types.h" @@ -50,7 +52,7 @@ #include "quiche/common/platform/api/quiche_bug_tracker.h" #include "quiche/common/platform/api/quiche_logging.h" #include "quiche/common/quiche_buffer_allocator.h" -#include "quiche/common/quiche_status_utils.h" +#include "quiche/common/quiche_mem_slice.h" #include "quiche/common/quiche_weak_ptr.h" #include "quiche/web_transport/web_transport.h" @@ -91,6 +93,7 @@ local_max_request_id_(parameters.max_request_id), alarm_factory_(std::move(alarm_factory)), weak_ptr_factory_(this), + weak_ptr_factory_for_publishers_(this), liveness_token_(std::make_shared<Empty>()) { if (parameters_.using_webtrans) { session_->SetOnDraining([this]() { @@ -165,6 +168,23 @@ void MoqtSession::OnIncomingBidirectionalStreamAvailable() { while (webtransport::Stream* stream = session_->AcceptIncomingBidirectionalStream()) { + if (sent_goaway_) { + // Immediately reject new requests with REQUEST_ERROR. If the stream + // cannot be written, just reset it. + if (!stream->CanWrite()) { + stream->ResetWithUserCode(kResetCodeSessionClosed); + continue; + } + webtransport::StreamWriteOptions options; + options.set_send_fin(true); + std::array write_vector = { + quiche::QuicheMemSlice(framer_.SerializeRequestError(MoqtRequestError{ + 0, RequestErrorCode::kGoingAway, std::nullopt, ""}))}; + if (!stream->Writev(absl::MakeSpan(write_vector), options).ok()) { + stream->ResetWithUserCode(kResetCodeSessionClosed); + }; + continue; + } auto bidi_stream = std::make_unique<UnknownBidiStream>(this, stream); stream->SetVisitor(std::move(bidi_stream)); stream->visitor()->OnCanRead(); @@ -196,7 +216,7 @@ << message.object_id << " priority " << message.publisher_priority << " length " << payload->size(); - SubscribeRemoteTrack* track = RemoteTrackByAlias(message.track_alias); + SubscribeRemoteTrack* track = SubscribeByAlias(message.track_alias); if (track == nullptr) { return; } @@ -444,10 +464,9 @@ } bool MoqtSession::Subscribe(const FullTrackName& name, - SubscribeVisitor* visitor, + SubscribeVisitor* absl_nonnull visitor, const MessageParameters& parameters) { QUICHE_DCHECK(name.IsValid()); - if (next_request_id_ >= peer_max_request_id_) { if (!last_requests_blocked_sent_.has_value() || peer_max_request_id_ > *last_requests_blocked_sent_) { @@ -471,24 +490,19 @@ QUIC_DLOG(INFO) << ENDPOINT << "Tried to send SUBSCRIBE after GOAWAY"; return false; } - MoqtSubscribe message(next_request_id_, name, parameters); - next_request_id_ += 2; - if (SupportsObjectAck() && visitor != nullptr) { - // Since we do not expose subscribe IDs directly in the API, instead wrap - // the session and subscribe ID in a callback. - visitor->OnCanAckObjects(absl::bind_front(&MoqtSession::SendObjectAck, this, - message.request_id)); - } else { - QUICHE_DLOG_IF(WARNING, message.parameters.oack_window_size.has_value()) - << "Attempting to set object_ack_window on a connection that does not " - "support it."; - message.parameters.oack_window_size = std::nullopt; + if (!session_->CanOpenNextOutgoingBidirectionalStream()) { + return false; // Do not retry opening a SUBSCRIBE stream. } - SendControlMessage(framer_.SerializeSubscribe(message)); - QUIC_DLOG(INFO) << ENDPOINT << "Sent SUBSCRIBE message for " - << message.full_track_name; - auto track = std::make_unique<SubscribeRemoteTrack>( - message, visitor, + auto stream_visitor = std::make_unique<MoqtSubscribeRequestStream>( + &framer_, ControlMessageParser(), NextRequestId(), + [weak_session = GetWeakPtr()](MoqtError code, absl::string_view reason) { + MoqtSessionInterface* session = weak_session.GetIfAvailable(); + if (session == nullptr) { + return; + } + session->Error(code, reason); + }, + name, visitor, parameters, [weakptr = GetWeakPtr()](SubscribeRemoteTrack* track) { MoqtSession* session = MoqtSessionFromWeakPtr(weakptr); if (session == nullptr || !track->track_alias().has_value()) { @@ -507,10 +521,18 @@ if (track->track_alias().has_value()) { session->subscribe_by_alias_.erase(*track->track_alias()); } - session->upstream_by_id_.erase(track->request_id()); - }); - subscribe_by_name_.emplace(message.full_track_name, track.get()); - upstream_by_id_.emplace(message.request_id, std::move(track)); + }, + callbacks_.clock, alarm_factory_.get()); + webtransport::Stream* stream = session_->OpenOutgoingBidirectionalStream(); + QUICHE_CHECK(stream != nullptr); + MoqtSubscribeRequestStream* stream_visitor_ptr = stream_visitor.get(); + stream->SetVisitor(std::move(stream_visitor)); + stream_visitor_ptr->BindStream(stream); + subscribe_by_name_[name] = stream_visitor_ptr->track(); + if (SupportsObjectAck()) { + visitor->OnCanAckObjects( + absl::bind_front(&MoqtSession::SendObjectAck, this, name)); + } return true; } @@ -518,32 +540,22 @@ const MessageParameters& parameters, MoqtResponseCallback response_callback) { QUICHE_DCHECK(name.IsValid()); - if (next_request_id_ >= peer_max_request_id_) { - if (!last_requests_blocked_sent_.has_value() || - peer_max_request_id_ > *last_requests_blocked_sent_) { - MoqtRequestsBlocked requests_blocked; - requests_blocked.max_request_id = peer_max_request_id_; - SendControlMessage(framer_.SerializeRequestsBlocked(requests_blocked)); - last_requests_blocked_sent_ = peer_max_request_id_; - } - QUIC_DLOG(INFO) << ENDPOINT << "Tried to send SUBSCRIBE with ID " - << next_request_id_ - << " which is greater than the maximum ID " - << peer_max_request_id_; - return false; - } auto it = subscribe_by_name_.find(name); if (it == subscribe_by_name_.end()) { return false; } - // TODO(martinduke): Support Update on PUBLISH streams. - pending_subscribe_updates_[next_request_id_] = {name, parameters, - std::move(response_callback)}; - MoqtRequestUpdate update{next_request_id_, it->second->request_id(), - parameters}; - next_request_id_ += 2; - SendControlMessage(framer_.SerializeRequestUpdate(update)); - return true; + // sending zero because related request ID is ignored for SUBSCRIBE. + return it->second->request_stream() + ->SendRequestUpdate(NextRequestId(), 0, parameters, + std::move(response_callback)) + .ok(); +} + +bool MoqtSession::PublishUpdate(const FullTrackName& name, + const MessageParameters& parameters, + MoqtResponseCallback response_callback) { + // TODO(martinduke): Implement this. + return false; } void MoqtSession::Unsubscribe(const FullTrackName& name) { @@ -551,16 +563,11 @@ return; } QUICHE_DCHECK(name.IsValid()); - SubscribeRemoteTrack* track = RemoteTrackByName(name); - if (track == nullptr) { + auto it = subscribe_by_name_.find(name); + if (it == subscribe_by_name_.end()) { return; } - QUICHE_DCHECK(name.IsValid()); - QUIC_DLOG(INFO) << ENDPOINT << "Sent UNSUBSCRIBE message for " << name; - MoqtUnsubscribe message; - message.request_id = track->request_id(); - SendControlMessage(framer_.SerializeUnsubscribe(message)); - track->Destroy(); + it->second->request_stream()->Reset(kResetCodeCancelled); } bool MoqtSession::Publish( @@ -576,10 +583,16 @@ if (!session_->CanOpenNextOutgoingBidirectionalStream()) { return false; // Do not retry opening a PUBLISH stream. } - if (!subscribed_track_names_.insert(name).second) { - QUICHE_DLOG(INFO) << ENDPOINT << "Tried to send PUBLISH for track " << name - << " which is already published"; - return false; + auto it = subscribed_track_names_.find(name); + if (it != subscribed_track_names_.end()) { + if (it->second->established()) { + QUICHE_DLOG(INFO) << ENDPOINT << "Tried to send PUBLISH for track " + << name << " which is already published"; + return false; + } + it->second->OnSubscribeRejected( + MoqtRequestErrorInfo{RequestErrorCode::kDuplicateSubscription, + std::nullopt, "PUBLISH is coming"}); } auto stream_visitor = std::make_unique<MoqtPublishPublisherStream>( &framer_, ControlMessageParser(), @@ -602,8 +615,8 @@ std::move(response_callback)); auto publish_state = std::make_unique<SubscriptionPublisher>( framer_, publisher, stream_visitor.get(), next_request_id_, - next_local_track_alias_, parameters, this, nullptr, callbacks_.clock, - trace_recorder_, true); + next_local_track_alias_, parameters, + weak_ptr_factory_for_publishers_.Create(), true); SubscriptionPublisher* publisher_ptr = publish_state.get(); stream_visitor->SetPublisher(std::move(publish_state)); webtransport::Stream* stream = session_->OpenOutgoingBidirectionalStream(); @@ -644,10 +657,10 @@ QUIC_DLOG(INFO) << ENDPOINT << "Sent FETCH message for " << name; auto fetch = std::make_unique<UpstreamFetch>( message, std::get<StandaloneFetch>(message.fetch), std::move(callback), - [this, id = message.request_id]() { // Deletion callback - upstream_by_id_.erase(id); + [this, id = message.request_id]() { + fetch_by_id_.erase(id); // Deletion callback }); - upstream_by_id_.emplace(message.request_id, std::move(fetch)); + fetch_by_id_.emplace(message.request_id, std::move(fetch)); return true; } @@ -658,15 +671,13 @@ QUICHE_DCHECK(name.IsValid()); return RelativeJoiningFetch( name, visitor, - [this, id = next_request_id_](std::unique_ptr<MoqtFetchTask> fetch_task) { + [this, track_name = name](std::unique_ptr<MoqtFetchTask> fetch_task) { // Move the fetch_task to the subscribe to plumb into its visitor. - RemoteTrack* track = RemoteTrackById(id); - if (track == nullptr || track->is_fetch()) { + SubscribeRemoteTrack* subscribe = SubscribeByName(track_name); + if (subscribe == nullptr || subscribe->is_fetch()) { fetch_task.release(); return; } - auto* subscribe = absl::down_cast<SubscribeRemoteTrack*>(track); - RemoteTrackByName(track->full_track_name()); subscribe->OnJoiningFetchReady(std::move(fetch_task)); }, num_previous_groups, parameters); @@ -702,9 +713,9 @@ auto upstream_fetch = std::make_unique<UpstreamFetch>( fetch, name, std::move(callback), /*Deletion callback=*/[this, id = fetch.request_id]() { - upstream_by_id_.erase(id); + fetch_by_id_.erase(id); }); - upstream_by_id_.emplace(fetch.request_id, std::move(upstream_fetch)); + fetch_by_id_.emplace(fetch.request_id, std::move(upstream_fetch)); return true; } @@ -733,19 +744,6 @@ "Peer did not close session after GOAWAY"); } -void MoqtSession::PublishIsDone(uint64_t request_id) { - if (is_closing_) { - return; - } - auto it = published_subscriptions_.find(request_id); - if (it == published_subscriptions_.end()) { - // If a PUBLISH, we will end up here. - return; - } - subscribed_track_names_.erase(it->second->publisher().GetTrackName()); - published_subscriptions_.erase(it); -} - void MoqtSession::UpdateTrackPriority( uint64_t request_id, std::optional<MoqtTrackPriority> old_priority, MoqtTrackPriority new_priority) { @@ -762,6 +760,25 @@ subscriptions_with_queued_streams_.emplace(new_priority, request_id); } +std::shared_ptr<MoqtTrackPublisher> MoqtSession::GetTrackPublisher( + const FullTrackName& name) { + if (publisher_ == nullptr) { + return nullptr; + } + return publisher_->GetTrack(name); +} + +MoqtPublishingMonitorInterface* MoqtSession::ReleaseMonitoringInterface( + const FullTrackName& name) { + auto it = monitoring_interfaces_for_published_tracks_.find(name); + if (it == monitoring_interfaces_for_published_tracks_.end()) { + return nullptr; + } + MoqtPublishingMonitorInterface* interface = it->second; + monitoring_interfaces_for_published_tracks_.erase(it); + return interface; +} + bool MoqtSession::OpenDataStream(PublishedFetch* fetch, webtransport::SendOrder send_order) { webtransport::Stream* new_stream = @@ -793,7 +810,7 @@ return true; } -SubscribeRemoteTrack* MoqtSession::RemoteTrackByAlias(uint64_t track_alias) { +SubscribeRemoteTrack* MoqtSession::SubscribeByAlias(uint64_t track_alias) { auto it = subscribe_by_alias_.find(track_alias); if (it == subscribe_by_alias_.end()) { return nullptr; @@ -801,24 +818,23 @@ return it->second; } -RemoteTrack* MoqtSession::RemoteTrackById(uint64_t request_id) { - auto it = upstream_by_id_.find(request_id); - if (it == upstream_by_id_.end()) { - return nullptr; - } - return it->second.get(); -} - -SubscribeRemoteTrack* MoqtSession::RemoteTrackByName( - const FullTrackName& name) { - QUICHE_DCHECK(name.IsValid()); - auto it = subscribe_by_name_.find(name); +SubscribeRemoteTrack* MoqtSession::SubscribeByName( + const FullTrackName& track_name) { + auto it = subscribe_by_name_.find(track_name); if (it == subscribe_by_name_.end()) { return nullptr; } return it->second; } +UpstreamFetch* MoqtSession::FetchById(uint64_t request_id) { + auto it = fetch_by_id_.find(request_id); + if (it == fetch_by_id_.end()) { + return nullptr; + } + return it->second.get(); +} + void MoqtSession::OnCanCreateNewOutgoingUnidirectionalStream() { while (!subscriptions_with_queued_streams_.empty() && session_->CanOpenNextOutgoingUnidirectionalStream()) { @@ -863,8 +879,9 @@ Error(MoqtError::kInvalidRequestId, "Request ID evenness incorrect"); return false; } - if (published_subscriptions_.contains(request_id) || - incoming_fetches_.contains(request_id) || + // TODO(martinduke): Write new checks for duplicate request IDs. It's + // probably best to track the largest observed plus a set of holes. + if (incoming_fetches_.contains(request_id) || incoming_track_status_.contains(request_id) || incoming_publish_namespaces_by_id_.contains(request_id)) { QUICHE_DLOG(INFO) << ENDPOINT << "Duplicate request ID"; @@ -978,7 +995,9 @@ auto it = session->subscribe_by_name_.find(track->full_track_name()); if (it != session->subscribe_by_name_.end()) { - // It's a pending SUBSCRIBE; kill it. + // It's a pending SUBSCRIBE; kill it, but use the parameters and + // visitor from the SUBSCRIBE. + track->Update(it->second->const_parameters()); track->set_visitor(it->second->ReleaseVisitor()); session->Unsubscribe(it->second->full_track_name()); } @@ -1005,6 +1024,51 @@ temp_stream->OnCanRead(); break; } + case MoqtMessageType::kSubscribe: { + auto subscribe_stream = std::make_unique<MoqtSubscribeResponseStream>( + &session_->framer_, session_->ControlMessageParser(), + session_->next_local_track_alias_++, + [weakptr = + session_->GetWeakPtr()](SubscriptionPublisher* subscription) { + MoqtSession* session = MoqtSessionFromWeakPtr(weakptr); + if (session == nullptr) { + return true; + } + auto [it, success] = session->published_subscriptions_.try_emplace( + subscription->request_id(), subscription); + if (!success) { + return false; + } + auto [it2, success2] = session->subscribed_track_names_.try_emplace( + subscription->publisher().GetTrackName(), subscription); + return success2; + }, + [weakptr = + session_->GetWeakPtr()](SubscriptionPublisher* subscription) { + MoqtSession* session = MoqtSessionFromWeakPtr(weakptr); + if (session == nullptr) { + return; + } + session->published_subscriptions_.erase(subscription->request_id()); + session->subscribed_track_names_.erase( + subscription->publisher().GetTrackName()); + }, + [weakptr = session_->GetWeakPtr()](MoqtError code, + absl::string_view reason) { + MoqtSession* session = MoqtSessionFromWeakPtr(weakptr); + if (session != nullptr) { + session->Error(code, reason); + } + }, + session_->weak_ptr_factory_for_publishers_.Create()); + subscribe_stream->BindStream(std::move(parser_)); + MoqtSubscribeResponseStream* temp_stream = subscribe_stream.get(); + stream_->SetVisitor(std::move(subscribe_stream)); + // The UnknownBidiStream object is deleted; no class access after this + // point. + temp_stream->OnCanRead(); + break; + } default: session_->Error(MoqtError::kProtocolViolation, "Unexpected message type received to start bidi stream"); @@ -1050,113 +1114,9 @@ } } -absl::Status MoqtSession::OnControlMessage(const MoqtSubscribe& message) { - if (!ValidateRequestId(message.request_id)) { - return absl::OkStatus(); - } - QUIC_DLOG(INFO) << ENDPOINT << "Received a SUBSCRIBE for " - << message.full_track_name; - if (sent_goaway_) { - QUIC_DLOG(INFO) << ENDPOINT << "Received a SUBSCRIBE after GOAWAY"; - SendRequestErrorOnControlStream(message.request_id, - RequestErrorCode::kUnauthorized, - std::nullopt, "SUBSCRIBE after GOAWAY"); - return absl::OkStatus(); - } - if (subscribed_track_names_.contains(message.full_track_name)) { - SendRequestErrorOnControlStream(message.request_id, - RequestErrorCode::kDuplicateSubscription, - std::nullopt, ""); - return absl::OkStatus(); - } - const FullTrackName& track_name = message.full_track_name; - std::shared_ptr<MoqtTrackPublisher> track_publisher = - publisher_->GetTrack(track_name); - if (track_publisher == nullptr) { - QUIC_DLOG(INFO) << ENDPOINT << "SUBSCRIBE for " << track_name - << " rejected by the application: does not exist"; - SendRequestErrorOnControlStream(message.request_id, - RequestErrorCode::kDoesNotExist, - std::nullopt, "not found"); - return absl::OkStatus(); - } - - MoqtPublishingMonitorInterface* monitoring = nullptr; - auto monitoring_it = - monitoring_interfaces_for_published_tracks_.find(track_name); - if (monitoring_it != monitoring_interfaces_for_published_tracks_.end()) { - monitoring = monitoring_it->second; - monitoring_interfaces_for_published_tracks_.erase(monitoring_it); - } - - MoqtTrackPublisher* track_publisher_ptr = track_publisher.get(); - auto subscription = std::make_unique<SubscriptionPublisher>( - framer_, track_publisher, GetControlStream(), message.request_id, - next_local_track_alias_++, message.parameters, this, monitoring, - callbacks_.clock, trace_recorder_, false); - SubscriptionPublisher* subscription_ptr = subscription.get(); - auto [it, success] = published_subscriptions_.emplace( - message.request_id, std::move(subscription)); - if (!success) { - QUICHE_NOTREACHED(); // ValidateRequestId() should have caught this. - } - subscribed_track_names_.insert(message.full_track_name); - track_publisher_ptr->AddObjectListener(subscription_ptr); - return absl::OkStatus(); -} - -absl::Status MoqtSession::OnControlMessage(const MoqtSubscribeOk& message) { - RemoteTrack* track = RemoteTrackById(message.request_id); - if (track == nullptr) { - QUIC_DLOG(INFO) << ENDPOINT << "Received the SUBSCRIBE_OK for " - << "request_id = " << message.request_id - << " but no track exists"; - // Subscription state might have been destroyed for internal reasons. - return absl::OkStatus(); - } - if (track->is_fetch()) { - return absl::InvalidArgumentError("Received SUBSCRIBE_OK for a FETCH"); - } - if (message.parameters.largest_object.has_value()) { - QUIC_DLOG(INFO) << ENDPOINT << "Received the SUBSCRIBE_OK for " - << "request_id = " << message.request_id << " " - << track->full_track_name() - << " largest_id = " << *message.parameters.largest_object; - } else { - QUIC_DLOG(INFO) << ENDPOINT << "Received the SUBSCRIBE_OK for " - << "request_id = " << message.request_id << " " - << track->full_track_name(); - } - SubscribeRemoteTrack* subscribe = - absl::down_cast<SubscribeRemoteTrack*>(track); - if (!subscribe->set_track_alias(message.track_alias)) { - return absl::AlreadyExistsError("Duplicate track alias"); - } - subscribe->OnObjectOrOk( - SubscribeOkData(message.parameters, message.extensions)); - return absl::OkStatus(); -} - absl::Status MoqtSession::OnControlMessage(const MoqtRequestOk& message) { - if (upstream_by_id_.contains(message.request_id)) { - return absl::InvalidArgumentError( - "Received REQUEST_OK for SUBSCRIBE, FETCH, or PUBLISH"); - } - // Response to REQUEST_UPDATE for a subscribe. - auto ru_it = pending_subscribe_updates_.find(message.request_id); - if (ru_it != pending_subscribe_updates_.end()) { - auto sub_it = subscribe_by_name_.find(ru_it->second.name); - if (sub_it == subscribe_by_name_.end()) { - std::move(ru_it->second.response_callback)( - MoqtRequestErrorInfo{RequestErrorCode::kDoesNotExist, std::nullopt, - "subscription does not exist anymore"}); - pending_subscribe_updates_.erase(ru_it); - return absl::OkStatus(); - } - sub_it->second->Update(ru_it->second.parameters); - std::move(ru_it->second.response_callback)(MessageParameters()); - pending_subscribe_updates_.erase(ru_it); - return absl::OkStatus(); + if (fetch_by_id_.contains(message.request_id)) { + return absl::InvalidArgumentError("Received REQUEST_OK for FETCH"); } // Response to PUBLISH_NAMESPACE. auto pn_it = publish_namespace_by_id_.find(message.request_id); @@ -1181,43 +1141,27 @@ MoqtRequestErrorInfo error_info{message.error_code, message.retry_interval, message.reason_phrase}; // TODO(martinduke): Do something with retry_interval. - RemoteTrack* track = RemoteTrackById(message.request_id); - if (track != nullptr) { - // It's in response to SUBSCRIBE or FETCH. - if (!track->ErrorIsAllowed()) { + UpstreamFetch* fetch = FetchById(message.request_id); + if (fetch != nullptr) { + // It's in response to FETCH. + if (!fetch->ErrorIsAllowed()) { return absl::InvalidArgumentError( "Received REQUEST_ERROR after REQUEST_OK or objects"); } QUIC_DLOG(INFO) << ENDPOINT << "Received the REQUEST_ERROR for " << "request_id = " << message.request_id << " (" - << track->full_track_name() << ")" + << fetch->full_track_name() << ")" << ", error = " << static_cast<uint64_t>(message.error_code) << " (" << message.reason_phrase << ")"; - if (track->is_fetch()) { - UpstreamFetch* fetch = absl::down_cast<UpstreamFetch*>(track); - absl::Status status = - RequestErrorCodeToStatus(message.error_code, message.reason_phrase); - fetch->OnFetchResult(Location(0, 0), status, nullptr); - } else { - SubscribeRemoteTrack* subscribe = - absl::down_cast<SubscribeRemoteTrack*>(track); - if (subscribe->visitor() != nullptr) { - subscribe->visitor()->OnReply(subscribe->full_track_name(), error_info); - } - } + absl::Status status = + RequestErrorCodeToStatus(message.error_code, message.reason_phrase); + fetch->OnFetchResult(Location(0, 0), status, nullptr); if (!is_closing_) { // The visitor might have closed the session. - track->Destroy(); + fetch->Destroy(); } return absl::OkStatus(); } - // Response to REQUEST_UPDATE for a subscribe. - auto ru_it = pending_subscribe_updates_.find(message.request_id); - if (ru_it != pending_subscribe_updates_.end()) { - std::move(ru_it->second.response_callback)(error_info); - pending_subscribe_updates_.erase(ru_it); - return absl::OkStatus(); - } // Response to PUBLISH_NAMESPACE. auto pn_it = publish_namespace_by_id_.find(message.request_id); if (pn_it != publish_namespace_by_id_.end()) { @@ -1238,41 +1182,7 @@ return absl::OkStatus(); } -absl::Status MoqtSession::OnControlMessage(const MoqtUnsubscribe& message) { - auto it = published_subscriptions_.find(message.request_id); - if (it == published_subscriptions_.end()) { - return absl::OkStatus(); - } - QUIC_DLOG(INFO) << ENDPOINT << "Received an UNSUBSCRIBE for " - << it->second->publisher().GetTrackName(); - PublishIsDone(message.request_id); - return absl::OkStatus(); -} - -absl::Status MoqtSession::OnControlMessage(const MoqtPublishDone& message) { - auto it = upstream_by_id_.find(message.request_id); - if (it == upstream_by_id_.end()) { - return absl::OkStatus(); - } - auto* subscribe = absl::down_cast<SubscribeRemoteTrack*>(it->second.get()); - QUIC_DLOG(INFO) << ENDPOINT << "Received a PUBLISH_DONE for " - << it->second->full_track_name(); - subscribe->OnPublishDone(message.stream_count, callbacks_.clock, - alarm_factory_.get()); - return absl::OkStatus(); -} - absl::Status MoqtSession::OnControlMessage(const MoqtRequestUpdate& message) { - auto it = published_subscriptions_.find(message.existing_request_id); - if (it != published_subscriptions_.end()) { - // It's updating SUBSCRIBE. - it->second->Update(message.parameters); - // TODO(martinduke): There should be an MoqtResponseCallback sent to the - // application, rather than automatic OK. - SendControlMessage(framer_.SerializeRequestOk( - MoqtRequestOk{.request_id = message.request_id})); - return absl::OkStatus(); - } auto pn_it = publish_namespace_by_id_.find(message.existing_request_id); if (pn_it != publish_namespace_by_id_.end()) { // It's updating PUBLISH_NAMESPACE. @@ -1622,7 +1532,7 @@ } absl::Status MoqtSession::OnControlMessage(const MoqtFetchOk& message) { - RemoteTrack* track = RemoteTrackById(message.request_id); + UpstreamFetch* track = FetchById(message.request_id); if (track == nullptr) { QUIC_DLOG(INFO) << ENDPOINT << "Received the FETCH_OK for " << "request_id = " << message.request_id @@ -1630,9 +1540,6 @@ // Subscription state might have been destroyed for internal reasons. return absl::OkStatus(); } - if (!track->is_fetch()) { - return absl::InvalidArgumentError("Received FETCH_OK for a SUBSCRIBE"); - } QUIC_DLOG(INFO) << ENDPOINT << "Received the FETCH_OK for request_id = " << message.request_id << " " << track->full_track_name(); UpstreamFetch* fetch = absl::down_cast<UpstreamFetch*>(track); @@ -1646,24 +1553,12 @@ return absl::OkStatus(); } -absl::Status MoqtSession::OnControlMessage(const MoqtPublish& message) { - if (!ValidateRequestId(message.request_id)) { - return absl::OkStatus(); - } - RequestErrorCode error_code = sent_goaway_ ? RequestErrorCode::kUnauthorized - : RequestErrorCode::kNotSupported; - absl::string_view error_reason = sent_goaway_ - ? "Received a PUBLISH after GOAWAY" - : "PUBLISH is not supported"; - SendRequestErrorOnControlStream(message.request_id, error_code, std::nullopt, - error_reason); - return absl::OkStatus(); -} - void MoqtSession::OnMalformedTrack(RemoteTrack* track) { if (!track->is_fetch()) { - absl::down_cast<SubscribeRemoteTrack*>(track)->visitor()->OnMalformedTrack( - track->full_track_name()); + auto* subscribe = absl::down_cast<SubscribeRemoteTrack*>(track); + if (subscribe->visitor() != nullptr) { + subscribe->visitor()->OnMalformedTrack(track->full_track_name()); + } Unsubscribe(track->full_track_name()); return; } @@ -1687,7 +1582,6 @@ // Incoming SUBSCRIBE_NAMESPACE is automatically cleaned up; the destroyed // session owns the webtransport stream, which owns the StreamVisitor, which // owns the task. Destroying the task notifies the application. - published_subscriptions_.clear(); for (auto& it : incoming_publish_namespaces_by_namespace_) { callbacks_.incoming_publish_namespace_callback(it.first, std::nullopt, nullptr); @@ -1696,8 +1590,17 @@ std::move(it.second.cancel_callback)(MoqtRequestErrorInfo{ RequestErrorCode::kUninterested, std::nullopt, "Session closed"}); } - while (!upstream_by_id_.empty()) { - upstream_by_id_.begin()->second->Destroy(); + while (!fetch_by_id_.empty()) { + fetch_by_id_.begin()->second->Destroy(); + } + for (auto& [track_name, subscriber] : subscribe_by_name_) { + // It's possible the application is going to destroy its visitor as early + // as session_deleted_callback is called. So call OnPublishDone() now and + // clear the visitor. + if (subscriber->visitor() != nullptr) { + subscriber->visitor()->OnPublishDone(subscriber->full_track_name()); + subscriber->ReleaseVisitor(); + } } } @@ -1705,8 +1608,8 @@ if (is_closing_) { return; } - auto it = upstream_by_id_.find(request_id); - if (it == upstream_by_id_.end()) { + auto it = fetch_by_id_.find(request_id); + if (it == fetch_by_id_.end()) { return; } it->second->Destroy();
diff --git a/quiche/quic/moqt/moqt_session.h b/quiche/quic/moqt/moqt_session.h index d0bcb21..fd45f9d 100644 --- a/quiche/quic/moqt/moqt_session.h +++ b/quiche/quic/moqt/moqt_session.h
@@ -87,11 +87,15 @@ MoqtSessionCallbacks& callbacks() override { return callbacks_; } void Error(MoqtError code, absl::string_view error) override; // Returns false if the SUBSCRIBE isn't sent. - bool Subscribe(const FullTrackName& name, SubscribeVisitor* visitor, + bool Subscribe(const FullTrackName& name, + SubscribeVisitor* absl_nonnull visitor, const MessageParameters& parameters) override; bool SubscribeUpdate(const FullTrackName& name, const MessageParameters& parameters, MoqtResponseCallback response_callback) override; + bool PublishUpdate(const FullTrackName& name, + const MessageParameters& parameters, + MoqtResponseCallback response_callback) override; void Unsubscribe(const FullTrackName& name) override; bool Publish(std::shared_ptr<MoqtTrackPublisher> absl_nonnull publisher, const MessageParameters& parameters, @@ -148,7 +152,12 @@ quic::QuicAlarmFactory* alarm_factory() override { return alarm_factory_.get(); } - void PublishIsDone(uint64_t request_id) override; + std::shared_ptr<MoqtTrackPublisher> GetTrackPublisher( + const FullTrackName& name) override; + MoqtPublishingMonitorInterface* ReleaseMonitoringInterface( + const FullTrackName& name) override; + const quic::QuicClock* clock() override { return callbacks_.clock; } + MoqtTraceRecorder& trace_recorder() override { return trace_recorder_; } webtransport::Session* session() override { return is_closing_ ? nullptr : session_; } @@ -161,16 +170,17 @@ // draft-ietf-moqt-moq-transport-12. Unsubscribe and notify the application so // the error can be propagated downstream, if necessary. void OnMalformedTrack(RemoteTrack* track); - quiche::QuicheWeakPtr<RemoteTrack> GetSubscribe(uint64_t track_alias) { - auto it = subscribe_by_alias_.find(track_alias); - if (it == subscribe_by_alias_.end()) { + quiche::QuicheWeakPtr<RemoteTrack> GetSubscribe( + uint64_t track_alias) override { + RemoteTrack* track = SubscribeByAlias(track_alias); + if (track == nullptr) { return quiche::QuicheWeakPtr<RemoteTrack>(); } - return it->second->weak_ptr(); + return track->weak_ptr(); } quiche::QuicheWeakPtr<RemoteTrack> GetFetch(uint64_t request_id) { - auto it = upstream_by_id_.find(request_id); - if (it == upstream_by_id_.end()) { + auto it = fetch_by_id_.find(request_id); + if (it == fetch_by_id_.end()) { return quiche::QuicheWeakPtr<RemoteTrack>(); } return it->second->weak_ptr(); @@ -208,8 +218,6 @@ void UseAlternateDeliveryTimeout() { alternate_delivery_timeout_ = true; } - MoqtTraceRecorder& trace_recorder() { return trace_recorder_; } - private: friend class ControlMessageDispatcher; friend class test::MoqtSessionPeer; @@ -254,18 +262,12 @@ } }), session_(session), - weak_ptr_factory_(this) { - this->set_control_stream(); - } - // TODO(martinduke): Remove constructor body once SUBSCRIBE moves to a bidi - // stream. + weak_ptr_factory_(this) {} void OnStreamBound() override; absl::Status OnRawControlMessage( const MoqtRawControlMessage& message) override; - // MoqtControlParserVisitor implementation. - // webtransport::StreamVisitor overrides void OnResetStreamReceived(webtransport::StreamErrorCode error) override { session_->Error(MoqtError::kProtocolViolation, @@ -338,7 +340,7 @@ MessageParameters parameters; parameters.expires = publisher_->expiration(); parameters.largest_object = publisher_->largest_location(); - MoqtBidiStreamBase* control_stream = session_->GetControlStream(); + ControlStream* control_stream = session_->GetControlStream(); if (control_stream != nullptr) { control_stream->CheckStatus( control_stream->SendRequestOk(request_id_, parameters)); @@ -348,7 +350,7 @@ } void OnSubscribeRejected(MoqtRequestErrorInfo info) override { - MoqtBidiStreamBase* control_stream = session_->GetControlStream(); + ControlStream* control_stream = session_->GetControlStream(); if (control_stream != nullptr) { control_stream->CheckStatus(control_stream->SendRequestError( request_id_, info.error_code, info.retry_interval, @@ -397,10 +399,9 @@ // Returns false if creation failed. [[nodiscard]] bool OpenDataStream(PublishedFetch* fetch, webtransport::SendOrder send_order); - - SubscribeRemoteTrack* RemoteTrackByAlias(uint64_t track_alias); - RemoteTrack* RemoteTrackById(uint64_t request_id); - SubscribeRemoteTrack* RemoteTrackByName(const FullTrackName& name); + SubscribeRemoteTrack* SubscribeByAlias(uint64_t track_alias); + SubscribeRemoteTrack* SubscribeByName(const FullTrackName& track_name); + UpstreamFetch* FetchById(uint64_t request_id); // Checks that a subscribe ID from a SUBSCRIBE or FETCH is valid, and throws // a session error if is not. @@ -409,20 +410,17 @@ void CancelFetch(uint64_t request_id); // Sends an OBJECT_ACK message for a specific subscribe ID. - void SendObjectAck(uint64_t subscribe_id, uint64_t group_id, + void SendObjectAck(FullTrackName track_name, uint64_t group_id, uint64_t object_id, quic::QuicTimeDelta delta_from_deadline) { if (!SupportsObjectAck()) { return; } - MoqtObjectAck ack; - ack.subscribe_id = subscribe_id; - ack.group_id = group_id; - ack.object_id = object_id; - ack.delta_from_deadline = delta_from_deadline; - SendControlMessage(framer_.SerializeObjectAck(ack)); + SubscribeRemoteTrack* track = SubscribeByName(track_name); + if (track != nullptr) { + track->SendObjectAck(group_id, object_id, delta_from_deadline); + } } - // Indicates if OBJECT_ACK is supported by both sides. bool SupportsObjectAck() const { return parameters_.support_object_acks && peer_supports_object_ack_; @@ -440,12 +438,11 @@ // Handlers for the control messages on the main control stream. absl::Status OnControlMessage(const MoqtSetup& message); + + // TODO(martinduke): All of these should be moved to bidi streams or + // deleted. absl::Status OnControlMessage(const MoqtRequestOk& message); absl::Status OnControlMessage(const MoqtRequestError& message); - absl::Status OnControlMessage(const MoqtSubscribe& message); - absl::Status OnControlMessage(const MoqtSubscribeOk& message); - absl::Status OnControlMessage(const MoqtUnsubscribe& message); - absl::Status OnControlMessage(const MoqtPublishDone& /*message*/); absl::Status OnControlMessage(const MoqtRequestUpdate& message); absl::Status OnControlMessage(const MoqtPublishNamespace& message); absl::Status OnControlMessage(const MoqtPublishNamespaceDone& /*message*/); @@ -459,15 +456,6 @@ } absl::Status OnControlMessage(const MoqtFetchOk& message); absl::Status OnControlMessage(const MoqtRequestsBlocked& message); - absl::Status OnControlMessage(const MoqtPublish& message); - absl::Status OnControlMessage(const MoqtObjectAck& message) { - auto subscription_it = published_subscriptions_.find(message.subscribe_id); - if (subscription_it == published_subscriptions_.end()) { - return absl::OkStatus(); - } - subscription_it->second->ProcessObjectAck(message); - return absl::OkStatus(); - } // TODO(vasilvv): remove this once all requests are moved into individual // streams. @@ -483,6 +471,12 @@ SendControlMessage(framer_.SerializeRequestError(request_error)); } + uint64_t NextRequestId() { + uint64_t id = next_request_id_; + next_request_id_ += 2; + return id; + } + bool is_closing_ = false; webtransport::Session* session_; MoqtSessionParameters parameters_; @@ -501,24 +495,12 @@ MoqtTraceRecorder trace_recorder_; - // Upstream SUBSCRIBE state. - // Upstream SUBSCRIBEs and FETCHes, indexed by subscribe_id. Do not erase - // directly, call RemoteTrack::Destroy(), except in deletion callbacks passed - // to RemoteTrack. - absl::flat_hash_map<uint64_t, std::unique_ptr<RemoteTrack>> upstream_by_id_; - // All SUBSCRIBEs, indexed by track_alias. + // Upstream FETCHes, indexed by request_id. Do not erase. + absl::flat_hash_map<uint64_t, std::unique_ptr<UpstreamFetch>> fetch_by_id_; + // All outgoing SUBSCRIBE and incoming PUBLISH, indexed by track_alias. absl::flat_hash_map<uint64_t, SubscribeRemoteTrack*> subscribe_by_alias_; - // All SUBSCRIBEs, indexed by track name. + // All outgoing SUBSCRIBE and incoming PUBLISH, indexed by track name. absl::flat_hash_map<FullTrackName, SubscribeRemoteTrack*> subscribe_by_name_; - struct SubscribeUpdateStatus { - FullTrackName name; - MessageParameters parameters; - MoqtResponseCallback response_callback; - }; - // Outgoing Subscribe Updates. We should not update parameters until a - // REQUEST_OK arrives. - absl::flat_hash_map<uint64_t, SubscribeUpdateStatus> - pending_subscribe_updates_; // The next subscribe ID that the local endpoint can send. uint64_t next_request_id_ = 0; @@ -528,15 +510,16 @@ // All open incoming subscriptions, indexed by track name, used to check for // duplicates. - absl::flat_hash_set<FullTrackName> subscribed_track_names_; + absl::flat_hash_map<FullTrackName, SubscriptionPublisher*> + subscribed_track_names_; // Application object representing the publisher for all of the tracks that // can be subscribed to via this connection. Must outlive this object. MoqtPublisher* publisher_; - // Subscriptions for local tracks by the remote peer, indexed by subscribe ID. - absl::flat_hash_map<uint64_t, std::unique_ptr<SubscriptionPublisher>> + // Subscriptions for local tracks by the remote peer, indexed by request ID. + absl::flat_hash_map<uint64_t, SubscriptionPublisher*> published_subscriptions_; - // Keeps track of all request IDs that have queued outgoing data streams. The - // first element is the highest priority (lowest integer). + // Keeps track of all request IDs that have queued outgoing data streams. + // The first element is the highest priority (lowest integer). absl::btree_multimap<MoqtTrackPriority, uint64_t> subscriptions_with_queued_streams_; // This is only used to check for track_alias collisions. @@ -587,6 +570,8 @@ bool alternate_delivery_timeout_ = false; quiche::QuicheWeakPtrFactory<MoqtSessionInterface> weak_ptr_factory_; + quiche::QuicheWeakPtrFactory<SessionToPublisherInterface> + weak_ptr_factory_for_publishers_; // Must be last. Token used to make sure that the streams do not call into // the session when the session has already been destroyed.
diff --git a/quiche/quic/moqt/moqt_session_interface.h b/quiche/quic/moqt/moqt_session_interface.h index e1455a6..72c5164 100644 --- a/quiche/quic/moqt/moqt_session_interface.h +++ b/quiche/quic/moqt/moqt_session_interface.h
@@ -97,13 +97,19 @@ virtual void Error(MoqtError code, absl::string_view error) = 0; // Return true if SUBSCRIBE was actually sent. - virtual bool Subscribe(const FullTrackName& name, SubscribeVisitor* visitor, + virtual bool Subscribe(const FullTrackName& name, + SubscribeVisitor* absl_nonnull visitor, const MessageParameters& parameters) = 0; // If a parameter is nullopt, there is no change to the current value. - // Returns false if the subscription is not found. + // Returns false if the subscription is not found. Used by the subscriber for + // a SUBSCRIBE or PUBLISH. virtual bool SubscribeUpdate(const FullTrackName& name, const MessageParameters& parameters, MoqtResponseCallback response_callback) = 0; + // Used by the publisher of a PUBLISH message. + virtual bool PublishUpdate(const FullTrackName& name, + const MessageParameters& parameters, + MoqtResponseCallback response_callback) = 0; // Sends an UNSUBSCRIBE message and removes all of the state related to the // subscription. Returns false if the subscription is not found.
diff --git a/quiche/quic/moqt/moqt_session_test.cc b/quiche/quic/moqt/moqt_session_test.cc index c327f5e..afff2a7 100644 --- a/quiche/quic/moqt/moqt_session_test.cc +++ b/quiche/quic/moqt/moqt_session_test.cc
@@ -7,7 +7,6 @@ #include <algorithm> #include <cstdint> #include <cstring> -#include <functional> #include <memory> #include <optional> #include <queue> @@ -17,7 +16,6 @@ #include <vector> #include "absl/base/casts.h" -#include "absl/memory/memory.h" #include "absl/status/status.h" #include "absl/strings/string_view.h" #include "absl/types/span.h" @@ -32,16 +30,13 @@ #include "quiche/quic/moqt/moqt_known_track_publisher.h" #include "quiche/quic/moqt/moqt_messages.h" #include "quiche/quic/moqt/moqt_names.h" -#include "quiche/quic/moqt/moqt_namespace_stream.h" #include "quiche/quic/moqt/moqt_object.h" -#include "quiche/quic/moqt/moqt_parser.h" #include "quiche/quic/moqt/moqt_priority.h" #include "quiche/quic/moqt/moqt_publisher.h" #include "quiche/quic/moqt/moqt_session_callbacks.h" #include "quiche/quic/moqt/moqt_session_interface.h" #include "quiche/quic/moqt/moqt_track.h" #include "quiche/quic/moqt/moqt_types.h" -#include "quiche/quic/moqt/session_namespace_tree.h" #include "quiche/quic/moqt/test_tools/moqt_framer_utils.h" #include "quiche/quic/moqt/test_tools/moqt_mock_visitor.h" #include "quiche/quic/moqt/test_tools/moqt_session_peer.h" @@ -51,7 +46,6 @@ #include "quiche/common/quiche_data_reader.h" #include "quiche/common/quiche_mem_slice.h" #include "quiche/common/quiche_weak_ptr.h" -#include "quiche/common/test_tools/quiche_test_utils.h" #include "quiche/web_transport/test_tools/in_memory_stream.h" #include "quiche/web_transport/test_tools/mock_web_transport.h" #include "quiche/web_transport/web_transport.h" @@ -146,7 +140,8 @@ session_.set_publisher(&publisher_); MoqtSessionPeer::set_peer_max_request_id(&session_, kDefaultInitialMaxRequestId); - ON_CALL(mock_session_, GetStreamById).WillByDefault(Return(&mock_stream_)); + ON_CALL(mock_session_, GetStreamById) + .WillByDefault(Return(&mock_bidi_stream_)); EXPECT_EQ(MoqtSessionPeer::GetImplementationString(&session_), kImplementationName); } @@ -168,6 +163,71 @@ ON_CALL(*publisher, largest_location()).WillByDefault(Return(largest_id)); } + // Opens an incoming request stream, determines the type based on first_byte + // and returns MoqtBidiStreamBase that can be downcast to the correct type. + // |wt_stream| is the underlying mock WebTransport stream; if nullptr, use + // mock_bidi_stream_. + static constexpr absl::string_view kSubscribeByte = "\x03"; + static constexpr absl::string_view kSubscribeNamespaceByte = "\x11"; + static constexpr absl::string_view kPublishByte = "\x1d"; + std::unique_ptr<MoqtBidiStreamBase> ResponseStream( + absl::string_view first_byte, + webtransport::test::MockStream* wt_stream = nullptr) { + webtransport::test::MockStream* stream = + wt_stream != nullptr ? wt_stream : &mock_bidi_stream_; + EXPECT_CALL(mock_session_, AcceptIncomingBidirectionalStream()) + .WillOnce(Return(stream)) + .WillOnce(Return(nullptr)); + std::unique_ptr<webtransport::StreamVisitor> unknown_bidi_stream; + std::unique_ptr<MoqtBidiStreamBase> final_stream; + EXPECT_CALL(*stream, SetVisitor) + .WillOnce([&](std::unique_ptr<webtransport::StreamVisitor> visitor) { + unknown_bidi_stream = std::move(visitor); + }) + .WillOnce([&](std::unique_ptr<webtransport::StreamVisitor> visitor) { + final_stream = std::unique_ptr<MoqtBidiStreamBase>( + absl::down_cast<MoqtBidiStreamBase*>(visitor.release())); + }); + EXPECT_CALL(*stream, visitor()) + .WillOnce([&]() { return unknown_bidi_stream.get(); }) + .WillRepeatedly([&]() { return final_stream.get(); }); + EXPECT_CALL(*stream, PeekNextReadableRegion) + .WillOnce( + Return(webtransport::Stream::PeekResult(first_byte, false, false))) + .WillRepeatedly( + Return(webtransport::Stream::PeekResult("", false, false))); + EXPECT_CALL(*stream, ReadableBytes()) + .WillOnce(Return(first_byte.length())) + .WillRepeatedly(Return(0)); + EXPECT_CALL(*stream, Read(::testing::An<absl::Span<char>>())) + .WillOnce([&](absl::Span<char> bytes_to_read) { + memcpy(bytes_to_read.data(), first_byte.data(), first_byte.length()); + return webtransport::Stream::ReadResult(first_byte.length(), false); + }); + session_.OnIncomingBidirectionalStreamAvailable(); + EXPECT_NE(final_stream, nullptr); + EXPECT_CALL(*stream, CanWrite).WillRepeatedly(Return(true)); + return final_stream; + } + + void PrepareRequestStream( + std::unique_ptr<MoqtBidiStreamTestWrapper>& stream_wrapper, + webtransport::test::MockStream* wt_stream = nullptr) { + webtransport::test::MockStream* stream = + wt_stream != nullptr ? wt_stream : &mock_bidi_stream_; + EXPECT_CALL(mock_session_, CanOpenNextOutgoingBidirectionalStream()) + .WillOnce(Return(true)); + EXPECT_CALL(mock_session_, OpenOutgoingBidirectionalStream()) + .WillOnce(Return(stream)); + EXPECT_CALL(*stream, SetVisitor) + .WillOnce([&](std::unique_ptr<webtransport::StreamVisitor> visitor) { + stream_wrapper = std::make_unique<MoqtBidiStreamTestWrapper>( + std::unique_ptr<MoqtBidiStreamBase>( + absl::down_cast<MoqtBidiStreamBase*>(visitor.release()))); + }); + EXPECT_CALL(*stream, CanWrite).WillRepeatedly(Return(true)); + } + // The publisher receives SUBSCRIBE and synchronously publishes namespaces it // supports. MoqtObjectListener* ReceiveSubscribeSynchronousOk( @@ -189,7 +249,8 @@ parameters, extensions, }; - EXPECT_CALL(mock_stream_, Writev(SerializedControlMessage(expected_ok), _)); + EXPECT_CALL(mock_bidi_stream_, + Writev(SerializedControlMessage(expected_ok), _)); control_parser->ReceiveMessage(subscribe); return listener_ptr; } @@ -258,12 +319,14 @@ } } - webtransport::test::MockStream mock_stream_, control_stream_; - MockSessionCallbacks session_callbacks_; - webtransport::test::MockSession mock_session_; MockSubscribeRemoteTrackVisitor remote_track_visitor_; - MoqtSession session_; MoqtKnownTrackPublisher publisher_; + webtransport::test::MockSession mock_session_; + MockSessionCallbacks session_callbacks_; + std::unique_ptr<MoqtBidiStreamTestWrapper> bidi_wrapper_; + MoqtSession session_; + webtransport::test::MockStream mock_bidi_stream_, mock_uni_stream_; + // std::shared_ptr<IncomingSubscribeInfo> last_incoming_subscribe_; }; TEST_F(MoqtSessionTest, Queries) { @@ -277,28 +340,28 @@ EXPECT_CALL(mock_session_, CanOpenNextOutgoingBidirectionalStream()) .WillOnce(Return(true)); EXPECT_CALL(mock_session_, OpenOutgoingBidirectionalStream()) - .WillOnce(Return(&mock_stream_)); - EXPECT_CALL(mock_stream_, CanWrite).WillRepeatedly(Return(true)); + .WillOnce(Return(&mock_bidi_stream_)); + EXPECT_CALL(mock_bidi_stream_, CanWrite).WillRepeatedly(Return(true)); std::unique_ptr<webtransport::StreamVisitor> visitor; // Save a reference to MoqtSession::Stream - EXPECT_CALL(mock_stream_, SetVisitor(_)) + EXPECT_CALL(mock_bidi_stream_, SetVisitor(_)) .WillOnce([&](std::unique_ptr<webtransport::StreamVisitor> new_visitor) { visitor = std::move(new_visitor); }); - EXPECT_CALL(mock_stream_, GetStreamId()) + EXPECT_CALL(mock_bidi_stream_, GetStreamId()) .WillRepeatedly(Return(webtransport::StreamId(4))); - EXPECT_CALL(mock_stream_, + EXPECT_CALL(mock_bidi_stream_, Writev(ControlMessageOfType(MoqtMessageType::kSetup), _)); session_.OnSessionReady(); // Receive SERVER_SETUP - std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input = + bidi_wrapper_ = MoqtSessionPeer::FetchParserVisitorFromWebtransportStreamVisitor( std::move(visitor)); // Handle the server setup MoqtSetup setup; // No fields are set. EXPECT_CALL(session_callbacks_.session_established_callback, Call()).Times(1); - stream_input->ReceiveMessage(setup); + bidi_wrapper_->ReceiveMessage(setup); } TEST_F(MoqtSessionTest, OnSessionReadyNoControlStream) { @@ -316,16 +379,16 @@ std::make_unique<quic::test::TestAlarmFactory>(), session_callbacks_.AsSessionCallbacks()); EXPECT_CALL(mock_session_, AcceptIncomingBidirectionalStream()) - .WillOnce(Return(&mock_stream_)) + .WillOnce(Return(&mock_bidi_stream_)) .WillOnce(Return(nullptr)); std::unique_ptr<webtransport::StreamVisitor> visitor; webtransport::test::MockStreamVisitor mock_stream_visitor; - EXPECT_CALL(mock_stream_, SetVisitor) + EXPECT_CALL(mock_bidi_stream_, SetVisitor) .WillOnce([&](std::unique_ptr<webtransport::StreamVisitor> new_visitor) { visitor = std::move(new_visitor); - EXPECT_CALL(mock_stream_, visitor).WillOnce(Return(visitor.get())); + EXPECT_CALL(mock_bidi_stream_, visitor).WillOnce(Return(visitor.get())); }); - EXPECT_CALL(mock_stream_, PeekNextReadableRegion()) + EXPECT_CALL(mock_bidi_stream_, PeekNextReadableRegion()) .WillRepeatedly(Return( webtransport::Stream::PeekResult(absl::string_view(), false, false))); server_session.OnIncomingBidirectionalStreamAvailable(); @@ -371,9 +434,10 @@ ::testing::InSequence seq; StrictMock<webtransport::test::MockStreamVisitor> mock_stream_visitor; EXPECT_CALL(mock_session_, AcceptIncomingBidirectionalStream()) - .WillOnce(Return(&mock_stream_)); - EXPECT_CALL(mock_stream_, SetVisitor(_)).Times(1); - EXPECT_CALL(mock_stream_, visitor()).WillOnce(Return(&mock_stream_visitor)); + .WillOnce(Return(&mock_bidi_stream_)); + EXPECT_CALL(mock_bidi_stream_, SetVisitor).Times(1); + EXPECT_CALL(mock_bidi_stream_, visitor()) + .WillOnce(Return(&mock_stream_visitor)); EXPECT_CALL(mock_stream_visitor, OnCanRead()).Times(1); EXPECT_CALL(mock_session_, AcceptIncomingBidirectionalStream()) .WillOnce(Return(nullptr)); @@ -384,9 +448,10 @@ ::testing::InSequence seq; StrictMock<webtransport::test::MockStreamVisitor> mock_stream_visitor; EXPECT_CALL(mock_session_, AcceptIncomingUnidirectionalStream()) - .WillOnce(Return(&mock_stream_)); - EXPECT_CALL(mock_stream_, SetVisitor(_)).Times(1); - EXPECT_CALL(mock_stream_, visitor()).WillOnce(Return(&mock_stream_visitor)); + .WillOnce(Return(&mock_uni_stream_)); + EXPECT_CALL(mock_uni_stream_, SetVisitor(_)).Times(1); + EXPECT_CALL(mock_uni_stream_, visitor()) + .WillOnce(Return(&mock_stream_visitor)); EXPECT_CALL(mock_stream_visitor, OnCanRead()).Times(1); EXPECT_CALL(mock_session_, AcceptIncomingUnidirectionalStream()) .WillOnce(Return(nullptr)); @@ -409,18 +474,20 @@ TEST_F(MoqtSessionTest, AddLocalTrack) { MoqtSubscribe request = DefaultSubscribe(); - std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input = - MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_); + bidi_wrapper_ = std::make_unique<MoqtBidiStreamTestWrapper>( + ResponseStream(kSubscribeByte)); // Request for track returns REQUEST_ERROR. - EXPECT_CALL(mock_stream_, + EXPECT_CALL(mock_bidi_stream_, Writev(ControlMessageOfType(MoqtMessageType::kRequestError), _)); - stream_input->ReceiveMessage(request); + bidi_wrapper_->ReceiveMessage(request); // Add the track. Now Subscribe should succeed. MockTrackPublisher* track = CreateTrackPublisher(); - std::make_shared<MockTrackPublisher>(request.full_track_name); request.request_id += 2; - ReceiveSubscribeSynchronousOk(track, request, stream_input.get()); + bidi_wrapper_ = std::make_unique<MoqtBidiStreamTestWrapper>( + ResponseStream(kSubscribeByte)); + ReceiveSubscribeSynchronousOk(track, request, bidi_wrapper_.get(), + MoqtSessionPeer::GetLastTrackAlias(&session_)); } TEST_F(MoqtSessionTest, IncomingPublishRejected) { @@ -431,22 +498,22 @@ .parameters = MessageParameters(), }; publish.parameters.largest_object = Location(4, 5); - std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input = - MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_); - // Request for track returns REQUEST_ERROR. - EXPECT_CALL(mock_stream_, + bidi_wrapper_ = + std::make_unique<MoqtBidiStreamTestWrapper>(ResponseStream(kPublishByte)); + // With the default incoming_publish_callbackm, will return REQUEST_ERROR. + EXPECT_CALL(mock_bidi_stream_, Writev(ControlMessageOfType(MoqtMessageType::kRequestError), _)); - stream_input->ReceiveMessage(publish); + bidi_wrapper_->ReceiveMessage(publish); } TEST_F(MoqtSessionTest, PublishNamespaceWithOkAndCancel) { testing::MockFunction<void( std::variant<MessageParameters, MoqtRequestErrorInfo> error_message)> publish_namespace_response_callback; - std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input = - MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_); + std::unique_ptr<MoqtBidiStreamTestWrapper> bidi_wrapper_ = + MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_); EXPECT_CALL( - mock_stream_, + mock_bidi_stream_, Writev(ControlMessageOfType(MoqtMessageType::kPublishNamespace), _)); MoqtRequestErrorInfo cancel_error_info; session_.PublishNamespace( @@ -460,14 +527,14 @@ [&](std::variant<MessageParameters, MoqtRequestErrorInfo> response) { EXPECT_TRUE(std::holds_alternative<MessageParameters>(response)); }); - stream_input->ReceiveMessage(ok); + bidi_wrapper_->ReceiveMessage(ok); MoqtPublishNamespaceCancel cancel = { /*request_id=*/0, RequestErrorCode::kInternalError, /*error_reason=*/"Test error", }; - stream_input->ReceiveMessage(cancel); + bidi_wrapper_->ReceiveMessage(cancel); EXPECT_EQ(cancel_error_info.error_code, RequestErrorCode::kInternalError); EXPECT_EQ(cancel_error_info.reason_phrase, "Test error"); // State is gone. @@ -478,10 +545,10 @@ testing::MockFunction<void( std::variant<MessageParameters, MoqtRequestErrorInfo>)> publish_namespace_resolved_callback; - std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input = - MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_); + std::unique_ptr<MoqtBidiStreamTestWrapper> bidi_wrapper_ = + MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_); EXPECT_CALL( - mock_stream_, + mock_bidi_stream_, Writev(ControlMessageOfType(MoqtMessageType::kPublishNamespace), _)); session_.PublishNamespace(TrackNamespace{"foo"}, MessageParameters(), publish_namespace_resolved_callback.AsStdFunction(), @@ -493,10 +560,10 @@ [&](std::variant<MessageParameters, MoqtRequestErrorInfo> response) { EXPECT_TRUE(std::holds_alternative<MessageParameters>(response)); }); - stream_input->ReceiveMessage(ok); + bidi_wrapper_->ReceiveMessage(ok); EXPECT_CALL( - mock_stream_, + mock_bidi_stream_, Writev(ControlMessageOfType(MoqtMessageType::kPublishNamespaceDone), _)); session_.PublishNamespaceDone(TrackNamespace{"foo"}); // State is gone. @@ -507,10 +574,10 @@ testing::MockFunction<void( std::variant<MessageParameters, MoqtRequestErrorInfo>)> publish_namespace_resolved_callback; - std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input = - MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_); + std::unique_ptr<MoqtBidiStreamTestWrapper> bidi_wrapper_ = + MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_); EXPECT_CALL( - mock_stream_, + mock_bidi_stream_, Writev(ControlMessageOfType(MoqtMessageType::kPublishNamespace), _)); session_.PublishNamespace(TrackNamespace{"foo"}, MessageParameters(), publish_namespace_resolved_callback.AsStdFunction(), @@ -527,23 +594,23 @@ EXPECT_EQ(error.error_code, RequestErrorCode::kInternalError); EXPECT_EQ(error.reason_phrase, "Test error"); }); - stream_input->ReceiveMessage(error); + bidi_wrapper_->ReceiveMessage(error); // State is gone. EXPECT_FALSE(session_.PublishNamespaceDone(TrackNamespace{"foo"})); } TEST_F(MoqtSessionTest, AsynchronousSubscribeReturnsOk) { - std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input = - MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_); + bidi_wrapper_ = std::make_unique<MoqtBidiStreamTestWrapper>( + ResponseStream(kSubscribeByte)); MoqtSubscribe request = DefaultSubscribe(); MockTrackPublisher* track = CreateTrackPublisher(); MoqtObjectListener* listener; EXPECT_CALL(*track, AddObjectListener) .WillOnce( [&](MoqtObjectListener* listener_ptr) { listener = listener_ptr; }); - stream_input->ReceiveMessage(request); + bidi_wrapper_->ReceiveMessage(request); - EXPECT_CALL(mock_stream_, + EXPECT_CALL(mock_bidi_stream_, Writev(ControlMessageOfType(MoqtMessageType::kSubscribeOk), _)); listener->OnSubscribeAccepted(); EXPECT_TRUE(MoqtSessionPeer::RequestIdIsSubscriptionPublisher( @@ -551,61 +618,66 @@ } TEST_F(MoqtSessionTest, AsynchronousSubscribeReturnsError) { - std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input = - MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_); + bidi_wrapper_ = std::make_unique<MoqtBidiStreamTestWrapper>( + ResponseStream(kSubscribeByte)); MoqtSubscribe request = DefaultSubscribe(); MockTrackPublisher* track = CreateTrackPublisher(); MoqtObjectListener* listener; EXPECT_CALL(*track, AddObjectListener) .WillOnce( [&](MoqtObjectListener* listener_ptr) { listener = listener_ptr; }); - stream_input->ReceiveMessage(request); - EXPECT_CALL(mock_stream_, - Writev(ControlMessageOfType(MoqtMessageType::kRequestError), _)); + bidi_wrapper_->ReceiveMessage(request); + EXPECT_CALL(mock_bidi_stream_, + Writev(ControlMessageOfType(MoqtMessageType::kRequestError), _)) + .WillOnce([](absl::Span<quiche::QuicheMemSlice>, + const webtransport::StreamWriteOptions& options) { + EXPECT_TRUE(options.send_fin()); + return absl::OkStatus(); + }); listener->OnSubscribeRejected(MoqtRequestErrorInfo( RequestErrorCode::kInternalError, std::nullopt, "Test error")); - EXPECT_FALSE(MoqtSessionPeer::RequestIdIsSubscriptionPublisher( - &session_, kDefaultPeerRequestId)); } TEST_F(MoqtSessionTest, SynchronousSubscribeReturnsError) { - std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input = - MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_); + bidi_wrapper_ = std::make_unique<MoqtBidiStreamTestWrapper>( + ResponseStream(kSubscribeByte)); MoqtSubscribe request = DefaultSubscribe(); MockTrackPublisher* track = CreateTrackPublisher(); EXPECT_CALL(*track, AddObjectListener) .WillOnce([&](MoqtObjectListener* listener) { - EXPECT_CALL( - mock_stream_, - Writev(ControlMessageOfType(MoqtMessageType::kRequestError), _)); EXPECT_CALL(*track, RemoveObjectListener); listener->OnSubscribeRejected(MoqtRequestErrorInfo( RequestErrorCode::kInternalError, std::nullopt, "Test error")); }); - stream_input->ReceiveMessage(request); - EXPECT_FALSE(MoqtSessionPeer::RequestIdIsSubscriptionPublisher( - &session_, kDefaultPeerRequestId)); + EXPECT_CALL(mock_bidi_stream_, + Writev(ControlMessageOfType(MoqtMessageType::kRequestError), _)) + .WillOnce([](absl::Span<quiche::QuicheMemSlice>, + const webtransport::StreamWriteOptions& options) { + EXPECT_TRUE(options.send_fin()); + return absl::OkStatus(); + }); + bidi_wrapper_->ReceiveMessage(request); } TEST_F(MoqtSessionTest, SubscribeForPast) { - std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input = - MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_); + bidi_wrapper_ = std::make_unique<MoqtBidiStreamTestWrapper>( + ResponseStream(kSubscribeByte)); MockTrackPublisher* track = CreateTrackPublisher(); SetLargestId(track, Location(10, 20)); MoqtSubscribe request = DefaultSubscribe(); - ReceiveSubscribeSynchronousOk(track, request, stream_input.get()); + ReceiveSubscribeSynchronousOk(track, request, bidi_wrapper_.get()); } TEST_F(MoqtSessionTest, SubscribeDoNotForward) { - std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input = - MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_); + bidi_wrapper_ = std::make_unique<MoqtBidiStreamTestWrapper>( + ResponseStream(kSubscribeByte)); MockTrackPublisher* track = CreateTrackPublisher(); MoqtSubscribe request = DefaultSubscribe(); request.parameters.set_forward(false); request.parameters.subscription_filter.emplace( MoqtFilterType::kLargestObject); MoqtObjectListener* listener = - ReceiveSubscribeSynchronousOk(track, request, stream_input.get()); + ReceiveSubscribeSynchronousOk(track, request, bidi_wrapper_.get()); // forward=false, so incoming objects are ignored. EXPECT_CALL(mock_session_, CanOpenNextOutgoingUnidirectionalStream()) .Times(0); @@ -613,13 +685,13 @@ } TEST_F(MoqtSessionTest, SubscribeAbsoluteStartNoDataYet) { - std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input = - MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_); + bidi_wrapper_ = std::make_unique<MoqtBidiStreamTestWrapper>( + ResponseStream(kSubscribeByte)); MockTrackPublisher* track = CreateTrackPublisher(); MoqtSubscribe request = DefaultSubscribe(); request.parameters.subscription_filter.emplace(Location(1, 0)); MoqtObjectListener* listener = - ReceiveSubscribeSynchronousOk(track, request, stream_input.get()); + ReceiveSubscribeSynchronousOk(track, request, bidi_wrapper_.get()); // Window was not set to (0, 0) by SUBSCRIBE acceptance. EXPECT_CALL(mock_session_, CanOpenNextOutgoingUnidirectionalStream()) .Times(0); @@ -627,15 +699,15 @@ } TEST_F(MoqtSessionTest, SubscribeNextGroup) { - std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input = - MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_); + bidi_wrapper_ = std::make_unique<MoqtBidiStreamTestWrapper>( + ResponseStream(kSubscribeByte)); MockTrackPublisher* track = CreateTrackPublisher(); MoqtSubscribe request = DefaultSubscribe(); request.parameters.subscription_filter.emplace( MoqtFilterType::kNextGroupStart); SetLargestId(track, Location(10, 20)); MoqtObjectListener* listener = - ReceiveSubscribeSynchronousOk(track, request, stream_input.get()); + ReceiveSubscribeSynchronousOk(track, request, bidi_wrapper_.get()); // Later objects in group 10 ignored. EXPECT_CALL(mock_session_, CanOpenNextOutgoingUnidirectionalStream()) .Times(0); @@ -648,88 +720,83 @@ } TEST_F(MoqtSessionTest, TwoSubscribesForTrack) { - std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input = - MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_); + bidi_wrapper_ = std::make_unique<MoqtBidiStreamTestWrapper>( + ResponseStream(kSubscribeByte)); MockTrackPublisher* track = CreateTrackPublisher(); MoqtSubscribe request = DefaultSubscribe(); - ReceiveSubscribeSynchronousOk(track, request, stream_input.get()); + ReceiveSubscribeSynchronousOk(track, request, bidi_wrapper_.get()); request.request_id = 3; request.parameters.subscription_filter.emplace(Location(12, 0)); - EXPECT_CALL(mock_stream_, + webtransport::test::MockStream bidi_stream_2; + auto bidi_wrapper_2 = std::make_unique<MoqtBidiStreamTestWrapper>( + ResponseStream(kSubscribeByte, &bidi_stream_2)); + EXPECT_CALL(bidi_stream_2, Writev(ControlMessageOfType(MoqtMessageType::kRequestError), _)); - stream_input->ReceiveMessage(request); + bidi_wrapper_2->ReceiveMessage(request); } TEST_F(MoqtSessionTest, UnsubscribeAllowsSecondSubscribe) { - std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input = - MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_); + bidi_wrapper_ = std::make_unique<MoqtBidiStreamTestWrapper>( + ResponseStream(kSubscribeByte)); MockTrackPublisher* track = CreateTrackPublisher(); MoqtSubscribe request = DefaultSubscribe(); - ReceiveSubscribeSynchronousOk(track, request, stream_input.get()); + ReceiveSubscribeSynchronousOk(track, request, bidi_wrapper_.get()); // Peer unsubscribes. - MoqtUnsubscribe unsubscribe = { - kDefaultPeerRequestId, - }; - stream_input->ReceiveMessage(unsubscribe); + bidi_wrapper_->stream().Reset(kResetCodeCancelled); + bidi_wrapper_ = nullptr; EXPECT_FALSE(MoqtSessionPeer::RequestIdIsSubscriptionPublisher(&session_, 1)); // Subscribe again, succeeds. request.request_id = 3; request.parameters.subscription_filter.emplace(Location(12, 0)); - ReceiveSubscribeSynchronousOk(track, request, stream_input.get(), + bidi_wrapper_ = std::make_unique<MoqtBidiStreamTestWrapper>( + ResponseStream(kSubscribeByte)); + ReceiveSubscribeSynchronousOk(track, request, bidi_wrapper_.get(), /*track_alias=*/1); } -TEST_F(MoqtSessionTest, RequestIdTooHigh) { - // Peer subscribes to (0, 0) - MoqtSubscribe request = DefaultSubscribe(); - request.request_id = kDefaultInitialMaxRequestId + 1; - - std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input = - MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_); - EXPECT_CALL(mock_session_, - CloseSession(static_cast<uint64_t>(MoqtError::kTooManyRequests), - "Received request with too large ID")); - stream_input->ReceiveMessage(request); -} - TEST_F(MoqtSessionTest, RequestIdWrongLsb) { // TODO(martinduke): Implement this test. } TEST_F(MoqtSessionTest, SubscribeIdNotIncreasing) { MoqtSubscribe request = DefaultSubscribe(); - std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input = - MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_); + bidi_wrapper_ = std::make_unique<MoqtBidiStreamTestWrapper>( + ResponseStream(kSubscribeByte)); MockTrackPublisher* track = CreateTrackPublisher(); EXPECT_CALL(*track, AddObjectListener); - stream_input->ReceiveMessage(request); + bidi_wrapper_->ReceiveMessage(request); // Second request is a protocol violation. request.full_track_name = FullTrackName({"dead", "beef"}); - EXPECT_CALL(mock_session_, - CloseSession(static_cast<uint64_t>(MoqtError::kInvalidRequestId), - "Duplicate request ID")); - stream_input->ReceiveMessage(request); + auto publisher = + std::make_shared<MockTrackPublisher>(request.full_track_name); + publisher_.Add(publisher); + webtransport::test::MockStream bidi_stream_2; + auto bidi_wrapper_2 = std::make_unique<MoqtBidiStreamTestWrapper>( + ResponseStream(kSubscribeByte, &bidi_stream_2)); + EXPECT_CALL(bidi_stream_2, + Writev(ControlMessageOfType(MoqtMessageType::kRequestError), _)); + bidi_wrapper_2->ReceiveMessage(request); } TEST_F(MoqtSessionTest, TooManySubscribes) { MoqtSessionPeer::set_next_request_id(&session_, kDefaultInitialMaxRequestId - 1); - std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input = - MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_); - EXPECT_CALL(mock_session_, GetStreamById(_)) - .WillRepeatedly(Return(&mock_stream_)); - EXPECT_CALL(mock_stream_, + PrepareRequestStream(bidi_wrapper_); + EXPECT_CALL(mock_bidi_stream_, Writev(ControlMessageOfType(MoqtMessageType::kSubscribe), _)); MessageParameters parameters(SubscribeForTest()); parameters.subscription_filter.emplace(MoqtFilterType::kLargestObject); EXPECT_TRUE(session_.Subscribe(FullTrackName("foo", "bar"), &remote_track_visitor_, parameters)); + webtransport::test::MockStream control_stream; + std::unique_ptr<MoqtBidiStreamTestWrapper> control_wrapper = + MoqtSessionPeer::CreateControlStream(&session_, &control_stream); EXPECT_CALL( - mock_stream_, + control_stream, Writev(ControlMessageOfType(MoqtMessageType::kRequestsBlocked), _)) .Times(1); EXPECT_FALSE(session_.Subscribe(FullTrackName("foo2", "bar2"), @@ -740,11 +807,8 @@ } TEST_F(MoqtSessionTest, SubscribeDuplicateTrackName) { - std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input = - MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_); - EXPECT_CALL(mock_session_, GetStreamById(_)) - .WillRepeatedly(Return(&mock_stream_)); - EXPECT_CALL(mock_stream_, + PrepareRequestStream(bidi_wrapper_); + EXPECT_CALL(mock_bidi_stream_, Writev(ControlMessageOfType(MoqtMessageType::kSubscribe), _)); MessageParameters parameters(SubscribeForTest()); EXPECT_TRUE(session_.Subscribe(FullTrackName("foo", "bar"), @@ -754,9 +818,8 @@ } TEST_F(MoqtSessionTest, SubscribeWithOk) { - std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input = - MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_); - EXPECT_CALL(mock_stream_, + PrepareRequestStream(bidi_wrapper_); + EXPECT_CALL(mock_bidi_stream_, Writev(ControlMessageOfType(MoqtMessageType::kSubscribe), _)); MessageParameters parameters(SubscribeForTest()); session_.Subscribe(FullTrackName("foo", "bar"), &remote_track_visitor_, @@ -775,16 +838,16 @@ EXPECT_EQ(ftn, FullTrackName("foo", "bar")); EXPECT_TRUE(std::holds_alternative<SubscribeOkData>(response)); }); - stream_input->ReceiveMessage(ok); + bidi_wrapper_->ReceiveMessage(ok); } TEST_F(MoqtSessionTest, SubscribeNextGroupWithOk) { - std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input = - MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_); + PrepareRequestStream(bidi_wrapper_); MoqtSubscribe subscribe = DefaultLocalSubscribe(); subscribe.parameters.subscription_filter.emplace( MoqtFilterType::kNextGroupStart); - EXPECT_CALL(mock_stream_, Writev(SerializedControlMessage(subscribe), _)); + EXPECT_CALL(mock_bidi_stream_, + Writev(SerializedControlMessage(subscribe), _)); session_.Subscribe(FullTrackName("foo", "bar"), &remote_track_visitor_, subscribe.parameters); @@ -801,15 +864,12 @@ EXPECT_EQ(ftn, FullTrackName("foo", "bar")); EXPECT_TRUE(std::holds_alternative<SubscribeOkData>(response)); }); - stream_input->ReceiveMessage(ok); + bidi_wrapper_->ReceiveMessage(ok); } TEST_F(MoqtSessionTest, OutgoingSubscribeUpdate) { - std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input = - MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_); - EXPECT_CALL(mock_session_, GetStreamById) - .WillRepeatedly(Return(&mock_stream_)); - EXPECT_CALL(mock_stream_, + PrepareRequestStream(bidi_wrapper_); + EXPECT_CALL(mock_bidi_stream_, Writev(ControlMessageOfType(MoqtMessageType::kSubscribe), _)); MessageParameters parameters(SubscribeForTest()); parameters.subscription_filter.emplace(Location(1, 0), 10); @@ -822,8 +882,8 @@ TrackExtensions(), }; EXPECT_CALL(remote_track_visitor_, OnReply); - stream_input->ReceiveMessage(ok); - EXPECT_CALL(mock_stream_, + bidi_wrapper_->ReceiveMessage(ok); + EXPECT_CALL(mock_bidi_stream_, Writev(ControlMessageOfType(MoqtMessageType::kRequestUpdate), _)); MessageParameters update_parameters; update_parameters.subscription_filter.emplace(Location(2, 1), 9); @@ -836,7 +896,7 @@ ASSERT_TRUE(std::holds_alternative<MessageParameters>(info)); EXPECT_EQ(std::get<MessageParameters>(info), MessageParameters()); })); - stream_input->ReceiveMessage(MoqtRequestOk{ + bidi_wrapper_->ReceiveMessage(MoqtRequestOk{ /*request_id=*/2, MessageParameters(), }); @@ -867,12 +927,11 @@ TEST_F(MoqtSessionTest, MaxRequestIdChangesResponse) { MoqtSessionPeer::set_next_request_id(&session_, kDefaultInitialMaxRequestId); - std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input = - MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_); - EXPECT_CALL(mock_session_, GetStreamById(_)) - .WillRepeatedly(Return(&mock_stream_)); + webtransport::test::MockStream control_stream; + std::unique_ptr<MoqtBidiStreamTestWrapper> control_wrapper = + MoqtSessionPeer::CreateControlStream(&session_, &control_stream); EXPECT_CALL( - mock_stream_, + control_stream, Writev(ControlMessageOfType(MoqtMessageType::kRequestsBlocked), _)); MessageParameters parameters(SubscribeForTest()); parameters.subscription_filter.emplace(MoqtFilterType::kLargestObject); @@ -881,9 +940,10 @@ MoqtMaxRequestId max_request_id = { /*max_request_id=*/kDefaultInitialMaxRequestId + 1, }; - stream_input->ReceiveMessage(max_request_id); + control_wrapper->ReceiveMessage(max_request_id); - EXPECT_CALL(mock_stream_, + PrepareRequestStream(bidi_wrapper_); + EXPECT_CALL(mock_bidi_stream_, Writev(ControlMessageOfType(MoqtMessageType::kSubscribe), _)); EXPECT_TRUE(session_.Subscribe(FullTrackName("foo", "bar"), &remote_track_visitor_, parameters)); @@ -893,32 +953,34 @@ MoqtMaxRequestId max_request_id = { /*max_request_id=*/kDefaultInitialMaxRequestId - 1, }; - std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input = - MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_); + bidi_wrapper_ = + MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_); EXPECT_CALL(mock_session_, CloseSession(static_cast<uint64_t>(MoqtError::kProtocolViolation), "MAX_REQUEST_ID has lower value than previous")) .Times(1); - stream_input->ReceiveMessage(max_request_id); + bidi_wrapper_->ReceiveMessage(max_request_id); } TEST_F(MoqtSessionTest, GrantMoreRequests) { - std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input = - MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_); - EXPECT_CALL(mock_stream_, + webtransport::test::MockStream control_stream; + std::unique_ptr<MoqtBidiStreamTestWrapper> control_wrapper_ = + MoqtSessionPeer::CreateControlStream(&session_, &control_stream); + EXPECT_CALL(control_stream, Writev(ControlMessageOfType(MoqtMessageType::kMaxRequestId), _)); session_.GrantMoreRequests(1); // Peer subscribes to (0, 0) + bidi_wrapper_ = std::make_unique<MoqtBidiStreamTestWrapper>( + ResponseStream(kSubscribeByte)); MoqtSubscribe request = DefaultSubscribe(); request.request_id = kDefaultInitialMaxRequestId + 1; MockTrackPublisher* track = CreateTrackPublisher(); - ReceiveSubscribeSynchronousOk(track, request, stream_input.get()); + ReceiveSubscribeSynchronousOk(track, request, bidi_wrapper_.get()); } TEST_F(MoqtSessionTest, SubscribeWithError) { - std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input = - MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_); - EXPECT_CALL(mock_stream_, + PrepareRequestStream(bidi_wrapper_); + EXPECT_CALL(mock_bidi_stream_, Writev(ControlMessageOfType(MoqtMessageType::kSubscribe), _)); MessageParameters parameters(SubscribeForTest()); parameters.subscription_filter.emplace(MoqtFilterType::kLargestObject); @@ -941,31 +1003,30 @@ std::get<MoqtRequestErrorInfo>(response).reason_phrase == "deadbeef"); }); - stream_input->ReceiveMessage(error); + EXPECT_CALL(mock_bidi_stream_, Writev(testing::IsEmpty(), _)); // FIN. + bidi_wrapper_->ReceiveMessage(error); } TEST_F(MoqtSessionTest, Unsubscribe) { - std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input = - MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_); + PrepareRequestStream(bidi_wrapper_); FullTrackName ftn = FullTrackName("foo", "bar"); - EXPECT_CALL(mock_stream_, + EXPECT_CALL(mock_bidi_stream_, Writev(ControlMessageOfType(MoqtMessageType::kSubscribe), _)); EXPECT_TRUE( session_.Subscribe(ftn, &remote_track_visitor_, MessageParameters())); - EXPECT_CALL(mock_stream_, - Writev(ControlMessageOfType(MoqtMessageType::kUnsubscribe), _)); + EXPECT_CALL(mock_bidi_stream_, ResetWithUserCode); EXPECT_CALL(remote_track_visitor_, OnPublishDone); session_.Unsubscribe(ftn); // Verify it was destroyed. - EXPECT_CALL(mock_stream_, Writev).Times(0); + EXPECT_CALL(mock_bidi_stream_, ResetWithUserCode).Times(0); EXPECT_CALL(remote_track_visitor_, OnPublishDone).Times(0); session_.Unsubscribe(ftn); } TEST_F(MoqtSessionTest, ReplyToPublishNamespaceWithOkThenPublishNamespaceDone) { TrackNamespace track_namespace{"foo"}; - std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input = - MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_); + bidi_wrapper_ = + MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_); MessageParameters parameters; parameters.authorization_tokens.emplace_back(AuthTokenType::kOutOfBand, "foo"); @@ -981,11 +1042,11 @@ MoqtResponseCallback callback) { std::move(callback)(MessageParameters()); }); - EXPECT_CALL(mock_stream_, + EXPECT_CALL(mock_bidi_stream_, Writev(SerializedControlMessage(MoqtRequestOk{ kDefaultPeerRequestId, MessageParameters()}), _)); - stream_input->ReceiveMessage(publish_namespace); + bidi_wrapper_->ReceiveMessage(publish_namespace); MoqtPublishNamespaceDone publish_namespace_done = { /*request_id=*/0, }; @@ -994,15 +1055,15 @@ .WillOnce( [](const TrackNamespace&, const std::optional<MessageParameters>&, MoqtResponseCallback callback) { EXPECT_EQ(callback, nullptr); }); - stream_input->ReceiveMessage(publish_namespace_done); + bidi_wrapper_->ReceiveMessage(publish_namespace_done); } TEST_F(MoqtSessionTest, ReplyToPublishNamespaceWithOkThenPublishNamespaceCancel) { TrackNamespace track_namespace{"foo"}; - std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input = - MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_); + bidi_wrapper_ = + MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_); MessageParameters parameters; parameters.authorization_tokens.emplace_back(AuthTokenType::kOutOfBand, "foo"); @@ -1018,12 +1079,12 @@ MoqtResponseCallback callback) { std::move(callback)(MessageParameters()); }); - EXPECT_CALL(mock_stream_, + EXPECT_CALL(mock_bidi_stream_, Writev(SerializedControlMessage(MoqtRequestOk{ kDefaultPeerRequestId, MessageParameters()}), _)); - stream_input->ReceiveMessage(publish_namespace); - EXPECT_CALL(mock_stream_, + bidi_wrapper_->ReceiveMessage(publish_namespace); + EXPECT_CALL(mock_bidi_stream_, Writev(SerializedControlMessage(MoqtPublishNamespaceCancel{ kDefaultPeerRequestId, RequestErrorCode::kInternalError, "deadbeef"}), @@ -1035,8 +1096,8 @@ TEST_F(MoqtSessionTest, ReplyToPublishNamespaceWithError) { TrackNamespace track_namespace{"foo"}; - std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input = - MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_); + bidi_wrapper_ = + MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_); MessageParameters parameters; parameters.authorization_tokens.emplace_back(AuthTokenType::kOutOfBand, "foo"); @@ -1055,31 +1116,20 @@ .WillOnce( [&](const TrackNamespace&, const std::optional<MessageParameters>&, MoqtResponseCallback callback) { std::move(callback)(error); }); - EXPECT_CALL(mock_stream_, + EXPECT_CALL(mock_bidi_stream_, Writev(SerializedControlMessage(MoqtRequestError{ kDefaultPeerRequestId, error.error_code, error.retry_interval, error.reason_phrase}), _)); - stream_input->ReceiveMessage(publish_namespace); + bidi_wrapper_->ReceiveMessage(publish_namespace); } TEST_F(MoqtSessionTest, SubscribeNamespaceLifeCycle) { TrackNamespace prefix({"foo"}); bool got_callback = false; - EXPECT_CALL(mock_session_, CanOpenNextOutgoingBidirectionalStream()) - .WillOnce(Return(true)); - EXPECT_CALL(mock_session_, OpenOutgoingBidirectionalStream()) - .WillOnce(Return(&mock_stream_)); - std::unique_ptr<MoqtNamespaceSubscriberStream> stream_input; - EXPECT_CALL(mock_stream_, SetVisitor) - .WillOnce([&](std::unique_ptr<webtransport::StreamVisitor> visitor) { - stream_input = absl::WrapUnique( - absl::down_cast<MoqtNamespaceSubscriberStream*>(visitor.release())); - ASSERT_NE(stream_input, nullptr); - }); - EXPECT_CALL(mock_stream_, CanWrite).WillRepeatedly(Return(true)); + PrepareRequestStream(bidi_wrapper_); EXPECT_CALL( - mock_stream_, + mock_bidi_stream_, Writev(ControlMessageOfType(MoqtMessageType::kSubscribeNamespace), _)); std::unique_ptr<MoqtNamespaceTask> task = session_.SubscribeNamespace( prefix, SubscribeNamespaceOption::kNamespace, MessageParameters(), @@ -1088,28 +1138,17 @@ EXPECT_TRUE(std::holds_alternative<MessageParameters>(response)); }); MoqtRequestOk ok = {kDefaultLocalRequestId, MessageParameters()}; - QUICHE_ASSERT_OK(stream_input->OnControlMessage(ok)); + bidi_wrapper_->ReceiveMessage(ok); EXPECT_TRUE(got_callback); - EXPECT_CALL(mock_stream_, ResetWithUserCode); + EXPECT_CALL(mock_bidi_stream_, ResetWithUserCode); } TEST_F(MoqtSessionTest, SubscribeNamespaceError) { TrackNamespace prefix({"foo"}); bool got_callback = false; - EXPECT_CALL(mock_session_, CanOpenNextOutgoingBidirectionalStream()) - .WillOnce(Return(true)); - EXPECT_CALL(mock_session_, OpenOutgoingBidirectionalStream()) - .WillOnce(Return(&mock_stream_)); - std::unique_ptr<MoqtNamespaceSubscriberStream> stream_input; - EXPECT_CALL(mock_stream_, SetVisitor) - .WillOnce([&](std::unique_ptr<webtransport::StreamVisitor> visitor) { - stream_input = std::unique_ptr<MoqtNamespaceSubscriberStream>( - absl::down_cast<MoqtNamespaceSubscriberStream*>(visitor.release())); - ASSERT_NE(stream_input, nullptr); - }); - EXPECT_CALL(mock_stream_, CanWrite).WillRepeatedly(Return(true)); + PrepareRequestStream(bidi_wrapper_); EXPECT_CALL( - mock_stream_, + mock_bidi_stream_, Writev(ControlMessageOfType(MoqtMessageType::kSubscribeNamespace), _)); std::unique_ptr<MoqtNamespaceTask> task = session_.SubscribeNamespace( prefix, SubscribeNamespaceOption::kNamespace, MessageParameters(), @@ -1124,7 +1163,7 @@ MoqtRequestError error = {kDefaultLocalRequestId, RequestErrorCode::kInvalidRange, std::nullopt, "deadbeef"}; - QUICHE_ASSERT_OK(stream_input->OnControlMessage(error)); + bidi_wrapper_->ReceiveMessage(error); EXPECT_TRUE(got_callback); } @@ -1136,17 +1175,8 @@ [&](std::variant<MessageParameters, MoqtRequestErrorInfo>) {}), nullptr); // kBoth is treated as kNamespace. - EXPECT_CALL(mock_session_, CanOpenNextOutgoingBidirectionalStream()) - .WillOnce(Return(true)); - EXPECT_CALL(mock_session_, OpenOutgoingBidirectionalStream()) - .WillOnce(Return(&mock_stream_)); - std::unique_ptr<webtransport::StreamVisitor> stream_visitor; - EXPECT_CALL(mock_stream_, SetVisitor) - .WillOnce([&](std::unique_ptr<webtransport::StreamVisitor> visitor) { - stream_visitor = std::move(visitor); - }); - EXPECT_CALL(mock_stream_, CanWrite).WillRepeatedly(Return(true)); - EXPECT_CALL(mock_stream_, + PrepareRequestStream(bidi_wrapper_); + EXPECT_CALL(mock_bidi_stream_, Writev(SerializedControlMessage(MoqtSubscribeNamespace{ 0, prefix, SubscribeNamespaceOption::kNamespace, MessageParameters()}), @@ -1158,9 +1188,7 @@ } TEST_F(MoqtSessionTest, SubscribeOkWithBadTrackAlias) { - webtransport::test::MockStream mock_control_stream; - std::unique_ptr<MoqtBidiStreamTestWrapper> control_stream = - MoqtSessionPeer::CreateControlStream(&session_, &mock_control_stream); + PrepareRequestStream(bidi_wrapper_); session_.Subscribe(FullTrackName("foo", "bar"), &remote_track_visitor_, MessageParameters()); MoqtSubscribeOk subscribe_ok = { @@ -1169,42 +1197,44 @@ MessageParameters(), TrackExtensions(), }; - control_stream->ReceiveMessage(subscribe_ok); + bidi_wrapper_->ReceiveMessage(subscribe_ok); // Second subscribe, but OK has the same track alias. + webtransport::test::MockStream bidi_stream_2; + std::unique_ptr<MoqtBidiStreamTestWrapper> bidi_wrapper_2 = + MoqtSessionPeer::CreateControlStream(&session_, &bidi_stream_2); + PrepareRequestStream(bidi_wrapper_2, &bidi_stream_2); session_.Subscribe(FullTrackName("foo2", "bar2"), &remote_track_visitor_, MessageParameters()); subscribe_ok.request_id += 2; EXPECT_CALL( mock_session_, CloseSession(static_cast<uint64_t>(MoqtError::kDuplicateTrackAlias), - "Duplicate track alias")); - control_stream->ReceiveMessage(subscribe_ok); + "Track alias already exists")); + bidi_wrapper_2->ReceiveMessage(subscribe_ok); } TEST_F(MoqtSessionTest, ReceiveUnsubscribe) { MockTrackPublisher* track = CreateTrackPublisher(); MoqtSubscribe request = DefaultSubscribe(); const MoqtPriority kLocalDefaultPriority = 0x20; - std::unique_ptr<MoqtBidiStreamTestWrapper> control_stream = - MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_); + bidi_wrapper_ = std::make_unique<MoqtBidiStreamTestWrapper>( + ResponseStream(kSubscribeByte)); TrackExtensions extensions(std::nullopt, std::nullopt, kLocalDefaultPriority, std::nullopt, std::nullopt, std::nullopt); EXPECT_CALL(*track, extensions) .WillRepeatedly(testing::ReturnRef(extensions)); MoqtObjectListener* listener = ReceiveSubscribeSynchronousOk( - track, request, control_stream.get(), /*track_alias=*/0, extensions); - MoqtUnsubscribe unsubscribe = {/*request_id=*/1}; + track, request, bidi_wrapper_.get(), /*track_alias=*/0, extensions); EXPECT_CALL(*track, RemoveObjectListener(listener)); - control_stream->ReceiveMessage(unsubscribe); + bidi_wrapper_->stream().OnResetStreamReceived(kResetCodeCancelled); } TEST_F(MoqtSessionTest, ReceiveDatagram) { FullTrackName ftn("foo", "bar"); const MoqtPriority kPeerDefaultPriority = 0x20; - std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input = - MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_); + PrepareRequestStream(bidi_wrapper_); std::string payload = "deadbeef"; - EXPECT_CALL(mock_stream_, + EXPECT_CALL(mock_bidi_stream_, Writev(ControlMessageOfType(MoqtMessageType::kSubscribe), _)); session_.Subscribe(ftn, &remote_track_visitor_, MessageParameters()); MoqtSubscribeOk ok; @@ -1214,7 +1244,7 @@ TrackExtensions(std::nullopt, std::nullopt, kPeerDefaultPriority, std::nullopt, std::nullopt, std::nullopt); EXPECT_CALL(remote_track_visitor_, OnReply); - stream_input->ReceiveMessage(ok); + bidi_wrapper_->ReceiveMessage(ok); MoqtObject object = { /*track_alias=*/2, @@ -1249,9 +1279,8 @@ TEST_F(MoqtSessionTest, UsePeerDefaultPriority) { FullTrackName ftn("foo", "bar"); const MoqtPriority kPeerDefaultPriority = 0x20; - std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input = - MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_); - EXPECT_CALL(mock_stream_, + PrepareRequestStream(bidi_wrapper_); + EXPECT_CALL(mock_bidi_stream_, Writev(ControlMessageOfType(MoqtMessageType::kSubscribe), _)); session_.Subscribe(ftn, &remote_track_visitor_, MessageParameters()); MoqtSubscribeOk ok; @@ -1261,7 +1290,7 @@ TrackExtensions(std::nullopt, std::nullopt, kPeerDefaultPriority, std::nullopt, std::nullopt, std::nullopt); EXPECT_CALL(remote_track_visitor_, OnReply); - stream_input->ReceiveMessage(ok); + bidi_wrapper_->ReceiveMessage(ok); // Omit priority from a datagram. char datagram[] = {0x0c, 0x02, 0x05, 0x64, 0x65, 0x61, 0x64, 0x62, 0x65, 0x65, 0x66}; @@ -1296,8 +1325,8 @@ TEST_F(MoqtSessionTest, OmitPublisherPriority) { MoqtSubscribe request = DefaultSubscribe(); const MoqtPriority kLocalDefaultPriority = 0x20; - std::unique_ptr<MoqtBidiStreamTestWrapper> control_stream = - MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_); + bidi_wrapper_ = std::make_unique<MoqtBidiStreamTestWrapper>( + ResponseStream(kSubscribeByte)); // Create the publisher and the SUBSCRIBE with kLocalDefaultPriority. MockTrackPublisher* track = CreateTrackPublisher(); std::make_shared<MockTrackPublisher>(request.full_track_name); @@ -1306,25 +1335,25 @@ EXPECT_CALL(*track, extensions) .WillRepeatedly(testing::ReturnRef(extensions)); MoqtObjectListener* listener = ReceiveSubscribeSynchronousOk( - track, request, control_stream.get(), /*track_alias=*/0, extensions); + track, request, bidi_wrapper_.get(), /*track_alias=*/0, extensions); // Deliver an object with kLocalDefaultPriority; stream_type will omit // the priority. EXPECT_CALL(mock_session_, CanOpenNextOutgoingUnidirectionalStream()) .WillOnce(Return(true)); EXPECT_CALL(mock_session_, OpenOutgoingUnidirectionalStream()) - .WillOnce(Return(&mock_stream_)); - EXPECT_CALL(mock_stream_, GetStreamId()).WillRepeatedly(Return(1)); + .WillOnce(Return(&mock_uni_stream_)); + EXPECT_CALL(mock_uni_stream_, GetStreamId()).WillRepeatedly(Return(1)); std::unique_ptr<webtransport::StreamVisitor> stream_visitor; - EXPECT_CALL(mock_stream_, SetVisitor) + EXPECT_CALL(mock_uni_stream_, SetVisitor) .WillOnce([&](std::unique_ptr<webtransport::StreamVisitor> visitor) { stream_visitor = std::move(visitor); }); - EXPECT_CALL(mock_stream_, SetPriority); - EXPECT_CALL(mock_stream_, visitor()).WillRepeatedly([&]() { + EXPECT_CALL(mock_uni_stream_, SetPriority); + EXPECT_CALL(mock_uni_stream_, visitor()).WillRepeatedly([&]() { return stream_visitor.get(); }); - EXPECT_CALL(mock_stream_, CanWrite).WillRepeatedly(Return(true)); + EXPECT_CALL(mock_uni_stream_, CanWrite).WillRepeatedly(Return(true)); EXPECT_CALL(*track, GetCachedObject(_, _, _, _)) .WillOnce(Return(PublishedObject{ PublishedObjectMetadata{ @@ -1332,7 +1361,7 @@ kLocalDefaultPriority, true, 8, MoqtSessionPeer::Now(&session_)}, PayloadFromString("deadbeef")})) .WillOnce(Return(std::nullopt)); - EXPECT_CALL(mock_stream_, Writev) + EXPECT_CALL(mock_uni_stream_, Writev) .WillOnce([&](absl::Span<quiche::QuicheMemSlice> data, const webtransport::StreamWriteOptions& options) { // The stream type omits the priority. @@ -1380,9 +1409,8 @@ TEST_F(MoqtSessionTest, DatagramOutOfWindow) { FullTrackName ftn("foo", "bar"); const MoqtPriority kPeerDefaultPriority = 0x20; - std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input = - MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_); - EXPECT_CALL(mock_stream_, + PrepareRequestStream(bidi_wrapper_); + EXPECT_CALL(mock_bidi_stream_, Writev(ControlMessageOfType(MoqtMessageType::kSubscribe), _)); MessageParameters params; params.subscription_filter.emplace(Location(1, 0)); @@ -1394,7 +1422,7 @@ TrackExtensions(std::nullopt, std::nullopt, kPeerDefaultPriority, std::nullopt, std::nullopt, std::nullopt); EXPECT_CALL(remote_track_visitor_, OnReply); - stream_input->ReceiveMessage(ok); + bidi_wrapper_->ReceiveMessage(ok); char datagram[] = {0x01, 0x02, 0x00, 0x00, 0x80, 0x00, 0x08, 0x64, 0x65, 0x61, 0x64, 0x62, 0x65, 0x65, 0x66}; EXPECT_CALL(remote_track_visitor_, OnObjectFragment).Times(0); @@ -1517,8 +1545,8 @@ // All callbacks are called asynchronously. TEST_F(MoqtSessionTest, ProcessFetchGetEverythingFromUpstream) { - std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input = - MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_); + bidi_wrapper_ = + MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_); MoqtFetch fetch = DefaultFetch(); MockTrackPublisher* track = CreateTrackPublisher(); @@ -1527,14 +1555,15 @@ MockFetchTask* fetch_task = fetch_task_ptr.get(); EXPECT_CALL(*track, StandaloneFetch) .WillOnce(Return(std::move(fetch_task_ptr))); - stream_input->ReceiveMessage(fetch); + bidi_wrapper_->ReceiveMessage(fetch); // Compose and send the FETCH_OK. MoqtFetchOk expected_ok; expected_ok.request_id = fetch.request_id; expected_ok.end_of_track = false; expected_ok.end_location = Location(1, 4); - EXPECT_CALL(mock_stream_, Writev(SerializedControlMessage(expected_ok), _)); + EXPECT_CALL(mock_bidi_stream_, + Writev(SerializedControlMessage(expected_ok), _)); fetch_task->CallFetchResponseCallback(expected_ok); // Data arrives. webtransport::test::MockStream data_stream; @@ -1549,8 +1578,8 @@ // All callbacks are called synchronously. All relevant data is cached (or this // is the original publisher). TEST_F(MoqtSessionTest, ProcessFetchWholeRangeIsPresent) { - std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input = - MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_); + bidi_wrapper_ = + MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_); MoqtFetch fetch = DefaultFetch(); MockTrackPublisher* track = CreateTrackPublisher(); @@ -1563,7 +1592,8 @@ MockFetchTask* fetch_task = fetch_task_ptr.get(); EXPECT_CALL(*track, StandaloneFetch) .WillOnce(Return(std::move(fetch_task_ptr))); - EXPECT_CALL(mock_stream_, Writev(SerializedControlMessage(expected_ok), _)); + EXPECT_CALL(mock_bidi_stream_, + Writev(SerializedControlMessage(expected_ok), _)); webtransport::test::MockStream data_stream; std::unique_ptr<webtransport::StreamVisitor> stream_visitor; ExpectStreamOpen(mock_session_, fetch_task, data_stream, stream_visitor); @@ -1572,13 +1602,13 @@ MoqtFetchTask::GetNextObjectResult::kPending); // Everything spins upon message receipt. FetchTask is generating the // necessary callbacks. - stream_input->ReceiveMessage(fetch); + bidi_wrapper_->ReceiveMessage(fetch); } TEST_F(MoqtSessionTest, SendFragmentedFetchObject) { using ::testing::ByMove; - std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input = - MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_); + bidi_wrapper_ = + MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_); MoqtFetch fetch = DefaultFetch(); fetch.request_id = 3; // Use an odd ID for peer request in client session. MockTrackPublisher* track = CreateTrackPublisher(); @@ -1591,27 +1621,27 @@ .WillOnce(Return(ByMove(std::move(fetch_task_ptr)))); // Receive FETCH, send FETCH_OK. - stream_input->ReceiveMessage(fetch); + bidi_wrapper_->ReceiveMessage(fetch); // FETCH_OK responding to the request. MoqtFetchOk expected_ok; expected_ok.request_id = fetch.request_id; expected_ok.end_of_track = false; expected_ok.end_location = Location(1, 0); - EXPECT_CALL(mock_stream_, Writev(SerializedControlMessage(expected_ok), _)); + EXPECT_CALL(mock_bidi_stream_, + Writev(SerializedControlMessage(expected_ok), _)); fetch_task->CallFetchResponseCallback(expected_ok); - webtransport::test::MockStream data_stream; std::unique_ptr<webtransport::StreamVisitor> stream_visitor; EXPECT_CALL(mock_session_, CanOpenNextOutgoingUnidirectionalStream) .WillRepeatedly(Return(true)); EXPECT_CALL(mock_session_, OpenOutgoingUnidirectionalStream()) - .WillOnce(Return(&data_stream)); - EXPECT_CALL(data_stream, SetVisitor) + .WillOnce(Return(&mock_uni_stream_)); + EXPECT_CALL(mock_uni_stream_, SetVisitor) .WillOnce([&](std::unique_ptr<webtransport::StreamVisitor> visitor) { stream_visitor = std::move(visitor); }); - EXPECT_CALL(data_stream, SetPriority); - EXPECT_CALL(data_stream, CanWrite).WillRepeatedly(Return(true)); + EXPECT_CALL(mock_uni_stream_, SetPriority); + EXPECT_CALL(mock_uni_stream_, CanWrite).WillRepeatedly(Return(true)); // Trigger stream opening (calls SetObjectAvailableCallback with lambda1). // Setting the stream visitor will cause a second call to the callback. PublishedObjectMetadata metadata = { @@ -1623,7 +1653,7 @@ return MoqtFetchTask::GetNextObjectResult::kSuccess; }) .WillOnce(Return(MoqtFetchTask::GetNextObjectResult::kPending)); - EXPECT_CALL(data_stream, Writev) + EXPECT_CALL(mock_uni_stream_, Writev) .WillOnce([&](absl::Span<const quiche::QuicheMemSlice> data, const webtransport::StreamWriteOptions& options) { EXPECT_EQ(data.size(), 2); @@ -1631,7 +1661,7 @@ return absl::OkStatus(); }); fetch_task->CallObjectsAvailableCallback(); - // lambda1 ran, data_stream captured, stream_visitor set. + // lambda1 ran, mock_uni_stream_ captured, stream_visitor set. ASSERT_NE(stream_visitor, nullptr); // The second fragment is available. @@ -1642,7 +1672,7 @@ return MoqtFetchTask::GetNextObjectResult::kSuccess; }) .WillRepeatedly(Return(MoqtFetchTask::GetNextObjectResult::kPending)); - EXPECT_CALL(data_stream, Writev) + EXPECT_CALL(mock_uni_stream_, Writev) .WillOnce([&](absl::Span<const quiche::QuicheMemSlice> data, const webtransport::StreamWriteOptions& options) { EXPECT_EQ(data.size(), 1); // No header. @@ -1655,8 +1685,8 @@ // The publisher has the first object locally, but has to go upstream to get // the rest. TEST_F(MoqtSessionTest, FetchReturnsObjectBeforeOk) { - std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input = - MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_); + bidi_wrapper_ = + MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_); MoqtFetch fetch = DefaultFetch(); MockTrackPublisher* track = CreateTrackPublisher(); @@ -1672,19 +1702,20 @@ ExpectSendObject(fetch_task, data_stream, MoqtObjectStatus::kNormal, Location(0, 0), "foo", MoqtFetchTask::GetNextObjectResult::kPending); - stream_input->ReceiveMessage(fetch); + bidi_wrapper_->ReceiveMessage(fetch); MoqtFetchOk expected_ok; expected_ok.request_id = fetch.request_id; expected_ok.end_of_track = false; expected_ok.end_location = Location(1, 4); - EXPECT_CALL(mock_stream_, Writev(SerializedControlMessage(expected_ok), _)); + EXPECT_CALL(mock_bidi_stream_, + Writev(SerializedControlMessage(expected_ok), _)); fetch_task->CallFetchResponseCallback(expected_ok); } TEST_F(MoqtSessionTest, FetchReturnsObjectBeforeError) { - std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input = - MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_); + bidi_wrapper_ = + MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_); MoqtFetch fetch = DefaultFetch(); MockTrackPublisher* track = CreateTrackPublisher(); @@ -1699,34 +1730,33 @@ ExpectSendObject(fetch_task, data_stream, MoqtObjectStatus::kNormal, Location(0, 0), "foo", MoqtFetchTask::GetNextObjectResult::kPending); - stream_input->ReceiveMessage(fetch); + bidi_wrapper_->ReceiveMessage(fetch); MoqtRequestError expected_error{ fetch.request_id, RequestErrorCode::kDoesNotExist, std::nullopt, "foo"}; - EXPECT_CALL(mock_stream_, + EXPECT_CALL(mock_bidi_stream_, Writev(SerializedControlMessage(expected_error), _)); fetch_task->CallFetchResponseCallback(expected_error); } TEST_F(MoqtSessionTest, InvalidFetch) { - webtransport::test::MockStream control_stream; - std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input = - MoqtSessionPeer::CreateControlStream(&session_, &control_stream); + bidi_wrapper_ = + MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_); MockTrackPublisher* track = CreateTrackPublisher(); MoqtFetch fetch = DefaultFetch(); EXPECT_CALL(*track, StandaloneFetch) .WillOnce(Return(std::make_unique<MockFetchTask>())); - stream_input->ReceiveMessage(fetch); + bidi_wrapper_->ReceiveMessage(fetch); EXPECT_CALL(mock_session_, CloseSession(static_cast<uint64_t>(MoqtError::kInvalidRequestId), "Duplicate request ID")) .Times(1); - stream_input->ReceiveMessage(fetch); + bidi_wrapper_->ReceiveMessage(fetch); } TEST_F(MoqtSessionTest, FetchFails) { - std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input = - MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_); + bidi_wrapper_ = + MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_); MoqtFetch fetch = DefaultFetch(); MockTrackPublisher* track = CreateTrackPublisher(); @@ -1736,14 +1766,14 @@ .WillOnce(Return(std::move(fetch_task_ptr))); EXPECT_CALL(*fetch_task, GetStatus()) .WillRepeatedly(Return(absl::Status(absl::StatusCode::kInternal, "foo"))); - EXPECT_CALL(mock_stream_, + EXPECT_CALL(mock_bidi_stream_, Writev(ControlMessageOfType(MoqtMessageType::kRequestError), _)); - stream_input->ReceiveMessage(fetch); + bidi_wrapper_->ReceiveMessage(fetch); } TEST_F(MoqtSessionTest, FullFetchDeliveryWithFlowControl) { - std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input = - MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_); + bidi_wrapper_ = + MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_); MoqtFetch fetch = DefaultFetch(); MockTrackPublisher* track = CreateTrackPublisher(); @@ -1753,7 +1783,7 @@ EXPECT_CALL(*track, StandaloneFetch) .WillOnce(Return(std::move(fetch_task_ptr))); - stream_input->ReceiveMessage(fetch); + bidi_wrapper_->ReceiveMessage(fetch); EXPECT_CALL(mock_session_, CanOpenNextOutgoingUnidirectionalStream()) .WillOnce(Return(false)); fetch_task->CallObjectsAvailableCallback(); @@ -1776,12 +1806,15 @@ // Give it the latest object filter. subscribe.parameters.subscription_filter.emplace( MoqtFilterType::kLargestObject); - std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input = - MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_); + bidi_wrapper_ = std::make_unique<MoqtBidiStreamTestWrapper>( + ResponseStream(kSubscribeByte)); MockTrackPublisher* track = CreateTrackPublisher(); SetLargestId(track, Location(4, 10)); - ReceiveSubscribeSynchronousOk(track, subscribe, stream_input.get()); + ReceiveSubscribeSynchronousOk(track, subscribe, bidi_wrapper_.get()); + webtransport::test::MockStream control_stream; + std::unique_ptr<MoqtBidiStreamTestWrapper> control_wrapper = + MoqtSessionPeer::CreateControlStream(&session_, &control_stream); ASSERT_TRUE(MoqtSessionPeer::RequestIdIsSubscriptionPublisher( &session_, subscribe.request_id)); MoqtFetch fetch = DefaultFetch(); @@ -1789,7 +1822,7 @@ fetch.fetch = JoiningFetchRelative(1, 2); EXPECT_CALL(*track, StandaloneFetch(Location(2, 0), Location(4, 10), _)) .WillOnce(Return(std::make_unique<MockFetchTask>())); - stream_input->ReceiveMessage(fetch); + control_wrapper->ReceiveMessage(fetch); } TEST_F(MoqtSessionTest, IncomingAbsoluteJoiningFetch) { @@ -1797,25 +1830,28 @@ // Give it the latest object filter. subscribe.parameters.subscription_filter.emplace( MoqtFilterType::kLargestObject); - std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input = - MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_); + bidi_wrapper_ = std::make_unique<MoqtBidiStreamTestWrapper>( + ResponseStream(kSubscribeByte)); MockTrackPublisher* track = CreateTrackPublisher(); SetLargestId(track, Location(4, 10)); - ReceiveSubscribeSynchronousOk(track, subscribe, stream_input.get()); + ReceiveSubscribeSynchronousOk(track, subscribe, bidi_wrapper_.get()); ASSERT_TRUE(MoqtSessionPeer::RequestIdIsSubscriptionPublisher( &session_, subscribe.request_id)); + webtransport::test::MockStream control_stream; + std::unique_ptr<MoqtBidiStreamTestWrapper> control_wrapper = + MoqtSessionPeer::CreateControlStream(&session_, &control_stream); MoqtFetch fetch = DefaultFetch(); fetch.request_id = 3; fetch.fetch = JoiningFetchAbsolute(1, 2); EXPECT_CALL(*track, StandaloneFetch(Location(2, 0), Location(4, 10), _)) .WillOnce(Return(std::make_unique<MockFetchTask>())); - stream_input->ReceiveMessage(fetch); + control_wrapper->ReceiveMessage(fetch); } TEST_F(MoqtSessionTest, IncomingJoiningFetchBadRequestId) { - std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input = - MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_); + bidi_wrapper_ = + MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_); MoqtFetch fetch = DefaultFetch(); fetch.fetch = JoiningFetchRelative(1, 2); MoqtRequestError expected_error = { @@ -1824,20 +1860,23 @@ /*retry_interval=*/std::nullopt, "Joining Fetch for non-existent request", }; - EXPECT_CALL(mock_stream_, + EXPECT_CALL(mock_bidi_stream_, Writev(SerializedControlMessage(expected_error), _)); - stream_input->ReceiveMessage(fetch); + bidi_wrapper_->ReceiveMessage(fetch); } TEST_F(MoqtSessionTest, IncomingJoiningFetchForwardZero) { MoqtSubscribe subscribe = DefaultSubscribe(); subscribe.parameters.set_forward(false); - std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input = - MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_); + bidi_wrapper_ = std::make_unique<MoqtBidiStreamTestWrapper>( + ResponseStream(kSubscribeByte)); MockTrackPublisher* track = CreateTrackPublisher(); SetLargestId(track, Location(2, 10)); - ReceiveSubscribeSynchronousOk(track, subscribe, stream_input.get()); + ReceiveSubscribeSynchronousOk(track, subscribe, bidi_wrapper_.get()); + webtransport::test::MockStream control_stream; + std::unique_ptr<MoqtBidiStreamTestWrapper> control_wrapper = + MoqtSessionPeer::CreateControlStream(&session_, &control_stream); MoqtFetch fetch = DefaultFetch(); fetch.request_id = 3; fetch.fetch = JoiningFetchRelative(1, 2); @@ -1845,14 +1884,14 @@ CloseSession(static_cast<uint64_t>(MoqtError::kProtocolViolation), "Joining Fetch for non-forwarding subscribe")) .Times(1); - stream_input->ReceiveMessage(fetch); + control_wrapper->ReceiveMessage(fetch); } TEST_F(MoqtSessionTest, SendJoiningFetch) { - std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input = - MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_); - EXPECT_CALL(mock_session_, GetStreamById(_)) - .WillRepeatedly(Return(&mock_stream_)); + PrepareRequestStream(bidi_wrapper_); + webtransport::test::MockStream control_stream; + std::unique_ptr<MoqtBidiStreamTestWrapper> control_wrapper = + MoqtSessionPeer::CreateControlStream(&session_, &control_stream); MoqtSubscribe expected_subscribe( 0, FullTrackName("foo", "bar"), MessageParameters(MoqtFilterType::kLargestObject)); @@ -1861,9 +1900,9 @@ /*fetch=*/JoiningFetchRelative(0, 1), MessageParameters(), }; - EXPECT_CALL(mock_stream_, + EXPECT_CALL(mock_bidi_stream_, Writev(SerializedControlMessage(expected_subscribe), _)); - EXPECT_CALL(mock_stream_, + EXPECT_CALL(control_stream, Writev(SerializedControlMessage(expected_fetch), _)); EXPECT_TRUE(session_.RelativeJoiningFetch(expected_subscribe.full_track_name, &remote_track_visitor_, nullptr, 1, @@ -1871,13 +1910,13 @@ } TEST_F(MoqtSessionTest, SendJoiningFetchNoFlowControl) { - std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input = - MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_); - EXPECT_CALL(mock_session_, GetStreamById(_)) - .WillRepeatedly(Return(&mock_stream_)); - EXPECT_CALL(mock_stream_, + PrepareRequestStream(bidi_wrapper_); + webtransport::test::MockStream control_stream; + std::unique_ptr<MoqtBidiStreamTestWrapper> control_wrapper = + MoqtSessionPeer::CreateControlStream(&session_, &control_stream); + EXPECT_CALL(mock_bidi_stream_, Writev(ControlMessageOfType(MoqtMessageType::kSubscribe), _)); - EXPECT_CALL(mock_stream_, + EXPECT_CALL(control_stream, Writev(ControlMessageOfType(MoqtMessageType::kFetch), _)); EXPECT_TRUE(session_.RelativeJoiningFetch(FullTrackName("foo", "bar"), &remote_track_visitor_, 0, @@ -1886,9 +1925,9 @@ EXPECT_CALL(remote_track_visitor_, OnReply).Times(1); MessageParameters parameters; parameters.largest_object = Location(2, 0); - stream_input->ReceiveMessage( + bidi_wrapper_->ReceiveMessage( MoqtSubscribeOk(0, 2, parameters, TrackExtensions())); - stream_input->ReceiveMessage(MoqtFetchOk( + control_wrapper->ReceiveMessage(MoqtFetchOk( 2, false, Location(2, 0), MessageParameters(), TrackExtensions())); // Packet arrives on FETCH stream. MoqtObject object = { @@ -1912,7 +1951,7 @@ data_stream.Receive(header.AsStringView(), false); EXPECT_CALL(remote_track_visitor_, OnObjectFragment).Times(1); // Last object of the FETCH causes FETCH_CANCEL. - EXPECT_CALL(mock_stream_, + EXPECT_CALL(control_stream, Writev(ControlMessageOfType(MoqtMessageType::kFetchCancel), _)); data_stream.Receive("foo", false); } @@ -1922,17 +1961,9 @@ MessageParameters parameters; parameters.authorization_tokens.emplace_back(AuthTokenType::kOutOfBand, "foo"); - auto bidi_stream = - std::make_unique<webtransport::test::InMemoryStreamWithWriteBuffer>(4); - MoqtFramer framer(true, quic::Perspective::IS_SERVER); + MoqtSubscribeNamespace subscribe_namespace = { /*request_id=*/1, prefix, SubscribeNamespaceOption::kBoth, parameters}; - bidi_stream->Receive( - framer.SerializeSubscribeNamespace(subscribe_namespace).AsStringView(), - /*fin=*/false); - EXPECT_CALL(mock_session_, AcceptIncomingBidirectionalStream()) - .WillOnce(Return(bidi_stream.get())) - .WillOnce(Return(nullptr)); quiche::QuicheWeakPtr<MockNamespaceTask> task; MoqtRequestOk expected_ok(/*request_id=*/1); expected_ok.parameters.expires = quic::QuicTimeDelta::FromSeconds(60); @@ -1946,10 +1977,12 @@ task = task_ptr->GetWeakPtr(); return task_ptr; }); - session_.OnIncomingBidirectionalStreamAvailable(); - EXPECT_EQ(bidi_stream->write_buffer(), - framer.SerializeRequestOk(expected_ok).AsStringView()); - bidi_stream->write_buffer().clear(); + bidi_wrapper_ = std::make_unique<MoqtBidiStreamTestWrapper>( + ResponseStream(kSubscribeNamespaceByte)); + EXPECT_CALL(mock_bidi_stream_, + Writev(SerializedControlMessage(expected_ok), _)) + .WillOnce(Return(absl::OkStatus())); + bidi_wrapper_->ReceiveMessage(subscribe_namespace); // Deliver a NAMESPACE ASSERT_TRUE(task.IsValid()); @@ -1960,14 +1993,16 @@ return GetNextResult::kSuccess; }) .WillOnce(Return(GetNextResult::kPending)); + MoqtNamespace expected_namespace = { + TrackNamespace({"bar"}), + }; + EXPECT_CALL(mock_bidi_stream_, + Writev(SerializedControlMessage(expected_namespace), _)) + .WillOnce(Return(absl::OkStatus())); task.GetIfAvailable()->InvokeCallback(); - char expected_data[] = {0x08, 0x00, 0x05, 0x01, 0x03, 'b', 'a', 'r'}; - absl::string_view expected_data_view(expected_data, sizeof(expected_data)); - EXPECT_EQ(expected_data_view, - bidi_stream->write_buffer().substr(0, expected_data_view.length())); // Unsubscribe - bidi_stream.reset(); + bidi_wrapper_.reset(); EXPECT_FALSE(task.IsValid()); } @@ -1976,16 +2011,10 @@ MessageParameters parameters; parameters.authorization_tokens.emplace_back(AuthTokenType::kOutOfBand, "foo"); - webtransport::test::InMemoryStreamWithWriteBuffer bidi_stream(4); - MoqtFramer framer(true, quic::Perspective::IS_SERVER); + bidi_wrapper_ = std::make_unique<MoqtBidiStreamTestWrapper>( + ResponseStream(kSubscribeNamespaceByte)); MoqtSubscribeNamespace subscribe_namespace = { /*request_id=*/1, prefix, SubscribeNamespaceOption::kBoth, parameters}; - bidi_stream.Receive( - framer.SerializeSubscribeNamespace(subscribe_namespace).AsStringView(), - /*fin=*/false); - EXPECT_CALL(mock_session_, AcceptIncomingBidirectionalStream()) - .WillOnce(Return(&bidi_stream)) - .WillOnce(Return(nullptr)); EXPECT_CALL(session_callbacks_.incoming_subscribe_namespace_callback, Call(prefix, SubscribeNamespaceOption::kBoth, parameters, _)) .WillOnce([&](const TrackNamespace&, SubscribeNamespaceOption, @@ -1995,10 +2024,14 @@ RequestErrorCode::kUnauthorized, std::nullopt, "foo"}); return nullptr; }); - session_.OnIncomingBidirectionalStreamAvailable(); - EXPECT_EQ(PeekControlMessageType(bidi_stream.write_buffer()), - MoqtMessageType::kRequestError); - EXPECT_TRUE(bidi_stream.fin_sent()); + EXPECT_CALL(mock_bidi_stream_, + Writev(ControlMessageOfType(MoqtMessageType::kRequestError), _)) + .WillOnce([](absl::Span<quiche::QuicheMemSlice>, + const webtransport::StreamWriteOptions& options) { + EXPECT_TRUE(options.send_fin()); + return absl::OkStatus(); + }); + bidi_wrapper_->ReceiveMessage(subscribe_namespace); } TEST_F(MoqtSessionTest, IncomingSubscribeNamespaceWithPrefixOverlap) { @@ -2006,17 +2039,10 @@ MessageParameters parameters; parameters.authorization_tokens.emplace_back(AuthTokenType::kOutOfBand, "foo"); - webtransport::test::InMemoryStreamWithWriteBuffer bidi_stream1(4), - bidi_stream2(8); - MoqtFramer framer(true, quic::Perspective::IS_SERVER); + bidi_wrapper_ = std::make_unique<MoqtBidiStreamTestWrapper>( + ResponseStream(kSubscribeNamespaceByte)); MoqtSubscribeNamespace subscribe_namespace = { /*request_id=*/1, foo, SubscribeNamespaceOption::kBoth, parameters}; - bidi_stream1.Receive( - framer.SerializeSubscribeNamespace(subscribe_namespace).AsStringView(), - /*fin=*/false); - EXPECT_CALL(mock_session_, AcceptIncomingBidirectionalStream()) - .WillOnce(Return(&bidi_stream1)) - .WillOnce(Return(nullptr)); EXPECT_CALL(session_callbacks_.incoming_subscribe_namespace_callback, Call(foo, SubscribeNamespaceOption::kBoth, parameters, _)) .WillOnce([&](const TrackNamespace& prefix, SubscribeNamespaceOption, @@ -2026,27 +2052,29 @@ auto task_ptr = std::make_unique<MockNamespaceTask>(prefix); return task_ptr; }); - session_.OnIncomingBidirectionalStreamAvailable(); - EXPECT_EQ(PeekControlMessageType(bidi_stream1.write_buffer()), - MoqtMessageType::kRequestOk); + EXPECT_CALL(mock_bidi_stream_, + Writev(ControlMessageOfType(MoqtMessageType::kRequestOk), _)) + .WillOnce(Return(absl::OkStatus())); + bidi_wrapper_->ReceiveMessage(subscribe_namespace); subscribe_namespace.request_id += 2; subscribe_namespace.track_namespace_prefix = foobar; - bidi_stream2.Receive( - framer.SerializeSubscribeNamespace(subscribe_namespace).AsStringView(), - /*fin=*/false); - EXPECT_CALL(mock_session_, AcceptIncomingBidirectionalStream()) - .WillOnce(Return(&bidi_stream2)) - .WillOnce(Return(nullptr)); - session_.OnIncomingBidirectionalStreamAvailable(); - EXPECT_EQ(PeekControlMessageType(bidi_stream2.write_buffer()), - MoqtMessageType::kRequestError); - EXPECT_TRUE(bidi_stream2.fin_sent()); + webtransport::test::MockStream bidi_stream_2; + auto bidi_wrapper_2 = std::make_unique<MoqtBidiStreamTestWrapper>( + ResponseStream(kSubscribeNamespaceByte, &bidi_stream_2)); + EXPECT_CALL(bidi_stream_2, + Writev(ControlMessageOfType(MoqtMessageType::kRequestError), _)) + .WillOnce([](absl::Span<quiche::QuicheMemSlice>, + const webtransport::StreamWriteOptions& options) { + EXPECT_TRUE(options.send_fin()); + return absl::OkStatus(); + }); + bidi_wrapper_2->ReceiveMessage(subscribe_namespace); } TEST_F(MoqtSessionTest, FetchThenOkThenCancel) { - std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input = - MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_); + bidi_wrapper_ = + MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_); std::unique_ptr<MoqtFetchTask> fetch_task; session_.Fetch( FullTrackName("foo", "bar"), @@ -2059,21 +2087,21 @@ /*end_of_track=*/false, Location(3, 25), MessageParameters(), TrackExtensions(), }; - stream_input->ReceiveMessage(ok); + bidi_wrapper_->ReceiveMessage(ok); ASSERT_NE(fetch_task, nullptr); EXPECT_TRUE(fetch_task->GetStatus().ok()); PublishedObject object; EXPECT_EQ(fetch_task->GetNextObject(object), MoqtFetchTask::GetNextObjectResult::kPending); // Cancel the fetch. - EXPECT_CALL(mock_stream_, + EXPECT_CALL(mock_bidi_stream_, Writev(ControlMessageOfType(MoqtMessageType::kFetchCancel), _)); fetch_task.reset(); } TEST_F(MoqtSessionTest, FetchThenError) { - std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input = - MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_); + bidi_wrapper_ = + MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_); std::unique_ptr<MoqtFetchTask> fetch_task; session_.Fetch( FullTrackName("foo", "bar"), @@ -2087,7 +2115,7 @@ /*retry_interval=*/std::nullopt, "No username provided", }; - stream_input->ReceiveMessage(error); + bidi_wrapper_->ReceiveMessage(error); ASSERT_NE(fetch_task, nullptr); EXPECT_TRUE(absl::IsPermissionDenied(fetch_task->GetStatus())); EXPECT_EQ(fetch_task->GetStatus().message(), "No username provided"); @@ -2095,8 +2123,8 @@ // The application takes objects as they arrive. TEST_F(MoqtSessionTest, IncomingFetchObjectsGreedyApp) { - std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input = - MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_); + bidi_wrapper_ = + MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_); std::unique_ptr<MoqtFetchTask> fetch_task; uint64_t expected_object_id = 0; session_.Fetch( @@ -2165,7 +2193,7 @@ MessageParameters(), TrackExtensions(), }; - stream_input->ReceiveMessage(ok); + bidi_wrapper_->ReceiveMessage(ok); ASSERT_NE(fetch_task, nullptr); EXPECT_EQ(expected_object_id, 2); @@ -2180,8 +2208,8 @@ } TEST_F(MoqtSessionTest, IncomingFetchObjectsSlowApp) { - std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input = - MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_); + bidi_wrapper_ = + MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_); std::unique_ptr<MoqtFetchTask> fetch_task; uint64_t expected_object_id = 0; bool objects_available = false; @@ -2237,7 +2265,7 @@ /*end_of_track=*/false, Location(3, 25), MessageParameters(), TrackExtensions(), }; - stream_input->ReceiveMessage(ok); + bidi_wrapper_->ReceiveMessage(ok); ASSERT_NE(fetch_task, nullptr); EXPECT_TRUE(objects_available); @@ -2278,10 +2306,10 @@ TEST_F(MoqtSessionTest, DeliveryTimeoutParameter) { MoqtSubscribe request = DefaultSubscribe(); request.parameters.delivery_timeout = quic::QuicTimeDelta::FromSeconds(1); - std::unique_ptr<MoqtBidiStreamTestWrapper> control_stream = - MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_); + bidi_wrapper_ = std::make_unique<MoqtBidiStreamTestWrapper>( + ResponseStream(kSubscribeByte)); MockTrackPublisher* track = CreateTrackPublisher(); - ReceiveSubscribeSynchronousOk(track, request, control_stream.get()); + ReceiveSubscribeSynchronousOk(track, request, bidi_wrapper_.get()); std::optional<quic::QuicTimeDelta> delivery_timeout = MoqtSessionPeer::GetDeliveryTimeout(&session_, request.request_id); EXPECT_TRUE(delivery_timeout.has_value() && @@ -2289,12 +2317,12 @@ } TEST_F(MoqtSessionTest, ReceiveGoAwayEnforcement) { - std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input = - MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_); + bidi_wrapper_ = + MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_); EXPECT_CALL(session_callbacks_.goaway_received_callback, Call("foo")); - stream_input->ReceiveMessage(MoqtGoAway("foo")); + bidi_wrapper_->ReceiveMessage(MoqtGoAway("foo")); // New requests not allowed. - EXPECT_CALL(mock_stream_, Writev).Times(0); + EXPECT_CALL(mock_bidi_stream_, Writev).Times(0); MessageParameters parameters = SubscribeForTest(); parameters.subscription_filter.emplace(MoqtFilterType::kLargestObject); EXPECT_FALSE(session_.Subscribe(FullTrackName("foo", "bar"), @@ -2324,50 +2352,52 @@ reported_error = true; EXPECT_EQ(error_message, "Received multiple GOAWAY messages"); }); - stream_input->ReceiveMessage(MoqtGoAway("foo")); + bidi_wrapper_->ReceiveMessage(MoqtGoAway("foo")); } TEST_F(MoqtSessionTest, SendGoAwayEnforcement) { - std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input = - MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_); + bidi_wrapper_ = + MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_); CreateTrackPublisher(); - EXPECT_CALL(mock_stream_, + EXPECT_CALL(mock_bidi_stream_, Writev(ControlMessageOfType(MoqtMessageType::kGoAway), _)); session_.GoAway(""); - EXPECT_CALL(mock_stream_, + + EXPECT_CALL(mock_bidi_stream_, Writev(ControlMessageOfType(MoqtMessageType::kRequestError), _)); - stream_input->ReceiveMessage(DefaultSubscribe()); - EXPECT_CALL(mock_stream_, - Writev(ControlMessageOfType(MoqtMessageType::kRequestError), _)); - stream_input->ReceiveMessage( + bidi_wrapper_->ReceiveMessage( MoqtPublishNamespace(3, TrackNamespace({"foo"}), MessageParameters())); - EXPECT_CALL(mock_stream_, + EXPECT_CALL(mock_bidi_stream_, Writev(ControlMessageOfType(MoqtMessageType::kRequestError), _)); MoqtFetch fetch = DefaultFetch(); fetch.request_id = 5; - stream_input->ReceiveMessage(fetch); + bidi_wrapper_->ReceiveMessage(fetch); - MoqtFramer framer(true, quic::Perspective::IS_CLIENT); - SessionNamespaceTree tree; - MoqtIncomingSubscribeNamespaceCallback callback = - DefaultIncomingSubscribeNamespaceCallback; - MoqtNamespacePublisherStream namespace_stream( - &framer, - MoqtControlMessageParser(kDefaultMoqtVersion, true, - quic::Perspective::IS_CLIENT), - [](const TrackNamespace&) { return true; }, nullptr, nullptr, callback); - namespace_stream.BindStream(&mock_stream_); - EXPECT_CALL(mock_stream_, - Writev(ControlMessageOfType(MoqtMessageType::kRequestError), _)); - QUICHE_ASSERT_OK( - namespace_stream.OnControlMessage(MoqtSubscribeNamespace(7))); - MoqtTrackStatus track_status = DefaultSubscribe(); - track_status.request_id = 7; - EXPECT_CALL(mock_stream_, - Writev(ControlMessageOfType(MoqtMessageType::kRequestError), _)); - stream_input->ReceiveMessage(track_status); - // Block all outgoing SUBSCRIBE, PUBLISH_NAMESPACE, GOAWAY,etc. - EXPECT_CALL(mock_stream_, Writev).Times(0); + // All new bidi streams types are immediately rejected. + webtransport::test::MockStream new_request_stream; + EXPECT_CALL(mock_session_, AcceptIncomingBidirectionalStream) + .WillOnce(Return(&new_request_stream)) + .WillOnce(Return(nullptr)); + EXPECT_CALL(new_request_stream, CanWrite()).WillOnce(Return(true)); + EXPECT_CALL(new_request_stream, + Writev(ControlMessageOfType(MoqtMessageType::kRequestError), _)) + .WillOnce([](absl::Span<quiche::QuicheMemSlice>, + const webtransport::StreamWriteOptions& options) { + EXPECT_TRUE(options.send_fin()); + return absl::OkStatus(); + }); + session_.OnIncomingBidirectionalStreamAvailable(); + + // If a new stream can't be written, reset it. + webtransport::test::MockStream new_request_stream_2; + EXPECT_CALL(mock_session_, AcceptIncomingBidirectionalStream) + .WillOnce(Return(&new_request_stream_2)) + .WillOnce(Return(nullptr)); + EXPECT_CALL(new_request_stream_2, CanWrite()).WillOnce(Return(false)); + EXPECT_CALL(new_request_stream_2, ResetWithUserCode); + session_.OnIncomingBidirectionalStreamAvailable(); + + // Block all outgoing PUBLISH_NAMESPACE, GOAWAY,etc. MessageParameters parameters = SubscribeForTest(); parameters.subscription_filter.emplace(MoqtFilterType::kLargestObject); EXPECT_FALSE(session_.Subscribe(FullTrackName({"foo"}, "bar"), @@ -2400,10 +2430,10 @@ TEST_F(MoqtSessionTest, ClientCannotSendNewSessionUri) { // session_ is a client session. - std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input = - MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_); + bidi_wrapper_ = + MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_); // Client GOAWAY not sent. - EXPECT_CALL(mock_stream_, Writev).Times(0); + EXPECT_CALL(mock_bidi_stream_, Writev).Times(0); session_.GoAway("foo"); } @@ -2413,8 +2443,8 @@ MoqtSessionParameters(quic::Perspective::IS_SERVER), std::make_unique<quic::test::TestAlarmFactory>(), session_callbacks_.AsSessionCallbacks()); - std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input = - MoqtSessionPeer::CreateControlStream(&session, &mock_stream_); + bidi_wrapper_ = + MoqtSessionPeer::CreateControlStream(&session, &mock_bidi_stream_); EXPECT_CALL( mock_session, CloseSession(static_cast<uint64_t>(MoqtError::kProtocolViolation), @@ -2427,14 +2457,13 @@ EXPECT_EQ(error_message, "Received GOAWAY with new_session_uri on the server"); }); - stream_input->ReceiveMessage(MoqtGoAway("foo")); + bidi_wrapper_->ReceiveMessage(MoqtGoAway("foo")); EXPECT_TRUE(reported_error); } TEST_F(MoqtSessionTest, IncomingTrackStatusThenSynchronousOk) { - webtransport::test::MockStream control_stream; - std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input = - MoqtSessionPeer::CreateControlStream(&session_, &control_stream); + bidi_wrapper_ = + MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_); auto* track = CreateTrackPublisher(); MoqtTrackStatus track_status = DefaultSubscribe(); @@ -2450,25 +2479,24 @@ expected_ok.parameters.expires = quic::QuicTimeDelta::FromMilliseconds(10000); expected_ok.parameters.largest_object = Location(5, 30); - EXPECT_CALL(control_stream, + EXPECT_CALL(mock_bidi_stream_, Writev(SerializedControlMessage(expected_ok), _)); EXPECT_CALL(*track, RemoveObjectListener); listener->OnSubscribeAccepted(); }); - stream_input->ReceiveMessage(track_status); + bidi_wrapper_->ReceiveMessage(track_status); } TEST_F(MoqtSessionTest, IncomingTrackStatusThenAsynchronousOk) { - webtransport::test::MockStream control_stream; - std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input = - MoqtSessionPeer::CreateControlStream(&session_, &control_stream); + bidi_wrapper_ = + MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_); auto* track = CreateTrackPublisher(); MoqtTrackStatus track_status = DefaultSubscribe(); MoqtObjectListener* listener = nullptr; EXPECT_CALL(*track, AddObjectListener) .WillOnce(testing::SaveArg<0>(&listener)); - stream_input->ReceiveMessage(track_status); + bidi_wrapper_->ReceiveMessage(track_status); ASSERT_NE(listener, nullptr); EXPECT_CALL(*track, expiration) .WillRepeatedly(Return(quic::QuicTimeDelta::FromMilliseconds(10000))); @@ -2477,15 +2505,15 @@ expected_ok.request_id = track_status.request_id; expected_ok.parameters.expires = quic::QuicTimeDelta::FromMilliseconds(10000); expected_ok.parameters.largest_object = Location(5, 30); - EXPECT_CALL(control_stream, Writev(SerializedControlMessage(expected_ok), _)); + EXPECT_CALL(mock_bidi_stream_, + Writev(SerializedControlMessage(expected_ok), _)); EXPECT_CALL(*track, RemoveObjectListener(listener)); listener->OnSubscribeAccepted(); } TEST_F(MoqtSessionTest, IncomingTrackStatusThenSynchronousError) { - webtransport::test::MockStream control_stream; - std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input = - MoqtSessionPeer::CreateControlStream(&session_, &control_stream); + bidi_wrapper_ = + MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_); auto* track = CreateTrackPublisher(); MoqtTrackStatus track_status = DefaultSubscribe(); @@ -2493,30 +2521,29 @@ EXPECT_CALL(*track, AddObjectListener) .WillOnce([&](MoqtObjectListener* listener) { EXPECT_CALL( - control_stream, + mock_bidi_stream_, Writev(ControlMessageOfType(MoqtMessageType::kRequestError), _)); EXPECT_CALL(*track, RemoveObjectListener); listener->OnSubscribeRejected(MoqtRequestErrorInfo( RequestErrorCode::kInternalError, std::nullopt, "Test error")); executed_AddObjectListener = true; }); - stream_input->ReceiveMessage(track_status); + bidi_wrapper_->ReceiveMessage(track_status); EXPECT_TRUE(executed_AddObjectListener); } TEST_F(MoqtSessionTest, IncomingTrackStatusThenAsynchronousError) { - webtransport::test::MockStream control_stream; - std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input = - MoqtSessionPeer::CreateControlStream(&session_, &control_stream); + bidi_wrapper_ = + MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_); auto* track = CreateTrackPublisher(); MoqtTrackStatus track_status = DefaultSubscribe(); MoqtObjectListener* listener; EXPECT_CALL(*track, AddObjectListener) .WillOnce(testing::SaveArg<0>(&listener)); - stream_input->ReceiveMessage(track_status); + bidi_wrapper_->ReceiveMessage(track_status); ASSERT_NE(listener, nullptr); - EXPECT_CALL(control_stream, + EXPECT_CALL(mock_bidi_stream_, Writev(ControlMessageOfType(MoqtMessageType::kRequestError), _)); EXPECT_CALL(*track, RemoveObjectListener(listener)); listener->OnSubscribeRejected(MoqtRequestErrorInfo( @@ -2524,11 +2551,8 @@ } TEST_F(MoqtSessionTest, FinReportedToVisitor) { - std::unique_ptr<MoqtBidiStreamTestWrapper> control_stream = - MoqtSessionPeer::CreateControlStream(&session_, &control_stream_); - EXPECT_CALL(mock_session_, GetStreamById) - .WillRepeatedly(Return(&control_stream_)); - EXPECT_CALL(control_stream_, + PrepareRequestStream(bidi_wrapper_); + EXPECT_CALL(mock_bidi_stream_, Writev(ControlMessageOfType(MoqtMessageType::kSubscribe), _)); MessageParameters parameters = SubscribeForTest(); parameters.subscription_filter.emplace(MoqtFilterType::kLargestObject); @@ -2543,7 +2567,7 @@ EXPECT_EQ(ftn, FullTrackName("foo", "bar")); EXPECT_TRUE(std::holds_alternative<SubscribeOkData>(response)); }); - control_stream->ReceiveMessage(ok); + bidi_wrapper_->ReceiveMessage(ok); MoqtObject object = { /*track_alias=*/2, /*group_id=*/0, @@ -2555,13 +2579,13 @@ /*first_object_in_subgroup=*/true, /*payload_length=*/0, }; - EXPECT_CALL(mock_stream_, GetStreamId()) + EXPECT_CALL(mock_uni_stream_, GetStreamId()) .WillRepeatedly(Return(kIncomingUniStreamId)); EXPECT_CALL(mock_session_, GetStreamById(kIncomingUniStreamId)) - .WillRepeatedly(Return(&mock_stream_)); + .WillRepeatedly(Return(&mock_uni_stream_)); std::unique_ptr<webtransport::StreamVisitor> data_stream; - DeliverObject(object, /*fin=*/true, mock_session_, &mock_stream_, data_stream, - &remote_track_visitor_); + DeliverObject(object, /*fin=*/true, mock_session_, &mock_uni_stream_, + data_stream, &remote_track_visitor_); // The data stream died and destroyed the visitor (IncomingDataStream). EXPECT_CALL(remote_track_visitor_, OnStreamFin(FullTrackName("foo", "bar"), DataStreamIndex(0, 0))); @@ -2569,11 +2593,8 @@ } TEST_F(MoqtSessionTest, ResetReportedToVisitor) { - std::unique_ptr<MoqtBidiStreamTestWrapper> control_stream = - MoqtSessionPeer::CreateControlStream(&session_, &control_stream_); - EXPECT_CALL(mock_session_, GetStreamById) - .WillRepeatedly(Return(&control_stream_)); - EXPECT_CALL(control_stream_, + PrepareRequestStream(bidi_wrapper_); + EXPECT_CALL(mock_bidi_stream_, Writev(ControlMessageOfType(MoqtMessageType::kSubscribe), _)); MessageParameters parameters = SubscribeForTest(); parameters.subscription_filter.emplace(MoqtFilterType::kLargestObject); @@ -2588,7 +2609,7 @@ EXPECT_EQ(ftn, FullTrackName("foo", "bar")); EXPECT_TRUE(std::holds_alternative<SubscribeOkData>(response)); }); - control_stream->ReceiveMessage(ok); + bidi_wrapper_->ReceiveMessage(ok); MoqtObject object = { /*track_alias=*/2, /*group_id=*/0, @@ -2600,12 +2621,12 @@ /*first_object_in_subgroup=*/true, /*payload_length=*/0, }; - EXPECT_CALL(mock_stream_, GetStreamId()) + EXPECT_CALL(mock_uni_stream_, GetStreamId()) .WillRepeatedly(Return(kIncomingUniStreamId)); EXPECT_CALL(mock_session_, GetStreamById(kIncomingUniStreamId)) - .WillRepeatedly(Return(&mock_stream_)); + .WillRepeatedly(Return(&mock_uni_stream_)); std::unique_ptr<webtransport::StreamVisitor> data_stream; - DeliverObject(object, /*fin=*/false, mock_session_, &mock_stream_, + DeliverObject(object, /*fin=*/false, mock_session_, &mock_uni_stream_, data_stream, &remote_track_visitor_); // The data stream died and destroyed the visitor (IncomingDataStream). data_stream->OnResetStreamReceived(kResetCodeCancelled); @@ -2615,9 +2636,8 @@ } TEST_F(MoqtSessionTest, IncomingPublishNamespaceCleanup) { - webtransport::test::MockStream control_stream; - std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input = - MoqtSessionPeer::CreateControlStream(&session_, &control_stream); + bidi_wrapper_ = + MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_); // Register two incoming PUBLISH_NAMESPACE. MoqtPublishNamespace publish_namespace{ /*request_id=*/1, TrackNamespace{"foo"}, MessageParameters()}; @@ -2630,8 +2650,9 @@ MoqtResponseCallback callback) { std::move(callback)(expected_ok.parameters); }); - EXPECT_CALL(control_stream, Writev(SerializedControlMessage(expected_ok), _)); - stream_input->ReceiveMessage(publish_namespace); + EXPECT_CALL(mock_bidi_stream_, + Writev(SerializedControlMessage(expected_ok), _)); + bidi_wrapper_->ReceiveMessage(publish_namespace); publish_namespace = MoqtPublishNamespace( /*request_id=*/3, TrackNamespace{"bar"}, MessageParameters()); @@ -2642,9 +2663,9 @@ MoqtResponseCallback callback) { std::move(callback)(MessageParameters()); }); - EXPECT_CALL(control_stream, + EXPECT_CALL(mock_bidi_stream_, Writev(ControlMessageOfType(MoqtMessageType::kRequestOk), _)); - stream_input->ReceiveMessage(publish_namespace); + bidi_wrapper_->ReceiveMessage(publish_namespace); // Revoke "bar" MoqtPublishNamespaceDone done{/*request_id=*/3}; @@ -2654,7 +2675,7 @@ .WillOnce( [](const TrackNamespace&, const std::optional<MessageParameters>&, MoqtResponseCallback callback) { EXPECT_EQ(callback, nullptr); }); - stream_input->ReceiveMessage(done); + bidi_wrapper_->ReceiveMessage(done); // Destroying the session should revoke "foo". EXPECT_CALL( @@ -2683,85 +2704,77 @@ session_.OnSessionReady(); } -TEST_F(MoqtSessionTest, SubscribeThenRequestOk) { - webtransport::test::MockStream control_stream; - std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input = - MoqtSessionPeer::CreateControlStream(&session_, &control_stream); - MessageParameters parameters = SubscribeForTest(); - parameters.subscription_filter.emplace(MoqtFilterType::kLargestObject); - session_.Subscribe(FullTrackName("foo", "bar"), &remote_track_visitor_, - parameters); - EXPECT_CALL(mock_session_, CloseSession); - EXPECT_CALL(session_callbacks_.session_terminated_callback, Call); - stream_input->ReceiveMessage(MoqtRequestOk{0, MessageParameters()}); -} TEST_F(MoqtSessionTest, ClientSetupNotAllowedOnControlStream) { // While technically on the Control stream, when it arrives, it's an // UnknownBidiStream - std::unique_ptr<MoqtBidiStreamTestWrapper> control_stream = - MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_); + bidi_wrapper_ = + MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_); EXPECT_CALL(mock_session_, CloseSession); EXPECT_CALL(session_callbacks_.session_terminated_callback, Call); - control_stream->ReceiveMessage( + bidi_wrapper_->ReceiveMessage( MoqtSetup(SetupParameters("/", "example.com", 0))); } TEST_F(MoqtSessionTest, NamespaceNotAllowedOnControlStream) { - std::unique_ptr<MoqtBidiStreamTestWrapper> control_stream = - MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_); + bidi_wrapper_ = + MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_); EXPECT_CALL(mock_session_, CloseSession); EXPECT_CALL(session_callbacks_.session_terminated_callback, Call); - control_stream->ReceiveMessage(MoqtNamespace()); + bidi_wrapper_->ReceiveMessage(MoqtNamespace()); } TEST_F(MoqtSessionTest, NamespaceDoneNotAllowedOnControlStream) { - std::unique_ptr<MoqtBidiStreamTestWrapper> control_stream = - MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_); + bidi_wrapper_ = + MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_); EXPECT_CALL(mock_session_, CloseSession); EXPECT_CALL(session_callbacks_.session_terminated_callback, Call); - control_stream->ReceiveMessage(MoqtNamespaceDone()); + bidi_wrapper_->ReceiveMessage(MoqtNamespaceDone()); } TEST_F(MoqtSessionTest, IncomingRequestUpdateTriggersRequestOk) { MoqtSubscribe subscribe = DefaultSubscribe(); MockTrackPublisher* track = CreateTrackPublisher(); - std::unique_ptr<MoqtBidiStreamTestWrapper> control_stream = - MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_); - ReceiveSubscribeSynchronousOk(track, subscribe, control_stream.get(), 0); - EXPECT_CALL(mock_stream_, + bidi_wrapper_ = std::make_unique<MoqtBidiStreamTestWrapper>( + ResponseStream(kSubscribeByte)); + ReceiveSubscribeSynchronousOk(track, subscribe, bidi_wrapper_.get(), 0); + EXPECT_CALL(mock_bidi_stream_, Writev(ControlMessageOfType(MoqtMessageType::kRequestOk), _)); - control_stream->ReceiveMessage(MoqtRequestUpdate{3, 1, MessageParameters()}); + bidi_wrapper_->ReceiveMessage(MoqtRequestUpdate{3, 1, MessageParameters()}); } TEST_F(MoqtSessionTest, IncomingRequestUpdateTriggersRequestError) { - std::unique_ptr<MoqtBidiStreamTestWrapper> control_stream = - MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_); - EXPECT_CALL(mock_stream_, + bidi_wrapper_ = + MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_); + EXPECT_CALL(mock_bidi_stream_, Writev(ControlMessageOfType(MoqtMessageType::kRequestError), _)); - control_stream->ReceiveMessage(MoqtRequestUpdate{3, 1, MessageParameters()}); + bidi_wrapper_->ReceiveMessage(MoqtRequestUpdate{3, 1, MessageParameters()}); } TEST_F(MoqtSessionTest, StopSendingBlocksSubgroup) { MoqtSubscribe subscribe = DefaultSubscribe(); MockTrackPublisher* track = CreateTrackPublisher(); - std::unique_ptr<MoqtBidiStreamTestWrapper> control_stream = - MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_); + bidi_wrapper_ = std::make_unique<MoqtBidiStreamTestWrapper>( + ResponseStream(kSubscribeByte)); MoqtObjectListener* listener = - ReceiveSubscribeSynchronousOk(track, subscribe, control_stream.get(), 0); + ReceiveSubscribeSynchronousOk(track, subscribe, bidi_wrapper_.get(), 0); EXPECT_CALL(mock_session_, CanOpenNextOutgoingUnidirectionalStream) .WillOnce(Return(true)); EXPECT_CALL(mock_session_, OpenOutgoingUnidirectionalStream) - .WillOnce(Return(&mock_stream_)); + .WillOnce(Return(&mock_uni_stream_)); std::unique_ptr<webtransport::StreamVisitor> data_stream_visitor; - EXPECT_CALL(mock_stream_, SetVisitor) + EXPECT_CALL(mock_uni_stream_, SetVisitor) .WillOnce([&](std::unique_ptr<webtransport::StreamVisitor> visitor) { data_stream_visitor = std::move(visitor); }); - EXPECT_CALL(mock_stream_, visitor).WillRepeatedly([&]() { + EXPECT_CALL(mock_uni_stream_, GetStreamId()) + .WillRepeatedly(Return(kOutgoingUniStreamId)); + EXPECT_CALL(mock_session_, GetStreamById(kOutgoingUniStreamId)) + .WillRepeatedly(Return(&mock_uni_stream_)); + EXPECT_CALL(mock_uni_stream_, visitor).WillRepeatedly([&]() { return data_stream_visitor.get(); }); - EXPECT_CALL(mock_stream_, CanWrite).WillRepeatedly(Return(true)); + EXPECT_CALL(mock_uni_stream_, CanWrite).WillRepeatedly(Return(true)); EXPECT_CALL(*track, GetCachedObject(0, Optional(1), 0, 0)) .WillOnce(Return(PublishedObject{ PublishedObjectMetadata{Location(0, 0), 1, "", @@ -2771,14 +2784,14 @@ EXPECT_CALL(*track, GetCachedObject(0, Optional(1), 1, 0)) .WillOnce(Return(std::nullopt)); SetLargestId(track, Location(0, 0)); - EXPECT_CALL(mock_stream_, Writev).WillOnce(Return(absl::OkStatus())); + EXPECT_CALL(mock_uni_stream_, Writev).WillOnce(Return(absl::OkStatus())); listener->OnNewObjectAvailable(Location(0, 0), 1, 0x80); - EXPECT_CALL(mock_stream_, ResetWithUserCode(kResetCodeCancelled)); + EXPECT_CALL(mock_uni_stream_, ResetWithUserCode(kResetCodeCancelled)); data_stream_visitor->OnStopSendingReceived(kResetCodeCancelled); // New object in the same subgroup should not be sent. EXPECT_CALL(*track, GetCachedObject).Times(0); - EXPECT_CALL(mock_stream_, Writev).Times(0); + EXPECT_CALL(mock_uni_stream_, Writev).Times(0); listener->OnNewObjectAvailable(Location(0, 1), 1, 0x80); } @@ -2786,21 +2799,10 @@ CreateTrackPublisher(); std::shared_ptr<MoqtTrackPublisher> track_publisher = publisher_.GetTrack(kDefaultTrackName()); - webtransport::test::MockStream mock_publish_stream; - EXPECT_CALL(mock_session_, CanOpenNextOutgoingBidirectionalStream()) - .WillOnce(Return(true)); - EXPECT_CALL(mock_session_, OpenOutgoingBidirectionalStream()) - .WillOnce(Return(&mock_publish_stream)); - std::unique_ptr<webtransport::StreamVisitor> publish_stream_visitor; - EXPECT_CALL(mock_publish_stream, SetVisitor) - .WillOnce([&](std::unique_ptr<webtransport::StreamVisitor> visitor) { - publish_stream_visitor = std::move(visitor); - }); - EXPECT_CALL(mock_publish_stream, CanWrite).WillRepeatedly(Return(true)); - + PrepareRequestStream(bidi_wrapper_); // Verify PUBLISH message is sent on the publish stream. - EXPECT_CALL(mock_publish_stream, + EXPECT_CALL(mock_bidi_stream_, Writev(ControlMessageOfType(MoqtMessageType::kPublish), _)) .WillOnce(Return(absl::OkStatus())); @@ -2810,17 +2812,11 @@ [&](std::variant<MessageParameters, MoqtRequestErrorInfo> resp) { response = resp; })); - ASSERT_NE(publish_stream_visitor, nullptr); - - std::unique_ptr<MoqtBidiStreamBase> bidi_stream( - absl::down_cast<MoqtBidiStreamBase*>(publish_stream_visitor.release())); - MoqtBidiStreamTestWrapper wrapper(std::move(bidi_stream)); MoqtRequestOk request_ok; request_ok.request_id = 0; request_ok.parameters.delivery_timeout = quic::QuicTimeDelta::FromSeconds(2); - - wrapper.ReceiveMessage(request_ok); + bidi_wrapper_->ReceiveMessage(request_ok); ASSERT_TRUE(response.has_value()); EXPECT_TRUE(std::holds_alternative<MessageParameters>(*response)); @@ -2840,11 +2836,11 @@ } TEST_F(MoqtSessionTest, PublishAfterGoaway) { - std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input = - MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_); + bidi_wrapper_ = + MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_); MoqtGoAway goaway; goaway.new_session_uri = ""; - stream_input->ReceiveMessage(goaway); + bidi_wrapper_->ReceiveMessage(goaway); CreateTrackPublisher(); std::shared_ptr<MoqtTrackPublisher> track_publisher = publisher_.GetTrack(kDefaultTrackName()); @@ -2854,10 +2850,8 @@ } TEST_F(MoqtSessionTest, IncomingPublishAbortsPendingSubscribe) { - // 1. Start a pending SUBSCRIBE. - std::unique_ptr<MoqtBidiStreamTestWrapper> control_stream = - MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_); - EXPECT_CALL(mock_stream_, + PrepareRequestStream(bidi_wrapper_); + EXPECT_CALL(mock_bidi_stream_, Writev(ControlMessageOfType(MoqtMessageType::kSubscribe), _)); MessageParameters parameters(SubscribeForTest()); session_.Subscribe(kDefaultTrackName(), &remote_track_visitor_, parameters); @@ -2873,110 +2867,24 @@ }; // Prepare PUBLISH message. - MoqtPublish publish; - publish.request_id = 0; // Matches pending SUBSCRIBE request ID - publish.full_track_name = kDefaultTrackName(); - publish.track_alias = 10; - publish.parameters.delivery_timeout = quic::QuicTimeDelta::FromSeconds(5); - - MoqtFramer peer_framer(true, quic::Perspective::IS_SERVER); - quiche::QuicheBuffer serialized_buffer = - peer_framer.SerializePublish(publish); - std::string serialized(serialized_buffer.data(), serialized_buffer.size()); - - // Setup mock_publish_stream. - webtransport::test::MockStream mock_publish_stream; - EXPECT_CALL(mock_session_, AcceptIncomingBidirectionalStream()) - .WillOnce(Return(&mock_publish_stream)) - .WillOnce(Return(nullptr)); - - // Setup mock_publish_stream to return serialized data. - size_t data_read = 0; - auto peek_lambda = [&data_read, - &serialized]() -> webtransport::Stream::PeekResult { - webtransport::Stream::PeekResult result; - result.peeked_data = absl::string_view(serialized.data() + data_read, - serialized.size() - data_read); - result.fin_next = (data_read == serialized.size()); - result.all_data_received = false; - return result; - }; - std::function<webtransport::Stream::PeekResult()> peek_action = peek_lambda; - - auto readable_bytes_lambda = [&data_read, &serialized]() -> size_t { - if (data_read >= serialized.size()) { - return 0; - } - return serialized.size() - data_read; - }; - std::function<size_t()> readable_bytes_action = readable_bytes_lambda; - - auto read_lambda = - [&data_read, &serialized]( - absl::Span<char> bytes_to_read) -> webtransport::Stream::ReadResult { - size_t read_size = - std::min(bytes_to_read.size(), serialized.size() - data_read); - memcpy(bytes_to_read.data(), serialized.data() + data_read, read_size); - data_read += read_size; - webtransport::Stream::ReadResult result; - result.bytes_read = read_size; - result.fin = (data_read == serialized.size()); - return result; - }; - std::function<webtransport::Stream::ReadResult(absl::Span<char>)> - read_action = read_lambda; - - auto skip_lambda = [&data_read, &serialized](size_t bytes) -> bool { - data_read += bytes; - return data_read == serialized.size(); - }; - std::function<bool(size_t)> skip_action = skip_lambda; - - EXPECT_CALL(mock_publish_stream, PeekNextReadableRegion()) - .WillRepeatedly(peek_action); - EXPECT_CALL(mock_publish_stream, ReadableBytes()) - .WillRepeatedly(readable_bytes_action); - EXPECT_CALL(mock_publish_stream, Read(testing::An<absl::Span<char>>())) - .WillRepeatedly(read_action); - EXPECT_CALL(mock_publish_stream, SkipBytes).WillRepeatedly(skip_action); - - // Capture SetVisitor calls and mock visitor(). - std::unique_ptr<webtransport::StreamVisitor> unknown_bidi_stream_visitor; - std::unique_ptr<webtransport::StreamVisitor> upgraded_visitor; - - { - testing::InSequence seq; - EXPECT_CALL(mock_publish_stream, SetVisitor) - .WillOnce([&](std::unique_ptr<webtransport::StreamVisitor> visitor) { - unknown_bidi_stream_visitor = std::move(visitor); - }) - .WillOnce([&](std::unique_ptr<webtransport::StreamVisitor> visitor) { - upgraded_visitor = std::move(visitor); - }); - } - - EXPECT_CALL(mock_publish_stream, visitor()) - .WillRepeatedly([&]() -> webtransport::StreamVisitor* { - return upgraded_visitor ? upgraded_visitor.get() - : unknown_bidi_stream_visitor.get(); - }); - - // The RemoteTrackVisitor is being reused, not destroyed. - EXPECT_CALL(remote_track_visitor_, - OnReply(kDefaultTrackName(), - testing::VariantWith<SubscribeOkData>( - testing::Field(&SubscribeOkData::parameters, - testing::Eq(publish.parameters))))); - EXPECT_CALL(mock_stream_, - Writev(ControlMessageOfType(MoqtMessageType::kUnsubscribe), _)) - .WillOnce(Return(absl::OkStatus())); - EXPECT_CALL(remote_track_visitor_, OnPublishDone(kDefaultTrackName())) - .Times(0); + MoqtPublish publish{1, kDefaultTrackName(), 10, MessageParameters(), + TrackExtensions()}; + webtransport::test::MockStream publish_stream; + std::unique_ptr<MoqtBidiStreamTestWrapper> publish_wrapper = + std::make_unique<MoqtBidiStreamTestWrapper>( + ResponseStream(kPublishByte, &publish_stream)); + EXPECT_CALL(mock_bidi_stream_, ResetWithUserCode(kResetCodeCancelled)); + MoqtRequestOk expected_request_ok; + expected_request_ok.request_id = publish.request_id; + expected_request_ok.parameters = parameters; // params from the SUBSCRIBE. + // group_order can be in SUBSCRIBE but not REQUEST_OK. + expected_request_ok.parameters.group_order = std::nullopt; + EXPECT_CALL(publish_stream, + Writev(SerializedControlMessage(expected_request_ok), _)); + // remote_track_visitor_ is reused, not destroyed. + EXPECT_CALL(remote_track_visitor_, OnReply); + publish_wrapper->ReceiveMessage(publish); EXPECT_FALSE(incoming_publish_callback_called); - - // Trigger the read by making incoming stream available. - session_.OnIncomingBidirectionalStreamAvailable(); - // Verify it was aborted immediately (not at teardown). EXPECT_TRUE( testing::Mock::VerifyAndClearExpectations(&remote_track_visitor_)); @@ -3048,22 +2956,27 @@ control_stream.write_buffer().clear(); // 4. Subscribe + webtransport::test::InMemoryStreamWithWriteBuffer sub_stream(5); + EXPECT_CALL(mock_session_, CanOpenNextOutgoingBidirectionalStream) + .WillOnce(Return(true)); + EXPECT_CALL(mock_session_, OpenOutgoingBidirectionalStream) + .WillOnce(Return(&sub_stream)); FullTrackName track_name1("namespace2", "track1"); bool s1 = session_.Subscribe(track_name1, &remote_track_visitor_, MessageParameters()); EXPECT_TRUE(s1); - EXPECT_EQ(get_request_id(control_stream), next_request_id); + EXPECT_EQ(get_request_id(sub_stream), next_request_id); next_request_id += 2; - control_stream.write_buffer().clear(); + sub_stream.write_buffer().clear(); // 5. SubscribeUpdate bool s_update = session_.SubscribeUpdate( track_name1, MessageParameters(), [](std::variant<MessageParameters, MoqtRequestErrorInfo>) {}); EXPECT_TRUE(s_update); - EXPECT_EQ(get_request_id(control_stream), next_request_id); + EXPECT_EQ(get_request_id(sub_stream), next_request_id); next_request_id += 2; - control_stream.write_buffer().clear(); + sub_stream.write_buffer().clear(); // 6. Fetch FullTrackName fetch_track("namespace2", "fetch_track");
diff --git a/quiche/quic/moqt/moqt_subscribe_stream.cc b/quiche/quic/moqt/moqt_subscribe_stream.cc new file mode 100644 index 0000000..473de89 --- /dev/null +++ b/quiche/quic/moqt/moqt_subscribe_stream.cc
@@ -0,0 +1,224 @@ +// Copyright (c) 2026 The Chromium Authors. All rights reserved. +// Use of this source code is governed by a BSD-style license that can be +// found in the LICENSE file. + +#include "quiche/quic/moqt/moqt_subscribe_stream.h" + +#include <cstdint> +#include <memory> +#include <optional> +#include <utility> +#include <variant> + +#include "absl/base/nullability.h" +#include "absl/status/status.h" +#include "quiche/quic/core/quic_alarm_factory.h" +#include "quiche/quic/core/quic_clock.h" +#include "quiche/quic/moqt/moqt_bidi_stream.h" +#include "quiche/quic/moqt/moqt_error.h" +#include "quiche/quic/moqt/moqt_framer.h" +#include "quiche/quic/moqt/moqt_key_value_pair.h" +#include "quiche/quic/moqt/moqt_messages.h" +#include "quiche/quic/moqt/moqt_names.h" +#include "quiche/quic/moqt/moqt_parser.h" +#include "quiche/quic/moqt/moqt_publisher.h" +#include "quiche/quic/moqt/moqt_session_callbacks.h" +#include "quiche/quic/moqt/moqt_subscription.h" +#include "quiche/quic/moqt/moqt_track.h" +#include "quiche/common/quiche_weak_ptr.h" + +namespace moqt { + +MoqtSubscribeRequestStream::MoqtSubscribeRequestStream( + MoqtFramer* absl_nonnull framer, + const MoqtControlMessageParser& message_parser, uint64_t request_id, + SessionErrorCallback session_error_callback, const FullTrackName& name, + SubscribeVisitor* absl_nonnull visitor, const MessageParameters& parameters, + SubscribeRemoteTrack::AddCallback add_callback, + SubscribeRemoteTrack::RemoveCallback remove_callback, + const quic::QuicClock* absl_nonnull clock, + quic::QuicAlarmFactory* absl_nonnull alarm_factory) + : MoqtBidiStreamBase(framer, message_parser, + std::move(session_error_callback)), + track_(std::make_unique<SubscribeRemoteTrack>( + MoqtSubscribe{request_id, name, parameters}, visitor, this)), + add_callback_(std::move(add_callback)), + remove_callback_(std::move(remove_callback)), + clock_(clock), + alarm_factory_(alarm_factory) {} + +void MoqtSubscribeRequestStream::OnStreamBound() { + stream_parser()->set_allow_fin(true); + SendOrBufferMessageOrFatal(framer()->SerializeSubscribe( + MoqtSubscribe{track_->request_id(), track_->full_track_name(), + track_->const_parameters()})); +} + +absl::Status MoqtSubscribeRequestStream::OnRawControlMessage( + const MoqtRawControlMessage& message) { + return ControlMessageDispatcher::DispatchControlMessage( + *this, message_parser(), message, "subscribe request"); +} + +absl::Status MoqtSubscribeRequestStream::OnControlMessage( + const MoqtSubscribeOk& message) { + if (message.request_id != track_->request_id()) { + return absl::InvalidArgumentError("SUBSCRIBE_OK request ID mismatch"); + } + if (add_callback_ == nullptr) { + return absl::InvalidArgumentError( + "Multiple SUBSCRIBE_OK on the same stream"); + } + track_->set_track_alias(message.track_alias); + if (!std::move(add_callback_)(track_.get())) { + add_callback_ = nullptr; + OnFatalError(absl::AlreadyExistsError("Track alias already exists")); + return absl::OkStatus(); + } + add_callback_ = nullptr; + + track_->OnObjectOrOk(SubscribeOkData(message.parameters, message.extensions)); + return absl::OkStatus(); +} + +absl::Status MoqtSubscribeRequestStream::OnControlMessage( + const MoqtRequestOk& message) { + if (!track_->track_alias().has_value()) { + // Not yet established. + OnFatalError( + absl::InvalidArgumentError("REQUEST_OK received before SUBSCRIBE_OK")); + return absl::OkStatus(); + } + auto status_or_params = PopParameters(); + if (status_or_params.ok()) { + MessageParameters parameters = status_or_params.value(); + // EXPIRES or LARGEST_OBJECT could be present in REQUEST_OK. + if (message.parameters.largest_object.has_value()) { + parameters.largest_object = message.parameters.largest_object; + } + if (message.parameters.expires.has_value()) { + parameters.expires = message.parameters.expires; + } + track_->Update(parameters); + } + return MoqtBidiStreamBase::OnControlMessage(message); +} + +absl::Status MoqtSubscribeRequestStream::OnControlMessage( + const MoqtRequestError& message) { + MoqtRequestErrorInfo error_info{message.error_code, message.retry_interval, + message.reason_phrase}; + if (track_->ErrorIsAllowed()) { + if (track_->visitor() != nullptr) { + track_->visitor()->OnReply(track_->full_track_name(), error_info); + } + Fin(); + return absl::OkStatus(); + } + // In response to REQUEST_UPDATE, utilize the ResponseCallback and do not + // update parameters. + return MoqtBidiStreamBase::OnControlMessage(message); +} + +absl::Status MoqtSubscribeRequestStream::OnControlMessage( + const MoqtPublishDone& message) { + if (track_ == nullptr) { + // PUBLISH_DONE can be sent before the subscriber rejects the track. + return absl::OkStatus(); + } + track_->OnPublishDone(message.stream_count, clock_, alarm_factory_); + return absl::OkStatus(); +} + +void MoqtSubscribeRequestStream::Detach() { + if (remove_callback_ != nullptr) { + SubscribeRemoteTrack::RemoveCallback remove_callback = + std::move(remove_callback_); + remove_callback_ = nullptr; + std::move(remove_callback)(track_.get()); + } + track_ = nullptr; +} + +MoqtSubscribeResponseStream::MoqtSubscribeResponseStream( + MoqtFramer* absl_nonnull framer, + const MoqtControlMessageParser& message_parser, uint64_t track_alias, + SubscriptionPublisher::AddCallback add_callback, + SubscriptionPublisher::RemoveCallback remove_callback, + SessionErrorCallback session_error_callback, + quiche::QuicheWeakPtr<SessionToPublisherInterface> session) + : MoqtBidiStreamBase(framer, message_parser, + std::move(session_error_callback)), + track_alias_(track_alias), + add_callback_(std::move(add_callback)), + remove_callback_(std::move(remove_callback)), + session_(std::move(session)) {} + +absl::Status MoqtSubscribeResponseStream::OnRawControlMessage( + const MoqtRawControlMessage& message) { + return ControlMessageDispatcher::DispatchControlMessage( + *this, message_parser(), message, "subscribe response"); +} + +absl::Status MoqtSubscribeResponseStream::OnControlMessage( + const MoqtSubscribe& message) { + if (subscription_ != nullptr) { + return absl::InvalidArgumentError( + "SUBSCRIBE received on stream that already has a subscription"); + } + QUIC_DLOG(INFO) << "Received a SUBSCRIBE for " << message.full_track_name; + if (session() == nullptr) { + return absl::OkStatus(); + } + std::shared_ptr<MoqtTrackPublisher> track_publisher = + session()->GetTrackPublisher(message.full_track_name); + if (track_publisher == nullptr) { + QUIC_DLOG(INFO) << "SUBSCRIBE for " << message.full_track_name + << " rejected by the application: does not exist"; + return SendRequestError(message.request_id, RequestErrorCode::kDoesNotExist, + std::nullopt, "not found", /*fin=*/true); + } + subscription_ = std::make_unique<SubscriptionPublisher>( + *framer(), track_publisher, this, message.request_id, track_alias_, + message.parameters, session_, false); + if (add_callback_ != nullptr) { + bool result = std::move(add_callback_)(subscription_.get()); + add_callback_ = nullptr; + if (!result) { + return SendRequestError(message.request_id, + RequestErrorCode::kDuplicateSubscription, + std::nullopt, "duplicate subscription", + /*fin=*/true); + } + } + // Don't add the publisher until we know it's successful. + track_publisher->AddObjectListener(subscription_.get()); + return absl::OkStatus(); +} + +absl::Status MoqtSubscribeResponseStream::OnControlMessage( + const MoqtRequestUpdate& message) { + if (subscription_ == nullptr) { + QUICHE_BUG(INFO) << "Received REQUEST_UPDATE, no subscription state"; + return SendRequestError(message.request_id, + RequestErrorCode::kInternalError, std::nullopt, + "no subscription", /*fin=*/true); + } + subscription_->Update(message.parameters); + return SendRequestOk(message.request_id, MessageParameters()); +} + +void MoqtSubscribeResponseStream::Detach() { + if (remove_callback_ != nullptr && subscription_ != nullptr) { + SubscriptionPublisher::RemoveCallback remove_callback = + std::move(remove_callback_); + remove_callback_ = nullptr; + std::move(remove_callback)(subscription_.get()); + } + if (subscription_ != nullptr) { + subscription_->ResetAllStreams(); + subscription_ = nullptr; + } +} + +} // namespace moqt
diff --git a/quiche/quic/moqt/moqt_subscribe_stream.h b/quiche/quic/moqt/moqt_subscribe_stream.h new file mode 100644 index 0000000..3b6f5aa --- /dev/null +++ b/quiche/quic/moqt/moqt_subscribe_stream.h
@@ -0,0 +1,114 @@ +// Copyright (c) 2026 The Chromium Authors. All rights reserved. +// Use of this source code is governed by a BSD-style license that can be +// found in the LICENSE file. + +#ifndef QUICHE_QUIC_MOQT_MOQT_SUBSCRIBE_STREAM_H_ +#define QUICHE_QUIC_MOQT_MOQT_SUBSCRIBE_STREAM_H_ + +#include <cstdint> +#include <memory> + +#include "absl/base/nullability.h" +#include "absl/status/status.h" +#include "quiche/quic/core/quic_alarm_factory.h" +#include "quiche/quic/core/quic_clock.h" +#include "quiche/quic/moqt/moqt_bidi_stream.h" +#include "quiche/quic/moqt/moqt_framer.h" +#include "quiche/quic/moqt/moqt_key_value_pair.h" +#include "quiche/quic/moqt/moqt_messages.h" +#include "quiche/quic/moqt/moqt_names.h" +#include "quiche/quic/moqt/moqt_parser.h" +#include "quiche/quic/moqt/moqt_session_callbacks.h" +#include "quiche/quic/moqt/moqt_subscription.h" +#include "quiche/quic/moqt/moqt_track.h" +#include "quiche/common/quiche_weak_ptr.h" + +namespace moqt { + +class MoqtSubscribeRequestStream : public MoqtBidiStreamBase { + public: + MoqtSubscribeRequestStream( + MoqtFramer* absl_nonnull framer, + const MoqtControlMessageParser& message_parser, uint64_t request_id, + SessionErrorCallback session_error_callback, const FullTrackName& name, + SubscribeVisitor* absl_nonnull visitor, + const MessageParameters& parameters, + SubscribeRemoteTrack::AddCallback add_callback, + SubscribeRemoteTrack::RemoveCallback remove_callback, + const quic::QuicClock* absl_nonnull clock, + quic::QuicAlarmFactory* absl_nonnull alarm_factory); + ~MoqtSubscribeRequestStream() { Detach(); } + + // StreamBase overrides. + void OnStreamBound() override; + absl::Status OnRawControlMessage( + const MoqtRawControlMessage& message) override; + absl::Status OnControlMessage(const MoqtRequestOk& message) override; + absl::Status OnControlMessage(const MoqtRequestError& message) override; + absl::Status OnControlMessage(const MoqtSubscribeOk& message); + absl::Status OnControlMessage(const MoqtPublishDone& message); + + SubscribeRemoteTrack* track() const { return track_.get(); } + void Detach() override; + + private: + std::unique_ptr<SubscribeRemoteTrack> track_; + SubscribeRemoteTrack::AddCallback add_callback_; + SubscribeRemoteTrack::RemoveCallback remove_callback_; + const quic::QuicClock* clock_; + quic::QuicAlarmFactory* alarm_factory_; +}; + +class MoqtSubscribeResponseStream : public MoqtBidiStreamBase { + public: + MoqtSubscribeResponseStream( + MoqtFramer* absl_nonnull framer, + const MoqtControlMessageParser& message_parser, uint64_t track_alias, + SubscriptionPublisher::AddCallback add_callback, + SubscriptionPublisher::RemoveCallback remove_callback, + SessionErrorCallback session_error_callback, + quiche::QuicheWeakPtr<SessionToPublisherInterface> session); + ~MoqtSubscribeResponseStream() { + if (subscription_ != nullptr) { + subscription_->IgnoreResetAllStreams(); + } + Detach(); + } + + // MoqtBidiStreamBase overrides. + void OnStreamBound() override { stream_parser()->set_allow_fin(true); } + absl::Status OnRawControlMessage( + const MoqtRawControlMessage& message) override; + absl::Status OnControlMessage(const MoqtRequestOk& message) override { + return absl::InvalidArgumentError( + "REQUEST_OK not allowed from Subscriber on SUBSCRIBE stream"); + } + absl::Status OnControlMessage(const MoqtRequestError& message) override { + return absl::InvalidArgumentError( + "REQUEST_ERROR not allowed from Subscriber on SUBSCRIBE stream"); + } + + absl::Status OnControlMessage(const MoqtSubscribe& message); + absl::Status OnControlMessage(const MoqtRequestUpdate& message); + absl::Status OnControlMessage(const MoqtObjectAck& message) { + subscription_->ProcessObjectAck(message); + return absl::OkStatus(); + } + void Detach() override; + + private: + // Returns nullptr if MoqtSession is gone. + SessionToPublisherInterface* absl_nullable session() const { + return session_.GetIfAvailable(); + } + + uint64_t track_alias_; + std::unique_ptr<SubscriptionPublisher> subscription_; + SubscriptionPublisher::AddCallback add_callback_; + SubscriptionPublisher::RemoveCallback remove_callback_; + quiche::QuicheWeakPtr<SessionToPublisherInterface> session_; +}; + +} // namespace moqt + +#endif // QUICHE_QUIC_MOQT_MOQT_SUBSCRIBE_STREAM_H_
diff --git a/quiche/quic/moqt/moqt_subscribe_stream_test.cc b/quiche/quic/moqt/moqt_subscribe_stream_test.cc new file mode 100644 index 0000000..300dafb --- /dev/null +++ b/quiche/quic/moqt/moqt_subscribe_stream_test.cc
@@ -0,0 +1,334 @@ +// Copyright (c) 2026 The Chromium Authors. All rights reserved. +// Use of this source code is governed by a BSD-style license that can be +// found in the LICENSE file. + +#include "quiche/quic/moqt/moqt_subscribe_stream.h" + +#include <cstdint> +#include <memory> +#include <optional> +#include <utility> +#include <variant> + +#include "absl/status/status.h" +#include "absl/strings/string_view.h" +#include "quiche/quic/core/quic_time.h" +#include "quiche/quic/core/quic_types.h" +#include "quiche/quic/moqt/moqt_error.h" +#include "quiche/quic/moqt/moqt_fetch_task.h" +#include "quiche/quic/moqt/moqt_framer.h" +#include "quiche/quic/moqt/moqt_key_value_pair.h" +#include "quiche/quic/moqt/moqt_messages.h" +#include "quiche/quic/moqt/moqt_names.h" +#include "quiche/quic/moqt/moqt_parser.h" +#include "quiche/quic/moqt/moqt_session_callbacks.h" +#include "quiche/quic/moqt/moqt_session_interface.h" +#include "quiche/quic/moqt/moqt_subscription.h" +#include "quiche/quic/moqt/moqt_trace_recorder.h" +#include "quiche/quic/moqt/moqt_track.h" +#include "quiche/quic/moqt/test_tools/mock_moqt_session.h" +#include "quiche/quic/moqt/test_tools/moqt_framer_utils.h" +#include "quiche/quic/moqt/test_tools/moqt_mock_visitor.h" +#include "quiche/quic/test_tools/mock_clock.h" +#include "quiche/quic/test_tools/quic_test_utils.h" +#include "quiche/common/platform/api/quiche_test.h" +#include "quiche/common/test_tools/quiche_test_utils.h" +#include "quiche/web_transport/test_tools/mock_web_transport.h" + +namespace moqt::test { +namespace { + +using ::testing::_; +using ::testing::Return; +using ::testing::StrictMock; + +class MoqtSubscribeRequestStreamTest : public quiche::test::QuicheTest { + public: + MoqtSubscribeRequestStreamTest() + : framer_(/*using_webtrans=*/true, quic::Perspective::IS_CLIENT), + message_parser_(kDefaultMoqtVersion, /*uses_web_transport=*/true, + quic::Perspective::IS_CLIENT), + track_name_("foo", "bar") { + stream_ = std::make_unique<MoqtSubscribeRequestStream>( + &framer_, message_parser_, kRequestId, error_callback_.AsStdFunction(), + track_name_, &mock_subscribe_visitor_, parameters_, + mock_add_callback_.AsStdFunction(), + mock_remove_callback_.AsStdFunction(), &mock_clock_, + &mock_alarm_factory_); + EXPECT_CALL(mock_stream_, CanWrite).WillRepeatedly(testing::Return(true)); + EXPECT_CALL(mock_subscribe_visitor_, OnPublishDone(track_name_)) + .Times(testing::AnyNumber()); + } + + MoqtFramer framer_; + MoqtControlMessageParser message_parser_; + uint64_t kRequestId = 1; + uint64_t kTrackAlias = 100; + FullTrackName track_name_; + MessageParameters parameters_; + testing::StrictMock<testing::MockFunction<void(MoqtError, absl::string_view)>> + error_callback_; + testing::MockFunction<bool(SubscribeRemoteTrack*)> mock_add_callback_; + testing::MockFunction<void(SubscribeRemoteTrack*)> mock_remove_callback_; + quic::MockClock mock_clock_; + quic::test::MockAlarmFactory mock_alarm_factory_; + StrictMock<MockSubscribeRemoteTrackVisitor> mock_subscribe_visitor_; + webtransport::test::MockStream mock_stream_; + std::unique_ptr<MoqtSubscribeRequestStream> stream_; +}; + +TEST_F(MoqtSubscribeRequestStreamTest, OnStreamBound) { + EXPECT_CALL(mock_stream_, + Writev(ControlMessageOfType(MoqtMessageType::kSubscribe), _)) + .WillOnce(Return(absl::OkStatus())); + stream_->BindStream(&mock_stream_); +} + +TEST_F(MoqtSubscribeRequestStreamTest, ReceiveSubscribeOk) { + EXPECT_CALL(mock_stream_, + Writev(ControlMessageOfType(MoqtMessageType::kSubscribe), _)) + .WillOnce(Return(absl::OkStatus())); + stream_->BindStream(&mock_stream_); + EXPECT_CALL(mock_add_callback_, Call(stream_->track())) + .WillOnce(Return(true)); + EXPECT_CALL(mock_subscribe_visitor_, OnReply(track_name_, _)); + MoqtSubscribeOk subscribe_ok; + subscribe_ok.request_id = kRequestId; + subscribe_ok.track_alias = kTrackAlias; + QUICHE_EXPECT_OK(stream_->OnControlMessage(subscribe_ok)); + EXPECT_EQ(stream_->track()->track_alias(), kTrackAlias); + // Test cleanup. + EXPECT_CALL(mock_remove_callback_, Call); +} + +TEST_F(MoqtSubscribeRequestStreamTest, ReceiveSubscribeOkAliasDuplicate) { + EXPECT_CALL(mock_stream_, + Writev(ControlMessageOfType(MoqtMessageType::kSubscribe), _)) + .WillOnce(Return(absl::OkStatus())); + stream_->BindStream(&mock_stream_); + EXPECT_CALL(mock_add_callback_, Call(stream_->track())) + .WillOnce(Return(false)); + EXPECT_CALL(error_callback_, Call(MoqtError::kDuplicateTrackAlias, _)); + MoqtSubscribeOk subscribe_ok; + subscribe_ok.request_id = kRequestId; + subscribe_ok.track_alias = kTrackAlias; + QUICHE_EXPECT_OK(stream_->OnControlMessage(subscribe_ok)); + // Test cleanup. + EXPECT_CALL(mock_remove_callback_, Call); +} + +TEST_F(MoqtSubscribeRequestStreamTest, RequestOkBeforeSubscribeOk) { + EXPECT_CALL(mock_stream_, + Writev(ControlMessageOfType(MoqtMessageType::kSubscribe), _)) + .WillOnce(Return(absl::OkStatus())); + stream_->BindStream(&mock_stream_); + EXPECT_CALL(error_callback_, Call(MoqtError::kProtocolViolation, _)); + MoqtRequestOk request_ok; + request_ok.request_id = kRequestId; + request_ok.parameters.expires = quic::QuicTimeDelta::FromSeconds(30); + QUICHE_EXPECT_OK(stream_->OnControlMessage(request_ok)); + // Test cleanup. + EXPECT_CALL(mock_remove_callback_, Call); +} + +TEST_F(MoqtSubscribeRequestStreamTest, ReceiveRequestOk) { + // SUBSCRIBE handshake. + EXPECT_CALL(mock_stream_, + Writev(ControlMessageOfType(MoqtMessageType::kSubscribe), _)) + .WillOnce(Return(absl::OkStatus())); + stream_->BindStream(&mock_stream_); + EXPECT_CALL(mock_add_callback_, Call(stream_->track())) + .WillOnce(Return(true)); + EXPECT_CALL(mock_subscribe_visitor_, OnReply(track_name_, _)); + MoqtSubscribeOk subscribe_ok; + subscribe_ok.request_id = kRequestId; + subscribe_ok.track_alias = kTrackAlias; + QUICHE_EXPECT_OK(stream_->OnControlMessage(subscribe_ok)); + // REQUEST_UPDATE. + EXPECT_CALL(mock_stream_, + Writev(ControlMessageOfType(MoqtMessageType::kRequestUpdate), _)) + .WillOnce(Return(absl::OkStatus())); + bool callback_called = false; + MoqtResponseCallback callback = + [&](std::variant<MessageParameters, MoqtRequestErrorInfo> res) { + callback_called = true; + ASSERT_TRUE(std::holds_alternative<MessageParameters>(res)); + EXPECT_EQ(std::get<MessageParameters>(res).expires, + quic::QuicTimeDelta::FromSeconds(30)); + }; + parameters_.subscriber_priority = 20; + QUICHE_EXPECT_OK(stream_->SendRequestUpdate( + kRequestId, kRequestId, parameters_, std::move(callback))); + // Params not yet updated. + EXPECT_EQ(stream_->track()->const_parameters().subscriber_priority, + std::nullopt); + MoqtRequestOk request_ok; + request_ok.request_id = kRequestId; + request_ok.parameters.expires = quic::QuicTimeDelta::FromSeconds(30); + QUICHE_EXPECT_OK(stream_->OnControlMessage(request_ok)); + EXPECT_EQ(stream_->track()->const_parameters().subscriber_priority, 20); + EXPECT_EQ(stream_->track()->const_parameters().expires, + quic::QuicTimeDelta::FromSeconds(30)); + EXPECT_TRUE(callback_called); + // Test cleanup. + EXPECT_CALL(mock_remove_callback_, Call); +} + +TEST_F(MoqtSubscribeRequestStreamTest, ReceiveRequestError) { + EXPECT_CALL(mock_stream_, + Writev(ControlMessageOfType(MoqtMessageType::kSubscribe), _)) + .WillOnce(Return(absl::OkStatus())); + stream_->BindStream(&mock_stream_); + EXPECT_CALL(mock_subscribe_visitor_, OnReply(track_name_, _)) + .WillOnce( + [](const FullTrackName&, + const std::variant<SubscribeOkData, MoqtRequestErrorInfo>& reply) { + ASSERT_TRUE(std::holds_alternative<MoqtRequestErrorInfo>(reply)); + EXPECT_EQ(std::get<MoqtRequestErrorInfo>(reply).error_code, + RequestErrorCode::kUnauthorized); + }); + EXPECT_CALL(mock_stream_, Writev(testing::IsEmpty(), _)) + .WillOnce(Return(absl::OkStatus())); + EXPECT_CALL(mock_remove_callback_, Call); + MoqtRequestError request_error; + request_error.request_id = kRequestId; + request_error.error_code = RequestErrorCode::kUnauthorized; + request_error.reason_phrase = "unauthorized"; + QUICHE_EXPECT_OK(stream_->OnControlMessage(request_error)); +} + +TEST_F(MoqtSubscribeRequestStreamTest, ReceivePublishDone) { + MoqtPublishDone publish_done; + publish_done.request_id = kRequestId; + publish_done.stream_count = 5; + EXPECT_CALL(mock_subscribe_visitor_, OnPublishDone(track_name_)); + QUICHE_EXPECT_OK(stream_->OnControlMessage(publish_done)); + EXPECT_CALL(mock_remove_callback_, Call); +} + +class MoqtSubscribeResponseStreamTest : public quiche::test::QuicheTest { + public: + MoqtSubscribeResponseStreamTest() + : framer_(/*using_webtrans=*/true, quic::Perspective::IS_SERVER), + message_parser_(kDefaultMoqtVersion, /*uses_web_transport=*/true, + quic::Perspective::IS_SERVER), + track_publisher_(std::make_shared<TestTrackPublisher>(kTrackName)) { + stream_ = std::make_unique<MoqtSubscribeResponseStream>( + &framer_, message_parser_, kTrackAlias, + mock_add_callback_.AsStdFunction(), + mock_remove_callback_.AsStdFunction(), error_callback_.AsStdFunction(), + visitor_.weak_ptr_factory_.Create()); + stream_->BindStream(&mock_stream_); + EXPECT_CALL(mock_stream_, CanWrite).WillRepeatedly(testing::Return(true)); + } + + MoqtFramer framer_; + MoqtControlMessageParser message_parser_; + uint64_t kRequestId = 1; + uint64_t kTrackAlias = 100; + FullTrackName kTrackName{"foo", "bar"}; + std::shared_ptr<TestTrackPublisher> track_publisher_; + testing::MockFunction<void(MoqtError, absl::string_view)> error_callback_; + testing::MockFunction<bool(SubscriptionPublisher*)> mock_add_callback_; + testing::MockFunction<void(SubscriptionPublisher*)> mock_remove_callback_; + MockSessionToPublisherInterface visitor_; + webtransport::test::MockSession webtrans_; + webtransport::test::MockStream mock_stream_; + MoqtTraceRecorder trace_recorder_; + std::unique_ptr<MoqtSubscribeResponseStream> stream_; +}; + +TEST_F(MoqtSubscribeResponseStreamTest, ReceiveSubscribeSuccess) { + EXPECT_CALL(visitor_, session).WillRepeatedly(Return(&webtrans_)); + EXPECT_CALL(visitor_, GetTrackPublisher(kTrackName)) + .WillOnce(Return(track_publisher_)); + EXPECT_CALL(mock_add_callback_, Call(testing::NotNull())) + .WillOnce(Return(true)); + MoqtSubscribe subscribe; + subscribe.request_id = kRequestId; + subscribe.full_track_name = kTrackName; + QUICHE_EXPECT_OK(stream_->OnControlMessage(subscribe)); + // Test cleanup. + EXPECT_CALL(mock_remove_callback_, Call); +} + +TEST_F(MoqtSubscribeResponseStreamTest, ReceiveSubscribeDoesNotExist) { + EXPECT_CALL(visitor_, session).WillRepeatedly(Return(&webtrans_)); + EXPECT_CALL(visitor_, GetTrackPublisher(kTrackName)) + .WillOnce(Return(nullptr)); + EXPECT_CALL(mock_stream_, + Writev(ControlMessageOfType(MoqtMessageType::kRequestError), _)) + .WillOnce(Return(absl::OkStatus())); + + MoqtSubscribe subscribe; + subscribe.request_id = kRequestId; + subscribe.full_track_name = kTrackName; + QUICHE_EXPECT_OK(stream_->OnControlMessage(subscribe)); +} + +TEST_F(MoqtSubscribeResponseStreamTest, ReceiveSubscribeDuplicate) { + EXPECT_CALL(visitor_, session).WillRepeatedly(Return(&webtrans_)); + EXPECT_CALL(visitor_, GetTrackPublisher(kTrackName)) + .WillOnce(Return(track_publisher_)); + EXPECT_CALL(mock_add_callback_, Call(testing::NotNull())) + .WillOnce(Return(false)); + EXPECT_CALL(mock_stream_, + Writev(ControlMessageOfType(MoqtMessageType::kRequestError), _)) + .WillOnce(Return(absl::OkStatus())); + MoqtSubscribe subscribe; + subscribe.request_id = kRequestId; + subscribe.full_track_name = kTrackName; + QUICHE_EXPECT_OK(stream_->OnControlMessage(subscribe)); +} + +TEST_F(MoqtSubscribeResponseStreamTest, ReceiveRequestUpdate) { + EXPECT_CALL(visitor_, session).WillRepeatedly(Return(&webtrans_)); + EXPECT_CALL(visitor_, GetTrackPublisher(kTrackName)) + .WillOnce(Return(track_publisher_)); + EXPECT_CALL(mock_add_callback_, Call(testing::NotNull())) + .WillOnce(Return(true)); + MoqtSubscribe subscribe; + subscribe.request_id = kRequestId; + subscribe.full_track_name = kTrackName; + QUICHE_EXPECT_OK(stream_->OnControlMessage(subscribe)); + EXPECT_CALL(mock_stream_, + Writev(ControlMessageOfType(MoqtMessageType::kRequestOk), _)) + .WillOnce(Return(absl::OkStatus())); + MoqtRequestUpdate update; + update.request_id = kRequestId; + update.parameters.subscriber_priority = 10; + QUICHE_EXPECT_OK(stream_->OnControlMessage(update)); + // Test cleanup. + EXPECT_CALL(mock_remove_callback_, Call(_)); +} + +TEST_F(MoqtSubscribeResponseStreamTest, ReceiveInvalidControlMessages) { + MoqtRequestOk request_ok; + EXPECT_FALSE(stream_->OnControlMessage(request_ok).ok()); + MoqtRequestError request_error; + EXPECT_FALSE(stream_->OnControlMessage(request_error).ok()); +} + +TEST_F(MoqtSubscribeResponseStreamTest, ReceiveObjectAck) { + EXPECT_CALL(visitor_, session).WillRepeatedly(Return(&webtrans_)); + EXPECT_CALL(visitor_, GetTrackPublisher(kTrackName)) + .WillOnce(Return(track_publisher_)); + EXPECT_CALL(mock_add_callback_, Call(testing::NotNull())) + .WillOnce(Return(true)); + MoqtSubscribe subscribe; + subscribe.request_id = kRequestId; + subscribe.full_track_name = kTrackName; + QUICHE_EXPECT_OK(stream_->OnControlMessage(subscribe)); + MoqtObjectAck ack; + ack.group_id = 1; + ack.object_id = 2; + ack.delta_from_deadline = quic::QuicTimeDelta::FromMilliseconds(100); + EXPECT_CALL(visitor_, trace_recorder()) + .WillOnce(testing::ReturnRef(trace_recorder_)); + QUICHE_EXPECT_OK(stream_->OnControlMessage(ack)); + // Test cleanup. + EXPECT_CALL(mock_remove_callback_, Call(_)); +} + +} // namespace +} // namespace moqt::test
diff --git a/quiche/quic/moqt/moqt_subscription.cc b/quiche/quic/moqt/moqt_subscription.cc index 9bf4513..fa83093 100644 --- a/quiche/quic/moqt/moqt_subscription.cc +++ b/quiche/quic/moqt/moqt_subscription.cc
@@ -30,6 +30,7 @@ #include "quiche/quic/moqt/moqt_types.h" #include "quiche/quic/moqt/moqt_uni_stream.h" #include "quiche/common/quiche_buffer_allocator.h" +#include "quiche/common/quiche_weak_ptr.h" #include "quiche/web_transport/web_transport.h" namespace moqt { @@ -38,10 +39,7 @@ MoqtFramer framer, std::shared_ptr<MoqtTrackPublisher> track_publisher, MoqtBidiStreamBase* absl_nonnull bidi_stream, uint64_t request_id, uint64_t track_alias, const MessageParameters& parameters, - SessionToPublisherInterface* absl_nonnull visitor, - MoqtPublishingMonitorInterface* monitoring_interface, - const quic::QuicClock* absl_nonnull clock, - MoqtTraceRecorder& trace_recorder, bool is_publish) + quiche::QuicheWeakPtr<SessionToPublisherInterface> visitor, bool is_publish) : track_publisher_(track_publisher), bidi_stream_(bidi_stream), visitor_(visitor), @@ -49,11 +47,14 @@ established_(is_publish), track_alias_(track_alias), framer_(framer), - trace_recorder_(trace_recorder), parameters_(parameters), - monitoring_interface_(monitoring_interface), - clock_(clock), weak_ptr_factory_(this) { + SessionToPublisherInterface* session_info = visitor_.GetIfAvailable(); + if (session_info == nullptr) { + return; + } + monitoring_interface_ = session_info->ReleaseMonitoringInterface( + track_publisher_->GetTrackName()); if (monitoring_interface_ != nullptr) { monitoring_interface_->OnObjectAckSupportKnown(parameters.oack_window_size); } @@ -66,13 +67,6 @@ if (track_publisher_ != nullptr) { track_publisher_->RemoveObjectListener(this); } - // Reset all streams. - for (const webtransport::StreamId stream_id : stream_map_.GetAllStreams()) { - webtransport::Stream* stream = GetStreamById(stream_id); - if (stream != nullptr) { - stream->ResetWithUserCode(kResetCodeCancelled); - } - } } void SubscriptionPublisher::Update(const MessageParameters& parameters) { @@ -103,13 +97,28 @@ pending_streams_.rbegin()->second.publisher_priority.value_or( track_publisher_->extensions().default_publisher_priority()); MoqtTrackPriority old_track_priority = {old_priority, publisher_priority}; - visitor_->UpdateTrackPriority( + if (visitor() == nullptr) { + return; + } + visitor()->UpdateTrackPriority( request_id_, old_track_priority, MoqtTrackPriority{new_priority, publisher_priority}); // Don't bother to update all the pending stream send orders. } } +void SubscriptionPublisher::ResetAllStreams() { + if (ignore_reset_all_streams_) { + return; + } + for (const webtransport::StreamId stream_id : stream_map_.GetAllStreams()) { + webtransport::Stream* stream = GetStreamById(stream_id); + if (stream != nullptr) { + stream->ResetWithUserCode(kResetCodeCancelled); + } + } +} + void SubscriptionPublisher::OnSubscribeAccepted() { if (established_) { return; // It's a PUBLISH. @@ -142,11 +151,9 @@ } void SubscriptionPublisher::OnSubscribeRejected(MoqtRequestErrorInfo info) { - bidi_stream_->CheckStatus(bidi_stream_->SendRequestError( - request_id_, info, /*fin=*/!bidi_stream_->is_control_stream())); - if (bidi_stream_->is_control_stream()) { - visitor_->PublishIsDone(request_id_); - } + bidi_stream_->CheckStatus(bidi_stream_->SendRequestError(request_id_, info, + /*fin=*/true)); + // Sending FIN will delete the class. } void SubscriptionPublisher::OnNewObjectAvailable( @@ -174,7 +181,11 @@ // TODO(vasilvv): This currently sends UINT64_MAX for datagram subgroups. // Maybe do something more satisfactory? - trace_recorder_.RecordNewObjectAvaliable( + SessionToPublisherInterface* session_info = visitor(); + if (session_info == nullptr) { + return; // Session is gone. + } + session_info->trace_recorder().RecordNewObjectAvaliable( track_alias_, *track_publisher_, location, subgroup.value_or(UINT64_MAX), publisher_priority); @@ -187,7 +198,7 @@ } stream_id = stream_map_.GetStreamFor(index); } - if (visitor_->alternate_delivery_timeout() && + if (session_info->alternate_delivery_timeout() && !delivery_timeout().IsInfinite() && largest_sent_.has_value() && location.group >= largest_sent_->group) { // Start the delivery timeout timer on all previous groups. @@ -201,7 +212,7 @@ } OutgoingSubgroupStream* stream = absl::down_cast<OutgoingSubgroupStream*>(raw_stream->visitor()); - stream->CreateAndSetAlarm(clock_->ApproximateNow() + + stream->CreateAndSetAlarm(session_info->clock()->ApproximateNow() + delivery_timeout()); } } @@ -229,7 +240,7 @@ if (raw_stream == nullptr) { StreamRank rank = StreamRankFor(parameters); if (pending_streams_.empty() || rank > pending_streams_.rbegin()->first) { - visitor_->UpdateTrackPriority( + session_info->UpdateTrackPriority( request_id_, /*old_priority=*/pending_streams_.empty() ? std::optional<MoqtTrackPriority>() @@ -267,6 +278,7 @@ OutgoingSubgroupStream* stream = absl::down_cast<OutgoingSubgroupStream*>(raw_stream->visitor()); stream->Fin(location); + // Sending FIN will delete the class. } void SubscriptionPublisher::OnSubgroupAbandoned( @@ -345,34 +357,39 @@ quiche::QuicheBuffer datagram = framer_.SerializeObjectDatagram( header, object->payload[0].AsStringView(), default_publisher_priority_.value_or(kDefaultPublisherPriority)); - if (visitor_->session() == nullptr) { + if (visitor() == nullptr) { return; } - visitor_->session()->SendOrQueueDatagram(datagram.AsStringView()); + visitor()->session()->SendOrQueueDatagram(datagram.AsStringView()); OnObjectSent(object->metadata.location); } void SubscriptionPublisher::ProcessObjectAck(const MoqtObjectAck& message) { - trace_recorder_.RecordObjectAck(track_alias_, - Location(message.group_id, message.object_id), - message.delta_from_deadline); - - if (monitoring_interface_ == nullptr) { + SessionToPublisherInterface* session_info = visitor(); + if (session_info == nullptr) { return; } - monitoring_interface_->OnObjectAckReceived( - Location(message.group_id, message.object_id), + session_info->trace_recorder().RecordObjectAck( + track_alias_, Location(message.group_id, message.object_id), message.delta_from_deadline); + if (monitoring_interface_ != nullptr) { + monitoring_interface_->OnObjectAckReceived( + Location(message.group_id, message.object_id), + message.delta_from_deadline); + } } webtransport::Stream* absl_nullable SubscriptionPublisher::OpenDataStream( const NewDataStreamParameters& parameters) { - if (visitor_->session() == nullptr || - !visitor_->session()->CanOpenNextOutgoingUnidirectionalStream()) { + SessionToPublisherInterface* session_info = visitor(); + if (session_info == nullptr) { + return nullptr; + } + if (!session_info->session()->CanOpenNextOutgoingUnidirectionalStream()) { return nullptr; } webtransport::Stream* new_stream = - visitor_->session()->OpenOutgoingUnidirectionalStream(); + session_info->session()->OpenOutgoingUnidirectionalStream(); if (new_stream == nullptr) { return nullptr; } @@ -380,7 +397,8 @@ new_stream->SetVisitor(std::make_unique<OutgoingSubgroupStream>( framer_, new_stream, parameters.index, parameters.first_object, weak_ptr_factory_.Create(), track_publisher_, - StreamPriorityFor(parameters), track_alias_, &trace_recorder_)); + StreamPriorityFor(parameters), track_alias_, + &session_info->trace_recorder())); ++streams_opened_; new_stream->visitor()->OnCanWrite(); return new_stream; @@ -398,17 +416,9 @@ // FIN naturally, where possible. QUICHE_DLOG(INFO) << "Sending PUBLISH_DONE message for " << track_publisher_->GetTrackName(); - // TODO(martinduke): For SUBSCRIBE, no FIN because it's the control stream. bidi_stream_->SendOrBufferMessageOrFatal( - framer_.SerializePublishDone(publish_done), - /*fin=*/!bidi_stream_->is_control_stream()); - if (bidi_stream_->is_control_stream()) { - visitor_->PublishIsDone(request_id_); - } else { - // Only detach immediately for PUBLISH flow. - track_publisher_->RemoveObjectListener(this); - track_publisher_ = nullptr; - } + framer_.SerializePublishDone(publish_done), /*fin=*/true); + // sending FIN will delete the class. } void SubscriptionPublisher::OnDataStreamDestroyed( @@ -417,8 +427,12 @@ } void SubscriptionPublisher::OnCanCreateNewUniStream() { - while (visitor_->session() != nullptr && - visitor_->session()->CanOpenNextOutgoingUnidirectionalStream()) { + SessionToPublisherInterface* session_info = visitor(); + if (session_info == nullptr) { + return; + } + while (visitor_.IsValid() && + session_info->session()->CanOpenNextOutgoingUnidirectionalStream()) { auto it = pending_streams_.rbegin(); while (it != pending_streams_.rend() && (it->second.index.group < first_active_group_ || @@ -434,7 +448,7 @@ } pending_streams_.erase(--(it.base())); if (!pending_streams_.empty()) { - visitor_->UpdateTrackPriority( + session_info->UpdateTrackPriority( request_id_, std::nullopt, MoqtTrackPriority{ subscriber_priority(),
diff --git a/quiche/quic/moqt/moqt_subscription.h b/quiche/quic/moqt/moqt_subscription.h index d32ed4c..d4af902 100644 --- a/quiche/quic/moqt/moqt_subscription.h +++ b/quiche/quic/moqt/moqt_subscription.h
@@ -21,6 +21,7 @@ #include "quiche/quic/moqt/moqt_framer.h" #include "quiche/quic/moqt/moqt_key_value_pair.h" #include "quiche/quic/moqt/moqt_messages.h" +#include "quiche/quic/moqt/moqt_names.h" #include "quiche/quic/moqt/moqt_priority.h" #include "quiche/quic/moqt/moqt_publisher.h" #include "quiche/quic/moqt/moqt_stream_map.h" @@ -76,17 +77,18 @@ virtual ~SessionToPublisherInterface() = default; virtual bool alternate_delivery_timeout() const = 0; // If |old_priority| is nullopt, the subscription does not have any pending - // streams. If it has a value, |old_priority| is the old value to be replaced - // by |new_priority|. + // streams. If it has a value, |old_priority| is the old value to be + // replaced by |new_priority|. virtual void UpdateTrackPriority( uint64_t request_id, std::optional<MoqtTrackPriority> old_priority, MoqtTrackPriority new_priority) = 0; virtual quic::QuicAlarmFactory* alarm_factory() = 0; - // Destroy any state associated with the subscription. It is OK destroy - // SubscriptionPublisher in this method. - // TODO(martinduke): Delete once SUBSCRIBE is on the bidi stream. - virtual void PublishIsDone(uint64_t request_id) = 0; - // Returns nullptr if MoqtSession is closing. + virtual std::shared_ptr<MoqtTrackPublisher> GetTrackPublisher( + const FullTrackName& name) = 0; + virtual MoqtPublishingMonitorInterface* ReleaseMonitoringInterface( + const FullTrackName& name) = 0; + virtual const quic::QuicClock* clock() = 0; + virtual MoqtTraceRecorder& trace_recorder() = 0; virtual webtransport::Session* session() = 0; }; @@ -95,20 +97,19 @@ class SubscriptionPublisher : public MoqtObjectListener, public SubscriptionPublisherInterface { public: - // The provider of this callback will delete whatever state it is tracking for - // the subscription. This will be used by both PUBLISH and SUBSCRIBE streams. - // For control stream SUBSCRIBE, visitor->PublishIsDone() does this instead. + // The provider of this callback will add/delete whatever state it is tracking + // for the subscription. This will be used by both PUBLISH and SUBSCRIBE + // streams. AddCallback returns |false| if the add fails because the key + // already exists. + using AddCallback = quiche::SingleUseCallback<bool(SubscriptionPublisher*)>; using RemoveCallback = quiche::SingleUseCallback<void(SubscriptionPublisher*)>; - SubscriptionPublisher(MoqtFramer framer, - std::shared_ptr<MoqtTrackPublisher> track_publisher, - MoqtBidiStreamBase* absl_nonnull bidi_stream, - uint64_t request_id, uint64_t track_alias, - const MessageParameters& parameters, - SessionToPublisherInterface* absl_nonnull visitor, - MoqtPublishingMonitorInterface* monitoring_interface, - const quic::QuicClock* absl_nonnull clock, - MoqtTraceRecorder& trace_recorder, bool is_publish); + SubscriptionPublisher( + MoqtFramer framer, std::shared_ptr<MoqtTrackPublisher> track_publisher, + MoqtBidiStreamBase* absl_nonnull bidi_stream, uint64_t request_id, + uint64_t track_alias, const MessageParameters& parameters, + quiche::QuicheWeakPtr<SessionToPublisherInterface> visitor, + bool is_publish); ~SubscriptionPublisher(); SubscriptionPublisher(const SubscriptionPublisher&) = delete; @@ -143,28 +144,39 @@ parameters_.subscription_filter->InWindow(location))); }; bool alternate_delivery_timeout() override { - return visitor_->alternate_delivery_timeout(); + if (visitor() == nullptr) { + return false; + } + return visitor()->alternate_delivery_timeout(); } - const quic::QuicClock* clock() override { return clock_; } + const quic::QuicClock* clock() override { + if (visitor() == nullptr) { + return nullptr; + } + return visitor()->clock(); + } quic::QuicTimeDelta delivery_timeout() override { return std::min( parameters_.delivery_timeout.value_or(kDefaultDeliveryTimeout), publisher_delivery_timeout_.value_or(kDefaultDeliveryTimeout)); } quic::QuicAlarmFactory* alarm_factory() override { - return visitor_->alarm_factory(); + if (visitor() == nullptr) { + return nullptr; + } + return visitor()->alarm_factory(); } void OnObjectSent(Location sequence) override; void OnStreamTimeout(DataStreamIndex index) override { reset_subgroups_.insert(index); - if (visitor_->alternate_delivery_timeout()) { + if (visitor()->alternate_delivery_timeout()) { first_active_group_ = std::max(first_active_group_, index.group + 1); } } // OnSubgroupAbandoned() is declared above with MoqtObjectListener. void OnDataStreamDestroyed(DataStreamIndex) override; - // Called by MoqtSession when this subscription can open a new stream. + // MoqtObjectPublisher implementation. void OnCanCreateNewUniStream(); // Called when the parameters_ needs an update. @@ -174,6 +186,17 @@ bool established() const { return established_; } + quiche::QuicheWeakPtr<SubscriptionPublisherInterface> GetWeakPtr() { + return weak_ptr_factory_.Create(); + } + + // Resets all unidirectional streams. + void ResetAllStreams(); + // Called when the bidi stream is being destroyed. If the result of a session + // teardown, it is not safe to access the streams via GetStreamById, and the + // uni streams will be destroyed anyway. + void IgnoreResetAllStreams() { ignore_reset_all_streams_ = true; } + private: friend class test::SubscriptionPublisherPeer; @@ -224,20 +247,25 @@ } webtransport::Stream* GetStreamById(webtransport::StreamId stream_id) { - return visitor_->session() == nullptr - ? nullptr - : visitor_->session()->GetStreamById(stream_id); + if (visitor() != nullptr) { + return visitor()->session()->GetStreamById(stream_id); + } + return nullptr; + } + + // If nullptr, MoqtSession is gone. + SessionToPublisherInterface* absl_nullable visitor() const { + return visitor_.GetIfAvailable(); } std::shared_ptr<MoqtTrackPublisher> track_publisher_; MoqtBidiStreamBase* absl_nonnull bidi_stream_; - SessionToPublisherInterface* absl_nonnull visitor_; + quiche::QuicheWeakPtr<SessionToPublisherInterface> visitor_; uint64_t request_id_; // Subscription is in the ESTABLISHED state. bool established_; const uint64_t track_alias_; MoqtFramer framer_; - MoqtTraceRecorder& trace_recorder_; // These are (mostly) the parameters from the SUBSCRIBE message. However, // group_order and largest_object may be updated by SUBSCRIBE_OK because // have no effect in a future REQUEST_UPDATE message. @@ -257,11 +285,11 @@ // Largest sequence number ever sent via this subscription. std::optional<Location> largest_sent_; SendStreamMap stream_map_; + bool ignore_reset_all_streams_ = false; // Store the StreamRank of queued outgoing data streams. High StreamRank is // highest priority, so use rbegin() to get the highest priority pending // stream. absl::btree_multimap<StreamRank, NewDataStreamParameters> pending_streams_; - const quic::QuicClock* absl_nonnull clock_; // Must be last. quiche::QuicheWeakPtrFactory<SubscriptionPublisherInterface> weak_ptr_factory_;
diff --git a/quiche/quic/moqt/moqt_subscription_test.cc b/quiche/quic/moqt/moqt_subscription_test.cc index 154725b..b35818d 100644 --- a/quiche/quic/moqt/moqt_subscription_test.cc +++ b/quiche/quic/moqt/moqt_subscription_test.cc
@@ -18,8 +18,8 @@ #include "absl/strings/match.h" #include "absl/strings/string_view.h" #include "absl/types/span.h" -#include "quiche/quic/core/quic_alarm_factory.h" #include "quiche/quic/core/quic_time.h" +#include "quiche/quic/core/quic_types.h" #include "quiche/quic/moqt/moqt_bidi_stream.h" #include "quiche/quic/moqt/moqt_error.h" #include "quiche/quic/moqt/moqt_framer.h" @@ -33,8 +33,7 @@ #include "quiche/quic/moqt/moqt_session_interface.h" #include "quiche/quic/moqt/moqt_trace_recorder.h" #include "quiche/quic/moqt/moqt_types.h" -#include "quiche/quic/moqt/moqt_uni_stream.h" -#include "quiche/quic/moqt/test_tools/moqt_framer_utils.h" +#include "quiche/quic/moqt/test_tools/mock_moqt_session.h" #include "quiche/quic/moqt/test_tools/moqt_mock_visitor.h" #include "quiche/quic/moqt/test_tools/moqt_session_peer.h" #include "quiche/quic/platform/api/quic_test.h" @@ -71,27 +70,13 @@ using ::webtransport::DatagramStatus; using ::webtransport::DatagramStatusCode; -class MockSessionToPublisherInterface : public SessionToPublisherInterface { - public: - ~MockSessionToPublisherInterface() override = default; - MOCK_METHOD(bool, alternate_delivery_timeout, (), (const, override)); - MOCK_METHOD(void, UpdateTrackPriority, - (uint64_t, std::optional<MoqtTrackPriority>, MoqtTrackPriority), - (override)); - MOCK_METHOD(quic::QuicAlarmFactory*, alarm_factory, (), (override)); - MOCK_METHOD(void, PublishIsDone, (uint64_t), (override)); - MOCK_METHOD(webtransport::Session*, session, (), (override)); -}; - class TestMoqtBidiStream : public MoqtBidiStreamBase { public: TestMoqtBidiStream(MoqtFramer* absl_nonnull framer, const MoqtControlMessageParser& message_parser, SessionErrorCallback session_error_callback) : MoqtBidiStreamBase(framer, message_parser, - std::move(session_error_callback)) { - set_control_stream(); // TODO(martinduke): Delete - } + std::move(session_error_callback)) {} ~TestMoqtBidiStream() override = default; void OnStreamBound() override {}; absl::Status OnRawControlMessage( @@ -134,21 +119,25 @@ EXPECT_CALL(monitoring_interface_, OnObjectAckSupportKnown) .Times(AtLeast(0)); EXPECT_CALL(visitor_, session).WillRepeatedly(Return(&webtrans_)); + ON_CALL(visitor_, ReleaseMonitoringInterface) + .WillByDefault(Return(&monitoring_interface_)); publisher_ = std::make_unique<SubscriptionPublisher>( framer_, track_publisher_, &bidi_stream_, kRequestId, kTrackAlias, - parameters_, &visitor_, &monitoring_interface_, &mock_clock_, - trace_recorder_, /*is_publish=*/false); + parameters_, visitor_.weak_ptr_factory_.Create(), + /*is_publish=*/false); ON_CALL(visitor_, alternate_delivery_timeout).WillByDefault(Return(false)); ON_CALL(webtrans_, GetStreamById(kStreamId)) .WillByDefault(Return(&mock_uni_stream_)); ON_CALL(visitor_, alarm_factory).WillByDefault(Return(&alarm_factory_)); + ON_CALL(visitor_, clock).WillByDefault(Return(&mock_clock_)); + ON_CALL(visitor_, trace_recorder).WillByDefault(ReturnRef(trace_recorder_)); } ~SubscriptionPublisherTest() override { + if (track_publisher_ == nullptr) { + return; + } EXPECT_CALL(*track_publisher_, RemoveObjectListener(publisher_.get())); - size_t num_open_streams = - SubscriptionPublisherPeer::num_open_streams(publisher_.get()); - EXPECT_CALL(mock_uni_stream_, ResetWithUserCode).Times(num_open_streams); } MoqtPriority subscriber_priority() const { @@ -308,7 +297,6 @@ EXPECT_CALL(mock_bidi_stream_, Writev(ControlMessageOfType(MoqtMessageType::kRequestError), _)) .WillOnce(Return(absl::OkStatus())); - EXPECT_CALL(visitor_, PublishIsDone(1)); publisher_->OnSubscribeRejected(MoqtRequestErrorInfo( RequestErrorCode::kDoesNotExist, std::nullopt, "reason")); } @@ -471,9 +459,15 @@ }; EXPECT_CALL(mock_bidi_stream_, CanWrite()).WillRepeatedly(Return(true)); EXPECT_CALL(mock_bidi_stream_, - Writev(SerializedControlMessage(expected_publish_done), _)); - EXPECT_CALL(visitor_, PublishIsDone(kRequestId)); + Writev(SerializedControlMessage(expected_publish_done), _)) + .WillOnce([&](absl::Span<quiche::QuicheMemSlice> data, + const webtransport::StreamWriteOptions& options) { + EXPECT_TRUE(options.send_fin()); + return absl::OkStatus(); + }); + EXPECT_CALL(*track_publisher_, RemoveObjectListener); publisher_->OnGroupAbandoned(5); + track_publisher_ = nullptr; } TEST_F(SubscriptionPublisherTest, OnCanCreateNewUniStreamPendingCleanup) { @@ -502,8 +496,9 @@ EXPECT_CALL(mock_bidi_stream_, Writev(ControlMessageOfType(MoqtMessageType::kPublishDone), _)) .WillOnce(Return(absl::OkStatus())); - EXPECT_CALL(visitor_, PublishIsDone(1)); + EXPECT_CALL(*track_publisher_, RemoveObjectListener); publisher_->OnTrackPublisherGone(); + track_publisher_ = nullptr; } TEST_F(SubscriptionPublisherTest, ProcessObjectAck) { @@ -701,6 +696,34 @@ publisher_->OnNewObjectAvailable(Location(5, 1), 0, 128); } +TEST_F(SubscriptionPublisherTest, OnNewFinAvailable) { + CreateStream( + Location(1, 0), 0, 127, + {0x51, static_cast<uint8_t>(kTrackAlias), 0x01, 0x7f, 0x00, 0x0a}); + EXPECT_CALL(mock_uni_stream_, Writev(testing::IsEmpty(), _)) + .WillOnce([](absl::Span<quiche::QuicheMemSlice> data, + const webtransport::StreamWriteOptions& options) { + EXPECT_TRUE(options.send_fin()); + return absl::OkStatus(); + }); + publisher_->OnNewFinAvailable(Location(1, 0), 0); +} + +TEST_F(SubscriptionPublisherTest, OnSubgroupAbandoned) { + CreateStream( + Location(1, 0), 0, 127, + {0x51, static_cast<uint8_t>(kTrackAlias), 0x01, 0x7f, 0x00, 0x0a}); + EXPECT_CALL(mock_uni_stream_, ResetWithUserCode(1234)); + publisher_->OnSubgroupAbandoned(1, 0, 1234); +} + +TEST_F(SubscriptionPublisherTest, OnSubgroupAbandonedOutsideWindow) { + parameters_.subscription_filter = SubscriptionFilter(Location(20, 0)); + publisher_->Update(parameters_); + EXPECT_CALL(mock_uni_stream_, ResetWithUserCode).Times(0); + publisher_->OnSubgroupAbandoned(1, 0, 1234); +} + } // namespace } // namespace moqt::test
diff --git a/quiche/quic/moqt/moqt_track.cc b/quiche/quic/moqt/moqt_track.cc index 36e06cf..23060f0 100644 --- a/quiche/quic/moqt/moqt_track.cc +++ b/quiche/quic/moqt/moqt_track.cc
@@ -43,14 +43,9 @@ if (publish_done_alarm_ != nullptr) { publish_done_alarm_->PermanentCancel(); } - if (remove_callback_ != nullptr) { - RemoveCallback callback = std::move(remove_callback_); - remove_callback_ = nullptr; - std::move(callback)(this); - } if (visitor_ != nullptr) { + // The application expects OnPublishDone even if the Session has gone away. visitor_->OnPublishDone(full_track_name()); - visitor_ = nullptr; } } @@ -62,10 +57,6 @@ publisher_delivery_timeout_ = data.extensions.delivery_timeout(); // TODO(martinduke): Is there anything to do with EXPIRES? default_publisher_priority_ = data.extensions.default_publisher_priority(); - if (!parameters().group_order.has_value()) { - // Use publisher default because the subscriber didn't care. - parameters().group_order = data.extensions.default_publisher_group_order(); - } dynamic_groups_ = data.extensions.dynamic_groups(); visitor_->OnReply(full_track_name(), data); OnObjectOrOk(); @@ -83,7 +74,7 @@ ++streams_closed_; --currently_open_streams_; QUICHE_DCHECK_GE(currently_open_streams_, -1); - if (index.has_value()) { + if (index.has_value() && visitor_ != nullptr) { // If index is nullopt, there was not an object received on the stream. if (fin_received) { visitor_->OnStreamFin(full_track_name(), *index); @@ -92,7 +83,8 @@ } } if (all_streams_closed()) { - Destroy(); + request_stream()->Fin(); + // Fin will destroy the class. return; } if (publish_done_alarm_ == nullptr) { @@ -107,11 +99,12 @@ total_streams_ = stream_count; clock_ = clock; if (all_streams_closed()) { - Destroy(); + request_stream()->Fin(); + // Fin will destroy the class. return; } publish_done_alarm_ = std::unique_ptr<quic::QuicAlarm>( - alarm_factory->CreateAlarm(new PublishDoneDelegate(this))); + alarm_factory->CreateAlarm(new PublishDoneDelegate(weak_ptr()))); MaybeSetPublishDoneAlarm(); } @@ -177,6 +170,14 @@ } } +void SubscribeRemoteTrack::SendObjectAck( + uint64_t group_id, uint64_t object_id, + quic::QuicTimeDelta delta_from_deadline) { + request_stream()->SendOrBufferMessageOrFatal( + request_stream()->framer()->SerializeObjectAck( + {group_id, object_id, delta_from_deadline})); +} + UpstreamFetch::~UpstreamFetch() { UpstreamFetchTask* task = task_.GetIfAvailable(); if (task != nullptr) { @@ -184,10 +185,7 @@ // If this has already been called, UpstreamFetchTask will ignore it. task->OnStreamAndFetchClosed(kResetCodeCancelled, ""); } - if (remove_callback_ != nullptr) { - std::move(remove_callback_)(); - remove_callback_ = nullptr; - } + task = nullptr; } void UpstreamFetch::OnFetchResult(Location largest_location,
diff --git a/quiche/quic/moqt/moqt_track.h b/quiche/quic/moqt/moqt_track.h index 54716d1..269ef2e 100644 --- a/quiche/quic/moqt/moqt_track.h +++ b/quiche/quic/moqt/moqt_track.h
@@ -12,12 +12,14 @@ #include <optional> #include <utility> +#include "absl/base/nullability.h" #include "absl/status/status.h" #include "absl/strings/string_view.h" #include "quiche/quic/core/quic_alarm.h" #include "quiche/quic/core/quic_alarm_factory.h" #include "quiche/quic/core/quic_time.h" #include "quiche/quic/moqt/moqt_bidi_stream.h" +#include "quiche/quic/moqt/moqt_error.h" #include "quiche/quic/moqt/moqt_fetch_task.h" #include "quiche/quic/moqt/moqt_key_value_pair.h" #include "quiche/quic/moqt/moqt_messages.h" @@ -69,8 +71,6 @@ virtual bool is_fetch() const = 0; - virtual void Destroy() = 0; - // A REQUEST_UPDATE changes any field that is present in |parameters|. void Update(const MessageParameters& parameters) { parameters_.Update(parameters); @@ -78,8 +78,9 @@ MoqtBidiStreamBase* request_stream() { return request_stream_; } - protected: const MessageParameters& const_parameters() const { return parameters_; } + + protected: MessageParameters& parameters() { return parameters_; } private: @@ -102,15 +103,11 @@ using AddCallback = quiche::SingleUseCallback<bool(SubscribeRemoteTrack*)>; using RemoveCallback = quiche::SingleUseCallback<void(SubscribeRemoteTrack*)>; SubscribeRemoteTrack(const MoqtSubscribe& subscribe, - SubscribeVisitor* visitor, AddCallback add_callback, - RemoveCallback remove_callback) + SubscribeVisitor* visitor, + MoqtBidiStreamBase* request_stream) : RemoteTrack(subscribe.full_track_name, subscribe.request_id, - subscribe.parameters, /*request_stream=*/nullptr), - visitor_(visitor), - add_callback_(std::move(add_callback)), - remove_callback_(std::move(remove_callback)) {} - // For incoming PUBLISH, all |callbacks| will be nullptr because it's handled - // by the stream. + subscribe.parameters, request_stream), + visitor_(visitor) {} SubscribeRemoteTrack(const MoqtPublish& publish, SubscribeVisitor* visitor, MoqtBidiStreamBase* request_stream) : RemoteTrack(publish.full_track_name, publish.request_id, @@ -127,12 +124,8 @@ std::optional<uint64_t> track_alias() const { return track_alias_; } // Returns false if the callback returns false, meaning the session has been // destroyed. - bool set_track_alias(uint64_t track_alias) { + void set_track_alias(uint64_t track_alias) { track_alias_.emplace(track_alias); - if (add_callback_ != nullptr) { - return std::move(add_callback_)(this); - } - return true; } void OnStreamOpened(); void OnStreamClosed(bool fin_received, std::optional<DataStreamIndex> index); @@ -169,19 +162,8 @@ } void set_visitor(SubscribeVisitor* visitor) { visitor_ = visitor; } - void Destroy() { - if (request_stream() != nullptr) { - request_stream()->Fin(); - } - if (remove_callback_) { - // Null the callback before calling, because the session owns this and - // the callback will call the destructor. When SUBSCRIBE moves to a stream - // this won't be a problem. - RemoveCallback callback = std::move(remove_callback_); - remove_callback_ = nullptr; - std::move(callback)(this); - } - } + void SendObjectAck(uint64_t group_id, uint64_t object_id, + quic::QuicTimeDelta delta_from_deadline); private: friend class test::MoqtSessionPeer; @@ -189,13 +171,19 @@ class PublishDoneDelegate : public quic::QuicAlarm::DelegateWithoutContext { public: - PublishDoneDelegate(SubscribeRemoteTrack* subscribe) + PublishDoneDelegate(quiche::QuicheWeakPtr<RemoteTrack> subscribe) : subscribe_(subscribe) {} - void OnAlarm() override { subscribe_->Destroy(); } + void OnAlarm() override { + RemoteTrack* subscribe = subscribe_.GetIfAvailable(); + if (subscribe == nullptr) { + return; + } + subscribe->request_stream()->Reset(kResetCodeCancelled); + } private: - SubscribeRemoteTrack* subscribe_; + quiche::QuicheWeakPtr<RemoteTrack> subscribe_; }; void MaybeSetPublishDoneAlarm(); @@ -216,9 +204,6 @@ int currently_open_streams_ = 0; // Every stream that has received FIN or RESET_STREAM. uint64_t streams_closed_ = 0; - // For PUBLISH (and later SUBSCRIBE), will be handled in the request stream. - AddCallback add_callback_ = nullptr; - RemoveCallback remove_callback_ = nullptr; // Value assigned on PUBLISH_DONE. Can destroy subscription state if // streams_closed_ == total_streams_. std::optional<uint64_t> total_streams_; @@ -292,7 +277,7 @@ // Called when the data stream is destroyed. void OnStreamClosed() { Destroy(); } - void Destroy() override { + void Destroy() { if (remove_callback_) { RemoveFetchCallback callback = std::move(remove_callback_); remove_callback_ = nullptr; @@ -421,7 +406,6 @@ // Initial values from Fetch() call. FetchResponseCallback ok_callback_; // Will be destroyed on FETCH_OK. - RemoveFetchCallback remove_callback_; };
diff --git a/quiche/quic/moqt/moqt_track_test.cc b/quiche/quic/moqt/moqt_track_test.cc index 010cda9..bb50891 100644 --- a/quiche/quic/moqt/moqt_track_test.cc +++ b/quiche/quic/moqt/moqt_track_test.cc
@@ -10,17 +10,20 @@ #include "absl/status/status.h" #include "quiche/quic/core/quic_alarm.h" +#include "quiche/quic/moqt/moqt_error.h" #include "quiche/quic/moqt/moqt_fetch_task.h" #include "quiche/quic/moqt/moqt_key_value_pair.h" #include "quiche/quic/moqt/moqt_messages.h" #include "quiche/quic/moqt/moqt_names.h" #include "quiche/quic/moqt/moqt_object.h" #include "quiche/quic/moqt/moqt_types.h" +#include "quiche/quic/moqt/test_tools/mock_moqt_session.h" #include "quiche/quic/moqt/test_tools/moqt_mock_visitor.h" #include "quiche/quic/platform/api/quic_test.h" #include "quiche/quic/test_tools/mock_clock.h" #include "quiche/quic/test_tools/quic_test_utils.h" #include "quiche/common/quiche_mem_slice.h" +#include "quiche/web_transport/test_tools/mock_web_transport.h" #include "quiche/web_transport/web_transport.h" namespace moqt { @@ -52,21 +55,16 @@ class SubscribeRemoteTrackTest : public quic::test::QuicTest { public: - SubscribeRemoteTrackTest() - : track_( - subscribe_, &visitor_, - [&](SubscribeRemoteTrack*) { - alias_registered_ = true; - return true; - }, - [this](SubscribeRemoteTrack*) { deleted_ = true; }) {} + SubscribeRemoteTrackTest() : track_(subscribe_, &visitor_, &stream_) { + stream_.BindStream(&wt_stream_); + } MockSubscribeRemoteTrackVisitor visitor_; MoqtSubscribe subscribe_ = {/*request_id=*/1, FullTrackName("foo", "bar"), MessageParameters(Location(2, 0))}; + MockBidiStream stream_; + webtransport::test::MockStream wt_stream_; SubscribeRemoteTrack track_; - bool alias_registered_ = false; - bool deleted_ = false; quic::MockClock clock_; quic::test::MockAlarmFactory alarm_factory_; }; @@ -77,9 +75,8 @@ EXPECT_FALSE(track_.track_alias().has_value()); EXPECT_EQ(track_.visitor(), &visitor_); EXPECT_FALSE(track_.is_fetch()); - EXPECT_TRUE(track_.set_track_alias(1)); + track_.set_track_alias(1); EXPECT_EQ(track_.track_alias(), 1); - EXPECT_TRUE(alias_registered_); } TEST_F(SubscribeRemoteTrackTest, AllowError) { @@ -97,34 +94,35 @@ track_.OnStreamOpened(); track_.OnStreamClosed(true, std::nullopt); EXPECT_CALL(visitor_, OnPublishDone); + ExpectFin(wt_stream_); track_.OnPublishDone(1, &clock_, &alarm_factory_); - EXPECT_TRUE(deleted_); } TEST_F(SubscribeRemoteTrackTest, OnPublishDoneAllStreamsCloseLater) { track_.OnStreamOpened(); EXPECT_CALL(visitor_, OnPublishDone).Times(0); + EXPECT_CALL(wt_stream_, Writev).Times(0); track_.OnPublishDone(2, &clock_, &alarm_factory_); track_.OnStreamClosed(true, std::nullopt); track_.OnStreamOpened(); + ExpectFin(wt_stream_); EXPECT_CALL(visitor_, OnPublishDone); track_.OnStreamClosed(true, std::nullopt); - EXPECT_TRUE(deleted_); } TEST_F(SubscribeRemoteTrackTest, OnPublishDoneTimesOut) { track_.OnStreamOpened(); EXPECT_CALL(visitor_, OnPublishDone).Times(0); + EXPECT_CALL(wt_stream_, Writev).Times(0); track_.OnPublishDone(2, &clock_, &alarm_factory_); - EXPECT_FALSE(deleted_); track_.OnStreamClosed(true, std::nullopt); // No streams are open; timer set. quic::QuicAlarm* alarm = SubscribeRemoteTrackPeer::GetPublishDoneAlarm(&track_); EXPECT_NE(alarm, nullptr); EXPECT_TRUE(alarm->IsSet()); EXPECT_CALL(visitor_, OnPublishDone); + EXPECT_CALL(wt_stream_, ResetWithUserCode(kResetCodeCancelled)); alarm_factory_.FireAlarm(alarm); - EXPECT_TRUE(deleted_); } TEST_F(SubscribeRemoteTrackTest, JoiningFetchMultiObject) {
diff --git a/quiche/quic/moqt/moqt_uni_stream_test.cc b/quiche/quic/moqt/moqt_uni_stream_test.cc index 9efcabd..65c8ad2 100644 --- a/quiche/quic/moqt/moqt_uni_stream_test.cc +++ b/quiche/quic/moqt/moqt_uni_stream_test.cc
@@ -491,15 +491,9 @@ subscribe_message_(1, ftn_, MessageParameters()) { EXPECT_CALL(session_, deliver_partial_objects()) .WillRepeatedly(Return(false)); - track_ = std::make_unique<SubscribeRemoteTrack>( - subscribe_message_, &visitor_, - [this](SubscribeRemoteTrack* track) { - alias_track_ = track; - alias_ = track->track_alias().value(); - return true; - }, - nullptr); - EXPECT_TRUE(track_->set_track_alias(2)); + track_ = std::make_unique<SubscribeRemoteTrack>(subscribe_message_, + &visitor_, nullptr); + track_->set_track_alias(2); CreateStream(); } @@ -519,8 +513,6 @@ EXPECT_CALL(session_, GetSubscribe(alias)) .WillOnce(Return(track_->weak_ptr())); stream_->OnCanRead(); - EXPECT_EQ(alias_, alias); - EXPECT_EQ(alias_track_, track_.get()); } webtransport::test::InMemoryStream mock_stream_; @@ -531,8 +523,6 @@ testing::NiceMock<MockSubscribeRemoteTrackVisitor> visitor_; std::unique_ptr<SubscribeRemoteTrack> track_; std::unique_ptr<IncomingDataStream> stream_; - uint64_t alias_ = 0; - SubscribeRemoteTrack* alias_track_ = nullptr; }; TEST_F(IncomingDataStreamTest, DestructorBeforeTrackAlias) {
diff --git a/quiche/quic/moqt/test_tools/mock_moqt_session.h b/quiche/quic/moqt/test_tools/mock_moqt_session.h index c09c8c6..8583e26 100644 --- a/quiche/quic/moqt/test_tools/mock_moqt_session.h +++ b/quiche/quic/moqt/test_tools/mock_moqt_session.h
@@ -8,22 +8,59 @@ #include <cstdint> #include <memory> #include <optional> +#include <utility> +#include "absl/base/nullability.h" +#include "absl/status/status.h" #include "absl/strings/string_view.h" +#include "absl/types/span.h" +#include "quiche/quic/core/quic_alarm_factory.h" +#include "quiche/quic/core/quic_clock.h" +#include "quiche/quic/core/quic_types.h" +#include "quiche/quic/moqt/moqt_bidi_stream.h" #include "quiche/quic/moqt/moqt_error.h" #include "quiche/quic/moqt/moqt_fetch_task.h" +#include "quiche/quic/moqt/moqt_framer.h" #include "quiche/quic/moqt/moqt_key_value_pair.h" +#include "quiche/quic/moqt/moqt_messages.h" #include "quiche/quic/moqt/moqt_names.h" +#include "quiche/quic/moqt/moqt_parser.h" +#include "quiche/quic/moqt/moqt_priority.h" +#include "quiche/quic/moqt/moqt_publisher.h" #include "quiche/quic/moqt/moqt_session_callbacks.h" #include "quiche/quic/moqt/moqt_session_interface.h" +#include "quiche/quic/moqt/moqt_subscription.h" +#include "quiche/quic/moqt/moqt_trace_recorder.h" #include "quiche/quic/moqt/moqt_types.h" #include "quiche/common/platform/api/quiche_test.h" -#include "quiche/common/quiche_callbacks.h" +#include "quiche/common/quiche_mem_slice.h" #include "quiche/common/quiche_weak_ptr.h" +#include "quiche/web_transport/test_tools/mock_web_transport.h" +#include "quiche/web_transport/web_transport.h" namespace moqt { namespace test { +class MockSessionToPublisherInterface : public SessionToPublisherInterface { + public: + MockSessionToPublisherInterface() : weak_ptr_factory_(this) {} + ~MockSessionToPublisherInterface() override = default; + MOCK_METHOD(bool, alternate_delivery_timeout, (), (const, override)); + MOCK_METHOD(void, UpdateTrackPriority, + (uint64_t, std::optional<MoqtTrackPriority>, MoqtTrackPriority), + (override)); + MOCK_METHOD(quic::QuicAlarmFactory*, alarm_factory, (), (override)); + MOCK_METHOD(std::shared_ptr<MoqtTrackPublisher>, GetTrackPublisher, + (const FullTrackName&), (override)); + MOCK_METHOD(MoqtPublishingMonitorInterface*, ReleaseMonitoringInterface, + (const FullTrackName&), (override)); + MOCK_METHOD(const quic::QuicClock*, clock, (), (override)); + MOCK_METHOD(MoqtTraceRecorder&, trace_recorder, (), (override)); + MOCK_METHOD(webtransport::Session*, session, (), (override)); + + quiche::QuicheWeakPtrFactory<SessionToPublisherInterface> weak_ptr_factory_; +}; + class MockMoqtSession : public MoqtSessionInterface { public: MOCK_METHOD(MoqtSessionCallbacks&, callbacks, (), (override)); @@ -37,6 +74,10 @@ (const FullTrackName&, const MessageParameters&, MoqtResponseCallback), (override)); + MOCK_METHOD(bool, PublishUpdate, + (const FullTrackName& name, const MessageParameters& parameters, + MoqtResponseCallback response_callback), + (override)); MOCK_METHOD(void, Unsubscribe, (const FullTrackName& name), (override)); MOCK_METHOD(bool, Publish, (std::shared_ptr<MoqtTrackPublisher> publisher, @@ -88,6 +129,45 @@ quiche::QuicheWeakPtrFactory<MoqtSessionInterface> weak_factory_{this}; }; +inline void ExpectFin(webtransport::test::MockStream& stream, + bool has_data = false) { + EXPECT_CALL(stream, Writev) + .WillOnce([&](absl::Span<quiche::QuicheMemSlice> data, + const webtransport::StreamWriteOptions& options) { + EXPECT_EQ(data.empty(), !has_data); + EXPECT_TRUE(options.send_fin()); + return absl::OkStatus(); + }); +} + +class MockBidiStream : public MoqtBidiStreamBase { + public: + MockBidiStream() + : MoqtBidiStreamBase(reinterpret_cast<MoqtFramer*>(0x1), + MoqtControlMessageParser( + "moqt-00", true, quic::Perspective::IS_CLIENT), + nullptr) {} + MockBidiStream(MoqtFramer* absl_nonnull framer, + const MoqtControlMessageParser& message_parser, + SessionErrorCallback session_error_callback) + : MoqtBidiStreamBase(framer, message_parser, + std::move(session_error_callback)) {} + + MOCK_METHOD(absl::Status, SendRequestUpdate, + (uint64_t request_id, uint64_t existing_request_id, + const MessageParameters& parameters, + MoqtResponseCallback callback), + (override)); + MOCK_METHOD(void, Detach, (), (override)); + MOCK_METHOD(void, OnStreamBound, (), (override)); + MOCK_METHOD(absl::Status, OnRawControlMessage, + (const MoqtRawControlMessage& message), (override)); + MOCK_METHOD(absl::Status, OnControlMessage, (const MoqtRequestOk& message), + (override)); + MOCK_METHOD(absl::Status, OnControlMessage, (const MoqtRequestError& message), + (override)); +}; + } // namespace test } // namespace moqt
diff --git a/quiche/quic/moqt/test_tools/moqt_framer_utils.cc b/quiche/quic/moqt/test_tools/moqt_framer_utils.cc index 7214426..a0e6d42 100644 --- a/quiche/quic/moqt/test_tools/moqt_framer_utils.cc +++ b/quiche/quic/moqt/test_tools/moqt_framer_utils.cc
@@ -31,9 +31,6 @@ quiche::QuicheBuffer operator()(const MoqtSubscribeOk& message) { return framer.SerializeSubscribeOk(message); } - quiche::QuicheBuffer operator()(const MoqtUnsubscribe& message) { - return framer.SerializeUnsubscribe(message); - } quiche::QuicheBuffer operator()(const MoqtPublishDone& message) { return framer.SerializePublishDone(message); }
diff --git a/quiche/quic/moqt/test_tools/moqt_framer_utils.h b/quiche/quic/moqt/test_tools/moqt_framer_utils.h index be2f36b..c5a5d4a 100644 --- a/quiche/quic/moqt/test_tools/moqt_framer_utils.h +++ b/quiche/quic/moqt/test_tools/moqt_framer_utils.h
@@ -19,13 +19,14 @@ namespace moqt::test { -using AnyMoqtControlMessage = std::variant< - MoqtSetup, MoqtRequestOk, MoqtRequestError, MoqtSubscribe, MoqtSubscribeOk, - MoqtUnsubscribe, MoqtPublishDone, MoqtRequestUpdate, MoqtPublishNamespace, - MoqtPublishNamespaceDone, MoqtPublishNamespaceCancel, MoqtTrackStatus, - MoqtGoAway, MoqtSubscribeNamespace, MoqtMaxRequestId, MoqtFetch, - MoqtFetchCancel, MoqtFetchOk, MoqtRequestsBlocked, MoqtPublish, - MoqtNamespace, MoqtNamespaceDone, MoqtObjectAck>; +using AnyMoqtControlMessage = + std::variant<MoqtSetup, MoqtRequestOk, MoqtRequestError, MoqtSubscribe, + MoqtSubscribeOk, MoqtPublishDone, MoqtRequestUpdate, + MoqtPublishNamespace, MoqtPublishNamespaceDone, + MoqtPublishNamespaceCancel, MoqtTrackStatus, MoqtGoAway, + MoqtSubscribeNamespace, MoqtMaxRequestId, MoqtFetch, + MoqtFetchCancel, MoqtFetchOk, MoqtRequestsBlocked, MoqtPublish, + MoqtNamespace, MoqtNamespaceDone, MoqtObjectAck>; std::string SerializeGenericMessage(const AnyMoqtControlMessage& frame, bool use_webtrans = false);
diff --git a/quiche/quic/moqt/test_tools/moqt_session_peer.h b/quiche/quic/moqt/test_tools/moqt_session_peer.h index 3bb847b..c7bdb73 100644 --- a/quiche/quic/moqt/test_tools/moqt_session_peer.h +++ b/quiche/quic/moqt/test_tools/moqt_session_peer.h
@@ -193,6 +193,10 @@ ? std::optional<uint64_t>() : session->subscriptions_with_queued_streams_.begin()->second; } + + static uint64_t GetLastTrackAlias(MoqtSession* session) { + return session->next_local_track_alias_ - 1; + } }; } // namespace moqt::test
diff --git a/quiche/quic/moqt/test_tools/moqt_test_message.h b/quiche/quic/moqt/test_tools/moqt_test_message.h index 4a01ac2..4e1d394 100644 --- a/quiche/quic/moqt/test_tools/moqt_test_message.h +++ b/quiche/quic/moqt/test_tools/moqt_test_message.h
@@ -112,15 +112,13 @@ public: virtual ~TestMessageBase() = default; - using MessageStructuredData = - std::variant<MoqtSetup, MoqtObject, MoqtRequestOk, MoqtRequestError, - MoqtSubscribe, MoqtSubscribeOk, MoqtUnsubscribe, - MoqtPublishDone, MoqtRequestUpdate, MoqtPublishNamespace, - MoqtPublishNamespaceDone, MoqtPublishNamespaceCancel, - MoqtTrackStatus, MoqtGoAway, MoqtSubscribeNamespace, - MoqtMaxRequestId, MoqtFetch, MoqtFetchCancel, MoqtFetchOk, - MoqtRequestsBlocked, MoqtPublish, MoqtNamespace, - MoqtNamespaceDone, MoqtObjectAck>; + using MessageStructuredData = std::variant< + MoqtSetup, MoqtObject, MoqtRequestOk, MoqtRequestError, MoqtSubscribe, + MoqtSubscribeOk, MoqtPublishDone, MoqtRequestUpdate, MoqtPublishNamespace, + MoqtPublishNamespaceDone, MoqtPublishNamespaceCancel, MoqtTrackStatus, + MoqtGoAway, MoqtSubscribeNamespace, MoqtMaxRequestId, MoqtFetch, + MoqtFetchCancel, MoqtFetchOk, MoqtRequestsBlocked, MoqtPublish, + MoqtNamespace, MoqtNamespaceDone, MoqtObjectAck>; // The total actual size of the message. size_t total_message_size() const { return wire_image_size_; } @@ -872,39 +870,6 @@ }; }; -class QUICHE_NO_EXPORT UnsubscribeMessage : public TestMessageBase { - public: - UnsubscribeMessage() : TestMessageBase() { - SetWireImage(raw_packet_, sizeof(raw_packet_)); - } - - bool EqualFieldValues(const MessageStructuredData& values) const override { - auto cast = std::get<MoqtUnsubscribe>(values); - if (cast.request_id != unsubscribe_.request_id) { - QUIC_LOG(INFO) << "UNSUBSCRIBE request ID mismatch"; - return false; - } - return true; - } - - void ExpandVarints() override { ExpandVarintsImpl("v"); } - - MessageStructuredData structured_data() const override { - return TestMessageBase::MessageStructuredData(unsubscribe_); - } - - private: - uint8_t raw_packet_[4] = { - 0x0a, - 0x00, - 0x01, - 0x03, // request_id = 3 - }; - - MoqtUnsubscribe unsubscribe_ = { - /*request_id=*/3, - }; -}; class QUICHE_NO_EXPORT PublishDoneMessage : public TestMessageBase { public: @@ -1764,10 +1729,6 @@ bool EqualFieldValues(const MessageStructuredData& values) const override { auto cast = std::get<MoqtObjectAck>(values); - if (cast.subscribe_id != object_ack_.subscribe_id) { - QUIC_LOG(INFO) << "OBJECT_ACK subscribe ID mismatch"; - return false; - } if (cast.group_id != object_ack_.group_id) { QUIC_LOG(INFO) << "OBJECT_ACK group ID mismatch"; return false; @@ -1783,21 +1744,20 @@ return true; } - void ExpandVarints() override { ExpandVarintsImpl("vvvv"); } + void ExpandVarints() override { ExpandVarintsImpl("vvv"); } MessageStructuredData structured_data() const override { return TestMessageBase::MessageStructuredData(object_ack_); } private: - uint8_t raw_packet_[8] = { - 0xb1, 0x84, 0x00, 0x04, // type - 0x01, 0x10, 0x20, // subscribe ID, group, object + uint8_t raw_packet_[7] = { + 0xb1, 0x84, 0x00, 0x03, // type + 0x10, 0x20, // group, object 0x20, // 0x10 time delta }; MoqtObjectAck object_ack_ = { - /*subscribe_id=*/0x01, /*group_id=*/0x10, /*object_id=*/0x20, /*delta_from_deadline=*/quic::QuicTimeDelta::FromMicroseconds(0x10), @@ -1817,8 +1777,6 @@ return std::make_unique<SubscribeMessage>(); case MoqtMessageType::kSubscribeOk: return std::make_unique<SubscribeOkMessage>(); - case MoqtMessageType::kUnsubscribe: - return std::make_unique<UnsubscribeMessage>(); case MoqtMessageType::kPublishDone: return std::make_unique<PublishDoneMessage>(); case MoqtMessageType::kRequestUpdate: