Move PUBLISH_NAMESPACE to a bidi stream. Remove the argument to PN's cancel_callback because it's a RESET_STREAM error code instead of the full Request Error registry. Change the std::optional& MessageParameters argument in IncomingPublishNamespaceCallback to MessageParameters* to avoid a copy. Calling MoqtSession::PublishNamespaceUpdate(), Done(), or Cancel() is now a QUICHE_BUG because the application provides callbacks to be notified when a PN closes. Changes the application API, slightly, and therefore touches many files. Part of draft-18 update. PiperOrigin-RevId: 963627080
diff --git a/build/source_list.bzl b/build/source_list.bzl index 64dfb01..d810ecd 100644 --- a/build/source_list.bzl +++ b/build/source_list.bzl
@@ -1610,6 +1610,7 @@ "quic/moqt/moqt_parser.h", "quic/moqt/moqt_priority.h", "quic/moqt/moqt_probe_manager.h", + "quic/moqt/moqt_publish_namespace_stream.h", "quic/moqt/moqt_publish_stream.h", "quic/moqt/moqt_publisher.h", "quic/moqt/moqt_quic_config.h", @@ -1652,6 +1653,7 @@ "quic/moqt/moqt_parser.cc", "quic/moqt/moqt_priority.cc", "quic/moqt/moqt_probe_manager.cc", + "quic/moqt/moqt_publish_namespace_stream.cc", "quic/moqt/moqt_publish_stream.cc", "quic/moqt/moqt_quic_config.cc", "quic/moqt/moqt_relay_publisher.cc", @@ -1691,6 +1693,7 @@ "quic/moqt/moqt_parser_test.cc", "quic/moqt/moqt_priority_test.cc", "quic/moqt/moqt_probe_manager_test.cc", + "quic/moqt/moqt_publish_namespace_stream_test.cc", "quic/moqt/moqt_publish_stream_test.cc", "quic/moqt/moqt_relay_publisher_test.cc", "quic/moqt/moqt_relay_track_publisher_test.cc",
diff --git a/build/source_list.gni b/build/source_list.gni index bd07125..b4ac905 100644 --- a/build/source_list.gni +++ b/build/source_list.gni
@@ -1615,6 +1615,7 @@ "src/quiche/quic/moqt/moqt_parser.h", "src/quiche/quic/moqt/moqt_priority.h", "src/quiche/quic/moqt/moqt_probe_manager.h", + "src/quiche/quic/moqt/moqt_publish_namespace_stream.h", "src/quiche/quic/moqt/moqt_publish_stream.h", "src/quiche/quic/moqt/moqt_publisher.h", "src/quiche/quic/moqt/moqt_quic_config.h", @@ -1657,6 +1658,7 @@ "src/quiche/quic/moqt/moqt_parser.cc", "src/quiche/quic/moqt/moqt_priority.cc", "src/quiche/quic/moqt/moqt_probe_manager.cc", + "src/quiche/quic/moqt/moqt_publish_namespace_stream.cc", "src/quiche/quic/moqt/moqt_publish_stream.cc", "src/quiche/quic/moqt/moqt_quic_config.cc", "src/quiche/quic/moqt/moqt_relay_publisher.cc", @@ -1697,6 +1699,7 @@ "src/quiche/quic/moqt/moqt_parser_test.cc", "src/quiche/quic/moqt/moqt_priority_test.cc", "src/quiche/quic/moqt/moqt_probe_manager_test.cc", + "src/quiche/quic/moqt/moqt_publish_namespace_stream_test.cc", "src/quiche/quic/moqt/moqt_publish_stream_test.cc", "src/quiche/quic/moqt/moqt_relay_publisher_test.cc", "src/quiche/quic/moqt/moqt_relay_track_publisher_test.cc",
diff --git a/build/source_list.json b/build/source_list.json index 1212ca4..f27cc10 100644 --- a/build/source_list.json +++ b/build/source_list.json
@@ -1614,6 +1614,7 @@ "quiche/quic/moqt/moqt_parser.h", "quiche/quic/moqt/moqt_priority.h", "quiche/quic/moqt/moqt_probe_manager.h", + "quiche/quic/moqt/moqt_publish_namespace_stream.h", "quiche/quic/moqt/moqt_publish_stream.h", "quiche/quic/moqt/moqt_publisher.h", "quiche/quic/moqt/moqt_quic_config.h", @@ -1656,6 +1657,7 @@ "quiche/quic/moqt/moqt_parser.cc", "quiche/quic/moqt/moqt_priority.cc", "quiche/quic/moqt/moqt_probe_manager.cc", + "quiche/quic/moqt/moqt_publish_namespace_stream.cc", "quiche/quic/moqt/moqt_publish_stream.cc", "quiche/quic/moqt/moqt_quic_config.cc", "quiche/quic/moqt/moqt_relay_publisher.cc", @@ -1696,6 +1698,7 @@ "quiche/quic/moqt/moqt_parser_test.cc", "quiche/quic/moqt/moqt_priority_test.cc", "quiche/quic/moqt/moqt_probe_manager_test.cc", + "quiche/quic/moqt/moqt_publish_namespace_stream_test.cc", "quiche/quic/moqt/moqt_publish_stream_test.cc", "quiche/quic/moqt/moqt_relay_publisher_test.cc", "quiche/quic/moqt/moqt_relay_track_publisher_test.cc",
diff --git a/quiche/quic/moqt/moqt_framer.cc b/quiche/quic/moqt/moqt_framer.cc index 81bd37a..fa3930a 100644 --- a/quiche/quic/moqt/moqt_framer.cc +++ b/quiche/quic/moqt/moqt_framer.cc
@@ -568,12 +568,6 @@ WireKeyValuePairList(message.parameters.ToKeyValuePairList())); } -quiche::QuicheBuffer MoqtFramer::SerializePublishNamespaceDone( - const MoqtPublishNamespaceDone& message) { - return SerializeControlMessage(MoqtMessageType::kPublishNamespaceDone, - WireMoqVarInt(message.request_id)); -} - quiche::QuicheBuffer MoqtFramer::SerializeNamespace( const MoqtNamespace& message) { return SerializeControlMessage( @@ -588,14 +582,6 @@ WireTrackNamespace(message.track_namespace_suffix)); } -quiche::QuicheBuffer MoqtFramer::SerializePublishNamespaceCancel( - const MoqtPublishNamespaceCancel& message) { - return SerializeControlMessage( - MoqtMessageType::kPublishNamespaceCancel, - WireMoqVarInt(message.request_id), WireMoqVarInt(message.error_code), - WireStringWithMoqVarIntLength(message.error_reason)); -} - quiche::QuicheBuffer MoqtFramer::SerializeTrackStatus( const MoqtTrackStatus& message) { return SerializeSubscribe(message, MoqtMessageType::kTrackStatus);
diff --git a/quiche/quic/moqt/moqt_framer.h b/quiche/quic/moqt/moqt_framer.h index e449986..c3beed3 100644 --- a/quiche/quic/moqt/moqt_framer.h +++ b/quiche/quic/moqt/moqt_framer.h
@@ -58,12 +58,8 @@ quiche::QuicheBuffer SerializeRequestUpdate(const MoqtRequestUpdate& message); quiche::QuicheBuffer SerializePublishNamespace( const MoqtPublishNamespace& message); - quiche::QuicheBuffer SerializePublishNamespaceDone( - const MoqtPublishNamespaceDone& message); quiche::QuicheBuffer SerializeNamespace(const MoqtNamespace& message); quiche::QuicheBuffer SerializeNamespaceDone(const MoqtNamespaceDone& message); - quiche::QuicheBuffer SerializePublishNamespaceCancel( - const MoqtPublishNamespaceCancel& message); quiche::QuicheBuffer SerializeTrackStatus(const MoqtTrackStatus& message); quiche::QuicheBuffer SerializeGoAway(const MoqtGoAway& message); quiche::QuicheBuffer SerializeSubscribeNamespace(
diff --git a/quiche/quic/moqt/moqt_framer_test.cc b/quiche/quic/moqt/moqt_framer_test.cc index 3a6992a..90463dc 100644 --- a/quiche/quic/moqt/moqt_framer_test.cc +++ b/quiche/quic/moqt/moqt_framer_test.cc
@@ -48,10 +48,8 @@ MoqtMessageType::kSubscribeOk, MoqtMessageType::kPublishDone, MoqtMessageType::kPublishNamespace, - MoqtMessageType::kPublishNamespaceDone, MoqtMessageType::kNamespace, MoqtMessageType::kNamespaceDone, - MoqtMessageType::kPublishNamespaceCancel, MoqtMessageType::kTrackStatus, MoqtMessageType::kGoAway, MoqtMessageType::kSubscribeNamespace, @@ -160,10 +158,6 @@ auto data = std::get<MoqtPublishNamespace>(structured_data); return framer_.SerializePublishNamespace(data); } - case MoqtMessageType::kPublishNamespaceDone: { - auto data = std::get<MoqtPublishNamespaceDone>(structured_data); - return framer_.SerializePublishNamespaceDone(data); - } case MoqtMessageType::kNamespace: { auto data = std::get<MoqtNamespace>(structured_data); return framer_.SerializeNamespace(data); @@ -172,10 +166,6 @@ auto data = std::get<MoqtNamespaceDone>(structured_data); return framer_.SerializeNamespaceDone(data); } - case moqt::MoqtMessageType::kPublishNamespaceCancel: { - auto data = std::get<MoqtPublishNamespaceCancel>(structured_data); - return framer_.SerializePublishNamespaceCancel(data); - } case moqt::MoqtMessageType::kTrackStatus: { auto data = std::get<MoqtTrackStatus>(structured_data); return framer_.SerializeTrackStatus(data);
diff --git a/quiche/quic/moqt/moqt_integration_test.cc b/quiche/quic/moqt/moqt_integration_test.cc index 4cfa3b7..934b902 100644 --- a/quiche/quic/moqt/moqt_integration_test.cc +++ b/quiche/quic/moqt/moqt_integration_test.cc
@@ -54,6 +54,7 @@ using ::testing::Assign; using ::testing::ElementsAre; using ::testing::IsNull; +using ::testing::NotNull; using ::testing::Return; class MoqtIntegrationTest : public quiche::test::QuicheTest { @@ -181,10 +182,10 @@ parameters.authorization_tokens.emplace_back(AuthTokenType::kOutOfBand, "foo"); EXPECT_CALL(server_callbacks_.incoming_publish_namespace_callback, - Call(TrackNamespace{"foo"}, std::make_optional(parameters), _)) - .WillOnce([](const TrackNamespace&, - const std::optional<MessageParameters>&, - MoqtResponseCallback callback) { + Call(TrackNamespace{"foo"}, NotNull(), _)) + .WillOnce([&](const TrackNamespace&, const MessageParameters* params, + MoqtResponseCallback callback) { + EXPECT_TRUE(params != nullptr && *params == parameters); std::move(callback)(MessageParameters()); }); testing::MockFunction<void( @@ -192,7 +193,7 @@ response_callback; client_->session()->PublishNamespace(TrackNamespace{"foo"}, parameters, response_callback.AsStdFunction(), - [](MoqtRequestErrorInfo) {}); + []() {}); bool matches = false; EXPECT_CALL(response_callback, Call) .WillOnce( @@ -204,11 +205,9 @@ test_harness_.RunUntilWithDefaultTimeout([&]() { return matches; }); EXPECT_TRUE(success); matches = false; - EXPECT_CALL( - server_callbacks_.incoming_publish_namespace_callback, - Call(TrackNamespace{"foo"}, std::optional<MessageParameters>(), IsNull())) - .WillOnce([&](const TrackNamespace&, - const std::optional<MessageParameters>&, + EXPECT_CALL(server_callbacks_.incoming_publish_namespace_callback, + Call(TrackNamespace{"foo"}, IsNull(), _)) + .WillOnce([&](const TrackNamespace&, const MessageParameters*, MoqtResponseCallback) { matches = true; }); EXPECT_TRUE(client_->session()->PublishNamespaceDone(TrackNamespace{"foo"})); success = test_harness_.RunUntilWithDefaultTimeout([&]() { return matches; }); @@ -221,16 +220,16 @@ parameters.authorization_tokens.emplace_back(AuthTokenType::kOutOfBand, "foo"); EXPECT_CALL(server_callbacks_.incoming_publish_namespace_callback, - Call(TrackNamespace{"foo"}, std::make_optional(parameters), _)) - .WillOnce([](const TrackNamespace&, - const std::optional<MessageParameters>&, - MoqtResponseCallback callback) { + Call(TrackNamespace{"foo"}, NotNull(), _)) + .WillOnce([&](const TrackNamespace&, const MessageParameters* params, + MoqtResponseCallback callback) { + EXPECT_TRUE(params != nullptr && *params == parameters); std::move(callback)(MessageParameters()); }); testing::MockFunction<void( std::variant<MessageParameters, MoqtRequestErrorInfo>)> response_callback; - testing::MockFunction<void(MoqtRequestErrorInfo)> cancel_callback; + testing::MockFunction<void()> cancel_callback; client_->session()->PublishNamespace(TrackNamespace{"foo"}, parameters, response_callback.AsStdFunction(), cancel_callback.AsStdFunction()); @@ -244,13 +243,12 @@ test_harness_.RunUntilWithDefaultTimeout([&]() { return matches; }); EXPECT_TRUE(success); matches = false; - EXPECT_CALL(cancel_callback, - Call(MoqtRequestErrorInfo{RequestErrorCode::kInternalError, - std::nullopt, "internal error"})) - .WillOnce([&](std::optional<MoqtRequestErrorInfo>) { matches = true; }); + EXPECT_CALL(cancel_callback, Call).WillOnce([&]() { matches = true; }); + // Resetting the stream will trigger the removal callback. + EXPECT_CALL(server_callbacks_.incoming_publish_namespace_callback, + Call(TrackNamespace{"foo"}, IsNull(), _)); server_->session()->PublishNamespaceCancel(TrackNamespace{"foo"}, - RequestErrorCode::kInternalError, - "internal error"); + kResetCodeCancelled); success = test_harness_.RunUntilWithDefaultTimeout([&]() { return matches; }); EXPECT_TRUE(success); } @@ -262,10 +260,11 @@ parameters.authorization_tokens.emplace_back(AuthTokenType::kOutOfBand, "foo"); EXPECT_CALL(server_callbacks_.incoming_publish_namespace_callback, - Call(TrackNamespace{"foo"}, std::make_optional(parameters), _)) + Call(TrackNamespace{"foo"}, NotNull(), _)) .WillOnce([&](const TrackNamespace& track_namespace, - const std::optional<MessageParameters>&, + const MessageParameters* params, MoqtResponseCallback callback) { + EXPECT_TRUE(params != nullptr && *params == parameters); std::move(callback)(MessageParameters()); absl::StatusOr<FullTrackName> track_name = FullTrackName::Create(track_namespace, "/catalog"); @@ -277,22 +276,21 @@ testing::MockFunction<void( std::variant<MessageParameters, MoqtRequestErrorInfo>)> response_callback; - client_->session()->PublishNamespace(prefix, parameters, - response_callback.AsStdFunction(), - [](MoqtRequestErrorInfo) {}); + client_->session()->PublishNamespace( + prefix, parameters, response_callback.AsStdFunction(), []() {}); bool matches = false; EXPECT_CALL(response_callback, Call) .WillOnce( - [&](std::variant<MessageParameters, MoqtRequestErrorInfo> error) { - EXPECT_TRUE(std::holds_alternative<MessageParameters>(error)); + [&](std::variant<MessageParameters, MoqtRequestErrorInfo> response) { + EXPECT_TRUE(std::holds_alternative<MessageParameters>(response)); }); EXPECT_CALL(subscribe_visitor_, OnReply).WillOnce([&]() { matches = true; }); bool success = test_harness_.RunUntilWithDefaultTimeout([&]() { return matches; }); EXPECT_TRUE(success); - // Session tears down PUBLISH_NAMESPACE. + // Teardown will invoke the close callbacks. EXPECT_CALL(server_callbacks_.incoming_publish_namespace_callback, - Call(prefix, std::optional<MessageParameters>(), IsNull())); + Call(TrackNamespace{"foo"}, IsNull(), _)); } TEST_F(MoqtIntegrationTest, PublishNamespaceSuccessSendDataInResponse) { @@ -304,9 +302,9 @@ parameters.authorization_tokens.emplace_back(AuthTokenType::kOutOfBand, "foo"); EXPECT_CALL(server_callbacks_.incoming_publish_namespace_callback, - Call(TrackNamespace{"test"}, std::make_optional(parameters), _)) + Call(TrackNamespace{"test"}, NotNull(), _)) .WillOnce([&](const TrackNamespace& track_namespace, - const std::optional<MessageParameters>&, + const MessageParameters* params, MoqtResponseCallback callback) { absl::StatusOr<FullTrackName> track_name = FullTrackName::Create(track_namespace, "data"); @@ -328,8 +326,7 @@ }); client_->session()->PublishNamespace( TrackNamespace{"test"}, parameters, - [](std::variant<MessageParameters, MoqtRequestErrorInfo>) {}, - [](MoqtRequestErrorInfo) {}); + [](std::variant<MessageParameters, MoqtRequestErrorInfo>) {}, []() {}); bool success = test_harness_.RunUntilWithDefaultTimeout( [&]() { return received_subscribe_ok; }); EXPECT_TRUE(success); @@ -352,10 +349,9 @@ success = test_harness_.RunUntilWithDefaultTimeout( [&]() { return received_object; }); EXPECT_TRUE(success); - // Session tears down PUBLISH_NAMESPACE. + // Teardown will invoke the close callbacks. EXPECT_CALL(server_callbacks_.incoming_publish_namespace_callback, - Call(TrackNamespace{"test"}, std::optional<MessageParameters>(), - IsNull())); + Call(TrackNamespace{"test"}, IsNull(), _)); } TEST_F(MoqtIntegrationTest, SendMultipleGroups) { @@ -520,10 +516,9 @@ testing::MockFunction<void( std::variant<MessageParameters, MoqtRequestErrorInfo>)> response_callback; - client_->session()->PublishNamespace(TrackNamespace{"foo"}, - MessageParameters(), - response_callback.AsStdFunction(), - [](MoqtRequestErrorInfo error_info) {}); + client_->session()->PublishNamespace( + TrackNamespace{"foo"}, MessageParameters(), + response_callback.AsStdFunction(), []() {}); bool matches = false; EXPECT_CALL(response_callback, Call) .WillOnce( @@ -1098,9 +1093,9 @@ EXPECT_CALL(relay1_callbacks.incoming_publish_namespace_callback, Call) .WillRepeatedly([&relay_publisher, &relay_endpoint1]( const TrackNamespace& track_namespace, - const std::optional<MessageParameters>& parameters, + const MessageParameters* parameters, MoqtResponseCallback callback) { - if (parameters.has_value()) { + if (parameters != nullptr) { relay_publisher.OnPublishNamespace(track_namespace, *parameters, relay_endpoint1.session(), std::move(callback)); @@ -1152,7 +1147,7 @@ }); client1.session()->PublishNamespace( TrackNamespace{"test"}, MessageParameters(), - publish_response_callback.AsStdFunction(), [](MoqtRequestErrorInfo) {}); + publish_response_callback.AsStdFunction(), []() {}); success = simulator.RunUntilOrTimeout( [&]() { return publish_namespace_ok; }, quic::simulator::TestHarness::kDefaultTimeout);
diff --git a/quiche/quic/moqt/moqt_messages.cc b/quiche/quic/moqt/moqt_messages.cc index d018263..c331e4f 100644 --- a/quiche/quic/moqt/moqt_messages.cc +++ b/quiche/quic/moqt/moqt_messages.cc
@@ -85,8 +85,6 @@ return "PUBLISH_DONE"; case MoqtMessageType::kRequestUpdate: return "REQUEST_UPDATE"; - case MoqtMessageType::kPublishNamespaceCancel: - return "PUBLISH_NAMESPACE_CANCEL"; case MoqtMessageType::kTrackStatus: return "TRACK_STATUS"; case MoqtMessageType::kPublishNamespace: @@ -97,8 +95,6 @@ return "NAMESPACE_DONE"; case MoqtMessageType::kRequestOk: return "REQUEST_OK"; - case MoqtMessageType::kPublishNamespaceDone: - return "PUBLISH_NAMESPACE_DONE"; case MoqtMessageType::kGoAway: return "GOAWAY"; case MoqtMessageType::kSubscribeNamespace:
diff --git a/quiche/quic/moqt/moqt_messages.h b/quiche/quic/moqt/moqt_messages.h index 98fa82e..f79dc04 100644 --- a/quiche/quic/moqt/moqt_messages.h +++ b/quiche/quic/moqt/moqt_messages.h
@@ -206,10 +206,8 @@ kPublishNamespace = 0x06, kRequestOk = 0x07, kNamespace = 0x08, - kPublishNamespaceDone = 0x09, kUnsubscribe = 0x0a, kPublishDone = 0x0b, - kPublishNamespaceCancel = 0x0c, kTrackStatus = 0x0d, kNamespaceDone = 0x0e, kGoAway = 0x10, @@ -411,16 +409,6 @@ MessageParameters parameters; }; -struct QUICHE_EXPORT MoqtPublishNamespaceDone { - uint64_t request_id; -}; - -struct QUICHE_EXPORT MoqtPublishNamespaceCancel { - uint64_t request_id; - RequestErrorCode error_code; - std::string error_reason; -}; - struct QUICHE_EXPORT MoqtTrackStatus : public MoqtSubscribe { MoqtTrackStatus() = default; MoqtTrackStatus(MoqtSubscribe subscribe) : MoqtSubscribe(subscribe) {}
diff --git a/quiche/quic/moqt/moqt_parser.cc b/quiche/quic/moqt/moqt_parser.cc index c7aa980..ab6babe 100644 --- a/quiche/quic/moqt/moqt_parser.cc +++ b/quiche/quic/moqt/moqt_parser.cc
@@ -791,35 +791,6 @@ return request_ok; } -absl::StatusOr<MoqtPublishNamespaceDone> -MoqtControlMessageParser::ProcessPublishNamespaceDone( - absl::string_view data) const { - quic::QuicDataReader reader(data); - MoqtPublishNamespaceDone pn_done; - if (!reader.ReadMoqVarInt(&pn_done.request_id)) { - return absl::InvalidArgumentError("Request ID missing"); - } - QUICHE_RETURN_IF_ERROR(CheckForTrailingData(reader)); - return pn_done; -} - -absl::StatusOr<MoqtPublishNamespaceCancel> -MoqtControlMessageParser::ProcessPublishNamespaceCancel( - absl::string_view data) const { - quic::QuicDataReader reader(data); - MoqtPublishNamespaceCancel publish_namespace_cancel; - uint64_t error_code; - if (!reader.ReadMoqVarInt(&publish_namespace_cancel.request_id) || - !reader.ReadMoqVarInt(&error_code) || - !reader.ReadStringMoqVarInt(publish_namespace_cancel.error_reason)) { - return absl::InvalidArgumentError("Message missing fields"); - } - publish_namespace_cancel.error_code = - static_cast<RequestErrorCode>(error_code); - QUICHE_RETURN_IF_ERROR(CheckForTrailingData(reader)); - return publish_namespace_cancel; -} - absl::StatusOr<MoqtTrackStatus> MoqtControlMessageParser::ProcessTrackStatus( absl::string_view data) const { return ProcessSubscribe(data);
diff --git a/quiche/quic/moqt/moqt_parser.h b/quiche/quic/moqt/moqt_parser.h index d28e203..0a2e46e 100644 --- a/quiche/quic/moqt/moqt_parser.h +++ b/quiche/quic/moqt/moqt_parser.h
@@ -161,13 +161,9 @@ absl::string_view data) const; absl::StatusOr<MoqtPublishNamespace> ProcessPublishNamespace( absl::string_view data) const; - absl::StatusOr<MoqtPublishNamespaceDone> ProcessPublishNamespaceDone( - absl::string_view data) const; absl::StatusOr<MoqtNamespace> ProcessNamespace(absl::string_view data) const; absl::StatusOr<MoqtNamespaceDone> ProcessNamespaceDone( absl::string_view data) const; - absl::StatusOr<MoqtPublishNamespaceCancel> ProcessPublishNamespaceCancel( - absl::string_view data) const; absl::StatusOr<MoqtTrackStatus> ProcessTrackStatus( absl::string_view data) const; absl::StatusOr<MoqtGoAway> ProcessGoAway(absl::string_view data) const; @@ -218,14 +214,10 @@ return parse(&MoqtControlMessageParser::ProcessRequestUpdate); case MoqtMessageType::kPublishNamespace: return parse(&MoqtControlMessageParser::ProcessPublishNamespace); - case MoqtMessageType::kPublishNamespaceDone: - return parse(&MoqtControlMessageParser::ProcessPublishNamespaceDone); case MoqtMessageType::kNamespace: return parse(&MoqtControlMessageParser::ProcessNamespace); case MoqtMessageType::kNamespaceDone: return parse(&MoqtControlMessageParser::ProcessNamespaceDone); - case MoqtMessageType::kPublishNamespaceCancel: - return parse(&MoqtControlMessageParser::ProcessPublishNamespaceCancel); case MoqtMessageType::kTrackStatus: return parse(&MoqtControlMessageParser::ProcessTrackStatus); case MoqtMessageType::kGoAway:
diff --git a/quiche/quic/moqt/moqt_parser_test.cc b/quiche/quic/moqt/moqt_parser_test.cc index 3407da4..c393367 100644 --- a/quiche/quic/moqt/moqt_parser_test.cc +++ b/quiche/quic/moqt/moqt_parser_test.cc
@@ -55,10 +55,8 @@ MoqtMessageType::kPublishDone, MoqtMessageType::kTrackStatus, MoqtMessageType::kPublishNamespace, - MoqtMessageType::kPublishNamespaceDone, MoqtMessageType::kNamespace, MoqtMessageType::kNamespaceDone, - MoqtMessageType::kPublishNamespaceCancel, MoqtMessageType::kGoAway, MoqtMessageType::kSubscribeNamespace, MoqtMessageType::kSubscribeTracks,
diff --git a/quiche/quic/moqt/moqt_publish_namespace_stream.cc b/quiche/quic/moqt/moqt_publish_namespace_stream.cc new file mode 100644 index 0000000..522fd0c --- /dev/null +++ b/quiche/quic/moqt/moqt_publish_namespace_stream.cc
@@ -0,0 +1,165 @@ +// 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_publish_namespace_stream.h" + +#include <optional> +#include <utility> +#include <variant> + +#include "absl/functional/overload.h" +#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_key_value_pair.h" +#include "quiche/quic/moqt/moqt_messages.h" +#include "quiche/quic/moqt/moqt_parser.h" + +namespace moqt { + +void MoqtPublishNamespaceRequestStream::OnStreamBound() { + // TODO(martinduke): Set the priority for this stream. + SendOrBufferMessageOrFatal( + framer()->SerializePublishNamespace( + MoqtPublishNamespace{request_id_, prefix_, parameters_}), + false); + QUIC_DLOG(INFO) << "Sent PUBLISH_NAMESPACE message for " << prefix_; +} + +absl::Status MoqtPublishNamespaceRequestStream::OnRawControlMessage( + const MoqtRawControlMessage& message) { + return ControlMessageDispatcher::DispatchControlMessage( + *this, message_parser(), message, "namespace publisher"); +} + +absl::Status MoqtPublishNamespaceRequestStream::OnControlMessage( + const MoqtRequestOk& message) { + if (response_callback_ != nullptr) { + // Response to the initial PUBLISH_NAMESPACE. + auto callback = std::move(response_callback_); + response_callback_ = nullptr; + std::move(callback)(message.parameters); + return absl::OkStatus(); + } + absl::StatusOr<MessageParameters> old_parameters = + request_update_queue().NextParameters(); + if (!old_parameters.ok()) { + return old_parameters.status(); + } + parameters_.Update(*old_parameters); + // Response to REQUEST_UPDATE. + return request_update_queue().OnControlMessage(message); +} + +absl::Status MoqtPublishNamespaceRequestStream::OnControlMessage( + const MoqtRequestError& message) { + if (response_callback_ != nullptr) { + // Response to the initial PUBLISH_NAMESPACE. + auto callback = std::move(response_callback_); + response_callback_ = nullptr; + Fin(); + std::move(callback)(MoqtRequestErrorInfo{ + message.error_code, message.retry_interval, message.reason_phrase}); + return absl::OkStatus(); + } + // The REQUEST_ERROR is a response to the REQUEST_UPDATE message. + absl::Status status = request_update_queue().OnControlMessage(message); + if (status.ok()) { + Fin(); + } + return status; +} + +void MoqtPublishNamespaceRequestStream::Detach() { + if (remove_callback_ != nullptr) { + RemovePublishNamespaceCallback callback = std::move(remove_callback_); + remove_callback_ = nullptr; + std::move(callback)(prefix_); + } +} + +absl::Status MoqtPublishNamespaceResponseStream::OnRawControlMessage( + const MoqtRawControlMessage& message) { + return ControlMessageDispatcher::DispatchControlMessage( + *this, message_parser(), message, "namespace publisher"); +} + +absl::Status MoqtPublishNamespaceResponseStream::OnControlMessage( + const MoqtPublishNamespace& message) { + if (add_callback_ == nullptr) { + return absl::InvalidArgumentError("Two PUBLISH_NAMESPACE on one stream"); + } + request_id_ = message.request_id; + if (!std::move(add_callback_)(message.track_namespace, this)) { + add_callback_ = nullptr; + return SendRequestError(request_id_, RequestErrorCode::kInternalError, + std::nullopt, "", /*fin=*/true); + } + add_callback_ = nullptr; + prefix_ = message.track_namespace; + application_( + *prefix_, &message.parameters, + [weakptr = weak_ptr_factory_.Create(), id = request_id_]( + 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( + id, error.error_code, error.retry_interval, + error.reason_phrase)); + }}, + response); + }); + return absl::OkStatus(); +} + +absl::Status MoqtPublishNamespaceResponseStream::OnControlMessage( + const MoqtRequestUpdate& message) { + if (!prefix_.has_value()) { + return absl::InvalidArgumentError( + "REQUEST_UPDATE before PUBLISH_NAMESPACE on a PN stream"); + } + application_( + *prefix_, &message.parameters, + [weakptr = weak_ptr_factory_.Create(), id = message.request_id]( + 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( + id, error.error_code, error.retry_interval, + error.reason_phrase)); + }}, + response); + }); + return absl::OkStatus(); +} + +void MoqtPublishNamespaceResponseStream::Detach() { + if (!prefix_.has_value()) { + return; + } + if (remove_callback_ != nullptr) { + RemovePublishNamespaceCallback callback = std::move(remove_callback_); + remove_callback_ = nullptr; + std::move(callback)(*prefix_); + application_(*prefix_, nullptr, nullptr); + } +} + +} // namespace moqt
diff --git a/quiche/quic/moqt/moqt_publish_namespace_stream.h b/quiche/quic/moqt/moqt_publish_namespace_stream.h new file mode 100644 index 0000000..1ad38d6 --- /dev/null +++ b/quiche/quic/moqt/moqt_publish_namespace_stream.h
@@ -0,0 +1,106 @@ +// 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_PUBLISH_NAMESPACE_STREAM_H_ +#define QUICHE_QUIC_MOQT_MOQT_PUBLISH_NAMESPACE_STREAM_H_ + +#include <cstdint> +#include <optional> +#include <utility> + +#include "absl/status/status.h" +#include "quiche/quic/moqt/moqt_bidi_stream.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/common/quiche_callbacks.h" +#include "quiche/common/quiche_weak_ptr.h" + +namespace moqt { + +using AddPublishNamespaceCallback = + quiche::SingleUseCallback<bool(const TrackNamespace&, MoqtBidiStreamBase*)>; +using RemovePublishNamespaceCallback = + quiche::SingleUseCallback<void(const TrackNamespace&)>; + +// This class will be owned by the webtransport stream. +class MoqtPublishNamespaceRequestStream : public MoqtBidiStreamBase { + public: + // Assumes the caller will send or queue the PUBLISH_NAMESPACE. + MoqtPublishNamespaceRequestStream( + const TrackNamespace& prefix, const MessageParameters& parameters, + MoqtFramer* framer, const MoqtControlMessageParser& message_parser, + uint64_t request_id, RemovePublishNamespaceCallback remove_callback, + SessionErrorCallback session_error_callback, + MoqtResponseCallback response_callback) + : MoqtBidiStreamBase(framer, message_parser, + std::move(session_error_callback)), + request_id_(request_id), + remove_callback_(std::move(remove_callback)), + response_callback_(std::move(response_callback)), + prefix_(prefix), + parameters_(parameters) {} + ~MoqtPublishNamespaceRequestStream() { Detach(); } + + // MoqtBidiStreamBase overrides. + void OnStreamBound() override; + absl::Status OnRawControlMessage( + const MoqtRawControlMessage& message) override; + absl::Status OnControlMessage(const MoqtRequestOk& message); + absl::Status OnControlMessage(const MoqtRequestError& message); + + void Detach() override; + + private: + const uint64_t request_id_; + RemovePublishNamespaceCallback remove_callback_; + MoqtResponseCallback response_callback_; + const TrackNamespace prefix_; + MessageParameters parameters_; +}; + +class MoqtPublishNamespaceResponseStream : public MoqtBidiStreamBase { + public: + // Constructor for the publisher side. + MoqtPublishNamespaceResponseStream( + MoqtFramer* framer, const MoqtControlMessageParser& message_parser, + AddPublishNamespaceCallback add_callback, + RemovePublishNamespaceCallback remove_callback, + SessionErrorCallback session_error_callback, + MoqtIncomingPublishNamespaceCallback application) + : MoqtBidiStreamBase(framer, message_parser, + std::move(session_error_callback)), + add_callback_(std::move(add_callback)), + remove_callback_(std::move(remove_callback)), + application_(std::move(application)), + weak_ptr_factory_(this) {} + ~MoqtPublishNamespaceResponseStream() { Detach(); } + + void OnStreamBound() override { + // TODO(martinduke): Set the priority for this stream. + } + absl::Status OnRawControlMessage( + const MoqtRawControlMessage& message) override; + absl::Status OnControlMessage(const MoqtPublishNamespace& message); + absl::Status OnControlMessage(const MoqtRequestUpdate& message); + + void Detach() override; + + private: + uint64_t request_id_; + std::optional<TrackNamespace> prefix_; + AddPublishNamespaceCallback add_callback_; + RemovePublishNamespaceCallback remove_callback_; + MoqtIncomingPublishNamespaceCallback application_; + quiche::QuicheWeakPtrFactory<MoqtPublishNamespaceResponseStream> + weak_ptr_factory_; +}; + +} // namespace moqt + +#endif // QUICHE_QUIC_MOQT_MOQT_PUBLISH_NAMESPACE_STREAM_H_
diff --git a/quiche/quic/moqt/moqt_publish_namespace_stream_test.cc b/quiche/quic/moqt/moqt_publish_namespace_stream_test.cc new file mode 100644 index 0000000..f244f80 --- /dev/null +++ b/quiche/quic/moqt/moqt_publish_namespace_stream_test.cc
@@ -0,0 +1,422 @@ +// 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_publish_namespace_stream.h" + +#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_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_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" +#include "quiche/web_transport/test_tools/mock_web_transport.h" + +namespace moqt::test { + +using ::testing::_; +using ::testing::Return; +using ::testing::StrictMock; + +class MoqtPublishNamespaceRequestStreamTest : public quiche::test::QuicheTest { + protected: + MoqtPublishNamespaceRequestStreamTest() + : framer_(true, quic::Perspective::IS_CLIENT), + remove_callback_(), + session_error_callback_(), + response_callback_() { + EXPECT_CALL(remove_callback_, Call(_)).Times(testing::AnyNumber()); + } + + std::unique_ptr<MoqtPublishNamespaceRequestStream> CreateAndBindStream() { + EXPECT_CALL(mock_stream_, CanWrite).WillRepeatedly(Return(true)); + EXPECT_CALL( + mock_stream_, + Writev(ControlMessageOfType(MoqtMessageType::kPublishNamespace), _)); + auto stream = std::make_unique<MoqtPublishNamespaceRequestStream>( + TrackNamespace({"foo"}), MessageParameters(), &framer_, + MoqtControlMessageParser(kDefaultMoqtVersion, /*webtransport=*/true, + quic::Perspective::IS_CLIENT), + /*request_id=*/10, remove_callback_.AsStdFunction(), + session_error_callback_.AsStdFunction(), + response_callback_.AsStdFunction()); + stream->BindStream(&mock_stream_); + return stream; + } + + MoqtFramer framer_; + StrictMock<testing::MockFunction<void(const TrackNamespace&)>> + remove_callback_; + StrictMock<testing::MockFunction<void(MoqtError, absl::string_view)>> + session_error_callback_; + StrictMock<testing::MockFunction<void( + std::variant<MessageParameters, MoqtRequestErrorInfo>)>> + response_callback_; + webtransport::test::MockStream mock_stream_; +}; + +TEST_F(MoqtPublishNamespaceRequestStreamTest, OnStreamBound) { + // Creating and binding the stream verifies OnStreamBound sends + // PUBLISH_NAMESPACE. + std::unique_ptr<MoqtPublishNamespaceRequestStream> request_stream = + CreateAndBindStream(); +} + +TEST_F(MoqtPublishNamespaceRequestStreamTest, DetachCallsRemoveCallback) { + std::unique_ptr<MoqtPublishNamespaceRequestStream> request_stream = + CreateAndBindStream(); + EXPECT_CALL(remove_callback_, Call(TrackNamespace({"foo"}))); + request_stream = nullptr; +} + +TEST_F(MoqtPublishNamespaceRequestStreamTest, OnControlMessageOk) { + std::unique_ptr<MoqtPublishNamespaceRequestStream> request_stream = + CreateAndBindStream(); + bool callback_called = false; + EXPECT_CALL(response_callback_, Call(_)) + .WillOnce([&](std::variant<MessageParameters, MoqtRequestErrorInfo> res) { + callback_called = true; + EXPECT_TRUE(std::holds_alternative<MessageParameters>(res)); + }); + + MoqtRequestOk message; + message.request_id = 10; + QUICHE_EXPECT_OK(request_stream->OnControlMessage(message)); + EXPECT_TRUE(callback_called); +} + +TEST_F(MoqtPublishNamespaceRequestStreamTest, OnControlMessageError) { + std::unique_ptr<MoqtPublishNamespaceRequestStream> request_stream = + CreateAndBindStream(); + bool callback_called = false; + EXPECT_CALL(response_callback_, Call(_)) + .WillOnce([&](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); + }); + ExpectFin(mock_stream_); + + MoqtRequestError message; + message.request_id = 10; + message.error_code = RequestErrorCode::kUnauthorized; + QUICHE_EXPECT_OK(request_stream->OnControlMessage(message)); + EXPECT_TRUE(callback_called); +} + +TEST_F(MoqtPublishNamespaceRequestStreamTest, SendRequestUpdateAndReceiveOk) { + std::unique_ptr<MoqtPublishNamespaceRequestStream> request_stream = + 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)); + + // Now send update. + EXPECT_CALL(mock_stream_, + Writev(ControlMessageOfType(MoqtMessageType::kRequestUpdate), _)); + + MessageParameters parameters; + parameters.subscriber_priority = 50; + bool update_callback_called = false; + MoqtResponseCallback update_callback = + [&](std::variant<MessageParameters, MoqtRequestErrorInfo> res) { + update_callback_called = true; + 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); +} + +class MoqtPublishNamespaceResponseStreamTest : public quiche::test::QuicheTest { + protected: + MoqtPublishNamespaceResponseStreamTest() + : framer_(true, quic::Perspective::IS_SERVER), + session_error_callback_(), + add_callback_(), + remove_callback_(), + application_() { + EXPECT_CALL(remove_callback_, Call(_)).Times(testing::AnyNumber()); + EXPECT_CALL(application_, Call(_, nullptr, _)).Times(testing::AnyNumber()); + } + + MoqtFramer framer_; + StrictMock<testing::MockFunction<void(MoqtError, absl::string_view)>> + session_error_callback_; + StrictMock< + testing::MockFunction<bool(const TrackNamespace&, MoqtBidiStreamBase*)>> + add_callback_; + StrictMock<testing::MockFunction<void(const TrackNamespace&)>> + remove_callback_; + StrictMock<testing::MockFunction<void( + const TrackNamespace&, const MessageParameters*, MoqtResponseCallback)>> + application_; + webtransport::test::MockStream mock_stream_; +}; + +TEST_F(MoqtPublishNamespaceResponseStreamTest, + OnPublishNamespaceAddCallbackFails) { + auto response_stream = std::make_unique<MoqtPublishNamespaceResponseStream>( + &framer_, + MoqtControlMessageParser(kDefaultMoqtVersion, /*webtransport=*/true, + quic::Perspective::IS_SERVER), + add_callback_.AsStdFunction(), remove_callback_.AsStdFunction(), + session_error_callback_.AsStdFunction(), application_.AsStdFunction()); + response_stream->BindStream(&mock_stream_); + + MoqtPublishNamespace message; + message.request_id = 5; + message.track_namespace = TrackNamespace({"foo"}); + + EXPECT_CALL(add_callback_, + Call(TrackNamespace({"foo"}), response_stream.get())) + .WillOnce(Return(false)); + EXPECT_CALL(mock_stream_, CanWrite).WillRepeatedly(Return(true)); + EXPECT_CALL(mock_stream_, + Writev(ControlMessageOfType(MoqtMessageType::kRequestError), _)); + + QUICHE_EXPECT_OK(response_stream->OnControlMessage(message)); +} + +TEST_F(MoqtPublishNamespaceResponseStreamTest, + OnPublishNamespaceSuccessAndApplicationAccepts) { + auto response_stream = std::make_unique<MoqtPublishNamespaceResponseStream>( + &framer_, + MoqtControlMessageParser(kDefaultMoqtVersion, /*webtransport=*/true, + quic::Perspective::IS_SERVER), + add_callback_.AsStdFunction(), remove_callback_.AsStdFunction(), + session_error_callback_.AsStdFunction(), application_.AsStdFunction()); + response_stream->BindStream(&mock_stream_); + + MoqtPublishNamespace message; + message.request_id = 5; + message.track_namespace = TrackNamespace({"foo"}); + + MoqtResponseCallback application_callback; + EXPECT_CALL(add_callback_, + Call(TrackNamespace({"foo"}), response_stream.get())) + .WillOnce(Return(true)); + EXPECT_CALL(application_, + Call(TrackNamespace({"foo"}), testing::NotNull(), _)) + .WillOnce([&](const TrackNamespace&, const MessageParameters*, + MoqtResponseCallback callback) { + application_callback = std::move(callback); + }); + + QUICHE_EXPECT_OK(response_stream->OnControlMessage(message)); + ASSERT_TRUE(application_callback != nullptr); + + // Application accepts the publish namespace. + EXPECT_CALL(mock_stream_, CanWrite).WillRepeatedly(Return(true)); + EXPECT_CALL(mock_stream_, + Writev(ControlMessageOfType(MoqtMessageType::kRequestOk), _)); + std::move(application_callback)(MessageParameters()); +} + +TEST_F(MoqtPublishNamespaceResponseStreamTest, + DetachCleansUpPrefixAndNotifiesApplication) { + auto response_stream = std::make_unique<MoqtPublishNamespaceResponseStream>( + &framer_, + MoqtControlMessageParser(kDefaultMoqtVersion, /*webtransport=*/true, + quic::Perspective::IS_SERVER), + add_callback_.AsStdFunction(), remove_callback_.AsStdFunction(), + session_error_callback_.AsStdFunction(), application_.AsStdFunction()); + response_stream->BindStream(&mock_stream_); + + MoqtPublishNamespace message; + message.request_id = 5; + message.track_namespace = TrackNamespace({"foo"}); + + MoqtResponseCallback application_callback; + EXPECT_CALL(add_callback_, + Call(TrackNamespace({"foo"}), response_stream.get())) + .WillOnce(Return(true)); + EXPECT_CALL(application_, + Call(TrackNamespace({"foo"}), testing::NotNull(), _)) + .WillOnce([&](const TrackNamespace&, const MessageParameters*, + MoqtResponseCallback callback) { + application_callback = std::move(callback); + }); + + QUICHE_EXPECT_OK(response_stream->OnControlMessage(message)); + + // Accept it first. + EXPECT_CALL(mock_stream_, CanWrite).WillRepeatedly(Return(true)); + EXPECT_CALL(mock_stream_, + Writev(ControlMessageOfType(MoqtMessageType::kRequestOk), _)); + std::move(application_callback)(MessageParameters()); + + // Since it was published (published_ = true), Detach() should call + // remove_callback_ and application_(*prefix_, nullptr, nullptr). + EXPECT_CALL(remove_callback_, Call(TrackNamespace({"foo"}))); + EXPECT_CALL(application_, Call(TrackNamespace({"foo"}), nullptr, nullptr)); + + response_stream = nullptr; // Destroys response_stream, triggering Detach(). +} + +TEST_F(MoqtPublishNamespaceResponseStreamTest, + DoublePublishNamespaceReturnsError) { + auto response_stream = std::make_unique<MoqtPublishNamespaceResponseStream>( + &framer_, + MoqtControlMessageParser(kDefaultMoqtVersion, /*webtransport=*/true, + quic::Perspective::IS_SERVER), + add_callback_.AsStdFunction(), remove_callback_.AsStdFunction(), + session_error_callback_.AsStdFunction(), application_.AsStdFunction()); + response_stream->BindStream(&mock_stream_); + + MoqtPublishNamespace message; + message.request_id = 5; + message.track_namespace = TrackNamespace({"foo"}); + + EXPECT_CALL(add_callback_, + Call(TrackNamespace({"foo"}), response_stream.get())) + .WillOnce(Return(true)); + EXPECT_CALL(application_, + Call(TrackNamespace({"foo"}), testing::NotNull(), _)); + + QUICHE_EXPECT_OK(response_stream->OnControlMessage(message)); + + // Second call must return absl::InvalidArgumentError + EXPECT_EQ(response_stream->OnControlMessage(message).code(), + absl::StatusCode::kInvalidArgument); +} + +TEST_F(MoqtPublishNamespaceResponseStreamTest, OnRequestUpdateSuccess) { + auto response_stream = std::make_unique<MoqtPublishNamespaceResponseStream>( + &framer_, + MoqtControlMessageParser(kDefaultMoqtVersion, /*webtransport=*/true, + quic::Perspective::IS_SERVER), + add_callback_.AsStdFunction(), remove_callback_.AsStdFunction(), + session_error_callback_.AsStdFunction(), application_.AsStdFunction()); + response_stream->BindStream(&mock_stream_); + + MoqtPublishNamespace message; + message.request_id = 5; + message.track_namespace = TrackNamespace({"foo"}); + + MoqtResponseCallback application_callback; + EXPECT_CALL(add_callback_, + Call(TrackNamespace({"foo"}), response_stream.get())) + .WillOnce(Return(true)); + EXPECT_CALL(application_, + Call(TrackNamespace({"foo"}), testing::NotNull(), _)) + .WillOnce([&](const TrackNamespace&, const MessageParameters*, + MoqtResponseCallback callback) { + application_callback = std::move(callback); + }); + + QUICHE_EXPECT_OK(response_stream->OnControlMessage(message)); + + // Accept it. + EXPECT_CALL(mock_stream_, CanWrite).WillRepeatedly(Return(true)); + EXPECT_CALL(mock_stream_, + Writev(ControlMessageOfType(MoqtMessageType::kRequestOk), _)); + std::move(application_callback)(MessageParameters()); + + // Now send a RequestUpdate. + MoqtRequestUpdate update; + update.request_id = 12; + update.parameters.subscriber_priority = 40; + + MoqtResponseCallback update_callback; + EXPECT_CALL(application_, + Call(TrackNamespace({"foo"}), testing::NotNull(), _)) + .WillOnce([&](const TrackNamespace&, const MessageParameters*, + MoqtResponseCallback callback) { + update_callback = std::move(callback); + }); + + QUICHE_EXPECT_OK(response_stream->OnControlMessage(update)); + ASSERT_TRUE(update_callback != nullptr); + + // Application accepts the update. + EXPECT_CALL(mock_stream_, + Writev(ControlMessageOfType(MoqtMessageType::kRequestOk), _)); + std::move(update_callback)(MessageParameters()); +} + +TEST_F(MoqtPublishNamespaceResponseStreamTest, OnRequestUpdateRejected) { + auto response_stream = std::make_unique<MoqtPublishNamespaceResponseStream>( + &framer_, + MoqtControlMessageParser(kDefaultMoqtVersion, /*webtransport=*/true, + quic::Perspective::IS_SERVER), + add_callback_.AsStdFunction(), remove_callback_.AsStdFunction(), + session_error_callback_.AsStdFunction(), application_.AsStdFunction()); + response_stream->BindStream(&mock_stream_); + + MoqtPublishNamespace message; + message.request_id = 5; + message.track_namespace = TrackNamespace({"foo"}); + + MoqtResponseCallback application_callback; + EXPECT_CALL(add_callback_, + Call(TrackNamespace({"foo"}), response_stream.get())) + .WillOnce(Return(true)); + EXPECT_CALL(application_, + Call(TrackNamespace({"foo"}), testing::NotNull(), _)) + .WillOnce([&](const TrackNamespace&, const MessageParameters*, + MoqtResponseCallback callback) { + application_callback = std::move(callback); + }); + + QUICHE_EXPECT_OK(response_stream->OnControlMessage(message)); + + // Accept it. + EXPECT_CALL(mock_stream_, CanWrite).WillRepeatedly(Return(true)); + EXPECT_CALL(mock_stream_, + Writev(ControlMessageOfType(MoqtMessageType::kRequestOk), _)); + std::move(application_callback)(MessageParameters()); + + // Now send a RequestUpdate. + MoqtRequestUpdate update; + update.request_id = 12; + update.parameters.subscriber_priority = 40; + + MoqtResponseCallback update_callback; + EXPECT_CALL(application_, + Call(TrackNamespace({"foo"}), testing::NotNull(), _)) + .WillOnce([&](const TrackNamespace&, const MessageParameters*, + MoqtResponseCallback callback) { + update_callback = std::move(callback); + }); + + QUICHE_EXPECT_OK(response_stream->OnControlMessage(update)); + + // Application rejects the update. + MoqtRequestErrorInfo error_info = { + RequestErrorCode::kUnauthorized, + /*retry_interval=*/std::nullopt, + "unauthorized", + }; + EXPECT_CALL(mock_stream_, + Writev(ControlMessageOfType(MoqtMessageType::kRequestError), _)); + std::move(update_callback)(error_info); +} + +} // namespace moqt::test
diff --git a/quiche/quic/moqt/moqt_session.cc b/quiche/quic/moqt/moqt_session.cc index 4b7c3a6..36a7056 100644 --- a/quiche/quic/moqt/moqt_session.cc +++ b/quiche/quic/moqt/moqt_session.cc
@@ -19,7 +19,6 @@ #include "absl/container/flat_hash_set.h" #include "absl/container/node_hash_map.h" #include "absl/functional/bind_front.h" -#include "absl/functional/overload.h" #include "absl/memory/memory.h" #include "absl/status/status.h" #include "absl/status/statusor.h" @@ -41,6 +40,7 @@ #include "quiche/quic/moqt/moqt_object_subscriber.h" #include "quiche/quic/moqt/moqt_parser.h" #include "quiche/quic/moqt/moqt_priority.h" +#include "quiche/quic/moqt/moqt_publish_namespace_stream.h" #include "quiche/quic/moqt/moqt_publish_stream.h" #include "quiche/quic/moqt/moqt_publisher.h" #include "quiche/quic/moqt/moqt_session_callbacks.h" @@ -53,6 +53,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_callbacks.h" #include "quiche/common/quiche_mem_slice.h" #include "quiche/common/quiche_weak_ptr.h" #include "quiche/web_transport/web_transport.h" @@ -377,25 +378,8 @@ bool MoqtSession::PublishNamespace( const TrackNamespace& track_namespace, const MessageParameters& parameters, MoqtResponseCallback response_callback, - quiche::SingleUseCallback<void(MoqtRequestErrorInfo)> cancel_callback) { - if (is_closing_) { - return false; - } - if (publish_namespace_by_namespace_.contains(track_namespace)) { - return false; - } - 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 PUBLISH_NAMESPACE with ID " - << next_request_id_ - << " which is greater than the maximum ID " - << peer_max_request_id_; + quiche::SingleUseCallback<void()> cancel_callback) { + if (is_closing_ || publish_namespace_requests_.contains(track_namespace)) { return false; } if (received_goaway_ || sent_goaway_) { @@ -403,18 +387,34 @@ << "Tried to send PUBLISH_NAMESPACE after GOAWAY"; return false; } - publish_namespace_by_namespace_[track_namespace] = next_request_id_; - publish_namespace_by_id_[next_request_id_] = - PublishNamespaceState{track_namespace, std::move(response_callback), - std::move(cancel_callback)}; - MoqtPublishNamespace message; - message.request_id = next_request_id_; - next_request_id_ += 2; - message.track_namespace = track_namespace; - message.parameters = parameters; - SendControlMessage(framer_.SerializePublishNamespace(message)); - QUIC_DLOG(INFO) << ENDPOINT << "Sent PUBLISH_NAMESPACE message for " - << message.track_namespace; + webtransport::Stream* stream = session_->OpenOutgoingBidirectionalStream(); + if (stream == nullptr) { + return false; + } + auto stream_visitor = std::make_unique<MoqtPublishNamespaceRequestStream>( + track_namespace, parameters, &framer_, ControlMessageParser(), + NextRequestId(), + [weakptr = GetWeakPtr(), callback = std::move(cancel_callback)]( + const TrackNamespace& prefix) mutable { + MoqtSession* session = MoqtSessionFromWeakPtr(weakptr); + if (session == nullptr) { + return; + } + session->publish_namespace_requests_.erase(prefix); + std::move(callback)(); + }, + [weakptr = GetWeakPtr()](MoqtError code, absl::string_view reason) { + MoqtSession* session = MoqtSessionFromWeakPtr(weakptr); + if (session == nullptr) { + return; + } + session->Error(code, reason); + }, + std::move(response_callback)); + MoqtPublishNamespaceRequestStream* stream_ptr = stream_visitor.get(); + publish_namespace_requests_.emplace(track_namespace, stream_ptr); + stream->SetVisitor(std::move(stream_visitor)); + stream_ptr->BindStream(stream); return true; } @@ -424,31 +424,15 @@ if (is_closing_) { return false; } - auto it = publish_namespace_by_namespace_.find(track_namespace); - if (it == publish_namespace_by_namespace_.end()) { - return false; // Could have been destroyed by PUBLISH_NAMESPACE_CANCEL. - } - 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 PUBLISH_NAMESPACE with ID " - << next_request_id_ - << " which is greater than the maximum ID " - << peer_max_request_id_; + auto it = publish_namespace_requests_.find(track_namespace); + if (it == publish_namespace_requests_.end()) { + QUICHE_BUG(quic_bug_publish_namespace_update_after_closure) + << "Tried to send PUBLISH_NAMESPACE_UPDATE for unknown namespace " + << track_namespace; return false; } - MoqtRequestUpdate message; - message.request_id = next_request_id_; - message.existing_request_id = it->second; - message.parameters = parameters; - publish_namespace_updates_[next_request_id_] = std::move(response_callback); - next_request_id_ += 2; - SendControlMessage(framer_.SerializeRequestUpdate(message)); + it->second->CheckStatus(it->second->SendRequestUpdate( + NextRequestId(), 0, parameters, std::move(response_callback))); return true; } @@ -456,33 +440,32 @@ if (is_closing_) { return false; } - auto it = publish_namespace_by_namespace_.find(track_namespace); - if (it == publish_namespace_by_namespace_.end()) { - return false; // Could have been destroyed by PUBLISH_NAMESPACE_CANCEL. + auto it = publish_namespace_requests_.find(track_namespace); + if (it == publish_namespace_requests_.end()) { + QUICHE_BUG(quic_bug_publish_namespace_update_after_closure) + << "Tried to reset PUBLISH_NAMESPACE for unknown namespace " + << track_namespace; + return false; } - MoqtPublishNamespaceDone message; - message.request_id = it->second; - SendControlMessage(framer_.SerializePublishNamespaceDone(message)); - QUIC_DLOG(INFO) << ENDPOINT << "Sent PUBLISH_NAMESPACE_DONE message for " + it->second->Reset(kResetCodeCancelled); + QUIC_DLOG(INFO) << ENDPOINT << "Revoked PUBLISH_NAMESPACE message for " << track_namespace; - publish_namespace_by_id_.erase(it->second); - publish_namespace_by_namespace_.erase(it); return true; } -bool MoqtSession::PublishNamespaceCancel(const TrackNamespace& track_namespace, - RequestErrorCode code, - absl::string_view reason) { - auto it = incoming_publish_namespaces_by_namespace_.find(track_namespace); - if (it == incoming_publish_namespaces_by_namespace_.end()) { - return false; // Could have been destroyed by PUBLISH_NAMESPACE_DONE. +bool MoqtSession::PublishNamespaceCancel( + const TrackNamespace& track_namespace, + webtransport::StreamErrorCode error_code) { + auto it = publish_namespace_responses_.find(track_namespace); + if (it == publish_namespace_responses_.end()) { + QUICHE_BUG(quic_bug_publish_namespace_update_after_closure) + << "Tried to reset PUBLISH_NAMESPACE for unknown namespace " + << track_namespace; + return false; } - MoqtPublishNamespaceCancel message{it->second, code, std::string(reason)}; - incoming_publish_namespaces_by_id_.erase(it->second); - incoming_publish_namespaces_by_namespace_.erase(it); - SendControlMessage(framer_.SerializePublishNamespaceCancel(message)); - QUIC_DLOG(INFO) << ENDPOINT << "Sent PUBLISH_NAMESPACE_CANCEL message for " - << track_namespace << " with reason " << reason; + it->second->Reset(error_code); + QUIC_DLOG(INFO) << ENDPOINT << "Signalled disinterest in PUBLISH_NAMESPACE " + << " for " << track_namespace; return true; } @@ -903,8 +886,7 @@ } // 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_publish_namespaces_by_id_.contains(request_id)) { + if (incoming_fetches_.contains(request_id)) { QUICHE_DLOG(INFO) << ENDPOINT << "Duplicate request ID"; Error(MoqtError::kInvalidRequestId, "Duplicate request ID"); return false; @@ -999,6 +981,54 @@ } break; } + case MoqtMessageType::kPublishNamespace: { + auto publish_namespace_stream = + std::make_unique<MoqtPublishNamespaceResponseStream>( + &session_->framer_, session_->ControlMessageParser(), + [weakptr = session_->GetWeakPtr()](const TrackNamespace& prefix, + MoqtBidiStreamBase* stream) { + MoqtSession* session = MoqtSessionFromWeakPtr(weakptr); + if (session == nullptr) { + return false; + } + auto [it, success] = + session->publish_namespace_responses_.try_emplace(prefix, + stream); + return success; + }, + [weakptr = session_->GetWeakPtr()](const TrackNamespace& prefix) { + MoqtSession* session = MoqtSessionFromWeakPtr(weakptr); + if (session != nullptr) { + session->publish_namespace_responses_.erase(prefix); + } + }, + [weakptr = session_->GetWeakPtr()](MoqtError code, + absl::string_view reason) { + MoqtSession* session = MoqtSessionFromWeakPtr(weakptr); + if (session != nullptr) { + session->Error(code, reason); + } + }, + [weakptr = session_->GetWeakPtr()]( + const TrackNamespace& track_namespace, + const MessageParameters* absl_nullable parameters, + MoqtResponseCallback callback) { + MoqtSession* session = MoqtSessionFromWeakPtr(weakptr); + if (session == nullptr) { + return; + } + session->callbacks_.incoming_publish_namespace_callback( + track_namespace, parameters, std::move(callback)); + }); + publish_namespace_stream->BindStream(std::move(parser_)); + MoqtPublishNamespaceResponseStream* temp_stream = + publish_namespace_stream.get(); + stream_->SetVisitor(std::move(publish_namespace_stream)); + // The UnknownBidiStream object is deleted; no class access after this + // point. + temp_stream->OnCanRead(); + break; + } case MoqtMessageType::kPublish: { auto publish_stream = std::make_unique<MoqtPublishResponseStream>( &session_->framer_, session_->ControlMessageParser(), @@ -1171,17 +1201,7 @@ 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); - if (pn_it != publish_namespace_by_id_.end()) { - if (pn_it->second.response_callback == nullptr) { - return absl::InvalidArgumentError( - "Multiple responses for PUBLISH_NAMESPACE"); - } - std::move(pn_it->second.response_callback)(MessageParameters()); - return absl::OkStatus(); - } - // Response to SUBSCRIBE_NAMESPACE is handled in the NamespaceStream. + // Response to PUBLISH/SUBSCRIBE_NAMESPACE is handled in the bidi stream.. // TRACK_STATUS response would go here, but we don't support upstream // TRACK_STATUS. // If it doesn't match any state, it might be because the local application @@ -1215,19 +1235,7 @@ } 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()) { - if (pn_it->second.response_callback == nullptr) { - return absl::InvalidArgumentError( - "Multiple responses for PUBLISH_NAMESPACE"); - } - std::move(pn_it->second.response_callback)(error_info); - publish_namespace_by_namespace_.erase(pn_it->second.track_namespace); - publish_namespace_by_id_.erase(pn_it); - return absl::OkStatus(); - } - // Response to SUBSCRIBE_NAMESPACE is handled in the NamespaceStream. + // Response to PUBLISH/SUBSCRIBE_NAMESPACE is handled in the bidi stream. // TRACK_STATUS response would go here, but we don't support upstream // TRACK_STATUS. // If it doesn't match any state, it might be because the local application @@ -1236,40 +1244,6 @@ } absl::Status MoqtSession::OnControlMessage(const MoqtRequestUpdate& message) { - 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. - quiche::QuicheWeakPtr<MoqtSessionInterface> session_weakptr = GetWeakPtr(); - TrackNamespace track_namespace = pn_it->second.track_namespace; - callbacks().incoming_publish_namespace_callback( - track_namespace, message.parameters, - [&](std::variant<MessageParameters, MoqtRequestErrorInfo> response) { - MoqtSession* session = - absl::down_cast<MoqtSession*>(session_weakptr.GetIfAvailable()); - if (session == nullptr) { - return; - } - std::visit( - absl::Overload{ - [this, request_id = message.request_id]( - const MessageParameters& parameters) { - // In draft-18, there are no useful parameters in - // PUBLISH_NAMESPACE_OK, but Issue #1639 would change that. - SendControlMessage(framer_.SerializeRequestOk(MoqtRequestOk{ - .request_id = request_id, .parameters = parameters})); - }, - [this, id = message.request_id, track_ns = track_namespace]( - const MoqtRequestErrorInfo& error_info) { - SendRequestErrorOnControlStream(id, error_info.error_code, - error_info.retry_interval, - error_info.reason_phrase); - incoming_publish_namespaces_by_id_.erase(id); - incoming_publish_namespaces_by_namespace_.erase(track_ns); - }}, - response); - }); - return absl::OkStatus(); - } // TODO(martinduke): Check all the request types. // Does not match any known request. SendRequestErrorOnControlStream(message.request_id, @@ -1278,88 +1252,6 @@ return absl::OkStatus(); } -absl::Status MoqtSession::OnControlMessage( - const MoqtPublishNamespace& message) { - if (!ValidateRequestId(message.request_id)) { - return absl::OkStatus(); - } - if (sent_goaway_) { - QUIC_DLOG(INFO) << ENDPOINT << "Received a PUBLISH_NAMESPACE after GOAWAY"; - SendRequestErrorOnControlStream( - message.request_id, RequestErrorCode::kUnauthorized, std::nullopt, - "PUBLISH_NAMESPACE after GOAWAY"); - return absl::OkStatus(); - } - QUIC_DLOG(INFO) << ENDPOINT << "Received a PUBLISH_NAMESPACE for " - << message.track_namespace; - auto [it, inserted] = incoming_publish_namespaces_by_namespace_.emplace( - message.track_namespace, message.request_id); - if (!inserted) { - SendRequestErrorOnControlStream( - message.request_id, RequestErrorCode::kDuplicateSubscription, - std::nullopt, "Duplicate PUBLISH_NAMESPACE"); - return absl::OkStatus(); - } - quiche::QuicheWeakPtr<MoqtSessionInterface> session_weakptr = GetWeakPtr(); - incoming_publish_namespaces_by_id_[message.request_id] = - message.track_namespace; - callbacks_.incoming_publish_namespace_callback( - message.track_namespace, message.parameters, - [&](std::variant<MessageParameters, MoqtRequestErrorInfo> response) { - MoqtSession* session = - absl::down_cast<MoqtSession*>(session_weakptr.GetIfAvailable()); - if (session == nullptr) { - return; - } - std::visit( - absl::Overload{ - [this, request_id = message.request_id]( - const MessageParameters& parameters) { - // In draft-18, there are no useful parameters in - // PUBLISH_NAMESPACE_OK, but Issue #1639 would change that. - SendControlMessage(framer_.SerializeRequestOk(MoqtRequestOk{ - .request_id = request_id, .parameters = parameters})); - }, - [this, id = message.request_id, - track_ns = message.track_namespace]( - const MoqtRequestErrorInfo& error_info) { - SendRequestErrorOnControlStream(id, error_info.error_code, - error_info.retry_interval, - error_info.reason_phrase); - incoming_publish_namespaces_by_id_.erase(id); - incoming_publish_namespaces_by_namespace_.erase(track_ns); - }}, - response); - }); - return absl::OkStatus(); -} - -absl::Status MoqtSession::OnControlMessage( - const MoqtPublishNamespaceDone& message) { - auto it = incoming_publish_namespaces_by_id_.find(message.request_id); - if (it == incoming_publish_namespaces_by_id_.end()) { - return absl::OkStatus(); - } - callbacks_.incoming_publish_namespace_callback(it->second, std::nullopt, - nullptr); - incoming_publish_namespaces_by_namespace_.erase(it->second); - incoming_publish_namespaces_by_id_.erase(it); - return absl::OkStatus(); -} - -absl::Status MoqtSession::OnControlMessage( - const MoqtPublishNamespaceCancel& message) { - auto it = publish_namespace_by_id_.find(message.request_id); - if (it == publish_namespace_by_id_.end()) { - return absl::OkStatus(); // State might have been destroyed due to - // PUBLISH_NAMESPACE_DONE. - } - std::move(it->second.cancel_callback)(MoqtRequestErrorInfo{ - message.error_code, std::nullopt, std::string(message.error_reason)}); - publish_namespace_by_namespace_.erase(it->second.track_namespace); - publish_namespace_by_id_.erase(it); - return absl::OkStatus(); -} absl::Status MoqtSession::OnControlMessage(const MoqtGoAway& message) { if (!message.new_session_uri.empty() && @@ -1604,17 +1496,22 @@ if (goaway_timeout_alarm_ != nullptr) { goaway_timeout_alarm_->PermanentCancel(); } - // 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. - for (auto& it : incoming_publish_namespaces_by_namespace_) { - callbacks_.incoming_publish_namespace_callback(it.first, std::nullopt, - nullptr); + // Although PUBLISH_NAMESPACE state will be cleaned up/ by the owning stream, + // the session can be destroyed first. In this case, the application callbacks + // will be inaccessible. Instead, invoke application callbacks now. + while (!publish_namespace_responses_.empty()) { + auto it = publish_namespace_responses_.begin(); + MoqtBidiStreamBase* stream = it->second; + publish_namespace_responses_.erase(it); + stream->Detach(); } - for (auto& it : publish_namespace_by_id_) { - std::move(it.second.cancel_callback)(MoqtRequestErrorInfo{ - RequestErrorCode::kUninterested, std::nullopt, "Session closed"}); + while (!publish_namespace_requests_.empty()) { + auto it = publish_namespace_requests_.begin(); + MoqtBidiStreamBase* stream = it->second; + publish_namespace_requests_.erase(it); + stream->Detach(); } + // WebTransport session, the incoming FETCHes are owned by this class. while (!fetch_by_id_.empty()) { fetch_by_id_.begin()->second->Destroy(); }
diff --git a/quiche/quic/moqt/moqt_session.h b/quiche/quic/moqt/moqt_session.h index 7f88f09..7e4904f 100644 --- a/quiche/quic/moqt/moqt_session.h +++ b/quiche/quic/moqt/moqt_session.h
@@ -115,18 +115,18 @@ FetchResponseCallback callback, uint64_t num_previous_groups, MessageParameters parameters) override; - bool PublishNamespace(const TrackNamespace& track_namespace, - const MessageParameters& parameters, - MoqtResponseCallback response_callback, - quiche::SingleUseCallback<void(MoqtRequestErrorInfo)> - cancel_callback) override; + bool PublishNamespace( + const TrackNamespace& track_namespace, + const MessageParameters& parameters, + MoqtResponseCallback response_callback, + quiche::SingleUseCallback<void()> cancel_callback) override; bool PublishNamespaceUpdate(const TrackNamespace& track_namespace, MessageParameters& parameters, MoqtResponseCallback response_callback) override; bool PublishNamespaceDone(const TrackNamespace& track_namespace) override; - bool PublishNamespaceCancel(const TrackNamespace& track_namespace, - RequestErrorCode error_code, - absl::string_view error_reason) override; + bool PublishNamespaceCancel( + const TrackNamespace& track_namespace, + webtransport::StreamErrorCode error_code) override; // TODO(martinduke): Support PUBLISH. For now, PUBLISH-only requests will be // rejected with nullptr, and kBoth requests will change to kNamespace. // After receiving MoqtNamespaceTask, call @@ -386,9 +386,6 @@ absl::Status OnControlMessage(const MoqtRequestOk& message); absl::Status OnControlMessage(const MoqtRequestError& message); absl::Status OnControlMessage(const MoqtRequestUpdate& message); - absl::Status OnControlMessage(const MoqtPublishNamespace& message); - absl::Status OnControlMessage(const MoqtPublishNamespaceDone& /*message*/); - absl::Status OnControlMessage(const MoqtPublishNamespaceCancel& message); absl::Status OnControlMessage(const MoqtGoAway& /*message*/); absl::Status OnControlMessage(const MoqtMaxRequestId& message); absl::Status OnControlMessage(const MoqtFetch& message); @@ -474,19 +471,10 @@ monitoring_interfaces_for_published_tracks_; // PUBLISH_NAMESPACE state. - struct PublishNamespaceState { - TrackNamespace track_namespace; - MoqtResponseCallback response_callback; - quiche::SingleUseCallback<void(MoqtRequestErrorInfo)> cancel_callback; - }; - absl::flat_hash_map<uint64_t, PublishNamespaceState> publish_namespace_by_id_; - absl::flat_hash_map<TrackNamespace, uint64_t> publish_namespace_by_namespace_; - absl::flat_hash_map<uint64_t, MoqtResponseCallback> - publish_namespace_updates_; - absl::flat_hash_map<TrackNamespace, uint64_t> - incoming_publish_namespaces_by_namespace_; - absl::flat_hash_map<uint64_t, TrackNamespace> - incoming_publish_namespaces_by_id_; + absl::flat_hash_map<TrackNamespace, MoqtBidiStreamBase*> + publish_namespace_requests_; + absl::flat_hash_map<TrackNamespace, MoqtBidiStreamBase*> + publish_namespace_responses_; // It's an error if the namespaces overlap, so keep track of them. SessionNamespaceTree incoming_subscribe_namespace_;
diff --git a/quiche/quic/moqt/moqt_session_callbacks.h b/quiche/quic/moqt/moqt_session_callbacks.h index 224ffac..c8b8356 100644 --- a/quiche/quic/moqt/moqt_session_callbacks.h +++ b/quiche/quic/moqt/moqt_session_callbacks.h
@@ -11,6 +11,7 @@ #include <utility> #include <variant> +#include "absl/base/nullability.h" #include "absl/strings/string_view.h" #include "quiche/quic/core/quic_clock.h" #include "quiche/quic/core/quic_default_clock.h" @@ -88,14 +89,14 @@ MoqtResponseCallback)>; // Called whenever a PUBLISH_NAMESPACE or PUBLISH_NAMESPACE_DONE message is -// received from the peer. PUBLISH_NAMESPACE sets a value for |parameters|, -// PUBLISH_NAMESPACE_DONE does not. This callback is not invoked by NAMESPACE or +// received from the peer. PUBLISH_NAMESPACE sets a pointer |parameters|, +// closing it does not. This callback is not invoked by NAMESPACE or // NAMESPACE_DONE messages that arrive on a SUBSCRIBE_NAMESPACE stream. // If the PUBLISH_NAMESPACE is updated, it will be called again, so be prepared // for duplicates. using MoqtIncomingPublishNamespaceCallback = quiche::MultiUseCallback<void( const TrackNamespace& track_namespace, - const std::optional<MessageParameters>& parameters, + const MessageParameters* absl_nullable parameters, MoqtResponseCallback callback)>; // Called whenever SUBSCRIBE_NAMESPACE is received from the peer. Unsubscribe @@ -111,7 +112,7 @@ MoqtResponseCallback response_callback)>; inline void DefaultIncomingPublishNamespaceCallback( - const TrackNamespace&, const std::optional<MessageParameters>&, + const TrackNamespace&, const MessageParameters* absl_nullable, MoqtResponseCallback callback) { if (callback == nullptr) { return;
diff --git a/quiche/quic/moqt/moqt_session_interface.h b/quiche/quic/moqt/moqt_session_interface.h index 1af02d8..c2f41f5 100644 --- a/quiche/quic/moqt/moqt_session_interface.h +++ b/quiche/quic/moqt/moqt_session_interface.h
@@ -25,6 +25,7 @@ #include "quiche/common/platform/api/quiche_export.h" #include "quiche/common/quiche_callbacks.h" #include "quiche/common/quiche_weak_ptr.h" +#include "quiche/web_transport/web_transport.h" namespace moqt { @@ -160,23 +161,22 @@ // Send a PUBLISH_NAMESPACE message for |track_namespace|, and call // |response_callback| when the response arrives. Will fail // immediately if there is already an unresolved PUBLISH_NAMESPACE for that - // namespace. Calls |cancel_callback| if the peer sends a - // PUBLISH_NAMESPACE_CANCEL. Returns true if the message was sent. + // namespace. Calls |cancel_callback| if the peer closes the stream. Returns + // true if the message was sent. virtual bool PublishNamespace( const TrackNamespace& track_namespace, const MessageParameters& parameters, MoqtResponseCallback response_callback, - quiche::SingleUseCallback<void(MoqtRequestErrorInfo)> - cancel_callback) = 0; + quiche::SingleUseCallback<void()> cancel_callback) = 0; virtual bool PublishNamespaceUpdate( const TrackNamespace& track_namespace, MessageParameters& parameters, MoqtResponseCallback response_callback) = 0; // Returns true if message was sent, false if there is no PUBLISH_NAMESPACE // that relates. virtual bool PublishNamespaceDone(const TrackNamespace& track_namespace) = 0; - virtual bool PublishNamespaceCancel(const TrackNamespace& track_namespace, - RequestErrorCode error_code, - absl::string_view error_reason) = 0; + virtual bool PublishNamespaceCancel( + const TrackNamespace& track_namespace, + webtransport::StreamErrorCode error_code) = 0; // Sends a SUBSCRIBE_NAMESPACE message for |prefix| and returns a // MoqtNamespaceTask that can be used to process the response.
diff --git a/quiche/quic/moqt/moqt_session_test.cc b/quiche/quic/moqt/moqt_session_test.cc index 436fc68..bd81199 100644 --- a/quiche/quic/moqt/moqt_session_test.cc +++ b/quiche/quic/moqt/moqt_session_test.cc
@@ -37,9 +37,11 @@ #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" #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" +#include "quiche/quic/platform/api/quic_expect_bug.h" #include "quiche/quic/platform/api/quic_test.h" #include "quiche/quic/test_tools/quic_test_utils.h" #include "quiche/common/quiche_buffer_allocator.h" @@ -171,6 +173,7 @@ static constexpr absl::string_view kSubscribeNamespaceByte = "\x50"; static constexpr absl::string_view kPublishByte = "\x1d"; static constexpr absl::string_view kTrackStatusByte = "\x0d"; + static constexpr absl::string_view kPublishNamespaceByte = "\x06"; std::unique_ptr<MoqtBidiStreamBase> ResponseStream( absl::string_view first_byte, webtransport::test::MockStream* wt_stream = nullptr) { @@ -216,8 +219,8 @@ 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)); + ON_CALL(mock_session_, CanOpenNextOutgoingBidirectionalStream()) + .WillByDefault(Return(true)); EXPECT_CALL(mock_session_, OpenOutgoingBidirectionalStream()) .WillOnce(Return(stream)); EXPECT_CALL(*stream, SetVisitor) @@ -324,8 +327,8 @@ MoqtKnownTrackPublisher publisher_; webtransport::test::MockSession mock_session_; MockSessionCallbacks session_callbacks_; - std::unique_ptr<MoqtBidiStreamTestWrapper> bidi_wrapper_; MoqtSession session_; + std::unique_ptr<MoqtBidiStreamTestWrapper> bidi_wrapper_; webtransport::test::MockStream mock_bidi_stream_, mock_uni_stream_; // std::shared_ptr<IncomingSubscribeInfo> last_incoming_subscribe_; }; @@ -511,16 +514,14 @@ testing::MockFunction<void( std::variant<MessageParameters, MoqtRequestErrorInfo> error_message)> publish_namespace_response_callback; - std::unique_ptr<MoqtBidiStreamTestWrapper> bidi_wrapper_ = - MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_); + PrepareRequestStream(bidi_wrapper_); EXPECT_CALL( mock_bidi_stream_, Writev(ControlMessageOfType(MoqtMessageType::kPublishNamespace), _)); - MoqtRequestErrorInfo cancel_error_info; - session_.PublishNamespace( - TrackNamespace({"foo"}), MessageParameters(), - publish_namespace_response_callback.AsStdFunction(), - [&](MoqtRequestErrorInfo info) { cancel_error_info = info; }); + bool cancel_called = false; + session_.PublishNamespace(TrackNamespace({"foo"}), MessageParameters(), + publish_namespace_response_callback.AsStdFunction(), + [&]() { cancel_called = true; }); MoqtRequestOk ok = {/*request_id=*/0, MessageParameters()}; EXPECT_CALL(publish_namespace_response_callback, Call) @@ -530,30 +531,25 @@ }); bidi_wrapper_->ReceiveMessage(ok); - MoqtPublishNamespaceCancel cancel = { - /*request_id=*/0, - RequestErrorCode::kInternalError, - /*error_reason=*/"Test error", - }; - bidi_wrapper_->ReceiveMessage(cancel); - EXPECT_EQ(cancel_error_info.error_code, RequestErrorCode::kInternalError); - EXPECT_EQ(cancel_error_info.reason_phrase, "Test error"); + bidi_wrapper_->stream().OnResetStreamReceived( + webtransport::StreamErrorCode(kResetCodeInternalError)); + EXPECT_TRUE(cancel_called); // State is gone. - EXPECT_FALSE(session_.PublishNamespaceDone(TrackNamespace({"foo"}))); + EXPECT_QUIC_BUG(session_.PublishNamespaceDone(TrackNamespace({"foo"})), + "Tried to reset PUBLISH_NAMESPACE for unknown namespace foo"); } TEST_F(MoqtSessionTest, PublishNamespaceWithOkAndPublishNamespaceDone) { testing::MockFunction<void( std::variant<MessageParameters, MoqtRequestErrorInfo>)> publish_namespace_resolved_callback; - std::unique_ptr<MoqtBidiStreamTestWrapper> bidi_wrapper_ = - MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_); + PrepareRequestStream(bidi_wrapper_); EXPECT_CALL( mock_bidi_stream_, Writev(ControlMessageOfType(MoqtMessageType::kPublishNamespace), _)); session_.PublishNamespace(TrackNamespace{"foo"}, MessageParameters(), publish_namespace_resolved_callback.AsStdFunction(), - [](MoqtRequestErrorInfo) {}); + []() {}); MoqtRequestOk ok = {/*request_id=*/0, MessageParameters()}; EXPECT_CALL(publish_namespace_resolved_callback, Call) @@ -563,26 +559,24 @@ }); bidi_wrapper_->ReceiveMessage(ok); - EXPECT_CALL( - mock_bidi_stream_, - Writev(ControlMessageOfType(MoqtMessageType::kPublishNamespaceDone), _)); + EXPECT_CALL(mock_bidi_stream_, ResetWithUserCode); session_.PublishNamespaceDone(TrackNamespace{"foo"}); // State is gone. - EXPECT_FALSE(session_.PublishNamespaceDone(TrackNamespace{"foo"})); + EXPECT_QUIC_BUG(session_.PublishNamespaceDone(TrackNamespace({"foo"})), + "Tried to reset PUBLISH_NAMESPACE for unknown namespace foo"); } TEST_F(MoqtSessionTest, PublishNamespaceWithError) { testing::MockFunction<void( std::variant<MessageParameters, MoqtRequestErrorInfo>)> publish_namespace_resolved_callback; - std::unique_ptr<MoqtBidiStreamTestWrapper> bidi_wrapper_ = - MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_); + PrepareRequestStream(bidi_wrapper_); EXPECT_CALL( mock_bidi_stream_, Writev(ControlMessageOfType(MoqtMessageType::kPublishNamespace), _)); session_.PublishNamespace(TrackNamespace{"foo"}, MessageParameters(), publish_namespace_resolved_callback.AsStdFunction(), - [](MoqtRequestErrorInfo) {}); + []() {}); MoqtRequestError error{/*request_id=*/0, RequestErrorCode::kInternalError, std::nullopt, "Test error"}; @@ -595,9 +589,11 @@ EXPECT_EQ(error.error_code, RequestErrorCode::kInternalError); EXPECT_EQ(error.reason_phrase, "Test error"); }); + ExpectFin(mock_bidi_stream_); bidi_wrapper_->ReceiveMessage(error); // State is gone. - EXPECT_FALSE(session_.PublishNamespaceDone(TrackNamespace{"foo"})); + EXPECT_QUIC_BUG(session_.PublishNamespaceDone(TrackNamespace({"foo"})), + "Tried to reset PUBLISH_NAMESPACE for unknown namespace foo"); } TEST_F(MoqtSessionTest, AsynchronousSubscribeReturnsOk) { @@ -1031,16 +1027,18 @@ MessageParameters parameters; parameters.authorization_tokens.emplace_back(AuthTokenType::kOutOfBand, "foo"); + bidi_wrapper_ = std::make_unique<MoqtBidiStreamTestWrapper>( + ResponseStream(kPublishNamespaceByte)); MoqtPublishNamespace publish_namespace = { kDefaultPeerRequestId, track_namespace, parameters, }; EXPECT_CALL(session_callbacks_.incoming_publish_namespace_callback, - Call(track_namespace, std::make_optional(parameters), _)) - .WillOnce([](const TrackNamespace&, - const std::optional<MessageParameters>&, - MoqtResponseCallback callback) { + Call(track_namespace, _, _)) + .WillOnce([&](const TrackNamespace&, const MessageParameters* params, + MoqtResponseCallback callback) { + EXPECT_TRUE(params != nullptr && *params == parameters); std::move(callback)(MessageParameters()); }); EXPECT_CALL(mock_bidi_stream_, @@ -1048,23 +1046,16 @@ kDefaultPeerRequestId, MessageParameters()}), _)); bidi_wrapper_->ReceiveMessage(publish_namespace); - MoqtPublishNamespaceDone publish_namespace_done = { - /*request_id=*/0, - }; EXPECT_CALL(session_callbacks_.incoming_publish_namespace_callback, - Call(track_namespace, std::optional<MessageParameters>(), _)) - .WillOnce( - [](const TrackNamespace&, const std::optional<MessageParameters>&, - MoqtResponseCallback callback) { EXPECT_EQ(callback, nullptr); }); - bidi_wrapper_->ReceiveMessage(publish_namespace_done); + Call(track_namespace, nullptr, nullptr)); + bidi_wrapper_->stream().OnResetStreamReceived(kResetCodeCancelled); } TEST_F(MoqtSessionTest, ReplyToPublishNamespaceWithOkThenPublishNamespaceCancel) { TrackNamespace track_namespace{"foo"}; - - bidi_wrapper_ = - MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_); + bidi_wrapper_ = std::make_unique<MoqtBidiStreamTestWrapper>( + ResponseStream(kPublishNamespaceByte)); MessageParameters parameters; parameters.authorization_tokens.emplace_back(AuthTokenType::kOutOfBand, "foo"); @@ -1074,10 +1065,10 @@ parameters, }; EXPECT_CALL(session_callbacks_.incoming_publish_namespace_callback, - Call(track_namespace, std::make_optional(parameters), _)) - .WillOnce([](const TrackNamespace&, - const std::optional<MessageParameters>&, - MoqtResponseCallback callback) { + Call(track_namespace, _, _)) + .WillOnce([&](const TrackNamespace&, const MessageParameters* params, + MoqtResponseCallback callback) { + EXPECT_TRUE(params != nullptr && *params == parameters); std::move(callback)(MessageParameters()); }); EXPECT_CALL(mock_bidi_stream_, @@ -1085,20 +1076,19 @@ kDefaultPeerRequestId, MessageParameters()}), _)); bidi_wrapper_->ReceiveMessage(publish_namespace); - EXPECT_CALL(mock_bidi_stream_, - Writev(SerializedControlMessage(MoqtPublishNamespaceCancel{ - kDefaultPeerRequestId, - RequestErrorCode::kInternalError, "deadbeef"}), - _)); - session_.PublishNamespaceCancel(track_namespace, - RequestErrorCode::kInternalError, "deadbeef"); + EXPECT_CALL(session_callbacks_.incoming_publish_namespace_callback, + Call(track_namespace, nullptr, nullptr)); + EXPECT_TRUE( + session_.PublishNamespaceCancel(track_namespace, kResetCodeCancelled)); + // State is gone. + EXPECT_QUIC_BUG(session_.PublishNamespaceDone(TrackNamespace({"foo"})), + "Tried to reset PUBLISH_NAMESPACE for unknown namespace foo"); } TEST_F(MoqtSessionTest, ReplyToPublishNamespaceWithError) { TrackNamespace track_namespace{"foo"}; - - bidi_wrapper_ = - MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_); + bidi_wrapper_ = std::make_unique<MoqtBidiStreamTestWrapper>( + ResponseStream(kPublishNamespaceByte)); MessageParameters parameters; parameters.authorization_tokens.emplace_back(AuthTokenType::kOutOfBand, "foo"); @@ -1113,10 +1103,17 @@ "deadbeef", }; EXPECT_CALL(session_callbacks_.incoming_publish_namespace_callback, - Call(track_namespace, std::make_optional(parameters), _)) - .WillOnce( - [&](const TrackNamespace&, const std::optional<MessageParameters>&, - MoqtResponseCallback callback) { std::move(callback)(error); }); + Call(track_namespace, _, _)) + .WillOnce([&](const TrackNamespace&, const MessageParameters* params, + MoqtResponseCallback callback) { + EXPECT_TRUE(params != nullptr && *params == parameters); + std::move(callback)(error); + }) + .WillOnce([&](const TrackNamespace&, const MessageParameters* params, + MoqtResponseCallback callback) { // Teardown + EXPECT_TRUE(params == nullptr); + EXPECT_TRUE(callback == nullptr); + }); EXPECT_CALL(mock_bidi_stream_, Writev(SerializedControlMessage(MoqtRequestError{ kDefaultPeerRequestId, error.error_code, @@ -2310,10 +2307,9 @@ prefix, MessageParameters(), +[](std::variant<MessageParameters, MoqtRequestErrorInfo>) {}), nullptr); - session_.PublishNamespace( + EXPECT_FALSE(session_.PublishNamespace( TrackNamespace{"foo"}, MessageParameters(), - +[](std::variant<MessageParameters, MoqtRequestErrorInfo>) {}, - +[](MoqtRequestErrorInfo) {}); + +[](std::variant<MessageParameters, MoqtRequestErrorInfo>) {}, +[]() {})); EXPECT_FALSE(session_.Fetch( FullTrackName{TrackNamespace({"foo"}), "bar"}, +[](std::unique_ptr<MoqtFetchTask>) {}, Location(0, 0), 5, std::nullopt, @@ -2342,10 +2338,6 @@ EXPECT_CALL(mock_bidi_stream_, Writev(ControlMessageOfType(MoqtMessageType::kRequestError), _)); - bidi_wrapper_->ReceiveMessage( - MoqtPublishNamespace(3, TrackNamespace({"foo"}), MessageParameters())); - EXPECT_CALL(mock_bidi_stream_, - Writev(ControlMessageOfType(MoqtMessageType::kRequestError), _)); MoqtFetch fetch = DefaultFetch(); fetch.request_id = 5; bidi_wrapper_->ReceiveMessage(fetch); @@ -2384,10 +2376,9 @@ prefix, MessageParameters(), +[](std::variant<MessageParameters, MoqtRequestErrorInfo>) {}), nullptr); - session_.PublishNamespace( + EXPECT_FALSE(session_.PublishNamespace( TrackNamespace{"foo"}, MessageParameters(), - +[](std::variant<MessageParameters, MoqtRequestErrorInfo>) {}, - +[](MoqtRequestErrorInfo) {}); + +[](std::variant<MessageParameters, MoqtRequestErrorInfo>) {}, +[]() {})); EXPECT_FALSE(session_.Fetch( FullTrackName(TrackNamespace({"foo"}), "bar"), +[](std::unique_ptr<MoqtFetchTask>) {}, Location(0, 0), 5, std::nullopt, @@ -2612,8 +2603,8 @@ } TEST_F(MoqtSessionTest, IncomingPublishNamespaceCleanup) { - bidi_wrapper_ = - MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_); + bidi_wrapper_ = std::make_unique<MoqtBidiStreamTestWrapper>( + ResponseStream(kPublishNamespaceByte)); // Register two incoming PUBLISH_NAMESPACE. MoqtPublishNamespace publish_namespace{ /*request_id=*/1, TrackNamespace{"foo"}, MessageParameters()}; @@ -2621,8 +2612,7 @@ expected_ok.parameters.expires = quic::QuicTimeDelta::FromSeconds(60); EXPECT_CALL(session_callbacks_.incoming_publish_namespace_callback, Call(TrackNamespace{"foo"}, _, _)) - .WillOnce([&](const TrackNamespace&, - const std::optional<MessageParameters>&, + .WillOnce([&](const TrackNamespace&, const MessageParameters*, MoqtResponseCallback callback) { std::move(callback)(expected_ok.parameters); }); @@ -2630,36 +2620,28 @@ Writev(SerializedControlMessage(expected_ok), _)); bidi_wrapper_->ReceiveMessage(publish_namespace); + auto bidi_wrapper_2 = std::make_unique<MoqtBidiStreamTestWrapper>( + ResponseStream(kPublishNamespaceByte)); publish_namespace = MoqtPublishNamespace( /*request_id=*/3, TrackNamespace{"bar"}, MessageParameters()); EXPECT_CALL(session_callbacks_.incoming_publish_namespace_callback, Call(TrackNamespace{"bar"}, _, _)) - .WillOnce([&](const TrackNamespace&, - const std::optional<MessageParameters>&, + .WillOnce([&](const TrackNamespace&, const MessageParameters*, MoqtResponseCallback callback) { std::move(callback)(MessageParameters()); }); EXPECT_CALL(mock_bidi_stream_, Writev(ControlMessageOfType(MoqtMessageType::kRequestOk), _)); - bidi_wrapper_->ReceiveMessage(publish_namespace); + bidi_wrapper_2->ReceiveMessage(publish_namespace); // Revoke "bar" - MoqtPublishNamespaceDone done{/*request_id=*/3}; - EXPECT_CALL( - session_callbacks_.incoming_publish_namespace_callback, - Call(TrackNamespace{"bar"}, std::optional<MessageParameters>(), _)) - .WillOnce( - [](const TrackNamespace&, const std::optional<MessageParameters>&, - MoqtResponseCallback callback) { EXPECT_EQ(callback, nullptr); }); - bidi_wrapper_->ReceiveMessage(done); + EXPECT_CALL(session_callbacks_.incoming_publish_namespace_callback, + Call(TrackNamespace{"bar"}, nullptr, nullptr)); + bidi_wrapper_2->stream().OnResetStreamReceived(kResetCodeCancelled); // Destroying the session should revoke "foo". - EXPECT_CALL( - session_callbacks_.incoming_publish_namespace_callback, - Call(TrackNamespace{"foo"}, std::optional<MessageParameters>(), _)) - .WillOnce( - [](const TrackNamespace&, const std::optional<MessageParameters>&, - MoqtResponseCallback callback) { EXPECT_EQ(callback, nullptr); }); + EXPECT_CALL(session_callbacks_.incoming_publish_namespace_callback, + Call(TrackNamespace{"foo"}, nullptr, nullptr)); // Test teardown will destroy session_, triggering removal of "foo". } @@ -2911,15 +2893,17 @@ next_request_id += 2; // 2. PublishNamespace + webtransport::test::InMemoryStreamWithWriteBuffer pub_ns_stream(5); + EXPECT_CALL(mock_session_, OpenOutgoingBidirectionalStream) + .WillOnce(Return(&pub_ns_stream)); TrackNamespace namespace2({"namespace2"}); bool p1 = session_.PublishNamespace( namespace2, MessageParameters(), - [](std::variant<MessageParameters, MoqtRequestErrorInfo>) {}, - [](MoqtRequestErrorInfo) {}); + [](std::variant<MessageParameters, MoqtRequestErrorInfo>) {}, []() {}); EXPECT_TRUE(p1); - EXPECT_EQ(next_request_id, get_request_id(control_stream)); + EXPECT_EQ(next_request_id, get_request_id(pub_ns_stream)); next_request_id += 2; - control_stream.write_buffer().clear(); + pub_ns_stream.write_buffer().clear(); // 3. PublishNamespaceUpdate MessageParameters params_update; @@ -2927,9 +2911,9 @@ namespace2, params_update, [](std::variant<MessageParameters, MoqtRequestErrorInfo>) {}); EXPECT_TRUE(p_update); - EXPECT_EQ(get_request_id(control_stream), next_request_id); + EXPECT_EQ(get_request_id(pub_ns_stream), next_request_id); next_request_id += 2; - control_stream.write_buffer().clear(); + pub_ns_stream.write_buffer().clear(); // 4. Subscribe webtransport::test::InMemoryStreamWithWriteBuffer sub_stream(5);
diff --git a/quiche/quic/moqt/test_tools/mock_moqt_session.h b/quiche/quic/moqt/test_tools/mock_moqt_session.h index e6e719d..de07a25 100644 --- a/quiche/quic/moqt/test_tools/mock_moqt_session.h +++ b/quiche/quic/moqt/test_tools/mock_moqt_session.h
@@ -100,13 +100,12 @@ FetchResponseCallback callback, uint64_t num_previous_groups, MessageParameters parameters), (override)); - MOCK_METHOD( - bool, PublishNamespace, - (const TrackNamespace& track_namespace, - const MessageParameters& parameters, - MoqtResponseCallback response_callback, - quiche::SingleUseCallback<void(MoqtRequestErrorInfo)> cancel_callback), - (override)); + MOCK_METHOD(bool, PublishNamespace, + (const TrackNamespace& track_namespace, + const MessageParameters& parameters, + MoqtResponseCallback response_callback, + quiche::SingleUseCallback<void()> cancel_callback), + (override)); MOCK_METHOD(bool, PublishNamespaceUpdate, (const TrackNamespace& track_namespace, MessageParameters& parameters, @@ -116,7 +115,7 @@ (const TrackNamespace& track_namespace), (override)); MOCK_METHOD(bool, PublishNamespaceCancel, (const TrackNamespace& track_namespace, - RequestErrorCode error_code, absl::string_view error_reason), + webtransport::StreamErrorCode error_code), (override)); MOCK_METHOD(std::unique_ptr<MoqtNamespaceTask>, SubscribeNamespace, (TrackNamespace&, const MessageParameters&, MoqtResponseCallback),
diff --git a/quiche/quic/moqt/test_tools/moqt_framer_utils.cc b/quiche/quic/moqt/test_tools/moqt_framer_utils.cc index f9c3052..aa45615 100644 --- a/quiche/quic/moqt/test_tools/moqt_framer_utils.cc +++ b/quiche/quic/moqt/test_tools/moqt_framer_utils.cc
@@ -46,18 +46,12 @@ quiche::QuicheBuffer operator()(const MoqtPublishNamespace& message) { return framer.SerializePublishNamespace(message); } - quiche::QuicheBuffer operator()(const MoqtPublishNamespaceDone& message) { - return framer.SerializePublishNamespaceDone(message); - } quiche::QuicheBuffer operator()(const MoqtNamespace& message) { return framer.SerializeNamespace(message); } quiche::QuicheBuffer operator()(const MoqtNamespaceDone& message) { return framer.SerializeNamespaceDone(message); } - quiche::QuicheBuffer operator()(const MoqtPublishNamespaceCancel& message) { - return framer.SerializePublishNamespaceCancel(message); - } quiche::QuicheBuffer operator()(const MoqtTrackStatus& message) { return framer.SerializeTrackStatus(message); }
diff --git a/quiche/quic/moqt/test_tools/moqt_framer_utils.h b/quiche/quic/moqt/test_tools/moqt_framer_utils.h index abea3ce..c0d9ee3 100644 --- a/quiche/quic/moqt/test_tools/moqt_framer_utils.h +++ b/quiche/quic/moqt/test_tools/moqt_framer_utils.h
@@ -23,8 +23,7 @@ using AnyMoqtControlMessage = std::variant<MoqtSetup, MoqtRequestOk, MoqtRequestError, MoqtSubscribe, MoqtSubscribeOk, MoqtPublishDone, MoqtRequestUpdate, - MoqtPublishNamespace, MoqtPublishNamespaceDone, - MoqtPublishNamespaceCancel, MoqtTrackStatus, MoqtGoAway, + MoqtPublishNamespace, MoqtTrackStatus, MoqtGoAway, MoqtSubscribeNamespace, MoqtSubscribeTracks, MoqtMaxRequestId, MoqtFetch, MoqtFetchCancel, MoqtFetchOk, MoqtRequestsBlocked, MoqtPublish, MoqtNamespace, MoqtNamespaceDone, MoqtObjectAck>;
diff --git a/quiche/quic/moqt/test_tools/moqt_mock_visitor.h b/quiche/quic/moqt/test_tools/moqt_mock_visitor.h index 69d8269..92a9e4c 100644 --- a/quiche/quic/moqt/test_tools/moqt_mock_visitor.h +++ b/quiche/quic/moqt/test_tools/moqt_mock_visitor.h
@@ -43,7 +43,7 @@ testing::MockFunction<void(absl::string_view)> session_terminated_callback; testing::MockFunction<void()> session_deleted_callback; testing::MockFunction<void(const TrackNamespace&, - const std::optional<MessageParameters>&, + const MessageParameters* absl_nullable, MoqtResponseCallback)> incoming_publish_namespace_callback; testing::MockFunction<std::unique_ptr<MoqtNamespaceTask>(
diff --git a/quiche/quic/moqt/test_tools/moqt_test_message.h b/quiche/quic/moqt/test_tools/moqt_test_message.h index 76b5c67..732a817 100644 --- a/quiche/quic/moqt/test_tools/moqt_test_message.h +++ b/quiche/quic/moqt/test_tools/moqt_test_message.h
@@ -112,13 +112,14 @@ public: virtual ~TestMessageBase() = default; - using MessageStructuredData = std::variant< - MoqtSetup, MoqtObject, MoqtRequestOk, MoqtRequestError, MoqtSubscribe, - MoqtSubscribeOk, MoqtPublishDone, MoqtRequestUpdate, MoqtPublishNamespace, - MoqtPublishNamespaceDone, MoqtPublishNamespaceCancel, MoqtTrackStatus, - MoqtGoAway, MoqtSubscribeNamespace, MoqtSubscribeTracks, MoqtMaxRequestId, - MoqtFetch, MoqtFetchCancel, MoqtFetchOk, MoqtRequestsBlocked, MoqtPublish, - MoqtNamespace, MoqtNamespaceDone, MoqtObjectAck>; + using MessageStructuredData = + std::variant<MoqtSetup, MoqtObject, MoqtRequestOk, MoqtRequestError, + MoqtSubscribe, MoqtSubscribeOk, MoqtPublishDone, + MoqtRequestUpdate, MoqtPublishNamespace, MoqtTrackStatus, + MoqtGoAway, MoqtSubscribeNamespace, MoqtSubscribeTracks, + 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_; } @@ -1120,83 +1121,6 @@ }; }; -class QUICHE_NO_EXPORT PublishNamespaceDoneMessage : public TestMessageBase { - public: - PublishNamespaceDoneMessage() : TestMessageBase() { - SetWireImage(raw_packet_, sizeof(raw_packet_)); - } - - bool EqualFieldValues(const MessageStructuredData& values) const override { - auto cast = std::get<MoqtPublishNamespaceDone>(values); - if (cast.request_id != publish_namespace_done_.request_id) { - QUIC_LOG(INFO) << "PUBLISH_NAMESPACE_DONE request ID mismatch"; - return false; - } - return true; - } - - void ExpandVarints() override { ExpandVarintsImpl("v"); } - - MessageStructuredData structured_data() const override { - return TestMessageBase::MessageStructuredData(publish_namespace_done_); - } - - private: - uint8_t raw_packet_[4] = { - 0x09, - 0x00, - 0x01, - 0x01, // request_id = 1 - }; - - MoqtPublishNamespaceDone publish_namespace_done_ = { - /*request_id=*/1, - }; -}; - -class QUICHE_NO_EXPORT PublishNamespaceCancelMessage : public TestMessageBase { - public: - PublishNamespaceCancelMessage() : TestMessageBase() { - SetWireImage(raw_packet_, sizeof(raw_packet_)); - } - - bool EqualFieldValues(const MessageStructuredData& values) const override { - auto cast = std::get<MoqtPublishNamespaceCancel>(values); - if (cast.request_id != publish_namespace_cancel_.request_id) { - QUIC_LOG(INFO) << "PUBLISH_NAMESPACE CANCEL request ID mismatch"; - return false; - } - if (cast.error_code != publish_namespace_cancel_.error_code) { - QUIC_LOG(INFO) << "PUBLISH_NAMESPACE CANCEL error code mismatch"; - return false; - } - if (cast.error_reason != publish_namespace_cancel_.error_reason) { - QUIC_LOG(INFO) << "PUBLISH_NAMESPACE CANCEL reason phrase mismatch"; - return false; - } - return true; - } - - void ExpandVarints() override { ExpandVarintsImpl("vvv---"); } - - MessageStructuredData structured_data() const override { - return TestMessageBase::MessageStructuredData(publish_namespace_cancel_); - } - - private: - uint8_t raw_packet_[9] = { - 0x0c, 0x00, 0x06, 0x02, // request_id = 2 - 0x03, // error_code = 3 - 0x03, 0x62, 0x61, 0x72, // error_reason = "bar" - }; - - MoqtPublishNamespaceCancel publish_namespace_cancel_ = { - /*request_id=*/2, - RequestErrorCode::kNotSupported, - /*error_reason=*/"bar", - }; -}; - class QUICHE_NO_EXPORT TrackStatusMessage : public SubscribeMessage { public: TrackStatusMessage() : SubscribeMessage() { @@ -1826,14 +1750,10 @@ return std::make_unique<RequestUpdateMessage>(); case MoqtMessageType::kPublishNamespace: return std::make_unique<PublishNamespaceMessage>(); - case MoqtMessageType::kPublishNamespaceDone: - return std::make_unique<PublishNamespaceDoneMessage>(); case MoqtMessageType::kNamespace: return std::make_unique<NamespaceMessage>(); case MoqtMessageType::kNamespaceDone: return std::make_unique<NamespaceDoneMessage>(); - case MoqtMessageType::kPublishNamespaceCancel: - return std::make_unique<PublishNamespaceCancelMessage>(); case MoqtMessageType::kTrackStatus: return std::make_unique<TrackStatusMessage>(); case MoqtMessageType::kGoAway:
diff --git a/quiche/quic/moqt/tools/chat_client.cc b/quiche/quic/moqt/tools/chat_client.cc index 86b2750..969f1a6 100644 --- a/quiche/quic/moqt/tools/chat_client.cc +++ b/quiche/quic/moqt/tools/chat_client.cc
@@ -51,22 +51,22 @@ void ChatClient::OnIncomingPublishNamespace( const moqt::TrackNamespace& track_namespace, - const std::optional<MessageParameters>& parameters, + const MessageParameters* absl_nullable parameters, moqt::MoqtResponseCallback absl_nullable callback) { if (!session_is_open_) { return; } if (track_namespace == GetUserNamespace(my_track_name_)) { // Ignore PUBLISH_NAMESPACE for my own track. - if (parameters.has_value() && callback != nullptr) { // callback exists. + if (parameters != nullptr && callback != nullptr) { // callback exists. std::move(callback)(MessageParameters()); } return; } std::optional<FullTrackName> track_name = ConstructTrackNameFromNamespace( track_namespace, GetChatId(my_track_name_)); - if (!parameters.has_value()) { - std::cout << "PUBLISH_NAMESPACE_DONE for " << track_namespace.ToString() + if (parameters == nullptr) { + std::cout << "PUBLISH_NAMESPACE done for " << track_namespace.ToString() << "\n"; if (track_name.has_value()) { other_users_.erase(*track_name); @@ -298,7 +298,7 @@ }}, response); }, - [](MoqtRequestErrorInfo) {}); + []() {}); // Send SUBSCRIBE_NAMESPACE. Pop 3 levels of namespace to get to // {moq-chat, chat-id} @@ -349,11 +349,10 @@ << "Error: received invalid suffix from namespace task\n"; return; } + MessageParameters parameters; OnIncomingPublishNamespace( *track_namespace, - (type == TransactionType::kAdd) - ? std::make_optional(MessageParameters()) - : std::nullopt, + (type == TransactionType::kAdd) ? ¶meters : nullptr, /*callback=*/nullptr); break; }
diff --git a/quiche/quic/moqt/tools/chat_client.h b/quiche/quic/moqt/tools/chat_client.h index b41f898..00f866b 100644 --- a/quiche/quic/moqt/tools/chat_client.h +++ b/quiche/quic/moqt/tools/chat_client.h
@@ -141,7 +141,7 @@ // a PUBLISH_NAMESPACE. void OnIncomingPublishNamespace( const moqt::TrackNamespace& track_namespace, - const std::optional<MessageParameters>& parameters, + const MessageParameters* absl_nullable parameters, moqt::MoqtResponseCallback absl_nullable callback); // Basic session information
diff --git a/quiche/quic/moqt/tools/moqt_ingestion_server_bin.cc b/quiche/quic/moqt/tools/moqt_ingestion_server_bin.cc index 8ad6337..753f2b5 100644 --- a/quiche/quic/moqt/tools/moqt_ingestion_server_bin.cc +++ b/quiche/quic/moqt/tools/moqt_ingestion_server_bin.cc
@@ -126,7 +126,7 @@ // TODO(martinduke): Handle when |publish_namespace| is false // (PUBLISH_NAMESPACE_DONE). void OnPublishNamespaceReceived(TrackNamespace track_namespace, - std::optional<MessageParameters>, + const MessageParameters*, MoqtResponseCallback callback) { if (!IsValidTrackNamespace(track_namespace) && !quiche::GetQuicheCommandLineFlag(
diff --git a/quiche/quic/moqt/tools/moqt_relay.cc b/quiche/quic/moqt/tools/moqt_relay.cc index 48d2a80..ebfa5d0 100644 --- a/quiche/quic/moqt/tools/moqt_relay.cc +++ b/quiche/quic/moqt/tools/moqt_relay.cc
@@ -10,6 +10,7 @@ #include <string> #include <utility> +#include "absl/base/nullability.h" #include "absl/status/statusor.h" #include "absl/strings/string_view.h" #include "quiche/quic/core/crypto/proof_source.h" @@ -23,7 +24,6 @@ #include "quiche/quic/moqt/moqt_session.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/tools/moqt_client.h" #include "quiche/quic/moqt/tools/moqt_server.h" #include "quiche/quic/platform/api/quic_default_proof_providers.h" @@ -115,12 +115,12 @@ void MoqtRelay::SetNamespaceCallbacks(MoqtSessionInterface* session) { session->callbacks().incoming_publish_namespace_callback = [this, session](const TrackNamespace& track_namespace, - const std::optional<MessageParameters>& parameters, + const MessageParameters* absl_nullable parameters, MoqtResponseCallback callback) { if (is_closing_) { return; } - if (parameters.has_value()) { + if (parameters != nullptr) { return publisher_.OnPublishNamespace(track_namespace, *parameters, session, std::move(callback)); } else {
diff --git a/quiche/quic/moqt/tools/moqt_relay_test.cc b/quiche/quic/moqt/tools/moqt_relay_test.cc index 2ea008e..0bc2584 100644 --- a/quiche/quic/moqt/tools/moqt_relay_test.cc +++ b/quiche/quic/moqt/tools/moqt_relay_test.cc
@@ -133,8 +133,7 @@ // relay_ publishes a namespace, so upstream_ will route to relay_. relay_.client_session()->PublishNamespace( TrackNamespace({"foo"}), MessageParameters(), - [](std::variant<MessageParameters, MoqtRequestErrorInfo>) {}, - [](MoqtRequestErrorInfo) {}); + [](std::variant<MessageParameters, MoqtRequestErrorInfo>) {}, []() {}); upstream_.RunOneEvent(); // There is now an upstream session for "Foo". std::shared_ptr<MoqtTrackPublisher> track = @@ -191,8 +190,7 @@ // hasn't been notified. downstream_.client_session()->PublishNamespace( foobar, MessageParameters(), - [](std::variant<MessageParameters, MoqtRequestErrorInfo>) {}, - [](MoqtRequestErrorInfo) {}); + [](std::variant<MessageParameters, MoqtRequestErrorInfo>) {}, []() {}); relay_.RunOneEvent(); upstream_.RunOneEvent(); EXPECT_THAT(relay_published_namespaces, ElementsAre(foobar)); @@ -223,8 +221,7 @@ // Downstream publishes another namespace. Everyone is notified. downstream_.client_session()->PublishNamespace( foobaz, MessageParameters(), - [](std::variant<MessageParameters, MoqtRequestErrorInfo>) {}, - [](MoqtRequestErrorInfo) {}); + [](std::variant<MessageParameters, MoqtRequestErrorInfo>) {}, []() {}); relay_.RunOneEvent(); upstream_.RunOneEvent(); EXPECT_THAT(relay_published_namespaces, ElementsAre(foobar, foobaz));