Add SubscribeNextGroup to the MoqtSession API PiperOrigin-RevId: 757866347
diff --git a/quiche/quic/moqt/moqt_integration_test.cc b/quiche/quic/moqt/moqt_integration_test.cc index 3ece1b6..8d97e6e 100644 --- a/quiche/quic/moqt/moqt_integration_test.cc +++ b/quiche/quic/moqt/moqt_integration_test.cc
@@ -520,7 +520,7 @@ EXPECT_TRUE(success); } -TEST_F(MoqtIntegrationTest, SubscribeCurrentGroupOk) { +TEST_F(MoqtIntegrationTest, SubscribeNextGroupOk) { EstablishSession(); FullTrackName full_track_name("foo", "bar"); @@ -539,8 +539,8 @@ }); EXPECT_CALL(client_visitor, OnReply(full_track_name, _, expected_reason)) .WillOnce([&]() { received_ok = true; }); - client_->session()->SubscribeCurrentObject(full_track_name, &client_visitor, - VersionSpecificParameters()); + client_->session()->SubscribeNextGroup(full_track_name, &client_visitor, + VersionSpecificParameters()); bool success = test_harness_.RunUntilWithDefaultTimeout([&]() { return received_ok; }); EXPECT_TRUE(success);
diff --git a/quiche/quic/moqt/moqt_session.cc b/quiche/quic/moqt/moqt_session.cc index 5645885..7ea870d 100644 --- a/quiche/quic/moqt/moqt_session.cc +++ b/quiche/quic/moqt/moqt_session.cc
@@ -57,10 +57,6 @@ using ::quic::Perspective; -constexpr MoqtPriority kDefaultSubscriberPriority = 0x80; -constexpr quic::QuicTimeDelta kDefaultGoAwayTimeout = - quic::QuicTime::Delta::FromSeconds(10); - // WebTransport lets applications split a session into multiple send groups // that have equal weight for scheduling. We don't have a use for that, so the // send group is always the same. @@ -384,6 +380,21 @@ return Subscribe(message, visitor); } +bool MoqtSession::SubscribeNextGroup(const FullTrackName& name, + SubscribeRemoteTrack::Visitor* visitor, + VersionSpecificParameters parameters) { + MoqtSubscribe message; + message.full_track_name = name; + message.subscriber_priority = kDefaultSubscriberPriority; + message.group_order = std::nullopt; + message.forward = true; + message.filter_type = MoqtFilterType::kNextGroupStart; + message.start = std::nullopt; + message.end_group = std::nullopt; + message.parameters = parameters; + return Subscribe(message, visitor); +} + void MoqtSession::Unsubscribe(const FullTrackName& name) { SubscribeRemoteTrack* track = RemoteTrackByName(name); if (track == nullptr) {
diff --git a/quiche/quic/moqt/moqt_session.h b/quiche/quic/moqt/moqt_session.h index bc4e80d..85e20ce 100644 --- a/quiche/quic/moqt/moqt_session.h +++ b/quiche/quic/moqt/moqt_session.h
@@ -44,6 +44,10 @@ class MoqtSessionPeer; } +inline constexpr MoqtPriority kDefaultSubscriberPriority = 0x80; +inline constexpr quic::QuicTimeDelta kDefaultGoAwayTimeout = + quic::QuicTime::Delta::FromSeconds(10); + struct SubscriptionWithQueuedStream { webtransport::SendOrder send_order; uint64_t subscription_id; @@ -126,7 +130,9 @@ bool SubscribeCurrentObject(const FullTrackName& name, SubscribeRemoteTrack::Visitor* visitor, VersionSpecificParameters parameters) override; - // TODO(martinduke): SubscribeNextGroup + bool SubscribeNextGroup(const FullTrackName& name, + SubscribeRemoteTrack::Visitor* visitor, + VersionSpecificParameters parameters) override; // Returns false if the subscription is not found. The session immediately // destroys all subscription state. void Unsubscribe(const FullTrackName& name);
diff --git a/quiche/quic/moqt/moqt_session_interface.h b/quiche/quic/moqt/moqt_session_interface.h index cdd197e..1bf2eb2 100644 --- a/quiche/quic/moqt/moqt_session_interface.h +++ b/quiche/quic/moqt/moqt_session_interface.h
@@ -61,6 +61,10 @@ virtual bool SubscribeCurrentObject(const FullTrackName& name, SubscribeRemoteTrack::Visitor* visitor, VersionSpecificParameters parameters) = 0; + // Start with the first group after the current Largest Group/Object ID. + virtual bool SubscribeNextGroup(const FullTrackName& name, + SubscribeRemoteTrack::Visitor* visitor, + VersionSpecificParameters parameters) = 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 1487e25..7e1b8df 100644 --- a/quiche/quic/moqt/moqt_session_test.cc +++ b/quiche/quic/moqt/moqt_session_test.cc
@@ -697,6 +697,43 @@ stream_input->OnSubscribeOkMessage(ok); } +TEST_F(MoqtSessionTest, SubscribeNextGroupWithOk) { + std::unique_ptr<MoqtControlParserVisitor> stream_input = + MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_); + MockSubscribeRemoteTrackVisitor remote_track_visitor; + EXPECT_CALL(mock_session_, GetStreamById(_)).WillOnce(Return(&mock_stream_)); + MoqtSubscribe subscribe = { + /*request_id=*/0, + /*track_alias=*/0, + FullTrackName("foo", "bar"), + kDefaultSubscriberPriority, + /*group_order=*/std::nullopt, + /*forward=*/true, + MoqtFilterType::kNextGroupStart, + std::nullopt, + std::nullopt, + VersionSpecificParameters(), + }; + subscribe.filter_type = MoqtFilterType::kNextGroupStart; + EXPECT_CALL(mock_stream_, Writev(SerializedControlMessage(subscribe), _)); + session_.SubscribeNextGroup(FullTrackName("foo", "bar"), + &remote_track_visitor, + VersionSpecificParameters()); + + MoqtSubscribeOk ok = { + /*request_id=*/0, + /*expires=*/quic::QuicTimeDelta::FromMilliseconds(0), + }; + EXPECT_CALL(remote_track_visitor, OnReply(_, _, _)) + .WillOnce([&](const FullTrackName& ftn, + std::optional<Location> /*largest_id*/, + std::optional<absl::string_view> error_message) { + EXPECT_EQ(ftn, FullTrackName("foo", "bar")); + EXPECT_FALSE(error_message.has_value()); + }); + stream_input->OnSubscribeOkMessage(ok); +} + TEST_F(MoqtSessionTest, MaxRequestIdChangesResponse) { MoqtSessionPeer::set_next_request_id(&session_, kDefaultInitialMaxRequestId); MockSubscribeRemoteTrackVisitor remote_track_visitor;
diff --git a/quiche/quic/moqt/test_tools/mock_moqt_session.h b/quiche/quic/moqt/test_tools/mock_moqt_session.h index 50405e4..7131181 100644 --- a/quiche/quic/moqt/test_tools/mock_moqt_session.h +++ b/quiche/quic/moqt/test_tools/mock_moqt_session.h
@@ -55,6 +55,11 @@ SubscribeRemoteTrack::Visitor* visitor, VersionSpecificParameters parameters), (override)); + MOCK_METHOD(bool, SubscribeNextGroup, + (const FullTrackName& name, + SubscribeRemoteTrack::Visitor* visitor, + VersionSpecificParameters parameters), + (override)); MOCK_METHOD(void, Unsubscribe, (const FullTrackName& name), (override)); MOCK_METHOD(bool, Fetch, (const FullTrackName& name, FetchResponseCallback callback,