Update SUBSCRIBE_OK and REQUEST_OK to draft-18. Update ALPN to moqt-18. Receipt of non-empty TrackExtensions on any REQUEST_OK except TRACK_STATUS_OK is a PROTOCOL_ERROR. MOQT is not in production. PiperOrigin-RevId: 983353332
diff --git a/quiche/quic/moqt/moqt_bidi_stream.cc b/quiche/quic/moqt/moqt_bidi_stream.cc index cb236c3..6d748a0 100644 --- a/quiche/quic/moqt/moqt_bidi_stream.cc +++ b/quiche/quic/moqt/moqt_bidi_stream.cc
@@ -59,9 +59,9 @@ } absl::Status MoqtBidiStreamBase::SendRequestOk( - uint64_t request_id, const MessageParameters& parameters, bool fin) { + const MessageParameters& parameters) { return SendOrBufferMessage( - framer_->SerializeRequestOk(MoqtRequestOk{request_id, parameters}), fin); + framer_->SerializeRequestOk(MoqtRequestOk(parameters)), false); } absl::Status MoqtBidiStreamBase::SendRequestError(
diff --git a/quiche/quic/moqt/moqt_bidi_stream.h b/quiche/quic/moqt/moqt_bidi_stream.h index a96b422..dd6ae15 100644 --- a/quiche/quic/moqt/moqt_bidi_stream.h +++ b/quiche/quic/moqt/moqt_bidi_stream.h
@@ -6,14 +6,12 @@ #define QUICHE_QUIC_MOQT_MOQT_BIDI_STREAM_H #include <cstdint> -#include <memory> #include <optional> #include <type_traits> #include <utility> #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" @@ -95,9 +93,9 @@ CheckStatus(SendOrBufferMessage(std::move(message), fin)); } - absl::Status SendRequestOk(uint64_t request_id, - const MessageParameters& parameters, - bool fin = false); + // Do not use for TRACK_STATUS_OK because that should also contain + // TrackProperties. + absl::Status SendRequestOk(const MessageParameters& parameters); absl::Status SendRequestError( RequestErrorCode error_code, std::optional<quic::QuicTimeDelta> retry_interval,
diff --git a/quiche/quic/moqt/moqt_bidi_stream_test.cc b/quiche/quic/moqt/moqt_bidi_stream_test.cc index 99e529a..bfc110e 100644 --- a/quiche/quic/moqt/moqt_bidi_stream_test.cc +++ b/quiche/quic/moqt/moqt_bidi_stream_test.cc
@@ -140,19 +140,8 @@ 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_); + QUICHE_EXPECT_OK(stream_->SendRequestOk(parameters)); + EXPECT_FALSE(stream_->detached_); // No FIN. } TEST_F(MoqtBidiStreamTest, SendRequestErrorOverload) { @@ -185,7 +174,6 @@ 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);
diff --git a/quiche/quic/moqt/moqt_fetch_stream.cc b/quiche/quic/moqt/moqt_fetch_stream.cc index 2d8e43c..724221e 100644 --- a/quiche/quic/moqt/moqt_fetch_stream.cc +++ b/quiche/quic/moqt/moqt_fetch_stream.cc
@@ -155,6 +155,11 @@ absl::Status MoqtFetchRequestStream::OnControlMessage( const MoqtRequestOk& message) { + if (!message.extensions.empty()) { + OnFatalError(absl::InvalidArgumentError( + "REQUEST_UPDATE_OK received with extensions")); + return absl::OkStatus(); + } absl::StatusOr<MessageParameters> old_parameters = request_update_queue().NextParameters(); if (!old_parameters.ok()) { @@ -367,7 +372,7 @@ message.parameters.subscriber_priority.has_value()) { data_stream_->UpdatePriority(*message.parameters.subscriber_priority); } - return SendRequestOk(message.request_id, MessageParameters()); + return SendRequestOk(MessageParameters()); } void MoqtFetchResponseStream::OnDataStreamOpen(
diff --git a/quiche/quic/moqt/moqt_fetch_stream_test.cc b/quiche/quic/moqt/moqt_fetch_stream_test.cc index 69ecec0..bd2e970 100644 --- a/quiche/quic/moqt/moqt_fetch_stream_test.cc +++ b/quiche/quic/moqt/moqt_fetch_stream_test.cc
@@ -301,7 +301,6 @@ // Receive REQUEST_OK for the update. MoqtRequestOk request_ok; - request_ok.request_id = 2; request_ok.parameters.subscriber_priority = 50; QUICHE_EXPECT_OK(stream->OnControlMessage(request_ok)); EXPECT_TRUE(update_callback_called); @@ -317,7 +316,6 @@ ok_message.end_of_track = true; QUICHE_EXPECT_OK(stream->OnControlMessage(ok_message)); MoqtRequestOk request_ok; - request_ok.request_id = 2; EXPECT_EQ(stream->OnControlMessage(request_ok).code(), absl::StatusCode::kFailedPrecondition); }
diff --git a/quiche/quic/moqt/moqt_framer.cc b/quiche/quic/moqt/moqt_framer.cc index 70cc872..e6cc15a 100644 --- a/quiche/quic/moqt/moqt_framer.cc +++ b/quiche/quic/moqt/moqt_framer.cc
@@ -503,8 +503,9 @@ quiche::QuicheBuffer MoqtFramer::SerializeRequestOk( const MoqtRequestOk& message) { return SerializeControlMessage( - MoqtMessageType::kRequestOk, WireMoqVarInt(message.request_id), - WireKeyValuePairList(message.parameters.ToKeyValuePairList())); + MoqtMessageType::kRequestOk, + WireKeyValuePairList(message.parameters.ToKeyValuePairList()), + WireKeyValuePairList(message.extensions, false)); } quiche::QuicheBuffer MoqtFramer::SerializeSubscribe( @@ -523,8 +524,7 @@ return quiche::QuicheBuffer(); } return SerializeControlMessage( - message_type, WireMoqVarInt(message.request_id), - WireMoqVarInt(message.track_alias), + message_type, WireMoqVarInt(message.track_alias), WireKeyValuePairList(message.parameters.ToKeyValuePairList()), WireKeyValuePairList(message.extensions, false)); }
diff --git a/quiche/quic/moqt/moqt_integration_test.cc b/quiche/quic/moqt/moqt_integration_test.cc index ee02ac3..01aea7f 100644 --- a/quiche/quic/moqt/moqt_integration_test.cc +++ b/quiche/quic/moqt/moqt_integration_test.cc
@@ -1379,10 +1379,10 @@ MessageParameters received_parameters; client_->session()->TrackStatus( track_name, MessageParameters(), - [&](std::variant<MessageParameters, MoqtRequestErrorInfo> response) { + [&](std::variant<TrackStatusOkData, MoqtRequestErrorInfo> response) { received_response = true; - ASSERT_TRUE(std::holds_alternative<MessageParameters>(response)); - received_parameters = std::get<MessageParameters>(response); + ASSERT_TRUE(std::holds_alternative<TrackStatusOkData>(response)); + received_parameters = std::get<TrackStatusOkData>(response).parameters; }); bool success = test_harness_.RunUntilWithDefaultTimeout( @@ -1403,7 +1403,7 @@ MoqtRequestErrorInfo received_error; client_->session()->TrackStatus( track_name, MessageParameters(), - [&](std::variant<MessageParameters, MoqtRequestErrorInfo> response) { + [&](std::variant<TrackStatusOkData, MoqtRequestErrorInfo> response) { received_response = true; ASSERT_TRUE(std::holds_alternative<MoqtRequestErrorInfo>(response)); received_error = std::get<MoqtRequestErrorInfo>(response);
diff --git a/quiche/quic/moqt/moqt_key_value_pair.h b/quiche/quic/moqt/moqt_key_value_pair.h index 90a9fdc..4d81519 100644 --- a/quiche/quic/moqt/moqt_key_value_pair.h +++ b/quiche/quic/moqt/moqt_key_value_pair.h
@@ -279,6 +279,7 @@ MoqtPriority default_publisher_priority() const; MoqtDeliveryOrder default_publisher_group_order() const; bool dynamic_groups() const; + bool empty() const { return size() == 0; } // Returns false if the extension list contains illegal values or illegally // duplicated extensions.
diff --git a/quiche/quic/moqt/moqt_live_publisher.cc b/quiche/quic/moqt/moqt_live_publisher.cc index fc91cd6..3a886e4 100644 --- a/quiche/quic/moqt/moqt_live_publisher.cc +++ b/quiche/quic/moqt/moqt_live_publisher.cc
@@ -133,7 +133,6 @@ parameters_.largest_object); } MoqtSubscribeOk subscribe_ok; - subscribe_ok.request_id = request_id_; subscribe_ok.track_alias = track_alias_; subscribe_ok.parameters.expires = track_publisher_->expiration(); subscribe_ok.parameters.largest_object = parameters_.largest_object;
diff --git a/quiche/quic/moqt/moqt_messages.h b/quiche/quic/moqt/moqt_messages.h index 2be9f7f..ebbf785 100644 --- a/quiche/quic/moqt/moqt_messages.h +++ b/quiche/quic/moqt/moqt_messages.h
@@ -371,7 +371,6 @@ }; struct QUICHE_EXPORT MoqtSubscribeOk { - uint64_t request_id; uint64_t track_alias; MessageParameters parameters; TrackExtensions extensions; @@ -396,10 +395,7 @@ MessageParameters parameters; }; -struct QUICHE_EXPORT MoqtRequestOk { - uint64_t request_id; - MessageParameters parameters; -}; +using MoqtRequestOk = TrackStatusOkData; struct QUICHE_EXPORT MoqtTrackStatus : public MoqtSubscribe { MoqtTrackStatus() = default;
diff --git a/quiche/quic/moqt/moqt_namespace_stream.cc b/quiche/quic/moqt/moqt_namespace_stream.cc index c7b6cb2..d5bf200 100644 --- a/quiche/quic/moqt/moqt_namespace_stream.cc +++ b/quiche/quic/moqt/moqt_namespace_stream.cc
@@ -47,13 +47,15 @@ absl::Status MoqtSubscribeNamespaceRequestStream::OnControlMessage( const MoqtRequestOk& message) { - if (message.request_id == request_id_) { - // Response to the initial SUBSCRIBE_NAMESPACE. - if (response_callback_ == nullptr) { - return absl::InvalidArgumentError("Two responses"); - } - std::move(response_callback_)(message.parameters); + if (!message.extensions.empty()) { + OnFatalError( + absl::InvalidArgumentError("REQUEST_OK received with extensions")); + return absl::OkStatus(); + } + if (response_callback_ != nullptr) { + MoqtResponseCallback callback = std::move(response_callback_); response_callback_ = nullptr; + std::move(callback)(message.parameters); return absl::OkStatus(); } NamespaceTask* task = task_.GetIfAvailable(); @@ -252,7 +254,7 @@ add_callback_ = nullptr; QUICHE_DCHECK(task_ == nullptr); task_ = application_(message.track_namespace_prefix, message.parameters, - ResponseCallback(request_id_)); + ResponseCallback()); if (task_ != nullptr) { task_->SetObjectsAvailableCallback([this]() { ProcessNamespaces(); }); } @@ -265,7 +267,7 @@ // This stream is dying. return absl::OkStatus(); } - task_->Update(message.parameters, ResponseCallback(message.request_id)); + task_->Update(message.parameters, ResponseCallback()); return absl::OkStatus(); } @@ -338,21 +340,19 @@ } } -MoqtResponseCallback MoqtSubscribeNamespaceResponseStream::ResponseCallback( - uint64_t request_id) { - return [this, request_id]( +MoqtResponseCallback MoqtSubscribeNamespaceResponseStream::ResponseCallback() { + return [this]( std::variant<MessageParameters, MoqtRequestErrorInfo> response) { - std::visit( - absl::Overload{[this, request_id](const MessageParameters& parameters) { - // In draft-18, there are no useful parameters in - // SUBSCRIBE_NAMESPACE_OK, but Issue #1639 would change - // that. - CheckStatus(SendRequestOk(request_id, parameters)); - }, - [this](const MoqtRequestErrorInfo& error_info) { - CheckStatus(SendRequestError(error_info)); - }}, - response); + std::visit(absl::Overload{[this](const MessageParameters& parameters) { + // In draft-18, there are no useful parameters + // in SUBSCRIBE_NAMESPACE_OK, but Issue #1639 + // would change that. + CheckStatus(SendRequestOk(parameters)); + }, + [this](const MoqtRequestErrorInfo& error_info) { + CheckStatus(SendRequestError(error_info)); + }}, + response); }; }
diff --git a/quiche/quic/moqt/moqt_namespace_stream.h b/quiche/quic/moqt/moqt_namespace_stream.h index 9820993..28972b5 100644 --- a/quiche/quic/moqt/moqt_namespace_stream.h +++ b/quiche/quic/moqt/moqt_namespace_stream.h
@@ -164,7 +164,7 @@ private: void ProcessNamespaces(); - MoqtResponseCallback ResponseCallback(uint64_t request_id); + MoqtResponseCallback ResponseCallback(); uint64_t request_id_; TrackNamespace prefix_;
diff --git a/quiche/quic/moqt/moqt_namespace_stream_test.cc b/quiche/quic/moqt/moqt_namespace_stream_test.cc index 6872e85..e7d017a 100644 --- a/quiche/quic/moqt/moqt_namespace_stream_test.cc +++ b/quiche/quic/moqt/moqt_namespace_stream_test.cc
@@ -87,7 +87,7 @@ EXPECT_CALL( response_callback_, Call(testing::VariantWith<MessageParameters>(Eq(MessageParameters())))); - ReceiveControlMessage(MoqtRequestOk{kRequestId}); + ReceiveControlMessage(MoqtRequestOk()); } TEST_F(MoqtSubscribeNamespaceRequestStreamTest, RequestError) { @@ -115,7 +115,7 @@ EXPECT_CALL( response_callback_, Call(testing::VariantWith<MessageParameters>(Eq(MessageParameters())))); - ReceiveControlMessage(MoqtRequestOk{kRequestId}); + ReceiveControlMessage(MoqtRequestOk()); ReceiveControlMessage(MoqtNamespace{TrackNamespace({"bar"})}); CheckNumberOfObjectsAvailable(1); TrackNamespace received_namespace; @@ -130,7 +130,7 @@ EXPECT_CALL( response_callback_, Call(testing::VariantWith<MessageParameters>(Eq(MessageParameters())))); - ReceiveControlMessage(MoqtRequestOk{kRequestId}); + ReceiveControlMessage(MoqtRequestOk()); ReceiveControlMessage(MoqtNamespace{TrackNamespace({"bar"})}); CheckNumberOfObjectsAvailable(1); ReceiveControlMessage(MoqtNamespaceDone{TrackNamespace({"bar"})}); @@ -150,7 +150,7 @@ EXPECT_CALL( response_callback_, Call(testing::VariantWith<MessageParameters>(Eq(MessageParameters())))); - ReceiveControlMessage(MoqtRequestOk{kRequestId}); + ReceiveControlMessage(MoqtRequestOk()); ReceiveControlMessage(MoqtNamespace{TrackNamespace({"bar"})}); CheckNumberOfObjectsAvailable(1); EXPECT_CALL(error_callback_, @@ -163,7 +163,7 @@ EXPECT_CALL( response_callback_, Call(testing::VariantWith<MessageParameters>(Eq(MessageParameters())))); - ReceiveControlMessage(MoqtRequestOk{kRequestId}); + ReceiveControlMessage(MoqtRequestOk()); EXPECT_CALL(error_callback_, Call(MoqtError::kProtocolViolation, "NAMESPACE_DONE with no active namespace")); ReceiveControlMessage(MoqtNamespaceDone{TrackNamespace({"bar"})}); @@ -173,7 +173,7 @@ EXPECT_CALL( response_callback_, Call(testing::VariantWith<MessageParameters>(Eq(MessageParameters())))); - ReceiveControlMessage(MoqtRequestOk{kRequestId}); + ReceiveControlMessage(MoqtRequestOk()); EXPECT_CALL(error_callback_, Call).Times(0); ReceiveControlMessage(MoqtNamespace{TrackNamespace({"bar"})}); CheckNumberOfObjectsAvailable(1); @@ -187,7 +187,7 @@ EXPECT_CALL( response_callback_, Call(testing::VariantWith<MessageParameters>(Eq(MessageParameters())))); - ReceiveControlMessage(MoqtRequestOk{kRequestId}); + ReceiveControlMessage(MoqtRequestOk()); ReceiveControlMessage(MoqtNamespace{TrackNamespace({"bar"})}); CheckNumberOfObjectsAvailable(1); ReceiveControlMessage(MoqtNamespace{TrackNamespace({"buzz"})}); @@ -225,7 +225,7 @@ EXPECT_CALL( response_callback_, Call(testing::VariantWith<MessageParameters>(Eq(MessageParameters())))); - QUICHE_EXPECT_OK(stream->OnControlMessage(MoqtRequestOk{kRequestId})); + QUICHE_EXPECT_OK(stream->OnControlMessage(MoqtRequestOk())); QUICHE_EXPECT_OK( stream->OnControlMessage(MoqtNamespace{TrackNamespace({"bar"})})); CheckNumberOfObjectsAvailable(1); @@ -243,7 +243,7 @@ EXPECT_CALL( response_callback_, Call(testing::VariantWith<MessageParameters>(Eq(MessageParameters())))); - ReceiveControlMessage(MoqtRequestOk{kRequestId}); + ReceiveControlMessage(MoqtRequestOk()); EXPECT_CALL(mock_stream_, Writev(ControlMessageOfType(MoqtMessageType::kRequestUpdate), _)); MessageParameters update_params; @@ -256,7 +256,7 @@ ok_params.expires = quic::QuicTimeDelta::FromSeconds(60); EXPECT_CALL(update_response_callback, Call(testing::VariantWith<MessageParameters>(Eq(ok_params)))); - ReceiveControlMessage(MoqtRequestOk{kRequestId + 2, ok_params}); + ReceiveControlMessage(MoqtRequestOk(ok_params)); } TEST_F(MoqtSubscribeNamespaceRequestStreamTest, UpdateAndRequestError) { @@ -264,7 +264,7 @@ ok_params.expires = quic::QuicTimeDelta::FromSeconds(60); EXPECT_CALL(response_callback_, Call(testing::VariantWith<MessageParameters>(Eq(ok_params)))); - ReceiveControlMessage(MoqtRequestOk{kRequestId, ok_params}); + ReceiveControlMessage(MoqtRequestOk(ok_params)); EXPECT_CALL(mock_stream_, Writev(ControlMessageOfType(MoqtMessageType::kRequestUpdate), _)); MessageParameters update_params; @@ -317,7 +317,7 @@ }; ObjectsAvailableCallback callback; MockNamespaceTask* task_ptr = nullptr; - MoqtRequestOk ok(kRequestId); + MoqtRequestOk ok; ok.parameters.expires = quic::QuicTimeDelta::FromSeconds(60); EXPECT_CALL(add_callback_, Call).WillOnce(Return(true)); EXPECT_CALL(mock_application_, Call) @@ -386,7 +386,7 @@ }; ObjectsAvailableCallback callback; MockNamespaceTask* task_ptr = nullptr; - MoqtRequestOk ok(kRequestId); + MoqtRequestOk ok; ok.parameters.expires = quic::QuicTimeDelta::FromSeconds(60); EXPECT_CALL(add_callback_, Call).WillOnce(Return(true)); EXPECT_CALL(mock_application_, Call) @@ -456,7 +456,7 @@ MessageParameters(), }; update_message.parameters.subscriber_priority = 10; - MoqtRequestOk ok_response(update_message.request_id); + MoqtRequestOk ok_response; ok_response.parameters.expires = quic::QuicTimeDelta::FromSeconds(60); EXPECT_CALL(*task_ptr, Update(_, _)) .WillOnce([&](const MessageParameters& params, MoqtResponseCallback cb) { @@ -529,7 +529,7 @@ MessageParameters(), }; EXPECT_CALL(add_callback_, Call).WillOnce(Return(true)); - MoqtRequestOk ok(kRequestId); + MoqtRequestOk ok; EXPECT_CALL(mock_stream_, Writev(SerializedControlMessage(ok), _)); EXPECT_CALL(mock_application_, Call) .WillOnce([&](const TrackNamespace&, const MessageParameters&, @@ -558,7 +558,7 @@ MessageParameters(), }; EXPECT_CALL(add_callback_, Call).WillOnce(Return(true)); - MoqtRequestOk ok1(kRequestId); + MoqtRequestOk ok1; EXPECT_CALL(mock_stream_, Writev(SerializedControlMessage(ok1), _)); EXPECT_CALL(mock_application_, Call) .WillOnce([&](const TrackNamespace&, const MessageParameters&,
diff --git a/quiche/quic/moqt/moqt_parser.cc b/quiche/quic/moqt/moqt_parser.cc index e5db23a..040104a 100644 --- a/quiche/quic/moqt/moqt_parser.cc +++ b/quiche/quic/moqt/moqt_parser.cc
@@ -665,9 +665,6 @@ absl::string_view data) const { quic::QuicDataReader reader(data); MoqtSubscribeOk subscribe_ok; - if (!reader.ReadMoqVarInt(&subscribe_ok.request_id)) { - return absl::InvalidArgumentError("Failed to read the request ID"); - } if (!reader.ReadMoqVarInt(&subscribe_ok.track_alias)) { return absl::InvalidArgumentError("Failed to read the track alias"); } @@ -774,11 +771,13 @@ absl::string_view data) const { quic::QuicDataReader reader(data); MoqtRequestOk request_ok; - if (!reader.ReadMoqVarInt(&request_ok.request_id)) { - return absl::InvalidArgumentError("Request ID missing"); - } QUICHE_RETURN_IF_ERROR( FillAndValidateMessageParameters(reader, request_ok.parameters)); + QUICHE_RETURN_IF_ERROR( + ParseKeyValuePairListWithNoPrefix(reader, request_ok.extensions)); + if (!request_ok.extensions.Validate()) { + return absl::InvalidArgumentError("Invalid REQUEST_OK track extensions"); + } QUICHE_RETURN_IF_ERROR(CheckForTrailingData(reader)); return request_ok; }
diff --git a/quiche/quic/moqt/moqt_parser_test.cc b/quiche/quic/moqt/moqt_parser_test.cc index 806ce6c..2067e06 100644 --- a/quiche/quic/moqt/moqt_parser_test.cc +++ b/quiche/quic/moqt/moqt_parser_test.cc
@@ -1383,8 +1383,8 @@ TEST_F(MoqtMessageSpecificTest, SubscribeOkExpirationIsZero) { char subscribe_ok[] = { - 0x04, 0x00, 0x05, 0x02, 0x01, // request_id = 2, track_alias = 1 - 0x01, 0x08, 0x00 // expires = 0 + 0x04, 0x00, 0x04, 0x01, // track_alias = 1 + 0x01, 0x08, 0x00 // expires = 0 }; absl::StatusOr<std::vector<AnyMoqtControlMessage>> parsed = ParseAllMessages(absl::string_view(subscribe_ok, sizeof(subscribe_ok)),
diff --git a/quiche/quic/moqt/moqt_publish_namespace_stream.cc b/quiche/quic/moqt/moqt_publish_namespace_stream.cc index dd0932e..227adc5 100644 --- a/quiche/quic/moqt/moqt_publish_namespace_stream.cc +++ b/quiche/quic/moqt/moqt_publish_namespace_stream.cc
@@ -36,6 +36,11 @@ absl::Status MoqtPublishNamespaceRequestStream::OnControlMessage( const MoqtRequestOk& message) { + if (!message.extensions.empty()) { + OnFatalError( + absl::InvalidArgumentError("REQUEST_OK received with extensions")); + return absl::OkStatus(); + } if (response_callback_ != nullptr) { // Response to the initial PUBLISH_NAMESPACE. auto callback = std::move(response_callback_); @@ -99,22 +104,20 @@ prefix_ = message.track_namespace; application_( *prefix_, &message.parameters, - [weakptr = weak_ptr_factory_.Create(), id = request_id_]( + [weakptr = weak_ptr_factory_.Create()]( std::variant<MessageParameters, MoqtRequestErrorInfo> response) { MoqtPublishNamespaceResponseStream* stream = weakptr.GetIfAvailable(); if (stream == nullptr) { return; } - std::visit( - absl::Overload{[&](const MessageParameters& parameters) { - stream->CheckStatus( - stream->SendRequestOk(id, parameters)); - }, - [&](const MoqtRequestErrorInfo& error) { - stream->CheckStatus( - stream->SendRequestError(error)); - }}, - response); + std::visit(absl::Overload{ + [&](const MessageParameters& parameters) { + stream->CheckStatus(stream->SendRequestOk(parameters)); + }, + [&](const MoqtRequestErrorInfo& error) { + stream->CheckStatus(stream->SendRequestError(error)); + }}, + response); }); return absl::OkStatus(); } @@ -127,22 +130,20 @@ } application_( *prefix_, &message.parameters, - [weakptr = weak_ptr_factory_.Create(), id = message.request_id]( + [weakptr = weak_ptr_factory_.Create()]( std::variant<MessageParameters, MoqtRequestErrorInfo> response) { MoqtPublishNamespaceResponseStream* stream = weakptr.GetIfAvailable(); if (stream == nullptr) { return; } - std::visit( - absl::Overload{[&](const MessageParameters& parameters) { - stream->CheckStatus( - stream->SendRequestOk(id, parameters)); - }, - [&](const MoqtRequestErrorInfo& error) { - stream->CheckStatus( - stream->SendRequestError(error)); - }}, - response); + std::visit(absl::Overload{ + [&](const MessageParameters& parameters) { + stream->CheckStatus(stream->SendRequestOk(parameters)); + }, + [&](const MoqtRequestErrorInfo& error) { + stream->CheckStatus(stream->SendRequestError(error)); + }}, + response); }); return absl::OkStatus(); }
diff --git a/quiche/quic/moqt/moqt_publish_namespace_stream_test.cc b/quiche/quic/moqt/moqt_publish_namespace_stream_test.cc index 4ee1c0e..b01507f 100644 --- a/quiche/quic/moqt/moqt_publish_namespace_stream_test.cc +++ b/quiche/quic/moqt/moqt_publish_namespace_stream_test.cc
@@ -94,10 +94,7 @@ callback_called = true; EXPECT_TRUE(std::holds_alternative<MessageParameters>(res)); }); - - MoqtRequestOk message; - message.request_id = 10; - QUICHE_EXPECT_OK(request_stream->OnControlMessage(message)); + QUICHE_EXPECT_OK(request_stream->OnControlMessage(MoqtRequestOk())); EXPECT_TRUE(callback_called); } @@ -123,14 +120,11 @@ CreateAndBindStream(); // Resolve initial response first. EXPECT_CALL(response_callback_, Call(_)); - MoqtRequestOk initial_ok; - initial_ok.request_id = 10; - QUICHE_EXPECT_OK(request_stream->OnControlMessage(initial_ok)); + QUICHE_EXPECT_OK(request_stream->OnControlMessage(MoqtRequestOk())); // Now send update. EXPECT_CALL(mock_stream_, Writev(ControlMessageOfType(MoqtMessageType::kRequestUpdate), _)); - MessageParameters parameters; parameters.subscriber_priority = 50; bool update_callback_called = false; @@ -140,13 +134,11 @@ ASSERT_TRUE(std::holds_alternative<MessageParameters>(res)); EXPECT_EQ(std::get<MessageParameters>(res).subscriber_priority, 50); }; - QUICHE_EXPECT_OK(request_stream->SendRequestUpdate( 11, 0, parameters, std::move(update_callback))); // Receive OK for update. MoqtRequestOk ok; - ok.request_id = 11; ok.parameters.subscriber_priority = 50; QUICHE_EXPECT_OK(request_stream->OnControlMessage(ok)); EXPECT_TRUE(update_callback_called);
diff --git a/quiche/quic/moqt/moqt_publish_stream.cc b/quiche/quic/moqt/moqt_publish_stream.cc index fcc6ebc..6c36172 100644 --- a/quiche/quic/moqt/moqt_publish_stream.cc +++ b/quiche/quic/moqt/moqt_publish_stream.cc
@@ -67,9 +67,10 @@ absl::Status MoqtPublishRequestStream::OnControlMessage( const MoqtRequestOk& message) { - if (message.request_id != publisher_->request_id()) { - return absl::InvalidArgumentError( - "REQUEST_OK does not match PUBLISH request ID"); + if (!message.extensions.empty()) { + OnFatalError( + absl::InvalidArgumentError("REQUEST_OK received with extensions")); + return absl::OkStatus(); } std::move(response_callback_)(message.parameters); publisher_->Update(message.parameters); @@ -91,7 +92,7 @@ out_parameters.largest_object); } publisher_->Update(in_parameters); - CheckStatus(SendRequestOk(message.request_id, MessageParameters())); + CheckStatus(SendRequestOk(MessageParameters())); return absl::OkStatus(); } @@ -136,7 +137,7 @@ // There was no existing SUBSCRIBE, so invoke the callback. subscriber_->set_visitor((*incoming_publish_callback_)( message.full_track_name, message.parameters, message.extensions, - [weakptr = weak_ptr_factory_.Create(), request_id = message.request_id]( + [weakptr = weak_ptr_factory_.Create()]( const std::variant<MessageParameters, MoqtRequestErrorInfo> response) { MoqtPublishResponseStream* stream = weakptr.GetIfAvailable(); @@ -144,22 +145,20 @@ return; } std::visit( - absl::Overload{[&](const MessageParameters& parameters) { - stream->subscriber_->Update(parameters); - stream->CheckStatus(stream->SendRequestOk( - request_id, parameters, /*fin=*/false)); - }, - [&](const MoqtRequestErrorInfo& error_info) { - stream->CheckStatus( - stream->SendRequestError(error_info)); - }}, + absl::Overload{ + [&](const MessageParameters& parameters) { + stream->subscriber_->Update(parameters); + stream->CheckStatus(stream->SendRequestOk(parameters)); + }, + [&](const MoqtRequestErrorInfo& error_info) { + stream->CheckStatus(stream->SendRequestError(error_info)); + }}, 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())); + CheckStatus(SendRequestOk(subscriber_->const_parameters())); } incoming_publish_callback_ = nullptr; if (subscriber_->visitor() == nullptr) { @@ -181,12 +180,17 @@ return absl::OkStatus(); } subscriber_->Update(message.parameters); - CheckStatus(SendRequestOk(message.request_id, MessageParameters())); + CheckStatus(SendRequestOk(MessageParameters())); return absl::OkStatus(); } absl::Status MoqtPublishResponseStream::OnControlMessage( const MoqtRequestOk& message) { + if (!message.extensions.empty()) { + OnFatalError( + absl::InvalidArgumentError("REQUEST_OK received with extensions")); + return absl::OkStatus(); + } // TODO(martinduke): Process REQUEST_OK parameters. return request_update_queue().OnControlMessage(message); }
diff --git a/quiche/quic/moqt/moqt_publish_stream_test.cc b/quiche/quic/moqt/moqt_publish_stream_test.cc index 304b471..b10f1c9 100644 --- a/quiche/quic/moqt/moqt_publish_stream_test.cc +++ b/quiche/quic/moqt/moqt_publish_stream_test.cc
@@ -131,7 +131,6 @@ stream_->BindStream(&mock_stream_); // Calls OnStreamBound MoqtRequestOk request_ok; - request_ok.request_id = kRequestId; request_ok.parameters.delivery_timeout = quic::QuicTimeDelta::FromSeconds(2); request_ok.parameters.group_order = MoqtDeliveryOrder::kDescending; QUICHE_EXPECT_OK(stream_->OnControlMessage(request_ok)); @@ -175,7 +174,7 @@ .WillOnce(Return(absl::OkStatus())); stream_->BindStream(&mock_stream_); - QUICHE_EXPECT_OK(stream_->OnControlMessage(MoqtRequestOk{kRequestId})); + QUICHE_EXPECT_OK(stream_->OnControlMessage(MoqtRequestOk())); // Set largest location on publisher track_publisher_->AddObject(Location(1, 2), 0, "payload", true);
diff --git a/quiche/quic/moqt/moqt_session.cc b/quiche/quic/moqt/moqt_session.cc index 2570852..fb5e4f6 100644 --- a/quiche/quic/moqt/moqt_session.cc +++ b/quiche/quic/moqt/moqt_session.cc
@@ -330,7 +330,7 @@ bool MoqtSession::TrackStatus(const FullTrackName& name, const MessageParameters& parameters, - MoqtResponseCallback response_callback) { + TrackStatusResponseCallback response_callback) { QUICHE_DCHECK(name.IsValid()); if (received_goaway_ || sent_goaway_) { QUIC_DLOG(INFO) << ENDPOINT << "Tried to send TRACK_STATUS after GOAWAY";
diff --git a/quiche/quic/moqt/moqt_session.h b/quiche/quic/moqt/moqt_session.h index 5320b39..557c0e7 100644 --- a/quiche/quic/moqt/moqt_session.h +++ b/quiche/quic/moqt/moqt_session.h
@@ -143,7 +143,7 @@ void UnsubscribeTracks(TrackNamespace& prefix) override; bool TrackStatus(const FullTrackName& name, const MessageParameters& parameters, - MoqtResponseCallback response_callback) override; + TrackStatusResponseCallback response_callback) override; quiche::QuicheWeakPtr<MoqtSessionInterface> GetWeakPtr() override { return weak_ptr_factory_.Create(); }
diff --git a/quiche/quic/moqt/moqt_session_callbacks.h b/quiche/quic/moqt/moqt_session_callbacks.h index 4221215..2a8f847 100644 --- a/quiche/quic/moqt/moqt_session_callbacks.h +++ b/quiche/quic/moqt/moqt_session_callbacks.h
@@ -76,6 +76,15 @@ using FetchResponseCallback = quiche::SingleUseCallback<void( std::variant<FetchOkData, MoqtRequestErrorInfo>)>; +struct TrackStatusOkData { + MessageParameters parameters; + TrackExtensions extensions; + bool operator==(const TrackStatusOkData& other) const = default; +}; + +using TrackStatusResponseCallback = quiche::SingleUseCallback<void( + std::variant<TrackStatusOkData, MoqtRequestErrorInfo>)>; + // Called when the SETUP message from the peer is received. using MoqtSessionEstablishedCallback = quiche::SingleUseCallback<void()>;
diff --git a/quiche/quic/moqt/moqt_session_interface.h b/quiche/quic/moqt/moqt_session_interface.h index 8d182cd..f39c82b 100644 --- a/quiche/quic/moqt/moqt_session_interface.h +++ b/quiche/quic/moqt/moqt_session_interface.h
@@ -29,12 +29,12 @@ namespace moqt { -inline constexpr absl::string_view kDraft16 = "moqt-16"; -inline constexpr absl::string_view kDefaultMoqtVersion = kDraft16; +inline constexpr absl::string_view kDraft18 = "moqt-18"; +inline constexpr absl::string_view kDefaultMoqtVersion = kDraft18; inline constexpr absl::string_view kUnrecognizedVersionForTests = "moqt-15"; inline constexpr absl::string_view kImplementationName = - "Google QUICHE MOQT draft 16"; + "Google QUICHE MOQT draft 18"; struct QUICHE_EXPORT MoqtSessionParameters { // TODO: support multiple versions. MoqtSessionParameters() = default; @@ -177,7 +177,7 @@ // `response_callback` will be eventually invoked if true. virtual bool TrackStatus(const FullTrackName& name, const MessageParameters& parameters, - MoqtResponseCallback response_callback) = 0; + TrackStatusResponseCallback response_callback) = 0; // TODO: Add RequestUpdate, PublishDone method. virtual quiche::QuicheWeakPtr<MoqtSessionInterface> GetWeakPtr() = 0; };
diff --git a/quiche/quic/moqt/moqt_session_test.cc b/quiche/quic/moqt/moqt_session_test.cc index debc29b..0786636 100644 --- a/quiche/quic/moqt/moqt_session_test.cc +++ b/quiche/quic/moqt/moqt_session_test.cc
@@ -245,12 +245,7 @@ MessageParameters parameters; parameters.expires = publisher->expiration(); parameters.largest_object = publisher->largest_location(); - MoqtSubscribeOk expected_ok = { - subscribe.request_id, - track_alias, - parameters, - extensions, - }; + MoqtSubscribeOk expected_ok(track_alias, parameters, extensions); EXPECT_CALL(mock_bidi_stream_, Writev(SerializedControlMessage(expected_ok), _)); control_parser->ReceiveMessage(subscribe); @@ -595,7 +590,7 @@ publish_namespace_response_callback.AsStdFunction(), [&]() { cancel_called = true; }); - MoqtRequestOk ok = {/*request_id=*/0, MessageParameters()}; + MoqtRequestOk ok; EXPECT_CALL(publish_namespace_response_callback, Call) .WillOnce( [&](std::variant<MessageParameters, MoqtRequestErrorInfo> response) { @@ -623,7 +618,7 @@ publish_namespace_resolved_callback.AsStdFunction(), []() {}); - MoqtRequestOk ok = {/*request_id=*/0, MessageParameters()}; + MoqtRequestOk ok; EXPECT_CALL(publish_namespace_resolved_callback, Call) .WillOnce( [&](std::variant<MessageParameters, MoqtRequestErrorInfo> response) { @@ -868,13 +863,7 @@ MessageParameters parameters(SubscribeForTest()); session_.Subscribe(FullTrackName("foo", "bar"), &remote_track_visitor_, parameters); - - MoqtSubscribeOk ok = { - /*request_id=*/0, - /*track_alias=*/2, - MessageParameters(), - TrackExtensions(), - }; + MoqtSubscribeOk ok(/*track_alias=*/2, MessageParameters(), TrackExtensions()); EXPECT_CALL(remote_track_visitor_, OnReply) .WillOnce( [&](const FullTrackName& ftn, @@ -894,13 +883,7 @@ Writev(SerializedControlMessage(subscribe), _)); session_.Subscribe(FullTrackName("foo", "bar"), &remote_track_visitor_, subscribe.parameters); - - MoqtSubscribeOk ok = { - /*request_id=*/0, - /*track_alias=*/2, - MessageParameters(), - TrackExtensions(), - }; + MoqtSubscribeOk ok(/*track_alias=*/2, MessageParameters(), TrackExtensions()); EXPECT_CALL(remote_track_visitor_, OnReply) .WillOnce( [&](const FullTrackName& ftn, @@ -919,12 +902,7 @@ parameters.subscription_filter.emplace(Location(1, 0), 10); session_.Subscribe(FullTrackName("foo", "bar"), &remote_track_visitor_, parameters); - MoqtSubscribeOk ok = { - /*request_id=*/0, - /*track_alias=*/2, - MessageParameters(), - TrackExtensions(), - }; + MoqtSubscribeOk ok(/*track_alias=*/2, MessageParameters(), TrackExtensions()); EXPECT_CALL(remote_track_visitor_, OnReply); bidi_wrapper_->ReceiveMessage(ok); EXPECT_CALL(mock_bidi_stream_, @@ -940,10 +918,7 @@ ASSERT_TRUE(std::holds_alternative<MessageParameters>(info)); EXPECT_EQ(std::get<MessageParameters>(info), MessageParameters()); })); - bidi_wrapper_->ReceiveMessage(MoqtRequestOk{ - /*request_id=*/2, - MessageParameters(), - }); + bidi_wrapper_->ReceiveMessage(MoqtRequestOk()); EXPECT_TRUE(got_response); // Check if window is functional by receiving datagrams. Type = 8, alias = 2, // Location = (2,0), payload = "foo". @@ -1030,9 +1005,7 @@ std::move(callback)(MessageParameters()); }); EXPECT_CALL(mock_bidi_stream_, - Writev(SerializedControlMessage(MoqtRequestOk{ - kDefaultPeerRequestId, MessageParameters()}), - _)); + Writev(SerializedControlMessage(MoqtRequestOk()), _)); bidi_wrapper_->ReceiveMessage(publish_namespace); EXPECT_CALL(session_callbacks_.incoming_publish_namespace_callback, Call(track_namespace, IsNull(), IsNull())); @@ -1060,9 +1033,7 @@ std::move(callback)(MessageParameters()); }); EXPECT_CALL(mock_bidi_stream_, - Writev(SerializedControlMessage(MoqtRequestOk{ - kDefaultPeerRequestId, MessageParameters()}), - _)); + Writev(SerializedControlMessage(MoqtRequestOk()), _)); bidi_wrapper_->ReceiveMessage(publish_namespace); EXPECT_CALL(session_callbacks_.incoming_publish_namespace_callback, Call(track_namespace, IsNull(), IsNull())); @@ -1119,8 +1090,7 @@ got_callback = true; EXPECT_TRUE(std::holds_alternative<MessageParameters>(response)); }); - MoqtRequestOk ok = {kDefaultLocalRequestId, MessageParameters()}; - bidi_wrapper_->ReceiveMessage(ok); + bidi_wrapper_->ReceiveMessage(MoqtRequestOk()); EXPECT_TRUE(got_callback); EXPECT_CALL(mock_bidi_stream_, ResetWithUserCode); } @@ -1154,12 +1124,8 @@ Writev(ControlMessageOfType(MoqtMessageType::kSubscribe), _)); session_.Subscribe(FullTrackName("foo", "bar"), &remote_track_visitor_, MessageParameters()); - MoqtSubscribeOk subscribe_ok = { - /*request_id=*/0, - /*track_alias=*/2, - MessageParameters(), - TrackExtensions(), - }; + MoqtSubscribeOk subscribe_ok(/*track_alias=*/2, MessageParameters(), + TrackExtensions()); bidi_wrapper_->ReceiveMessage(subscribe_ok); // Second subscribe, but OK has the same track alias. webtransport::test::MockStream bidi_stream_2; @@ -1169,7 +1135,6 @@ Writev(ControlMessageOfType(MoqtMessageType::kSubscribe), _)); 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), @@ -1201,12 +1166,10 @@ EXPECT_CALL(mock_bidi_stream_, Writev(ControlMessageOfType(MoqtMessageType::kSubscribe), _)); session_.Subscribe(ftn, &remote_track_visitor_, MessageParameters()); - MoqtSubscribeOk ok; - ok.request_id = 0; - ok.track_alias = 2; - ok.extensions = + MoqtSubscribeOk ok( + 2, MessageParameters(), TrackExtensions(std::nullopt, std::nullopt, kPeerDefaultPriority, - std::nullopt, std::nullopt, std::nullopt); + std::nullopt, std::nullopt, std::nullopt)); EXPECT_CALL(remote_track_visitor_, OnReply); bidi_wrapper_->ReceiveMessage(ok); @@ -1247,12 +1210,10 @@ EXPECT_CALL(mock_bidi_stream_, Writev(ControlMessageOfType(MoqtMessageType::kSubscribe), _)); session_.Subscribe(ftn, &remote_track_visitor_, MessageParameters()); - MoqtSubscribeOk ok; - ok.request_id = 0; - ok.track_alias = 2; - ok.extensions = + MoqtSubscribeOk ok( + 2, MessageParameters(), TrackExtensions(std::nullopt, std::nullopt, kPeerDefaultPriority, - std::nullopt, std::nullopt, std::nullopt); + std::nullopt, std::nullopt, std::nullopt)); EXPECT_CALL(remote_track_visitor_, OnReply); bidi_wrapper_->ReceiveMessage(ok); // Omit priority from a datagram. @@ -1379,12 +1340,10 @@ MessageParameters params; params.subscription_filter.emplace(Location(1, 0)); session_.Subscribe(ftn, &remote_track_visitor_, params); - MoqtSubscribeOk ok; - ok.request_id = 0; - ok.track_alias = 2; - ok.extensions = + MoqtSubscribeOk ok( + 2, MessageParameters(), TrackExtensions(std::nullopt, std::nullopt, kPeerDefaultPriority, - std::nullopt, std::nullopt, std::nullopt); + std::nullopt, std::nullopt, std::nullopt)); EXPECT_CALL(remote_track_visitor_, OnReply); bidi_wrapper_->ReceiveMessage(ok); char datagram[] = {0x01, 0x02, 0x00, 0x00, 0x80, 0x00, 0x08, 0x64, @@ -1903,7 +1862,7 @@ MessageParameters parameters; parameters.largest_object = Location(2, 0); subscribe_wrapper->ReceiveMessage( - MoqtSubscribeOk(0, 2, parameters, TrackExtensions())); + MoqtSubscribeOk(0, parameters, TrackExtensions())); bidi_wrapper_->ReceiveMessage(MoqtFetchOk( false, Location(2, 0), MessageParameters(), TrackExtensions())); // Packet arrives on FETCH stream. @@ -1939,7 +1898,7 @@ MoqtSubscribeNamespace subscribe_namespace = {/*request_id=*/1, prefix, parameters}; quiche::QuicheWeakPtr<MockNamespaceTask> task; - MoqtRequestOk expected_ok(/*request_id=*/1); + MoqtRequestOk expected_ok; expected_ok.parameters.expires = quic::QuicTimeDelta::FromSeconds(60); EXPECT_CALL(session_callbacks_.incoming_subscribe_namespace_callback, Call(prefix, parameters, _)) @@ -2466,7 +2425,6 @@ EXPECT_CALL(*track, largest_location) .WillRepeatedly(Return(Location(5, 30))); MoqtRequestOk expected_ok; - expected_ok.request_id = track_status.request_id; expected_ok.parameters.expires = quic::QuicTimeDelta::FromMilliseconds(10000); expected_ok.parameters.largest_object = Location(5, 30); @@ -2493,7 +2451,6 @@ .WillRepeatedly(Return(quic::QuicTimeDelta::FromMilliseconds(10000))); EXPECT_CALL(*track, largest_location).WillRepeatedly(Return(Location(5, 30))); MoqtRequestOk expected_ok; - 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(mock_bidi_stream_, @@ -2549,8 +2506,7 @@ parameters.subscription_filter.emplace(MoqtFilterType::kLargestObject); EXPECT_TRUE(session_.Subscribe(FullTrackName("foo", "bar"), &remote_track_visitor_, parameters)); - MoqtSubscribeOk ok = {/*request_id=*/0, /*track_alias=*/2, - MessageParameters(), TrackExtensions()}; + MoqtSubscribeOk ok(/*track_alias=*/2, MessageParameters(), TrackExtensions()); EXPECT_CALL(remote_track_visitor_, OnReply) .WillOnce( [&](const FullTrackName& ftn, @@ -2588,8 +2544,7 @@ parameters.subscription_filter.emplace(MoqtFilterType::kLargestObject); EXPECT_TRUE(session_.Subscribe(FullTrackName("foo", "bar"), &remote_track_visitor_, parameters)); - MoqtSubscribeOk ok = {/*request_id=*/0, /*track_alias=*/2, - MessageParameters(), TrackExtensions()}; + MoqtSubscribeOk ok(/*track_alias=*/2, MessageParameters(), TrackExtensions()); EXPECT_CALL(remote_track_visitor_, OnReply) .WillOnce( [&](const FullTrackName& ftn, @@ -2629,7 +2584,7 @@ // Register two incoming PUBLISH_NAMESPACE. MoqtPublishNamespace publish_namespace{ /*request_id=*/1, TrackNamespace{"foo"}, MessageParameters()}; - MoqtRequestOk expected_ok = {/*request_id=*/1, MessageParameters()}; + MoqtRequestOk expected_ok; expected_ok.parameters.expires = quic::QuicTimeDelta::FromSeconds(60); EXPECT_CALL(session_callbacks_.incoming_publish_namespace_callback, Call(TrackNamespace{"foo"}, _, _)) @@ -2783,7 +2738,6 @@ })); MoqtRequestOk request_ok; - request_ok.request_id = 0; request_ok.parameters.delivery_timeout = quic::QuicTimeDelta::FromSeconds(2); bidi_wrapper_->ReceiveMessage(request_ok); @@ -2844,7 +2798,6 @@ 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;
diff --git a/quiche/quic/moqt/moqt_subscribe_stream.cc b/quiche/quic/moqt/moqt_subscribe_stream.cc index 35578d8..b6186e8 100644 --- a/quiche/quic/moqt/moqt_subscribe_stream.cc +++ b/quiche/quic/moqt/moqt_subscribe_stream.cc
@@ -62,9 +62,6 @@ 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"); @@ -89,6 +86,11 @@ absl::InvalidArgumentError("REQUEST_OK received before SUBSCRIBE_OK")); return absl::OkStatus(); } + if (!message.extensions.empty()) { + OnFatalError(absl::InvalidArgumentError( + "REQUEST_UPDATE_OK received with extensions")); + return absl::OkStatus(); + } absl::StatusOr<MessageParameters> old_parameters = request_update_queue().NextParameters(); if (!old_parameters.ok()) { @@ -211,7 +213,7 @@ "no subscription"); } subscription_->Update(message.parameters); - return SendRequestOk(message.request_id, MessageParameters()); + return SendRequestOk(MessageParameters()); } void MoqtSubscribeResponseStream::Detach() {
diff --git a/quiche/quic/moqt/moqt_subscribe_stream_test.cc b/quiche/quic/moqt/moqt_subscribe_stream_test.cc index a380b38..301ee16 100644 --- a/quiche/quic/moqt/moqt_subscribe_stream_test.cc +++ b/quiche/quic/moqt/moqt_subscribe_stream_test.cc
@@ -92,9 +92,7 @@ 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; + MoqtSubscribeOk subscribe_ok(kTrackAlias); QUICHE_EXPECT_OK(stream_->OnControlMessage(subscribe_ok)); EXPECT_EQ(stream_->track()->track_alias(), kTrackAlias); // Test cleanup. @@ -109,9 +107,7 @@ 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; + MoqtSubscribeOk subscribe_ok(kTrackAlias); QUICHE_EXPECT_OK(stream_->OnControlMessage(subscribe_ok)); // Test cleanup. EXPECT_CALL(mock_remove_callback_, Call); @@ -124,13 +120,32 @@ 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, RequestOkWithExtensions) { + 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(kTrackAlias); + QUICHE_EXPECT_OK(stream_->OnControlMessage(subscribe_ok)); + EXPECT_CALL(error_callback_, Call(MoqtError::kProtocolViolation, _)); + MoqtRequestOk request_ok( + MessageParameters(), + TrackExtensions(quic::QuicTimeDelta::FromSeconds(5), std::nullopt, + std::nullopt, std::nullopt, std::nullopt, std::nullopt)); + 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_, @@ -140,9 +155,7 @@ 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; + MoqtSubscribeOk subscribe_ok(kTrackAlias); QUICHE_EXPECT_OK(stream_->OnControlMessage(subscribe_ok)); // REQUEST_UPDATE. EXPECT_CALL(mock_stream_, @@ -163,7 +176,6 @@ 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);
diff --git a/quiche/quic/moqt/moqt_track_status_stream.cc b/quiche/quic/moqt/moqt_track_status_stream.cc index 8e4863b..3af18fa 100644 --- a/quiche/quic/moqt/moqt_track_status_stream.cc +++ b/quiche/quic/moqt/moqt_track_status_stream.cc
@@ -20,6 +20,7 @@ #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/common/platform/api/quiche_logging.h" #include "quiche/common/quiche_weak_ptr.h" @@ -30,7 +31,7 @@ const MoqtControlMessageParser& message_parser, uint64_t request_id, const FullTrackName& full_track_name, const MessageParameters& parameters, SessionErrorCallback session_error_callback, - MoqtResponseCallback response_callback) + TrackStatusResponseCallback response_callback) : MoqtBidiStreamBase(framer, message_parser, std::move(session_error_callback)), request_id_(request_id), @@ -58,12 +59,10 @@ if (response_callback_ == nullptr) { return absl::InvalidArgumentError("Duplicate REQUEST_OK"); } - MoqtResponseCallback callback = std::move(response_callback_); + TrackStatusResponseCallback callback = std::move(response_callback_); response_callback_ = nullptr; Fin(); - // `message.request_id` is ignored, since request IDs in REQUEST_OK are - // deprecated and not present in draft-18. - std::move(callback)(message.parameters); + std::move(callback)(message); return absl::OkStatus(); } @@ -72,7 +71,7 @@ if (response_callback_ == nullptr) { return absl::InvalidArgumentError("Duplicate REQUEST_ERROR"); } - MoqtResponseCallback callback = std::move(response_callback_); + TrackStatusResponseCallback callback = std::move(response_callback_); response_callback_ = nullptr; Fin(); // `message.request_id` is ignored, since request IDs in REQUEST_ERROR are @@ -84,7 +83,7 @@ void MoqtTrackStatusRequestStream::Detach() { if (response_callback_ != nullptr) { - MoqtResponseCallback callback = std::move(response_callback_); + TrackStatusResponseCallback callback = std::move(response_callback_); response_callback_ = nullptr; std::move(callback)(MoqtRequestErrorInfo{RequestErrorCode::kInternalError, std::nullopt, "Stream closed"}); @@ -135,7 +134,7 @@ parameters.expires = publisher_->expiration(); parameters.largest_object = publisher_->largest_location(); // Since `fin` is true, this will also reset `publisher_`. - CheckStatus(SendRequestOk(*request_id_, parameters, /*fin=*/true)); + CheckStatus(SendRequestOk(parameters, publisher_->extensions())); } void MoqtTrackStatusResponseStream::OnSubscribeRejected( @@ -162,4 +161,11 @@ } } +absl::Status MoqtTrackStatusResponseStream::SendRequestOk( + const MessageParameters& parameters, const TrackExtensions& extensions) { + return SendOrBufferMessage( + framer()->SerializeRequestOk(MoqtRequestOk(parameters, extensions)), + /*fin=*/true); +} + } // namespace moqt
diff --git a/quiche/quic/moqt/moqt_track_status_stream.h b/quiche/quic/moqt/moqt_track_status_stream.h index 16f9886..95d0f09 100644 --- a/quiche/quic/moqt/moqt_track_status_stream.h +++ b/quiche/quic/moqt/moqt_track_status_stream.h
@@ -13,7 +13,6 @@ #include "absl/status/status.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_live_publisher.h" @@ -22,6 +21,7 @@ #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_types.h" #include "quiche/common/quiche_weak_ptr.h" #include "quiche/web_transport/web_transport.h" @@ -37,7 +37,7 @@ const FullTrackName& full_track_name, const MessageParameters& parameters, SessionErrorCallback session_error_callback, - MoqtResponseCallback response_callback); + TrackStatusResponseCallback response_callback); ~MoqtTrackStatusRequestStream() { Detach(); } // MoqtBidiStreamBase overrides. @@ -53,7 +53,7 @@ const uint64_t request_id_; const FullTrackName full_track_name_; const MessageParameters parameters_; - MoqtResponseCallback response_callback_; + TrackStatusResponseCallback response_callback_; }; // MoqtTrackStatusResponseStream represents an incoming TRACK_STATUS request. @@ -87,6 +87,9 @@ void Detach() override; private: + // Unlike other REQUEST_OK, TRACK_STATUS_OK has extensions and a FIN. + absl::Status SendRequestOk(const MessageParameters& parameters, + const TrackExtensions& extensions); SessionToPublisherInterface* absl_nullable session() const { return session_.GetIfAvailable(); }
diff --git a/quiche/quic/moqt/moqt_track_status_stream_test.cc b/quiche/quic/moqt/moqt_track_status_stream_test.cc index 274fc94..1da3741 100644 --- a/quiche/quic/moqt/moqt_track_status_stream_test.cc +++ b/quiche/quic/moqt/moqt_track_status_stream_test.cc
@@ -15,13 +15,13 @@ #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_publisher.h" +#include "quiche/quic/moqt/moqt_session_callbacks.h" #include "quiche/quic/moqt/moqt_session_interface.h" #include "quiche/quic/moqt/moqt_types.h" #include "quiche/quic/moqt/test_tools/mock_moqt_session.h" @@ -67,7 +67,7 @@ StrictMock<testing::MockFunction<void(MoqtError, absl::string_view)>> session_error_callback_; StrictMock<testing::MockFunction<void( - std::variant<MessageParameters, MoqtRequestErrorInfo>)>> + std::variant<TrackStatusOkData, MoqtRequestErrorInfo>)>> response_callback_; StrictMock<webtransport::test::MockStream> mock_stream_; }; @@ -78,7 +78,7 @@ EXPECT_CALL(mock_stream_, Writev(ControlMessageOfType(MoqtMessageType::kTrackStatus), _)); EXPECT_CALL(response_callback_, Call) - .WillOnce([&](std::variant<MessageParameters, MoqtRequestErrorInfo> v) { + .WillOnce([&](std::variant<TrackStatusOkData, MoqtRequestErrorInfo> v) { ASSERT_TRUE(std::holds_alternative<MoqtRequestErrorInfo>(v)); auto info = std::get<MoqtRequestErrorInfo>(v); EXPECT_EQ(info.error_code, RequestErrorCode::kInternalError); @@ -99,7 +99,7 @@ EXPECT_CALL(mock_stream_, Writev(SerializedControlMessage(expected_message), _)); EXPECT_CALL(response_callback_, Call) - .WillOnce([&](std::variant<MessageParameters, MoqtRequestErrorInfo> v) { + .WillOnce([&](std::variant<TrackStatusOkData, MoqtRequestErrorInfo> v) { ASSERT_TRUE(std::holds_alternative<MoqtRequestErrorInfo>(v)); auto info = std::get<MoqtRequestErrorInfo>(v); EXPECT_EQ(info.error_code, RequestErrorCode::kInternalError); @@ -115,23 +115,20 @@ Writev(ControlMessageOfType(MoqtMessageType::kTrackStatus), _)); stream.BindStream(&mock_stream_); - MessageParameters parameters; - parameters.expires = quic::QuicTimeDelta::FromSeconds(10); - parameters.largest_object = Location(1, 2); - + MoqtRequestOk ok( + MessageParameters(), + TrackExtensions(quic::QuicTimeDelta::FromSeconds(5), + quic::QuicTimeDelta::FromSeconds(10), std::nullopt, + std::nullopt, std::nullopt, std::nullopt)); + ok.parameters.expires = quic::QuicTimeDelta::FromSeconds(10); + ok.parameters.largest_object = Location(1, 2); EXPECT_CALL(response_callback_, Call) - .WillOnce([&](std::variant<MessageParameters, MoqtRequestErrorInfo> v) { - ASSERT_TRUE(std::holds_alternative<MessageParameters>(v)); - auto params = std::get<MessageParameters>(v); - EXPECT_EQ(params.expires, parameters.expires); - EXPECT_EQ(params.largest_object, parameters.largest_object); + .WillOnce([&](std::variant<TrackStatusOkData, MoqtRequestErrorInfo> v) { + ASSERT_TRUE(std::holds_alternative<TrackStatusOkData>(v)); + auto data = std::get<TrackStatusOkData>(v); + EXPECT_EQ(data, ok); }); EXPECT_CALL(mock_stream_, Writev(testing::IsEmpty(), _)); - - MoqtRequestOk ok; - ok.request_id = kRequestId; - ok.parameters = parameters; - QUICHE_EXPECT_OK( stream.OnRawControlMessage(GenericMessageToRawControlMessage(ok))); } @@ -145,7 +142,7 @@ MoqtRequestError error(RequestErrorCode::kDoesNotExist, std::nullopt, "Track does not exist"); EXPECT_CALL(response_callback_, Call) - .WillOnce([&](std::variant<MessageParameters, MoqtRequestErrorInfo> v) { + .WillOnce([&](std::variant<TrackStatusOkData, MoqtRequestErrorInfo> v) { ASSERT_TRUE(std::holds_alternative<MoqtRequestErrorInfo>(v)); EXPECT_EQ(std::get<MoqtRequestErrorInfo>(v), error); }); @@ -168,11 +165,8 @@ EXPECT_CALL(mock_stream_, Writev(testing::IsEmpty(), _)); MoqtRequestOk ok; - ok.request_id = kRequestId; - QUICHE_EXPECT_OK( stream.OnRawControlMessage(GenericMessageToRawControlMessage(ok))); - EXPECT_THAT( stream.OnRawControlMessage(GenericMessageToRawControlMessage(ok)), StatusIs(absl::StatusCode::kInvalidArgument, "Duplicate REQUEST_OK"));
diff --git a/quiche/quic/moqt/test_tools/mock_moqt_session.h b/quiche/quic/moqt/test_tools/mock_moqt_session.h index 6d5bc15..06889b3 100644 --- a/quiche/quic/moqt/test_tools/mock_moqt_session.h +++ b/quiche/quic/moqt/test_tools/mock_moqt_session.h
@@ -128,7 +128,7 @@ MOCK_METHOD(void, UnsubscribeTracks, (TrackNamespace&), (override)); MOCK_METHOD(bool, TrackStatus, (const FullTrackName&, const MessageParameters&, - MoqtResponseCallback), + TrackStatusResponseCallback), (override)); quiche::QuicheWeakPtr<MoqtSessionInterface> GetWeakPtr() override {
diff --git a/quiche/quic/moqt/test_tools/moqt_test_message.h b/quiche/quic/moqt/test_tools/moqt_test_message.h index 79c6c30..04b9b34 100644 --- a/quiche/quic/moqt/test_tools/moqt_test_message.h +++ b/quiche/quic/moqt/test_tools/moqt_test_message.h
@@ -757,10 +757,6 @@ bool EqualFieldValues(const MessageStructuredData& values) const override { auto cast = std::get<MoqtSubscribeOk>(values); - if (cast.request_id != subscribe_ok_.request_id) { - QUIC_LOG(INFO) << "SUBSCRIBE OK subscribe ID mismatch"; - return false; - } if (cast.track_alias != subscribe_ok_.track_alias) { QUIC_LOG(INFO) << "SUBSCRIBE OK track alias mismatch"; return false; @@ -776,21 +772,20 @@ return true; } - void ExpandVarints() override { ExpandVarintsImpl("vvvvvvv--v--v--vv"); } + void ExpandVarints() override { ExpandVarintsImpl("vvvvvv--v--v--vv"); } MessageStructuredData structured_data() const override { return TestMessageBase::MessageStructuredData(subscribe_ok_); } void SetInvalidDeliveryOrder() { - raw_packet_[19] = 0x10; + raw_packet_[18] = 0x10; SetWireImage(raw_packet_, sizeof(raw_packet_)); } protected: // This is protected so that TrackStatusOk can edit the fields. MoqtSubscribeOk subscribe_ok_ = { - /*request_id=*/1, /*track_alias=*/2, MessageParameters(), // Set in the constructor. TrackExtensions( @@ -803,10 +798,10 @@ }; private: - uint8_t raw_packet_[20] = { - 0x04, 0x00, 0x11, 0x01, 0x02, 0x02, // request_id, alias, 2 params - 0x08, 0x03, // expires = 3 - 0x01, 0x02, 0x0c, 0x14, // largest_location = (12, 20) + uint8_t raw_packet_[19] = { + 0x04, 0x00, 0x10, 0x02, 0x02, // alias, 2 params + 0x08, 0x03, // expires = 3 + 0x01, 0x02, 0x0c, 0x14, // largest_location = (12, 20) // Extensions 0x02, 0xa7, 0x10, // delivery_timeout = 10000 0x02, 0xa7, 0x10, // max_cache_duration = 10000 @@ -1079,33 +1074,43 @@ bool EqualFieldValues(const MessageStructuredData& values) const override { auto cast = std::get<MoqtRequestOk>(values); - if (cast.request_id != request_ok_.request_id) { - QUIC_LOG(INFO) << "REQUEST_OK request ID mismatch"; - return false; - } if (cast.parameters != request_ok_.parameters) { QUIC_LOG(INFO) << "REQUEST_OK parameter mismatch"; return false; } + if (cast.extensions != request_ok_.extensions) { + QUIC_LOG(INFO) << "REQUEST_OK extensions mismatch"; + return false; + } return true; } - void ExpandVarints() override { ExpandVarintsImpl("vvvv--"); } + void ExpandVarints() override { ExpandVarintsImpl("vvv--v--v--vv"); } MessageStructuredData structured_data() const override { return TestMessageBase::MessageStructuredData(request_ok_); } private: - uint8_t raw_packet_[9] = { - 0x07, 0x00, 0x06, 0x01, // request_id = 1 + uint8_t raw_packet_[16] = { + 0x07, 0x00, 0x0d, 0x01, // 1 parameter 0x09, 0x02, 0x05, 0x01, // Largest Object = (5, 1) + // Extensions + 0x02, 0xa7, 0x10, // delivery_timeout = 10000 + 0x02, 0xa7, 0x10, // max_cache_duration = 10000 + 0x1e, 0x02 // default_publisher_group_order = 2 }; MoqtRequestOk request_ok_ = { - /*request_id=*/1, MessageParameters(), // Set in the constructor. + TrackExtensions( + /*delivery_timeout=*/quic::QuicTimeDelta::FromMilliseconds(10000), + /*max_cache_duration=*/quic::QuicTimeDelta::FromMilliseconds(10000), + /*publisher_priority=*/std::nullopt, + /*group_order=*/MoqtDeliveryOrder::kDescending, + /*dynamic_groups=*/std::nullopt, + /*immutable_extensions=*/std::nullopt), }; };