Move PUBLISH to a Bidi stream part of draft-17/18 update. Unfortunately, as subscription ownership is moving from the session to the bidi stream there are some differences in how objects are deleted, and whether or not to send a FIN with PUBLISH_DONE. This necessitates a (temporary) is_control_stream() boolean in MoqtBidiStreamBase to have the correct behavior if a Subscription sits on the control stream or its bidi stream. PiperOrigin-RevId: 936330335
diff --git a/build/source_list.bzl b/build/source_list.bzl index 1a78372..9d12049 100644 --- a/build/source_list.bzl +++ b/build/source_list.bzl
@@ -1595,6 +1595,7 @@ "quic/moqt/moqt_parser.h", "quic/moqt/moqt_priority.h", "quic/moqt/moqt_probe_manager.h", + "quic/moqt/moqt_publish_stream.h", "quic/moqt/moqt_publisher.h", "quic/moqt/moqt_quic_config.h", "quic/moqt/moqt_relay_publisher.h", @@ -1633,6 +1634,7 @@ "quic/moqt/moqt_parser.cc", "quic/moqt/moqt_priority.cc", "quic/moqt/moqt_probe_manager.cc", + "quic/moqt/moqt_publish_stream.cc", "quic/moqt/moqt_quic_config.cc", "quic/moqt/moqt_relay_publisher.cc", "quic/moqt/moqt_relay_track_publisher.cc", @@ -1668,6 +1670,7 @@ "quic/moqt/moqt_parser_test.cc", "quic/moqt/moqt_priority_test.cc", "quic/moqt/moqt_probe_manager_test.cc", + "quic/moqt/moqt_publish_stream_test.cc", "quic/moqt/moqt_relay_publisher_test.cc", "quic/moqt/moqt_relay_track_publisher_test.cc", "quic/moqt/moqt_session_test.cc",
diff --git a/build/source_list.gni b/build/source_list.gni index 65f0259..fdccb89 100644 --- a/build/source_list.gni +++ b/build/source_list.gni
@@ -1599,6 +1599,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_stream.h", "src/quiche/quic/moqt/moqt_publisher.h", "src/quiche/quic/moqt/moqt_quic_config.h", "src/quiche/quic/moqt/moqt_relay_publisher.h", @@ -1637,6 +1638,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_stream.cc", "src/quiche/quic/moqt/moqt_quic_config.cc", "src/quiche/quic/moqt/moqt_relay_publisher.cc", "src/quiche/quic/moqt/moqt_relay_track_publisher.cc", @@ -1673,6 +1675,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_stream_test.cc", "src/quiche/quic/moqt/moqt_relay_publisher_test.cc", "src/quiche/quic/moqt/moqt_relay_track_publisher_test.cc", "src/quiche/quic/moqt/moqt_session_test.cc",
diff --git a/build/source_list.json b/build/source_list.json index ef8b9de..ba05903 100644 --- a/build/source_list.json +++ b/build/source_list.json
@@ -1598,6 +1598,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_stream.h", "quiche/quic/moqt/moqt_publisher.h", "quiche/quic/moqt/moqt_quic_config.h", "quiche/quic/moqt/moqt_relay_publisher.h", @@ -1636,6 +1637,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_stream.cc", "quiche/quic/moqt/moqt_quic_config.cc", "quiche/quic/moqt/moqt_relay_publisher.cc", "quiche/quic/moqt/moqt_relay_track_publisher.cc", @@ -1672,6 +1674,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_stream_test.cc", "quiche/quic/moqt/moqt_relay_publisher_test.cc", "quiche/quic/moqt/moqt_relay_track_publisher_test.cc", "quiche/quic/moqt/moqt_session_test.cc",
diff --git a/quiche/quic/moqt/moqt_bidi_stream.cc b/quiche/quic/moqt/moqt_bidi_stream.cc index 46a4d78..df12569 100644 --- a/quiche/quic/moqt/moqt_bidi_stream.cc +++ b/quiche/quic/moqt/moqt_bidi_stream.cc
@@ -90,10 +90,20 @@ } std::optional<MoqtError> error_code = GetMoqtErrorForStatus(status); if (!error_code.has_value()) { - error_code = absl::IsInvalidArgument(status) ? MoqtError::kProtocolViolation - : MoqtError::kInternalError; + switch (status.code()) { + case absl::StatusCode::kInvalidArgument: + error_code = MoqtError::kProtocolViolation; + break; + case absl::StatusCode::kAlreadyExists: + error_code = MoqtError::kDuplicateTrackAlias; + break; + default: + error_code = MoqtError::kInternalError; + break; + } } std::move(session_error_callback_)(*error_code, status.message()); + session_error_callback_ = nullptr; } } // namespace moqt
diff --git a/quiche/quic/moqt/moqt_bidi_stream.h b/quiche/quic/moqt/moqt_bidi_stream.h index 2c9def3..23b1d4e 100644 --- a/quiche/quic/moqt/moqt_bidi_stream.h +++ b/quiche/quic/moqt/moqt_bidi_stream.h
@@ -114,6 +114,11 @@ } } + // TODO(martinduke): Remove once SUBSCRIBE moves to a bidi stream. This is + // only needed to check whether or not to FIN the bidi stream. + bool is_control_stream() const { return control_stream_; } + void set_control_stream() { control_stream_ = true; } + protected: // Called when a WebTransport stream has been associated with the object. // Should be used to set the priority for the stream. @@ -139,6 +144,8 @@ private: friend class test::MoqtBidiStreamTestWrapper; + // TODO(martinduke): Remove once SUBSCRIBE moves to a bidi stream. + bool control_stream_ = false; MoqtFramer* absl_nonnull framer_; std::unique_ptr<MoqtControlStreamParser> absl_nullable stream_parser_; MoqtControlMessageParser message_parser_;
diff --git a/quiche/quic/moqt/moqt_error.h b/quiche/quic/moqt/moqt_error.h index 2b08442..30b1a58 100644 --- a/quiche/quic/moqt/moqt_error.h +++ b/quiche/quic/moqt/moqt_error.h
@@ -67,6 +67,7 @@ kNotSupported = 0x3, kMalformedAuthToken = 0x4, kExpiredAuthToken = 0x5, + kGoingAway = 0x6, kDoesNotExist = 0x10, kInvalidRange = 0x11, kMalformedTrack = 0x12,
diff --git a/quiche/quic/moqt/moqt_integration_test.cc b/quiche/quic/moqt/moqt_integration_test.cc index 0ac6e81..c00cd60 100644 --- a/quiche/quic/moqt/moqt_integration_test.cc +++ b/quiche/quic/moqt/moqt_integration_test.cc
@@ -953,6 +953,90 @@ EXPECT_EQ(objects_enqueued, 1); } +TEST_F(MoqtIntegrationTest, ClientPublishServerSubscribe) { + EstablishSession(); + FullTrackName full_track_name("foo", "bar"); + + // Server registers incoming publish callback. + bool server_received_publish = false; + MoqtResponseCallback server_response_callback; + server_->session()->callbacks().incoming_publish_callback = + [&](const FullTrackName& name, const MessageParameters& parameters, + const TrackExtensions& extensions, MoqtResponseCallback callback) { + EXPECT_EQ(name, full_track_name); + server_response_callback = std::move(callback); + server_received_publish = true; + return &subscribe_visitor_; + }; + + // Client publishes. + auto queue = std::make_shared<TestTrackPublisher>(full_track_name); + bool client_publish_completed = false; + bool client_publish_success = false; + MoqtResponseCallback client_publish_callback = + [&](std::variant<MessageParameters, MoqtRequestErrorInfo> response) { + client_publish_completed = true; + client_publish_success = + std::holds_alternative<MessageParameters>(response); + }; + + MessageParameters publish_parameters; + TrackExtensions publish_extensions; + bool publish_submitted = + client_->session()->Publish(queue, publish_parameters, publish_extensions, + std::move(client_publish_callback)); + ASSERT_TRUE(publish_submitted); + + // Run until server receives PUBLISH. + bool success = test_harness_.RunUntilWithDefaultTimeout( + [&]() { return server_received_publish; }); + ASSERT_TRUE(success); + + // Server responds with REQUEST_OK. + std::move(server_response_callback)(MessageParameters()); + + // Run until client receives REQUEST_OK. + success = test_harness_.RunUntilWithDefaultTimeout( + [&]() { return client_publish_completed; }); + ASSERT_TRUE(success); + EXPECT_TRUE(client_publish_success); + + // Deliver objects. + queue->AddObject(Location(0, 0), 0, "object0", false); + queue->AddObject(Location(0, 1), 0, "object1", true); + + int received_objects = 0; + EXPECT_CALL(subscribe_visitor_, OnObjectFragment) + .Times(2) + .WillRepeatedly([&](const FullTrackName& name, + const PublishedObjectMetadata& metadata, + absl::string_view object, uint64_t offset) { + EXPECT_EQ(name, full_track_name); + if (received_objects == 0) { + EXPECT_EQ(metadata.location, Location(0, 0)); + EXPECT_EQ(object, "object0"); + } else if (received_objects == 1) { + EXPECT_EQ(metadata.location, Location(0, 1)); + EXPECT_EQ(object, "object1"); + } + ++received_objects; + }); + + success = test_harness_.RunUntilWithDefaultTimeout( + [&]() { return received_objects == 2; }); + EXPECT_TRUE(success); + + // Destroy the publisher to induce PUBLISH_DONE. + queue->RemoveAllSubscriptions(); + bool publish_done = false; + EXPECT_CALL(subscribe_visitor_, OnPublishDone).WillOnce([&]() { + publish_done = true; + }); + success = + test_harness_.RunUntilWithDefaultTimeout([&]() { return publish_done; }); + EXPECT_TRUE(success); +} + } // namespace } // namespace moqt::test
diff --git a/quiche/quic/moqt/moqt_key_value_pair.cc b/quiche/quic/moqt/moqt_key_value_pair.cc index 65ddda4..c8520b8 100644 --- a/quiche/quic/moqt/moqt_key_value_pair.cc +++ b/quiche/quic/moqt/moqt_key_value_pair.cc
@@ -91,9 +91,7 @@ if (other.subscription_filter.has_value()) { subscription_filter = other.subscription_filter; } - if (other.group_order.has_value()) { - group_order = other.group_order; - } + // Group order cannot be updated. if (other.new_group_request.has_value()) { new_group_request = other.new_group_request; }
diff --git a/quiche/quic/moqt/moqt_key_value_pair_test.cc b/quiche/quic/moqt/moqt_key_value_pair_test.cc index 56f54bb..d91793c 100644 --- a/quiche/quic/moqt/moqt_key_value_pair_test.cc +++ b/quiche/quic/moqt/moqt_key_value_pair_test.cc
@@ -254,7 +254,7 @@ AuthToken(AuthTokenType::kOutOfBand, "token")); EXPECT_TRUE(p1.forward()); EXPECT_EQ(p1.subscriber_priority, 100); - EXPECT_EQ(p1.group_order, MoqtDeliveryOrder::kDescending); + EXPECT_EQ(p1.group_order, std::nullopt); EXPECT_EQ(p1.new_group_request, 1); }
diff --git a/quiche/quic/moqt/moqt_publish_stream.cc b/quiche/quic/moqt/moqt_publish_stream.cc new file mode 100644 index 0000000..1e10548 --- /dev/null +++ b/quiche/quic/moqt/moqt_publish_stream.cc
@@ -0,0 +1,233 @@ +// 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_stream.h" + +#include <memory> +#include <optional> +#include <utility> +#include <variant> + +#include "absl/base/nullability.h" +#include "absl/functional/overload.h" +#include "absl/status/status.h" +#include "quiche/quic/core/quic_alarm_factory.h" +#include "quiche/quic/core/quic_clock.h" +#include "quiche/quic/moqt/moqt_bidi_stream.h" +#include "quiche/quic/moqt/moqt_error.h" +#include "quiche/quic/moqt/moqt_fetch_task.h" +#include "quiche/quic/moqt/moqt_framer.h" +#include "quiche/quic/moqt/moqt_key_value_pair.h" +#include "quiche/quic/moqt/moqt_messages.h" +#include "quiche/quic/moqt/moqt_parser.h" +#include "quiche/quic/moqt/moqt_session_callbacks.h" +#include "quiche/quic/moqt/moqt_subscription.h" +#include "quiche/quic/moqt/moqt_track.h" + +namespace moqt { + +MoqtPublishPublisherStream::MoqtPublishPublisherStream( + MoqtFramer* absl_nonnull framer, + const MoqtControlMessageParser& message_parser, + BidiStreamDeletedCallback stream_deleted_callback, + SessionErrorCallback session_error_callback, + MoqtResponseCallback response_callback) + : MoqtBidiStreamBase(framer, message_parser, + std::move(stream_deleted_callback), + std::move(session_error_callback)), + response_callback_(std::move(response_callback)) {} + +MoqtPublishPublisherStream::~MoqtPublishPublisherStream() {} + +void MoqtPublishPublisherStream::OnStreamBound() { + stream_parser()->set_allow_fin(true); + publisher_->parameters().largest_object = + publisher_->publisher().largest_location(); + publisher_->parameters().expires = publisher_->publisher().expiration(); + SendOrBufferMessageOrFatal(framer()->SerializePublish(MoqtPublish{ + publisher_->request_id(), publisher_->publisher().GetTrackName(), + publisher_->track_alias(), publisher_->parameters(), + publisher_->publisher().extensions()})); + // Use the default group order. + publisher_->parameters().group_order = + publisher_->publisher().extensions().default_publisher_group_order(); +} + +absl::Status MoqtPublishPublisherStream::OnRawControlMessage( + const MoqtRawControlMessage& message) { + return ControlMessageDispatcher::DispatchControlMessage( + *this, message_parser(), message, "publish publisher"); +} + +// TODO(martinduke): When we allow the publisher to send REQUEST_UPDATE, +// REQUEST_OK and REQUEST_ERROR processing need to check the request ID. +absl::Status MoqtPublishPublisherStream::OnControlMessage( + const MoqtRequestOk& message) { + if (message.request_id != publisher_->request_id()) { + OnFatalError(absl::InvalidArgumentError( + "REQUEST_OK does not match PUBLISH request ID")); + return absl::OkStatus(); + } + std::move(response_callback_)(message.parameters); + publisher_->Update(message.parameters); + // TODO(martinduke): Update() will not update group order because that is not + // allowed in REQUEST_UPDATE. PUBLISH_OK therefore needs to explicitly + // change the group order, but this would require reordering all streams by + // priority, and might create edge cases. + return absl::OkStatus(); +} + +absl::Status MoqtPublishPublisherStream::OnControlMessage( + const MoqtRequestError& message) { + if (message.request_id != publisher_->request_id()) { + OnFatalError(absl::InvalidArgumentError( + "REQUEST_OK does not match PUBLISH request ID")); + return absl::OkStatus(); + } + std::move(response_callback_)(MoqtRequestErrorInfo{ + message.error_code, message.retry_interval, message.reason_phrase}); + return absl::OkStatus(); +} + +absl::Status MoqtPublishPublisherStream::OnControlMessage( + const MoqtRequestUpdate& message) { + MessageParameters in_parameters = message.parameters, out_parameters; + out_parameters.largest_object = publisher_->publisher().largest_location(); + if (in_parameters.subscription_filter.has_value()) { + in_parameters.subscription_filter->OnLargestObject( + out_parameters.largest_object); + } + publisher_->Update(in_parameters); + CheckStatus(SendRequestOk(message.request_id, MessageParameters())); + return absl::OkStatus(); +} + +MoqtPublishSubscriberStream::MoqtPublishSubscriberStream( + MoqtFramer* absl_nonnull framer, + const MoqtControlMessageParser& message_parser, + const quic::QuicClock* absl_nonnull clock, + quic::QuicAlarmFactory* absl_nonnull alarm_factory, + SessionErrorCallback session_error_callback, + const MoqtIncomingPublishCallback* absl_nonnull incoming_publish_callback, + SubscribeRemoteTrack::SubscribeCallbacks callbacks) + : MoqtBidiStreamBase( + framer, message_parser, + /*stream_deleted_callback=*/+[]() {}, + std::move(session_error_callback)), + clock_(clock), + alarm_factory_(alarm_factory), + incoming_publish_callback_(incoming_publish_callback), + callbacks_(std::move(callbacks)), + weak_ptr_factory_(this) {} + +MoqtPublishSubscriberStream::~MoqtPublishSubscriberStream() { + in_destructor_ = true; +} + +absl::Status MoqtPublishSubscriberStream::OnRawControlMessage( + const MoqtRawControlMessage& message) { + return ControlMessageDispatcher::DispatchControlMessage( + *this, message_parser(), message, "publish subscriber"); +} + +absl::Status MoqtPublishSubscriberStream::OnControlMessage( + const MoqtPublish& message) { + if (incoming_publish_callback_ == nullptr) { + // Two PUBLISH messages for the same stream. + return absl::InvalidArgumentError("Multiple PUBLISH on the same stream"); + } + SubscribeVisitor* visitor = nullptr; + SubscribeRemoteTrack* existing_track = + std::move(callbacks_.query_name)(message.full_track_name); + callbacks_.query_name = nullptr; + if (existing_track != nullptr) { + // Track already exists. + if (!existing_track->ErrorIsAllowed()) { + // It's not a pending SUBSCRIBE; refuse this PUBLISH. + return SendRequestError(message.request_id, + RequestErrorCode::kDuplicateSubscription, + /*retry_interval=*/std::nullopt, "", + /*fin=*/true); + } + // It's a pending SUBSCRIBE. Transition it and accept the PUBLISH. + visitor = existing_track->ReleaseVisitor(); + existing_track->Destroy(); + } else { + // No existing SUBSCRIBE, get a new visitor from the application callback. + visitor = (*incoming_publish_callback_)( + message.full_track_name, message.parameters, message.extensions, + [weakptr = weak_ptr_factory_.Create(), request_id = message.request_id]( + const std::variant<MessageParameters, MoqtRequestErrorInfo> + response) { + MoqtPublishSubscriberStream* stream = weakptr.GetIfAvailable(); + if (stream == nullptr) { + return; + } + std::visit( + absl::Overload{[&](const MessageParameters& parameters) { + stream->subscriber_->Update(parameters); + stream->CheckStatus(stream->SendRequestOk( + request_id, parameters)); + }, + [&](const MoqtRequestErrorInfo& error_info) { + stream->CheckStatus(stream->SendRequestError( + request_id, error_info)); + }}, + response); + }); + } + incoming_publish_callback_ = nullptr; + if (visitor == nullptr) { + CheckStatus(SendRequestError(message.request_id, + RequestErrorCode::kUninterested, + /*retry_interval=*/std::nullopt, "", + /*fin=*/true)); + return absl::OkStatus(); + } + subscriber_ = std::make_unique<SubscribeRemoteTrack>( + message, visitor, + [this]() { + if (!in_destructor_) { + subscriber_.reset(); + stream()->ResetWithUserCode(kResetCodeCancelled); + } + }, + std::move(callbacks_)); + bool success = subscriber_->set_track_alias(message.track_alias); + if (!success) { + OnFatalError(absl::AlreadyExistsError("")); + } + return absl::OkStatus(); +} + +absl::Status MoqtPublishSubscriberStream::OnControlMessage( + const MoqtRequestUpdate& message) { + subscriber_->Update(message.parameters); + CheckStatus(SendRequestOk(message.request_id, MessageParameters())); + return absl::OkStatus(); +} + +absl::Status MoqtPublishSubscriberStream::OnControlMessage( + const MoqtRequestOk& message) { + // TODO(martinduke): Implement REQUEST_UPDATE. + return absl::OkStatus(); +} + +absl::Status MoqtPublishSubscriberStream::OnControlMessage( + const MoqtRequestError& message) { + // TODO(martinduke): Implement REQUEST_UPDATE. + return absl::OkStatus(); +} + +absl::Status MoqtPublishSubscriberStream::OnControlMessage( + const MoqtPublishDone& message) { + if (subscriber_ == nullptr) { + // PUBLISH_DONE can be sent before the subscriber rejects the track. + return absl::OkStatus(); + } + subscriber_->OnPublishDone(message.stream_count, clock_, alarm_factory_); + return absl::OkStatus(); +} + +} // namespace moqt
diff --git a/quiche/quic/moqt/moqt_publish_stream.h b/quiche/quic/moqt/moqt_publish_stream.h new file mode 100644 index 0000000..b10f5fd --- /dev/null +++ b/quiche/quic/moqt/moqt_publish_stream.h
@@ -0,0 +1,101 @@ +// 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_STREAM_H_ +#define QUICHE_QUIC_MOQT_MOQT_PUBLISH_STREAM_H_ + +#include <cstdint> +#include <memory> +#include <utility> + +#include "absl/base/nullability.h" +#include "absl/container/flat_hash_map.h" +#include "absl/status/status.h" +#include "quiche/quic/core/quic_alarm_factory.h" +#include "quiche/quic/core/quic_clock.h" +#include "quiche/quic/moqt/moqt_bidi_stream.h" +#include "quiche/quic/moqt/moqt_fetch_task.h" +#include "quiche/quic/moqt/moqt_framer.h" +#include "quiche/quic/moqt/moqt_messages.h" +#include "quiche/quic/moqt/moqt_parser.h" +#include "quiche/quic/moqt/moqt_session_callbacks.h" +#include "quiche/quic/moqt/moqt_subscription.h" +#include "quiche/quic/moqt/moqt_track.h" +#include "quiche/common/quiche_weak_ptr.h" + +namespace moqt { + +class MoqtPublishPublisherStream : public MoqtBidiStreamBase { + public: + // Order of operations: + // 1. Call this constructor + // 2. Call SetPublisher() + // 3. Call Webtransport::Stream::SetVisitor() + // 4. Call this::BindStream() + MoqtPublishPublisherStream(MoqtFramer* absl_nonnull framer, + const MoqtControlMessageParser& message_parser, + BidiStreamDeletedCallback stream_deleted_callback, + SessionErrorCallback session_error_callback, + MoqtResponseCallback response_callback); + ~MoqtPublishPublisherStream(); + + // MoqtBidiStreamBase overrides. + void OnStreamBound() override; + absl::Status OnRawControlMessage( + const MoqtRawControlMessage& message) override; + absl::Status OnControlMessage(const MoqtRequestOk& message); + absl::Status OnControlMessage(const MoqtRequestError& message); + absl::Status OnControlMessage(const MoqtRequestUpdate& message); + + void SetPublisher(std::unique_ptr<SubscriptionPublisher> publisher) { + publisher_ = std::move(publisher); + } + + private: + MoqtResponseCallback response_callback_; + std::unique_ptr<SubscriptionPublisher> publisher_; + absl::flat_hash_map<uint64_t, MoqtResponseCallback> pending_updates_; +}; + +class MoqtPublishSubscriberStream : public MoqtBidiStreamBase { + public: + MoqtPublishSubscriberStream( + MoqtFramer* absl_nonnull framer, + const MoqtControlMessageParser& message_parser, + const quic::QuicClock* absl_nonnull clock, + quic::QuicAlarmFactory* absl_nonnull alarm_factory, + SessionErrorCallback session_error_callback, + const MoqtIncomingPublishCallback* absl_nonnull incoming_publish_callback, + SubscribeRemoteTrack::SubscribeCallbacks callbacks); + ~MoqtPublishSubscriberStream(); + + // MoqtBidiStreamBase overrides. + void OnStreamBound() override { + stream_parser()->set_allow_fin(true); + // TODO(martinduke): Set the priority for this stream. + } + absl::Status OnRawControlMessage( + const MoqtRawControlMessage& message) override; + absl::Status OnControlMessage(const MoqtPublish& message); + absl::Status OnControlMessage(const MoqtRequestUpdate& message); + absl::Status OnControlMessage(const MoqtRequestOk& message); + absl::Status OnControlMessage(const MoqtRequestError& message); + absl::Status OnControlMessage(const MoqtPublishDone& message); + + private: + uint64_t request_id_; + SubscribeVisitor* absl_nullable subscribe_visitor_ = nullptr; + bool in_destructor_ = false; + std::unique_ptr<SubscribeRemoteTrack> subscriber_; + absl::flat_hash_map<uint64_t, MoqtResponseCallback> pending_updates_; + const quic::QuicClock* clock_; + quic::QuicAlarmFactory* alarm_factory_; + const MoqtIncomingPublishCallback* incoming_publish_callback_; + SubscribeRemoteTrack::SubscribeCallbacks callbacks_; + quiche::QuicheWeakPtrFactory<MoqtPublishSubscriberStream> weak_ptr_factory_; +}; + +} // namespace moqt + +#endif // QUICHE_QUIC_MOQT_MOQT_PUBLISH_STREAM_H_
diff --git a/quiche/quic/moqt/moqt_publish_stream_test.cc b/quiche/quic/moqt/moqt_publish_stream_test.cc new file mode 100644 index 0000000..b55e429 --- /dev/null +++ b/quiche/quic/moqt/moqt_publish_stream_test.cc
@@ -0,0 +1,778 @@ +// 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_stream.h" + +#include <cstdint> +#include <memory> +#include <optional> +#include <string> +#include <utility> +#include <variant> + +#include "absl/status/status.h" +#include "absl/status/statusor.h" +#include "absl/strings/string_view.h" +#include "quiche/quic/core/quic_alarm_factory.h" +#include "quiche/quic/core/quic_time.h" +#include "quiche/quic/core/quic_types.h" +#include "quiche/quic/moqt/moqt_error.h" +#include "quiche/quic/moqt/moqt_fetch_task.h" +#include "quiche/quic/moqt/moqt_framer.h" +#include "quiche/quic/moqt/moqt_key_value_pair.h" +#include "quiche/quic/moqt/moqt_messages.h" +#include "quiche/quic/moqt/moqt_names.h" +#include "quiche/quic/moqt/moqt_parser.h" +#include "quiche/quic/moqt/moqt_priority.h" +#include "quiche/quic/moqt/moqt_session_callbacks.h" +#include "quiche/quic/moqt/moqt_session_interface.h" +#include "quiche/quic/moqt/moqt_subscription.h" +#include "quiche/quic/moqt/moqt_trace_recorder.h" +#include "quiche/quic/moqt/moqt_track.h" +#include "quiche/quic/moqt/moqt_types.h" +#include "quiche/quic/moqt/test_tools/moqt_mock_visitor.h" +#include "quiche/quic/test_tools/mock_clock.h" +#include "quiche/quic/test_tools/quic_test_utils.h" +#include "quiche/common/platform/api/quiche_test.h" +#include "quiche/web_transport/test_tools/in_memory_stream.h" +#include "quiche/web_transport/test_tools/mock_web_transport.h" +#include "quiche/web_transport/web_transport.h" + +namespace moqt::test { + +class SubscriptionPublisherPeer { + public: + static const MessageParameters& parameters( + const SubscriptionPublisher& publisher) { + return publisher.parameters_; + } +}; + +class SubscribeRemoteTrackPeer { + public: + static const MessageParameters& parameters( + const SubscribeRemoteTrack& track) { + return track.const_parameters(); + } +}; + +namespace { + +using ::testing::_; +using ::testing::NotNull; +using ::testing::Return; +using ::testing::StrictMock; + +class MockSessionToPublisherInterface : public SessionToPublisherInterface { + public: + ~MockSessionToPublisherInterface() override = default; + MOCK_METHOD(bool, alternate_delivery_timeout, (), (const, override)); + MOCK_METHOD(void, UpdateTrackPriority, + (uint64_t, std::optional<MoqtTrackPriority>, MoqtTrackPriority), + (override)); + MOCK_METHOD(quic::QuicAlarmFactory*, alarm_factory, (), (override)); + MOCK_METHOD(void, PublishIsDone, (uint64_t), (override)); + MOCK_METHOD(webtransport::Session*, session, (), (override)); +}; + +class BidiStreamWithReset + : public webtransport::test::InMemoryStreamWithWriteBuffer { + public: + using InMemoryStreamWithWriteBuffer::InMemoryStreamWithWriteBuffer; + void ResetWithUserCode(webtransport::StreamErrorCode error) override { + last_reset_code_ = error; + } + std::optional<webtransport::StreamErrorCode> last_reset_code() const { + return last_reset_code_; + } + + private: + std::optional<webtransport::StreamErrorCode> last_reset_code_; +}; + +constexpr uint64_t kRequestId = 1; +constexpr uint64_t kTrackAlias = 10; +const FullTrackName kTrackName("foo", "bar"); + +class MoqtPublishPublisherStreamTest : public quiche::test::QuicheTest { + public: + MoqtPublishPublisherStreamTest() + : framer_(/*using_webtrans=*/true, quic::Perspective::IS_CLIENT), + message_parser_(kDefaultMoqtVersion, /*uses_web_transport=*/true, + quic::Perspective::IS_CLIENT), + track_publisher_(std::make_shared<TestTrackPublisher>(kTrackName)) { + // Construct the stream visitor. + stream_visitor_ = std::make_unique<MoqtPublishPublisherStream>( + &framer_, message_parser_, deleted_callback_.AsStdFunction(), + error_callback_.AsStdFunction(), + [this](std::variant<MessageParameters, MoqtRequestErrorInfo> response) { + response_ = response; + }); + + // Construct the SubscriptionPublisher. + parameters_.set_forward(true); + parameters_.delivery_timeout = quic::QuicTimeDelta::FromSeconds(1); + parameters_.group_order = MoqtDeliveryOrder::kAscending; + + EXPECT_CALL(visitor_, session).WillRepeatedly(Return(&webtrans_)); + auto publisher = std::make_unique<SubscriptionPublisher>( + framer_, track_publisher_, stream_visitor_.get(), kRequestId, + kTrackAlias, parameters_, &visitor_, /*monitoring_interface=*/nullptr, + &mock_clock_, trace_recorder_, /*is_publish=*/true); + + publisher_ = publisher.get(); // Keep raw pointer for testing + stream_visitor_->SetPublisher(std::move(publisher)); + } + + MoqtFramer framer_; + MoqtControlMessageParser message_parser_; + std::shared_ptr<TestTrackPublisher> track_publisher_; + testing::MockFunction<void()> deleted_callback_; + testing::StrictMock<testing::MockFunction<void(MoqtError, absl::string_view)>> + error_callback_; + MockSessionToPublisherInterface visitor_; + webtransport::test::MockSession webtrans_; + quic::MockClock mock_clock_; + MoqtTraceRecorder trace_recorder_; + MessageParameters parameters_; + + std::unique_ptr<MoqtPublishPublisherStream> stream_visitor_; + SubscriptionPublisher* publisher_; // Raw pointer + std::optional<std::variant<MessageParameters, MoqtRequestErrorInfo>> + response_; +}; + +TEST_F(MoqtPublishPublisherStreamTest, OnStreamBoundSendsPublish) { + webtransport::test::InMemoryStreamWithWriteBuffer stream(0); + stream_visitor_->BindStream(&stream); // Calls OnStreamBound + + // Verify PUBLISH message was sent. + std::string& written = stream.write_buffer(); + MoqtControlStreamParser parser(&stream); + // Feed the written data back to a parser to verify it. + webtransport::test::InMemoryStream read_stream(0); + read_stream.Receive(written); + MoqtControlStreamParser read_parser(&read_stream); + absl::StatusOr<MoqtRawControlMessage> message = read_parser.ReadNextMessage(); + ASSERT_TRUE(message.ok()); + EXPECT_EQ(message->type, MoqtMessageType::kPublish); + + MoqtControlMessageParser cmp(kDefaultMoqtVersion, true, + quic::Perspective::IS_CLIENT); + absl::StatusOr<MoqtPublish> publish = cmp.ProcessPublish(message->payload); + ASSERT_TRUE(publish.ok()); + EXPECT_EQ(publish->request_id, kRequestId); + EXPECT_EQ(publish->full_track_name, kTrackName); + EXPECT_EQ(publish->track_alias, kTrackAlias); + EXPECT_EQ(publish->parameters.delivery_timeout, parameters_.delivery_timeout); + EXPECT_EQ(publish->parameters.group_order, parameters_.group_order); +} + +TEST_F(MoqtPublishPublisherStreamTest, ReceiveRequestOk) { + webtransport::test::InMemoryStreamWithWriteBuffer stream(0); + stream_visitor_->BindStream(&stream); + + MoqtRequestOk request_ok; + request_ok.request_id = kRequestId; + request_ok.parameters.delivery_timeout = quic::QuicTimeDelta::FromSeconds(2); + request_ok.parameters.group_order = MoqtDeliveryOrder::kDescending; + + stream.Receive(framer_.SerializeRequestOk(request_ok).AsStringView()); + stream_visitor_->OnCanRead(); + + // Verify response callback was called. + ASSERT_TRUE(response_.has_value()); + ASSERT_TRUE(std::holds_alternative<MessageParameters>(*response_)); + MessageParameters resp_params = std::get<MessageParameters>(*response_); + EXPECT_EQ(resp_params.delivery_timeout, + request_ok.parameters.delivery_timeout); + EXPECT_EQ(resp_params.group_order, request_ok.parameters.group_order); + + // Verify publisher parameters were updated. + const MessageParameters& pub_params = + SubscriptionPublisherPeer::parameters(*publisher_); + EXPECT_EQ(pub_params.delivery_timeout, + request_ok.parameters.delivery_timeout); + // Group order cannot be updated. + EXPECT_EQ(pub_params.group_order, parameters_.group_order); +} + +TEST_F(MoqtPublishPublisherStreamTest, ReceiveRequestError) { + webtransport::test::InMemoryStreamWithWriteBuffer stream(0); + stream_visitor_->BindStream(&stream); + + MoqtRequestError request_error; + request_error.request_id = kRequestId; + request_error.error_code = RequestErrorCode::kUnauthorized; + request_error.retry_interval = quic::QuicTimeDelta::FromSeconds(5); + request_error.reason_phrase = "Unauthorized"; + + stream.Receive(framer_.SerializeRequestError(request_error).AsStringView()); + stream_visitor_->OnCanRead(); + + // Verify response callback was called with error. + ASSERT_TRUE(response_.has_value()); + ASSERT_TRUE(std::holds_alternative<MoqtRequestErrorInfo>(*response_)); + MoqtRequestErrorInfo resp_error = std::get<MoqtRequestErrorInfo>(*response_); + EXPECT_EQ(resp_error.error_code, request_error.error_code); + EXPECT_EQ(resp_error.retry_interval, request_error.retry_interval); + EXPECT_EQ(resp_error.reason_phrase, request_error.reason_phrase); +} + +TEST_F(MoqtPublishPublisherStreamTest, ReceiveRequestUpdate) { + webtransport::test::InMemoryStreamWithWriteBuffer stream(0); + stream_visitor_->BindStream(&stream); + stream.write_buffer().clear(); // Clear initial PUBLISH + + // Set largest location on publisher + track_publisher_->AddObject(Location(1, 2), 0, "payload", true); + + MoqtRequestUpdate request_update; + request_update.request_id = kRequestId + 2; + request_update.existing_request_id = kRequestId; + request_update.parameters.delivery_timeout = + quic::QuicTimeDelta::FromSeconds(3); + request_update.parameters.subscriber_priority = 5; + request_update.parameters.subscription_filter.emplace( + MoqtFilterType::kLargestObject); + + stream.Receive(framer_.SerializeRequestUpdate(request_update).AsStringView()); + stream_visitor_->OnCanRead(); + + // Verify publisher parameters were updated. + const MessageParameters& pub_params = + SubscriptionPublisherPeer::parameters(*publisher_); + EXPECT_EQ(pub_params.delivery_timeout, + request_update.parameters.delivery_timeout); + EXPECT_EQ(pub_params.subscriber_priority, + request_update.parameters.subscriber_priority); + + // Verify filter was updated based on largest location (1, 2) -> (1, 3) + // AbsoluteStart + ASSERT_TRUE(pub_params.subscription_filter.has_value()); + EXPECT_EQ(pub_params.subscription_filter->type(), + MoqtFilterType::kAbsoluteStart); + EXPECT_EQ(pub_params.subscription_filter->start(), Location(1, 3)); + + // Verify REQUEST_OK response was sent. + std::string& written = stream.write_buffer(); + webtransport::test::InMemoryStream read_stream(0); + read_stream.Receive(written); + MoqtControlStreamParser read_parser(&read_stream); + absl::StatusOr<MoqtRawControlMessage> message = read_parser.ReadNextMessage(); + ASSERT_TRUE(message.ok()); + EXPECT_EQ(message->type, MoqtMessageType::kRequestOk); + + MoqtControlMessageParser cmp(kDefaultMoqtVersion, true, + quic::Perspective::IS_CLIENT); + absl::StatusOr<MoqtRequestOk> request_ok = + cmp.ProcessRequestOk(message->payload); + ASSERT_TRUE(request_ok.ok()); + EXPECT_EQ(request_ok->request_id, request_update.request_id); +} + +TEST_F(MoqtPublishPublisherStreamTest, ReceiveRequestOkMismatchedId) { + BidiStreamWithReset stream(0); + stream_visitor_->BindStream(&stream); + + // Receive REQUEST_OK with mismatched ID. + MoqtRequestOk request_ok; + request_ok.request_id = kRequestId + 1; // Mismatched + EXPECT_CALL(error_callback_, + Call(MoqtError::kProtocolViolation, + "REQUEST_OK does not match PUBLISH request ID")); + stream.Receive(framer_.SerializeRequestOk(request_ok).AsStringView()); + stream_visitor_->OnCanRead(); +} + +TEST_F(MoqtPublishPublisherStreamTest, ReceiveRequestErrorMismatchedId) { + BidiStreamWithReset stream(0); + stream_visitor_->BindStream(&stream); + + // Receive REQUEST_ERROR with mismatched ID. + MoqtRequestError request_error; + request_error.request_id = kRequestId + 1; // Mismatched + request_error.error_code = RequestErrorCode::kUninterested; + EXPECT_CALL(error_callback_, + Call(MoqtError::kProtocolViolation, + "REQUEST_OK does not match PUBLISH request ID")); + stream.Receive(framer_.SerializeRequestError(request_error).AsStringView()); + stream_visitor_->OnCanRead(); +} + +class MoqtPublishSubscriberStreamTest : public quiche::test::QuicheTest { + public: + MoqtPublishSubscriberStreamTest() + : framer_(/*using_webtrans=*/true, quic::Perspective::IS_SERVER), + message_parser_(kDefaultMoqtVersion, /*uses_web_transport=*/true, + quic::Perspective::IS_SERVER), + incoming_publish_callback_( + incoming_publish_callback_mock_.AsStdFunction()) { + SubscribeRemoteTrack::SubscribeCallbacks callbacks; + callbacks.query_name = [this](const FullTrackName& name) { + return query_name_mock_.Call(name); + }; + callbacks.register_name = [this](const FullTrackName& name, + SubscribeRemoteTrack* track) { + register_name_mock_.Call(name, track); + }; + callbacks.register_alias = [this](uint64_t alias, + SubscribeRemoteTrack* track) { + return register_alias_mock_.Call(alias, track); + }; + callbacks.unregister = [this](const FullTrackName& name, + std::optional<uint64_t> alias) { + unregister_mock_.Call(name, alias); + }; + + stream_visitor_ = std::make_unique<MoqtPublishSubscriberStream>( + &framer_, message_parser_, &mock_clock_, &mock_alarm_factory_, + error_callback_.AsStdFunction(), &incoming_publish_callback_, + std::move(callbacks)); + } + + void ExpectSubscriberDestruction() { + EXPECT_CALL(unregister_mock_, + Call(kTrackName, std::optional<uint64_t>(kTrackAlias))) + .Times(1); + } + + MoqtFramer framer_; + MoqtControlMessageParser message_parser_; + quic::MockClock mock_clock_; + quic::test::MockAlarmFactory mock_alarm_factory_; + testing::StrictMock<testing::MockFunction<void(MoqtError, absl::string_view)>> + error_callback_; + + testing::MockFunction<SubscribeVisitor*( + const FullTrackName&, const MessageParameters&, const TrackExtensions&, + MoqtResponseCallback)> + incoming_publish_callback_mock_; + MoqtIncomingPublishCallback incoming_publish_callback_; + + testing::MockFunction<SubscribeRemoteTrack*(const FullTrackName&)> + query_name_mock_; + testing::MockFunction<void(const FullTrackName&, SubscribeRemoteTrack*)> + register_name_mock_; + testing::MockFunction<bool(uint64_t, SubscribeRemoteTrack*)> + register_alias_mock_; + testing::MockFunction<void(const FullTrackName&, std::optional<uint64_t>)> + unregister_mock_; + + StrictMock<MockSubscribeRemoteTrackVisitor> mock_subscribe_visitor_; + MoqtResponseCallback captured_response_callback_; + std::unique_ptr<MoqtPublishSubscriberStream> stream_visitor_; +}; + +TEST_F(MoqtPublishSubscriberStreamTest, ReceivePublishAndAccept) { + BidiStreamWithReset stream(0); + stream_visitor_->BindStream(&stream); + + EXPECT_CALL(mock_subscribe_visitor_, OnReply(kTrackName, _)) + .WillOnce( + [](const FullTrackName&, + const std::variant<SubscribeOkData, MoqtRequestErrorInfo>& reply) { + EXPECT_TRUE(std::holds_alternative<SubscribeOkData>(reply)); + }); + EXPECT_CALL(mock_subscribe_visitor_, OnPublishDone(kTrackName)); + ExpectSubscriberDestruction(); + + MoqtPublish publish; + publish.request_id = kRequestId; + publish.full_track_name = kTrackName; + publish.track_alias = kTrackAlias; + publish.parameters.delivery_timeout = quic::QuicTimeDelta::FromSeconds(1); + + EXPECT_CALL(incoming_publish_callback_mock_, Call(kTrackName, _, _, _)) + .WillOnce([this](const FullTrackName&, const MessageParameters&, + const TrackExtensions&, MoqtResponseCallback callback) { + captured_response_callback_ = std::move(callback); + return &mock_subscribe_visitor_; + }); + + EXPECT_CALL(query_name_mock_, Call(kTrackName)).WillOnce(Return(nullptr)); + EXPECT_CALL(register_name_mock_, Call(kTrackName, NotNull())).Times(1); + SubscribeRemoteTrack* captured_subscriber = nullptr; + EXPECT_CALL(register_alias_mock_, Call(kTrackAlias, NotNull())) + .WillOnce([&](uint64_t, SubscribeRemoteTrack* track) { + captured_subscriber = track; + return true; + }); + + stream.Receive(framer_.SerializePublish(publish).AsStringView()); + stream_visitor_->OnCanRead(); + + // Verify subscriber is created. + ASSERT_NE(captured_subscriber, nullptr); + EXPECT_EQ(captured_subscriber->track_alias(), kTrackAlias); + EXPECT_EQ(captured_subscriber->visitor(), &mock_subscribe_visitor_); + + // Now call the response callback with success. + stream.write_buffer().clear(); + MessageParameters response_parameters; + response_parameters.delivery_timeout = quic::QuicTimeDelta::FromSeconds(2); + std::move(captured_response_callback_)(response_parameters); + + // Verify REQUEST_OK response was sent. + std::string& written = stream.write_buffer(); + webtransport::test::InMemoryStream read_stream(0); + read_stream.Receive(written); + MoqtControlStreamParser read_parser(&read_stream); + absl::StatusOr<MoqtRawControlMessage> message = read_parser.ReadNextMessage(); + ASSERT_TRUE(message.ok()); + EXPECT_EQ(message->type, MoqtMessageType::kRequestOk); + + MoqtControlMessageParser cmp(kDefaultMoqtVersion, true, + quic::Perspective::IS_SERVER); + absl::StatusOr<MoqtRequestOk> request_ok = + cmp.ProcessRequestOk(message->payload); + ASSERT_TRUE(request_ok.ok()); + EXPECT_EQ(request_ok->request_id, kRequestId); + EXPECT_EQ(request_ok->parameters.delivery_timeout, + response_parameters.delivery_timeout); + + // Verify subscriber parameters were updated. + const MessageParameters& sub_params = + SubscribeRemoteTrackPeer::parameters(*captured_subscriber); + EXPECT_EQ(sub_params.delivery_timeout, response_parameters.delivery_timeout); +} + +TEST_F(MoqtPublishSubscriberStreamTest, ReceivePublishAndReject) { + BidiStreamWithReset stream(0); + stream_visitor_->BindStream(&stream); + + MoqtPublish publish; + publish.request_id = kRequestId; + publish.full_track_name = kTrackName; + publish.track_alias = kTrackAlias; + + // Callback returns nullptr (rejection). + EXPECT_CALL(incoming_publish_callback_mock_, Call(kTrackName, _, _, _)) + .WillOnce(Return(nullptr)); + + stream.Receive(framer_.SerializePublish(publish).AsStringView()); + stream_visitor_->OnCanRead(); + + // Verify REQUEST_ERROR was sent. + std::string& written = stream.write_buffer(); + webtransport::test::InMemoryStream read_stream(0); + read_stream.Receive(written); + MoqtControlStreamParser read_parser(&read_stream); + absl::StatusOr<MoqtRawControlMessage> message = read_parser.ReadNextMessage(); + ASSERT_TRUE(message.ok()); + EXPECT_EQ(message->type, MoqtMessageType::kRequestError); + + MoqtControlMessageParser cmp(kDefaultMoqtVersion, true, + quic::Perspective::IS_SERVER); + absl::StatusOr<MoqtRequestError> request_error = + cmp.ProcessRequestError(message->payload); + ASSERT_TRUE(request_error.ok()); + EXPECT_EQ(request_error->request_id, kRequestId); + EXPECT_EQ(request_error->error_code, RequestErrorCode::kUninterested); + EXPECT_TRUE(stream.fin_sent()); +} + +TEST_F(MoqtPublishSubscriberStreamTest, ReceivePublishDuplicate) { + BidiStreamWithReset stream(0); + stream_visitor_->BindStream(&stream); + + MoqtPublish publish; + publish.request_id = kRequestId; + publish.full_track_name = kTrackName; + publish.track_alias = kTrackAlias; + + EXPECT_CALL(incoming_publish_callback_mock_, Call(kTrackName, _, _, _)) + .WillOnce(Return(&mock_subscribe_visitor_)); + EXPECT_CALL(query_name_mock_, Call(kTrackName)).WillOnce(Return(nullptr)); + EXPECT_CALL(register_name_mock_, Call(kTrackName, NotNull())).Times(1); + EXPECT_CALL(register_alias_mock_, Call(kTrackAlias, NotNull())) + .WillOnce(Return(true)); + EXPECT_CALL(mock_subscribe_visitor_, OnReply(kTrackName, _)) + .WillOnce( + [](const FullTrackName&, + const std::variant<SubscribeOkData, MoqtRequestErrorInfo>& reply) { + EXPECT_TRUE(std::holds_alternative<SubscribeOkData>(reply)); + }); + EXPECT_CALL(mock_subscribe_visitor_, OnPublishDone(kTrackName)); + ExpectSubscriberDestruction(); + + stream.Receive(framer_.SerializePublish(publish).AsStringView()); + stream_visitor_->OnCanRead(); + + // Send second PUBLISH on same stream. + stream.Receive(framer_.SerializePublish(publish).AsStringView()); + + // It should return InvalidArgumentError, which calls OnFatalError in + // MoqtBidiStreamBase. + EXPECT_CALL(error_callback_, Call(MoqtError::kProtocolViolation, + "Multiple PUBLISH on the same stream")); + stream_visitor_->OnCanRead(); +} + +TEST_F(MoqtPublishSubscriberStreamTest, ReceivePublishDuplicateName) { + BidiStreamWithReset stream(0); + stream_visitor_->BindStream(&stream); + + MoqtPublish publish; + publish.request_id = kRequestId; + publish.full_track_name = kTrackName; + publish.track_alias = kTrackAlias; + + // Simulate an existing established track. + MoqtSubscribe sub; + sub.full_track_name = kTrackName; + StrictMock<MockSubscribeRemoteTrackVisitor> existing_visitor; + SubscribeRemoteTrack existing_track(sub, &existing_visitor, []() {}, {}); + existing_track.OnObjectOrOk(); + + EXPECT_CALL(existing_visitor, OnPublishDone(kTrackName)).Times(1); + + EXPECT_CALL(query_name_mock_, Call(kTrackName)) + .WillOnce(Return(&existing_track)); + + stream.Receive(framer_.SerializePublish(publish).AsStringView()); + stream_visitor_->OnCanRead(); + + // Verify REQUEST_ERROR was sent. + std::string& written = stream.write_buffer(); + webtransport::test::InMemoryStream read_stream(0); + read_stream.Receive(written); + MoqtControlStreamParser read_parser(&read_stream); + absl::StatusOr<MoqtRawControlMessage> message = read_parser.ReadNextMessage(); + ASSERT_TRUE(message.ok()); + EXPECT_EQ(message->type, MoqtMessageType::kRequestError); + + MoqtControlMessageParser cmp(kDefaultMoqtVersion, true, + quic::Perspective::IS_SERVER); + absl::StatusOr<MoqtRequestError> request_error = + cmp.ProcessRequestError(message->payload); + ASSERT_TRUE(request_error.ok()); + EXPECT_EQ(request_error->request_id, kRequestId); + EXPECT_EQ(request_error->error_code, + RequestErrorCode::kDuplicateSubscription); + EXPECT_TRUE(stream.fin_sent()); +} + +TEST_F(MoqtPublishSubscriberStreamTest, ReceivePublishDuplicateAlias) { + BidiStreamWithReset stream(0); + stream_visitor_->BindStream(&stream); + + MoqtPublish publish; + publish.request_id = kRequestId; + publish.full_track_name = kTrackName; + publish.track_alias = kTrackAlias; + + EXPECT_CALL(incoming_publish_callback_mock_, Call(kTrackName, _, _, _)) + .WillOnce(Return(&mock_subscribe_visitor_)); + + EXPECT_CALL(query_name_mock_, Call(kTrackName)).WillOnce(Return(nullptr)); + EXPECT_CALL(register_name_mock_, Call(kTrackName, NotNull())).Times(1); + // Return duplicate alias error (false) from alias callback. + EXPECT_CALL(register_alias_mock_, Call(kTrackAlias, NotNull())) + .WillOnce(Return(false)); + EXPECT_CALL(mock_subscribe_visitor_, OnReply(kTrackName, _)) + .WillOnce( + [](const FullTrackName&, + const std::variant<SubscribeOkData, MoqtRequestErrorInfo>& reply) { + EXPECT_TRUE(std::holds_alternative<SubscribeOkData>(reply)); + }); + EXPECT_CALL(mock_subscribe_visitor_, OnPublishDone(kTrackName)); + ExpectSubscriberDestruction(); + + // It should call OnFatalError, which calls error_callback_. + // Note: The error message is now empty because we pass + // AlreadyExistsError(""). + EXPECT_CALL(error_callback_, Call(MoqtError::kDuplicateTrackAlias, "")); + + stream.Receive(framer_.SerializePublish(publish).AsStringView()); + stream_visitor_->OnCanRead(); +} + +TEST_F(MoqtPublishSubscriberStreamTest, ReceiveRequestUpdate) { + // First, establish subscription. + BidiStreamWithReset stream(0); + stream_visitor_->BindStream(&stream); + + MoqtPublish publish; + publish.request_id = kRequestId; + publish.full_track_name = kTrackName; + publish.track_alias = kTrackAlias; + + EXPECT_CALL(incoming_publish_callback_mock_, Call(kTrackName, _, _, _)) + .WillOnce(Return(&mock_subscribe_visitor_)); + + EXPECT_CALL(query_name_mock_, Call(kTrackName)).WillOnce(Return(nullptr)); + EXPECT_CALL(register_name_mock_, Call(kTrackName, NotNull())).Times(1); + SubscribeRemoteTrack* captured_subscriber = nullptr; + EXPECT_CALL(register_alias_mock_, Call(kTrackAlias, NotNull())) + .WillOnce([&](uint64_t, SubscribeRemoteTrack* track) { + captured_subscriber = track; + return true; + }); + EXPECT_CALL(mock_subscribe_visitor_, OnReply(kTrackName, _)) + .WillOnce( + [](const FullTrackName&, + const std::variant<SubscribeOkData, MoqtRequestErrorInfo>& reply) { + EXPECT_TRUE(std::holds_alternative<SubscribeOkData>(reply)); + }); + EXPECT_CALL(mock_subscribe_visitor_, OnPublishDone(kTrackName)); + ExpectSubscriberDestruction(); + + stream.Receive(framer_.SerializePublish(publish).AsStringView()); + stream_visitor_->OnCanRead(); + stream.write_buffer().clear(); + + // Now receive REQUEST_UPDATE. + MoqtRequestUpdate request_update; + request_update.request_id = kRequestId + 2; + request_update.existing_request_id = kRequestId; + request_update.parameters.delivery_timeout = + quic::QuicTimeDelta::FromSeconds(3); + + stream.Receive(framer_.SerializeRequestUpdate(request_update).AsStringView()); + stream_visitor_->OnCanRead(); + + // Verify subscriber parameters were updated. + ASSERT_NE(captured_subscriber, nullptr); + const MessageParameters& sub_params = + SubscribeRemoteTrackPeer::parameters(*captured_subscriber); + EXPECT_EQ(sub_params.delivery_timeout, + request_update.parameters.delivery_timeout); + + // Verify REQUEST_OK response was sent. + std::string& written = stream.write_buffer(); + webtransport::test::InMemoryStream read_stream(0); + read_stream.Receive(written); + MoqtControlStreamParser read_parser(&read_stream); + absl::StatusOr<MoqtRawControlMessage> message = read_parser.ReadNextMessage(); + ASSERT_TRUE(message.ok()); + EXPECT_EQ(message->type, MoqtMessageType::kRequestOk); +} + +TEST_F(MoqtPublishSubscriberStreamTest, ReceivePublishDone) { + // First, establish subscription. + BidiStreamWithReset stream(0); + stream_visitor_->BindStream(&stream); + + MoqtPublish publish; + publish.request_id = kRequestId; + publish.full_track_name = kTrackName; + publish.track_alias = kTrackAlias; + + EXPECT_CALL(incoming_publish_callback_mock_, Call(kTrackName, _, _, _)) + .WillOnce(Return(&mock_subscribe_visitor_)); + EXPECT_CALL(query_name_mock_, Call(kTrackName)).WillOnce(Return(nullptr)); + EXPECT_CALL(register_name_mock_, Call(kTrackName, NotNull())).Times(1); + EXPECT_CALL(register_alias_mock_, Call(kTrackAlias, NotNull())) + .WillOnce(Return(true)); + EXPECT_CALL(mock_subscribe_visitor_, OnReply(kTrackName, _)) + .WillOnce( + [](const FullTrackName&, + const std::variant<SubscribeOkData, MoqtRequestErrorInfo>& reply) { + EXPECT_TRUE(std::holds_alternative<SubscribeOkData>(reply)); + }); + EXPECT_CALL(mock_subscribe_visitor_, OnPublishDone(kTrackName)); + ExpectSubscriberDestruction(); + + stream.Receive(framer_.SerializePublish(publish).AsStringView()); + stream_visitor_->OnCanRead(); + + // Now receive PUBLISH_DONE. + MoqtPublishDone publish_done; + publish_done.request_id = kRequestId; + publish_done.status_code = PublishDoneCode::kTrackEnded; + publish_done.stream_count = 0; // Trigger immediate Destroy + + stream.Receive(framer_.SerializePublishDone(publish_done).AsStringView()); + stream_visitor_->OnCanRead(); + + EXPECT_EQ(stream.last_reset_code(), kResetCodeCancelled); +} + +TEST_F(MoqtPublishSubscriberStreamTest, ReceivePublishAndRejectCallback) { + BidiStreamWithReset stream(0); + stream_visitor_->BindStream(&stream); + + MoqtPublish publish; + publish.request_id = kRequestId; + publish.full_track_name = kTrackName; + publish.track_alias = kTrackAlias; + + EXPECT_CALL(incoming_publish_callback_mock_, Call(kTrackName, _, _, _)) + .WillOnce([this](const FullTrackName&, const MessageParameters&, + const TrackExtensions&, MoqtResponseCallback callback) { + captured_response_callback_ = std::move(callback); + return &mock_subscribe_visitor_; + }); + + EXPECT_CALL(query_name_mock_, Call(kTrackName)).WillOnce(Return(nullptr)); + EXPECT_CALL(register_name_mock_, Call(kTrackName, NotNull())).Times(1); + EXPECT_CALL(register_alias_mock_, Call(kTrackAlias, NotNull())) + .WillOnce(Return(true)); + EXPECT_CALL(mock_subscribe_visitor_, OnReply(kTrackName, _)) + .WillOnce( + [](const FullTrackName&, + const std::variant<SubscribeOkData, MoqtRequestErrorInfo>& reply) { + EXPECT_TRUE(std::holds_alternative<SubscribeOkData>(reply)); + }); + EXPECT_CALL(mock_subscribe_visitor_, OnPublishDone(kTrackName)); + ExpectSubscriberDestruction(); + + stream.Receive(framer_.SerializePublish(publish).AsStringView()); + stream_visitor_->OnCanRead(); + + // Now call the response callback with error (reject). + stream.write_buffer().clear(); + MoqtRequestErrorInfo error_info{RequestErrorCode::kUninterested, std::nullopt, + "rejected by app"}; + std::move(captured_response_callback_)(error_info); + + // Verify REQUEST_ERROR response was sent. + std::string& written = stream.write_buffer(); + webtransport::test::InMemoryStream read_stream(0); + read_stream.Receive(written); + MoqtControlStreamParser read_parser(&read_stream); + absl::StatusOr<MoqtRawControlMessage> message = read_parser.ReadNextMessage(); + ASSERT_TRUE(message.ok()); + EXPECT_EQ(message->type, MoqtMessageType::kRequestError); + + MoqtControlMessageParser cmp(kDefaultMoqtVersion, true, + quic::Perspective::IS_SERVER); + absl::StatusOr<MoqtRequestError> request_error = + cmp.ProcessRequestError(message->payload); + ASSERT_TRUE(request_error.ok()); + EXPECT_EQ(request_error->request_id, kRequestId); + EXPECT_EQ(request_error->error_code, RequestErrorCode::kUninterested); + EXPECT_EQ(request_error->reason_phrase, "rejected by app"); +} + +TEST_F(MoqtPublishSubscriberStreamTest, ReceivePublishDoneOnRejectedStream) { + BidiStreamWithReset stream(0); + stream_visitor_->BindStream(&stream); + + MoqtPublish publish; + publish.request_id = kRequestId; + publish.full_track_name = kTrackName; + publish.track_alias = kTrackAlias; + + // Callback returns nullptr (rejection). + EXPECT_CALL(incoming_publish_callback_mock_, Call(kTrackName, _, _, _)) + .WillOnce(Return(nullptr)); + + stream.Receive(framer_.SerializePublish(publish).AsStringView()); + stream_visitor_->OnCanRead(); + + // Now receive PUBLISH_DONE. + MoqtPublishDone publish_done; + publish_done.request_id = kRequestId; + publish_done.status_code = PublishDoneCode::kTrackEnded; + publish_done.stream_count = 0; + + EXPECT_CALL(mock_subscribe_visitor_, OnPublishDone(kTrackName)).Times(0); + stream.Receive(framer_.SerializePublishDone(publish_done).AsStringView()); + stream_visitor_->OnCanRead(); +} + +} // namespace +} // namespace moqt::test
diff --git a/quiche/quic/moqt/moqt_session.cc b/quiche/quic/moqt/moqt_session.cc index 8db77b3..9a7c7e4 100644 --- a/quiche/quic/moqt/moqt_session.cc +++ b/quiche/quic/moqt/moqt_session.cc
@@ -38,6 +38,7 @@ #include "quiche/quic/moqt/moqt_object.h" #include "quiche/quic/moqt/moqt_parser.h" #include "quiche/quic/moqt/moqt_priority.h" +#include "quiche/quic/moqt/moqt_publish_stream.h" #include "quiche/quic/moqt/moqt_publisher.h" #include "quiche/quic/moqt/moqt_session_callbacks.h" #include "quiche/quic/moqt/moqt_session_interface.h" @@ -487,26 +488,8 @@ << message.full_track_name; auto track = std::make_unique<SubscribeRemoteTrack>( message, visitor, - [this, request_id = message.request_id, ftn = name]() { - // Deletion callback - subscribe_by_name_.erase(ftn); - upstream_by_id_.erase(request_id); - }, - [this](uint64_t alias, SubscribeRemoteTrack* track) { - // Track alias registry callback. - if (is_closing_) { - return true; - } - if (track == nullptr) { - subscribe_by_alias_.erase(alias); - return true; - } - auto [it, success] = subscribe_by_alias_.try_emplace(alias, track); - if (!success) { - Error(MoqtError::kDuplicateTrackAlias, ""); - } - return success; - }); + [this, id = message.request_id]() { upstream_by_id_.erase(id); }, + GetSubscribeCallbacks()); subscribe_by_name_.emplace(message.full_track_name, track.get()); upstream_by_id_.emplace(message.request_id, std::move(track)); return true; @@ -534,6 +517,7 @@ if (it == subscribe_by_name_.end()) { return false; } + // TODO(martinduke): Support Update on PUBLISH streams. pending_subscribe_updates_[next_request_id_] = {name, parameters, std::move(response_callback)}; MoqtRequestUpdate update{next_request_id_, it->second->request_id(), @@ -560,6 +544,59 @@ track->Destroy(); } +bool MoqtSession::Publish( + std::shared_ptr<MoqtTrackPublisher> absl_nonnull publisher, + const MessageParameters& parameters, const TrackExtensions& extensions, + MoqtResponseCallback response_callback) { + if (received_goaway_ || sent_goaway_) { + QUICHE_DLOG(INFO) << ENDPOINT << "Tried to send PUBLISH after GOAWAY"; + return false; + } + const FullTrackName& name = publisher->GetTrackName(); + QUICHE_DCHECK(name.IsValid()); + if (!session_->CanOpenNextOutgoingBidirectionalStream()) { + return false; // Do not retry opening a PUBLISH stream. + } + if (!subscribed_track_names_.insert(name).second) { + QUICHE_DLOG(INFO) << ENDPOINT << "Tried to send PUBLISH for track " << name + << " which is already published"; + return false; + } + auto stream_visitor = std::make_unique<MoqtPublishPublisherStream>( + &framer_, ControlMessageParser(), + [weak_session = GetWeakPtr(), track_name = name]() { + // Stream deleted callback. + MoqtSession* session = + absl::down_cast<MoqtSession*>(weak_session.GetIfAvailable()); + if (session == nullptr) { + return; + } + session->subscribed_track_names_.erase(track_name); + }, + [weak_session = GetWeakPtr()](MoqtError code, absl::string_view reason) { + MoqtSessionInterface* session = weak_session.GetIfAvailable(); + if (session == nullptr) { + return; + } + session->Error(code, reason); + }, + std::move(response_callback)); + auto publish_state = std::make_unique<SubscriptionPublisher>( + framer_, publisher, stream_visitor.get(), next_request_id_, + next_local_track_alias_, parameters, this, nullptr, callbacks_.clock, + trace_recorder_, true); + SubscriptionPublisher* publisher_ptr = publish_state.get(); + stream_visitor->SetPublisher(std::move(publish_state)); + webtransport::Stream* stream = session_->OpenOutgoingBidirectionalStream(); + MoqtPublishPublisherStream* stream_visitor_ptr = stream_visitor.get(); + stream->SetVisitor(std::move(stream_visitor)); + stream_visitor_ptr->BindStream(stream); + next_request_id_ += 2; + ++next_local_track_alias_; + publisher->AddObjectListener(publisher_ptr); + return true; +} + bool MoqtSession::Fetch(const FullTrackName& name, FetchResponseCallback callback, Location start, uint64_t end_group, std::optional<uint64_t> end_object, @@ -683,6 +720,7 @@ } auto it = published_subscriptions_.find(request_id); if (it == published_subscriptions_.end()) { + // If a PUBLISH, we will end up here. return; } subscribed_track_names_.erase(it->second->publisher().GetTrackName()); @@ -762,6 +800,35 @@ return it->second; } +SubscribeRemoteTrack::SubscribeCallbacks MoqtSession::GetSubscribeCallbacks() { + return { + [this](const FullTrackName& name) { // query_name_ + return RemoteTrackByName(name); + }, + [this](const FullTrackName& name, SubscribeRemoteTrack* track) { + // register_name_ + subscribe_by_name_[name] = track; + }, + [this](const uint64_t alias, SubscribeRemoteTrack* track) { + // register_alias_ + return subscribe_by_alias_.try_emplace(alias, track).second; + }, + [weaksession = GetWeakPtr()](const FullTrackName& name, + std::optional<uint64_t> alias) { + // unregister_ + MoqtSession* session = + absl::down_cast<MoqtSession*>(weaksession.GetIfAvailable()); + if (session == nullptr) { + return; + } + session->subscribe_by_name_.erase(name); + if (alias.has_value()) { + session->subscribe_by_alias_.erase(*alias); + } + }, + }; +} + void MoqtSession::OnCanCreateNewOutgoingUnidirectionalStream() { while (!subscriptions_with_queued_streams_.empty() && session_->CanOpenNextOutgoingUnidirectionalStream()) { @@ -870,6 +937,23 @@ temp_stream->OnCanRead(); break; } + case MoqtMessageType::kPublish: { + auto publish_stream = std::make_unique<MoqtPublishSubscriberStream>( + &session_->framer_, session_->ControlMessageParser(), + session_->callbacks_.clock, session_->alarm_factory(), + [session = session_](MoqtError code, absl::string_view reason) { + session->Error(code, reason); + }, + &session_->callbacks_.incoming_publish_callback, + session_->GetSubscribeCallbacks()); + publish_stream->BindStream(std::move(parser_)); + MoqtPublishSubscriberStream* temp_stream = publish_stream.get(); + stream_->SetVisitor(std::move(publish_stream)); + // The UnknownBidiStream object is deleted; no class access after this + // point. + temp_stream->OnCanRead(); + break; + } default: session_->Error(MoqtError::kProtocolViolation, "Unexpected message type received to start bidi stream"); @@ -959,7 +1043,7 @@ auto subscription = std::make_unique<SubscriptionPublisher>( session_->framer_, track_publisher, this, message.request_id, session_->next_local_track_alias_++, message.parameters, session_, - monitoring, session_->callbacks_.clock, session_->trace_recorder_); + monitoring, session_->callbacks_.clock, session_->trace_recorder_, false); SubscriptionPublisher* subscription_ptr = subscription.get(); auto [it, success] = session_->published_subscriptions_.emplace( message.request_id, std::move(subscription)); @@ -996,9 +1080,8 @@ } SubscribeRemoteTrack* subscribe = absl::down_cast<SubscribeRemoteTrack*>(track); - if (!subscribe->set_track_alias(message.track_alias)) { - // A duplicate track alias could destroy the session. + OnFatalError(absl::AlreadyExistsError("")); return absl::OkStatus(); } subscribe->OnObjectOrOk(
diff --git a/quiche/quic/moqt/moqt_session.h b/quiche/quic/moqt/moqt_session.h index 0e42ce2..81eb53c 100644 --- a/quiche/quic/moqt/moqt_session.h +++ b/quiche/quic/moqt/moqt_session.h
@@ -92,6 +92,10 @@ const MessageParameters& parameters, MoqtResponseCallback response_callback) override; void Unsubscribe(const FullTrackName& name) override; + bool Publish(std::shared_ptr<MoqtTrackPublisher> absl_nonnull publisher, + const MessageParameters& parameters, + const TrackExtensions& extensions, + MoqtResponseCallback response_callback) override; bool Fetch(const FullTrackName& name, FetchResponseCallback callback, Location start, uint64_t end_group, std::optional<uint64_t> end_object, @@ -251,7 +255,11 @@ } }), session_(session), - weak_ptr_factory_(this) {} + weak_ptr_factory_(this) { + this->set_control_stream(); + } + // TODO(martinduke): Remove constructor body once SUBSCRIBE moves to a bidi + // stream. void OnStreamBound() override; absl::Status OnRawControlMessage( @@ -422,6 +430,8 @@ RemoteTrack* RemoteTrackById(uint64_t request_id); SubscribeRemoteTrack* RemoteTrackByName(const FullTrackName& name); + SubscribeRemoteTrack::SubscribeCallbacks GetSubscribeCallbacks(); + // Checks that a subscribe ID from a SUBSCRIBE or FETCH is valid, and throws // a session error if is not. bool ValidateRequestId(uint64_t request_id);
diff --git a/quiche/quic/moqt/moqt_session_callbacks.h b/quiche/quic/moqt/moqt_session_callbacks.h index 46a0dde..5b2070b 100644 --- a/quiche/quic/moqt/moqt_session_callbacks.h +++ b/quiche/quic/moqt/moqt_session_callbacks.h
@@ -5,22 +5,65 @@ #ifndef QUICHE_QUIC_MOQT_MOQT_SESSION_CALLBACKS_H_ #define QUICHE_QUIC_MOQT_MOQT_SESSION_CALLBACKS_H_ +#include <cstdint> #include <memory> #include <optional> #include <utility> +#include <variant> #include "absl/strings/string_view.h" #include "quiche/quic/core/quic_clock.h" #include "quiche/quic/core/quic_default_clock.h" +#include "quiche/quic/core/quic_time.h" #include "quiche/quic/moqt/moqt_error.h" #include "quiche/quic/moqt/moqt_fetch_task.h" #include "quiche/quic/moqt/moqt_key_value_pair.h" #include "quiche/quic/moqt/moqt_names.h" +#include "quiche/quic/moqt/moqt_object.h" #include "quiche/quic/moqt/moqt_types.h" #include "quiche/common/quiche_callbacks.h" namespace moqt { +using MoqtObjectAckFunction = + quiche::MultiUseCallback<void(uint64_t group_id, uint64_t object_id, + quic::QuicTimeDelta delta_from_deadline)>; + +struct SubscribeOkData { + MessageParameters parameters; + TrackExtensions extensions; +}; + +class SubscribeVisitor { + public: + virtual ~SubscribeVisitor() = default; + // Called when the session receives a response to the SUBSCRIBE. + virtual void OnReply( + const FullTrackName& full_track_name, + std::variant<SubscribeOkData, MoqtRequestErrorInfo> response) = 0; + // Called when the subscription process is far enough that it is possible to + // send OBJECT_ACK messages; provides a callback to do so. The callback is + // valid for as long as the session is valid. + virtual void OnCanAckObjects(MoqtObjectAckFunction ack_function) = 0; + // Called when an object fragment (or an entire object) is received. + virtual void OnObjectFragment(const FullTrackName& full_track_name, + const PublishedObjectMetadata& metadata, + absl::string_view object, uint64_t offset) = 0; + // Called when the subscription state goes away, regardless of whether or not + // there was a PUBLISH_DONE message. + virtual void OnPublishDone(FullTrackName full_track_name) = 0; + // Called when the track is malformed per Section 2.5 of + // draft-ietf-moqt-moq-transport-12. If the application is a relay, it MUST + // terminate downstream delivery of the track. + virtual void OnMalformedTrack(const FullTrackName& full_track_name) = 0; + + // End user applications might not care about stream state, but relays will. + virtual void OnStreamFin(const FullTrackName& full_track_name, + DataStreamIndex stream) = 0; + virtual void OnStreamReset(const FullTrackName& full_track_name, + DataStreamIndex stream) = 0; +}; + // Called when the SETUP message from the peer is received. using MoqtSessionEstablishedCallback = quiche::SingleUseCallback<void()>; @@ -35,6 +78,15 @@ // Called from the session destructor. using MoqtSessionDeletedCallback = quiche::SingleUseCallback<void()>; +// Called when a PUBLISH message is received from the peer. Returns a visitor +// for the subscription. If the returned visitor is nullptr, the session will +// immediately reject the PUBLISH. Otherwise, it will deliver objects for the +// track until either MoqtResponseCallback returns with an error or the +// application calls Unsubscribe. +using MoqtIncomingPublishCallback = quiche::MultiUseCallback<SubscribeVisitor*( + const FullTrackName&, const MessageParameters&, const TrackExtensions&, + 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 @@ -77,6 +129,12 @@ return nullptr; } +inline SubscribeVisitor* DefaultIncomingPublishCallback( + const FullTrackName&, const MessageParameters&, const TrackExtensions&, + MoqtResponseCallback) { + return nullptr; +} + // Callbacks for session-level events. struct MoqtSessionCallbacks { MoqtSessionEstablishedCallback session_established_callback = +[] {}; @@ -90,6 +148,8 @@ DefaultIncomingPublishNamespaceCallback; MoqtIncomingSubscribeNamespaceCallback incoming_subscribe_namespace_callback = DefaultIncomingSubscribeNamespaceCallback; + MoqtIncomingPublishCallback incoming_publish_callback = + DefaultIncomingPublishCallback; const quic::QuicClock* clock = quic::QuicDefaultClock::Get(); };
diff --git a/quiche/quic/moqt/moqt_session_interface.h b/quiche/quic/moqt/moqt_session_interface.h index 7b34a1d..e1455a6 100644 --- a/quiche/quic/moqt/moqt_session_interface.h +++ b/quiche/quic/moqt/moqt_session_interface.h
@@ -10,17 +10,16 @@ #include <optional> #include <string> #include <utility> -#include <variant> #include <vector> +#include "absl/base/nullability.h" #include "absl/strings/string_view.h" -#include "quiche/quic/core/quic_time.h" #include "quiche/quic/core/quic_types.h" #include "quiche/quic/moqt/moqt_error.h" #include "quiche/quic/moqt/moqt_fetch_task.h" #include "quiche/quic/moqt/moqt_key_value_pair.h" #include "quiche/quic/moqt/moqt_names.h" -#include "quiche/quic/moqt/moqt_object.h" +#include "quiche/quic/moqt/moqt_publisher.h" #include "quiche/quic/moqt/moqt_session_callbacks.h" #include "quiche/quic/moqt/moqt_types.h" #include "quiche/common/platform/api/quiche_export.h" @@ -78,44 +77,6 @@ void ToSetupParameters(SetupParameters& out) const; }; -using MoqtObjectAckFunction = - quiche::MultiUseCallback<void(uint64_t group_id, uint64_t object_id, - quic::QuicTimeDelta delta_from_deadline)>; - -struct SubscribeOkData { - MessageParameters parameters; - TrackExtensions extensions; -}; - -class SubscribeVisitor { - public: - virtual ~SubscribeVisitor() = default; - // Called when the session receives a response to the SUBSCRIBE. - virtual void OnReply( - const FullTrackName& full_track_name, - std::variant<SubscribeOkData, MoqtRequestErrorInfo> response) = 0; - // Called when the subscription process is far enough that it is possible to - // send OBJECT_ACK messages; provides a callback to do so. The callback is - // valid for as long as the session is valid. - virtual void OnCanAckObjects(MoqtObjectAckFunction ack_function) = 0; - // Called when an object fragment (or an entire object) is received. - virtual void OnObjectFragment(const FullTrackName& full_track_name, - const PublishedObjectMetadata& metadata, - absl::string_view object, uint64_t offset) = 0; - // Called when the subscription state goes away, regardless of whether or not - // there was a PUBLISH_DONE message. - virtual void OnPublishDone(FullTrackName full_track_name) = 0; - // Called when the track is malformed per Section 2.5 of - // draft-ietf-moqt-moq-transport-12. If the application is a relay, it MUST - // terminate downstream delivery of the track. - virtual void OnMalformedTrack(const FullTrackName& full_track_name) = 0; - - // End user applications might not care about stream state, but relays will. - virtual void OnStreamFin(const FullTrackName& full_track_name, - DataStreamIndex stream) = 0; - virtual void OnStreamReset(const FullTrackName& full_track_name, - DataStreamIndex stream) = 0; -}; // MoqtSession calls this when a FETCH_OK or REQUEST_ERROR is received. The // destination of the callback owns |fetch_task| and MoqtSession will react @@ -148,6 +109,14 @@ // subscription. Returns false if the subscription is not found. virtual void Unsubscribe(const FullTrackName& name) = 0; + // Returns false if the PUBLISH cannot be sent due stream flow control + // limitations (which spawns PUBLISH_BLOCKED in namespace streams). Any other + // failure will be covered by |response_callback|. + virtual bool Publish( + std::shared_ptr<MoqtTrackPublisher> absl_nonnull publisher, + const MessageParameters& parameters, const TrackExtensions& extensions, + MoqtResponseCallback response_callback) = 0; + // Sends a FETCH for a pre-specified object range. Once a FETCH_OK or a // FETCH_ERROR is received, `callback` is called with a MoqtFetchTask that can // be used to process the FETCH further. To cancel a FETCH, simply destroy
diff --git a/quiche/quic/moqt/moqt_session_test.cc b/quiche/quic/moqt/moqt_session_test.cc index 15eb900..2b6bbb9 100644 --- a/quiche/quic/moqt/moqt_session_test.cc +++ b/quiche/quic/moqt/moqt_session_test.cc
@@ -7,6 +7,7 @@ #include <algorithm> #include <cstdint> #include <cstring> +#include <functional> #include <memory> #include <optional> #include <queue> @@ -23,6 +24,7 @@ #include "quiche/quic/core/quic_data_reader.h" #include "quiche/quic/core/quic_time.h" #include "quiche/quic/core/quic_types.h" +#include "quiche/quic/moqt/moqt_bidi_stream.h" #include "quiche/quic/moqt/moqt_error.h" #include "quiche/quic/moqt/moqt_fetch_task.h" #include "quiche/quic/moqt/moqt_framer.h" @@ -1156,8 +1158,6 @@ } TEST_F(MoqtSessionTest, SubscribeOkWithBadTrackAlias) { - // Create open subscription. We cannot use CreateRemoteTrack because that - // skips the code that sets the track alias callbacks. webtransport::test::MockStream mock_control_stream; std::unique_ptr<MoqtBidiStreamTestWrapper> control_stream = MoqtSessionPeer::CreateControlStream(&session_, &mock_control_stream); @@ -2781,6 +2781,203 @@ listener->OnNewObjectAvailable(Location(0, 1), 1, 0x80); } +TEST_F(MoqtSessionTest, PublishSuccess) { + CreateTrackPublisher(); + std::shared_ptr<MoqtTrackPublisher> track_publisher = + publisher_.GetTrack(kDefaultTrackName()); + webtransport::test::MockStream mock_publish_stream; + EXPECT_CALL(mock_session_, CanOpenNextOutgoingBidirectionalStream()) + .WillOnce(Return(true)); + EXPECT_CALL(mock_session_, OpenOutgoingBidirectionalStream()) + .WillOnce(Return(&mock_publish_stream)); + + std::unique_ptr<webtransport::StreamVisitor> publish_stream_visitor; + EXPECT_CALL(mock_publish_stream, SetVisitor) + .WillOnce([&](std::unique_ptr<webtransport::StreamVisitor> visitor) { + publish_stream_visitor = std::move(visitor); + }); + EXPECT_CALL(mock_publish_stream, CanWrite).WillRepeatedly(Return(true)); + + // Verify PUBLISH message is sent on the publish stream. + EXPECT_CALL(mock_publish_stream, + Writev(ControlMessageOfType(MoqtMessageType::kPublish), _)) + .WillOnce(Return(absl::OkStatus())); + + std::optional<std::variant<MessageParameters, MoqtRequestErrorInfo>> response; + ASSERT_TRUE(session_.Publish( + track_publisher, MessageParameters(), TrackExtensions(), + [&](std::variant<MessageParameters, MoqtRequestErrorInfo> resp) { + response = resp; + })); + ASSERT_NE(publish_stream_visitor, nullptr); + + std::unique_ptr<MoqtBidiStreamBase> bidi_stream( + absl::down_cast<MoqtBidiStreamBase*>(publish_stream_visitor.release())); + MoqtBidiStreamTestWrapper wrapper(std::move(bidi_stream)); + + MoqtRequestOk request_ok; + request_ok.request_id = 0; + request_ok.parameters.delivery_timeout = quic::QuicTimeDelta::FromSeconds(2); + + wrapper.ReceiveMessage(request_ok); + + ASSERT_TRUE(response.has_value()); + EXPECT_TRUE(std::holds_alternative<MessageParameters>(*response)); + EXPECT_EQ(std::get<MessageParameters>(*response).delivery_timeout, + quic::QuicTimeDelta::FromSeconds(2)); +} + +TEST_F(MoqtSessionTest, PublishCannotOpenStream) { + CreateTrackPublisher(); + std::shared_ptr<MoqtTrackPublisher> track_publisher = + publisher_.GetTrack(kDefaultTrackName()); + EXPECT_CALL(mock_session_, CanOpenNextOutgoingBidirectionalStream()) + .WillOnce(Return(false)); + EXPECT_FALSE(session_.Publish( + track_publisher, MessageParameters(), TrackExtensions(), + [&](std::variant<MessageParameters, MoqtRequestErrorInfo> response) {})); +} + +TEST_F(MoqtSessionTest, PublishAfterGoaway) { + std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input = + MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_); + MoqtGoAway goaway; + goaway.new_session_uri = ""; + stream_input->ReceiveMessage(goaway); + CreateTrackPublisher(); + std::shared_ptr<MoqtTrackPublisher> track_publisher = + publisher_.GetTrack(kDefaultTrackName()); + EXPECT_FALSE(session_.Publish( + track_publisher, MessageParameters(), TrackExtensions(), + [&](std::variant<MessageParameters, MoqtRequestErrorInfo>) {})); +} + +TEST_F(MoqtSessionTest, IncomingPublishAbortsPendingSubscribe) { + // 1. Start a pending SUBSCRIBE. + std::unique_ptr<MoqtBidiStreamTestWrapper> control_stream = + MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_); + EXPECT_CALL(mock_stream_, + Writev(ControlMessageOfType(MoqtMessageType::kSubscribe), _)); + MessageParameters parameters(SubscribeForTest()); + session_.Subscribe(kDefaultTrackName(), &remote_track_visitor_, parameters); + + // Configure incoming publish callback to accept. Return nullptr because this + // should never be called. + bool incoming_publish_callback_called = false; + session_.callbacks().incoming_publish_callback = + [&](const FullTrackName&, const MessageParameters&, + const TrackExtensions&, MoqtResponseCallback callback) { + incoming_publish_callback_called = true; + return nullptr; + }; + + // Prepare PUBLISH message. + MoqtPublish publish; + publish.request_id = 0; // Matches pending SUBSCRIBE request ID + publish.full_track_name = kDefaultTrackName(); + publish.track_alias = 10; + publish.parameters.delivery_timeout = quic::QuicTimeDelta::FromSeconds(5); + + MoqtFramer peer_framer(true, quic::Perspective::IS_SERVER); + quiche::QuicheBuffer serialized_buffer = + peer_framer.SerializePublish(publish); + std::string serialized(serialized_buffer.data(), serialized_buffer.size()); + + // Setup mock_publish_stream. + webtransport::test::MockStream mock_publish_stream; + EXPECT_CALL(mock_session_, AcceptIncomingBidirectionalStream()) + .WillOnce(Return(&mock_publish_stream)) + .WillOnce(Return(nullptr)); + + // Setup mock_publish_stream to return serialized data. + size_t data_read = 0; + auto peek_lambda = [&data_read, + &serialized]() -> webtransport::Stream::PeekResult { + webtransport::Stream::PeekResult result; + result.peeked_data = absl::string_view(serialized.data() + data_read, + serialized.size() - data_read); + result.fin_next = (data_read == serialized.size()); + result.all_data_received = false; + return result; + }; + std::function<webtransport::Stream::PeekResult()> peek_action = peek_lambda; + + auto readable_bytes_lambda = [&data_read, &serialized]() -> size_t { + if (data_read >= serialized.size()) { + return 0; + } + return serialized.size() - data_read; + }; + std::function<size_t()> readable_bytes_action = readable_bytes_lambda; + + auto read_lambda = + [&data_read, &serialized]( + absl::Span<char> bytes_to_read) -> webtransport::Stream::ReadResult { + size_t read_size = + std::min(bytes_to_read.size(), serialized.size() - data_read); + memcpy(bytes_to_read.data(), serialized.data() + data_read, read_size); + data_read += read_size; + webtransport::Stream::ReadResult result; + result.bytes_read = read_size; + result.fin = (data_read == serialized.size()); + return result; + }; + std::function<webtransport::Stream::ReadResult(absl::Span<char>)> + read_action = read_lambda; + + auto skip_lambda = [&data_read, &serialized](size_t bytes) -> bool { + data_read += bytes; + return data_read == serialized.size(); + }; + std::function<bool(size_t)> skip_action = skip_lambda; + + EXPECT_CALL(mock_publish_stream, PeekNextReadableRegion()) + .WillRepeatedly(peek_action); + EXPECT_CALL(mock_publish_stream, ReadableBytes()) + .WillRepeatedly(readable_bytes_action); + EXPECT_CALL(mock_publish_stream, Read(testing::An<absl::Span<char>>())) + .WillRepeatedly(read_action); + EXPECT_CALL(mock_publish_stream, SkipBytes).WillRepeatedly(skip_action); + + // Capture SetVisitor calls and mock visitor(). + std::unique_ptr<webtransport::StreamVisitor> unknown_bidi_stream_visitor; + std::unique_ptr<webtransport::StreamVisitor> upgraded_visitor; + + { + testing::InSequence seq; + EXPECT_CALL(mock_publish_stream, SetVisitor) + .WillOnce([&](std::unique_ptr<webtransport::StreamVisitor> visitor) { + unknown_bidi_stream_visitor = std::move(visitor); + }) + .WillOnce([&](std::unique_ptr<webtransport::StreamVisitor> visitor) { + upgraded_visitor = std::move(visitor); + }); + } + + EXPECT_CALL(mock_publish_stream, visitor()) + .WillRepeatedly([&]() -> webtransport::StreamVisitor* { + return upgraded_visitor ? upgraded_visitor.get() + : unknown_bidi_stream_visitor.get(); + }); + + // The RemoteTrackVisitor is being reused, not destroyed. + EXPECT_CALL(remote_track_visitor_, + OnReply(kDefaultTrackName(), + testing::VariantWith<SubscribeOkData>( + testing::Field(&SubscribeOkData::parameters, + testing::Eq(publish.parameters))))); + EXPECT_CALL(remote_track_visitor_, OnPublishDone(kDefaultTrackName())) + .Times(0); + EXPECT_FALSE(incoming_publish_callback_called); + + // Trigger the read by making incoming stream available. + session_.OnIncomingBidirectionalStreamAvailable(); + + // Verify it was aborted immediately (not at teardown). + EXPECT_TRUE( + testing::Mock::VerifyAndClearExpectations(&remote_track_visitor_)); +} + } // namespace test } // namespace moqt
diff --git a/quiche/quic/moqt/moqt_subscription.cc b/quiche/quic/moqt/moqt_subscription.cc index c62d3e0..b3d5416 100644 --- a/quiche/quic/moqt/moqt_subscription.cc +++ b/quiche/quic/moqt/moqt_subscription.cc
@@ -41,11 +41,12 @@ SessionToPublisherInterface* absl_nonnull visitor, MoqtPublishingMonitorInterface* monitoring_interface, const quic::QuicClock* absl_nonnull clock, - MoqtTraceRecorder& trace_recorder) + MoqtTraceRecorder& trace_recorder, bool is_publish) : track_publisher_(track_publisher), bidi_stream_(bidi_stream), visitor_(visitor), request_id_(request_id), + established_(is_publish), track_alias_(track_alias), framer_(framer), trace_recorder_(trace_recorder), @@ -62,7 +63,9 @@ } SubscriptionPublisher::~SubscriptionPublisher() { - track_publisher_->RemoveObjectListener(this); + if (track_publisher_ != nullptr) { + track_publisher_->RemoveObjectListener(this); + } // Reset all streams. for (const webtransport::StreamId stream_id : stream_map_.GetAllStreams()) { webtransport::Stream* stream = GetStreamById(stream_id); @@ -108,7 +111,9 @@ } void SubscriptionPublisher::OnSubscribeAccepted() { - QUICHE_DCHECK(!established_); + if (established_) { + return; // It's a PUBLISH. + } established_ = true; parameters_.largest_object = track_publisher_->largest_location(); if (parameters_.subscription_filter.has_value()) { @@ -137,9 +142,11 @@ } void SubscriptionPublisher::OnSubscribeRejected(MoqtRequestErrorInfo info) { - bidi_stream_->CheckStatus(bidi_stream_->SendRequestError(request_id_, info)); - visitor_->PublishIsDone(request_id_); - // No class access below this line! + bidi_stream_->CheckStatus(bidi_stream_->SendRequestError(request_id_, info, + /*fin=*/true)); + if (bidi_stream_->is_control_stream()) { + visitor_->PublishIsDone(request_id_); + } } void SubscriptionPublisher::OnNewObjectAvailable( @@ -238,7 +245,7 @@ } void SubscriptionPublisher::OnTrackPublisherGone() { - PublishIsDone(request_id_, PublishDoneCode::kGoingAway, "Publisher is gone"); + PublishIsDone(PublishDoneCode::kGoingAway, "Publisher is gone"); } // TODO(martinduke): Revise to check if the last object has been delivered. @@ -296,7 +303,7 @@ stream_map_.GetStreamsForGroup(group_id); if (delivery_timeout().IsInfinite() && largest_sent_.has_value() && largest_sent_->group <= group_id) { - PublishIsDone(request_id_, PublishDoneCode::kTooFarBehind, ""); + PublishIsDone(PublishDoneCode::kTooFarBehind, ""); // No class access below this line! return; } @@ -379,23 +386,29 @@ return new_stream; } -void SubscriptionPublisher::PublishIsDone(uint64_t request_id, - PublishDoneCode code, +void SubscriptionPublisher::PublishIsDone(PublishDoneCode code, absl::string_view error_reason) { MoqtPublishDone publish_done; - publish_done.request_id = request_id; + publish_done.request_id = request_id_; publish_done.status_code = code; publish_done.stream_count = streams_opened_; publish_done.error_reason = error_reason; - // TODO(martinduke): It is technically correct, but not good, to simply - // reset all the streams in order to send PUBLISH_DONE. It's better to wait - // until streams FIN naturally, where possible. + // TODO(martinduke): It is technically correct, but not good, to reset all the + // streams in order to send PUBLISH_DONE. It's better to wait until streams + // FIN naturally, where possible. QUICHE_DLOG(INFO) << "Sending PUBLISH_DONE message for " << track_publisher_->GetTrackName(); + // TODO(martinduke): For SUBSCRIBE, no FIN because it's the control stream. bidi_stream_->SendOrBufferMessageOrFatal( - framer_.SerializePublishDone(publish_done)); - visitor_->PublishIsDone(request_id_); - // No class access below this line! + framer_.SerializePublishDone(publish_done), + /*fin=*/!bidi_stream_->is_control_stream()); + if (bidi_stream_->is_control_stream()) { + visitor_->PublishIsDone(request_id_); + } else { + // Only detach immediately for PUBLISH flow. + track_publisher_->RemoveObjectListener(this); + track_publisher_ = nullptr; + } } void SubscriptionPublisher::OnDataStreamDestroyed(
diff --git a/quiche/quic/moqt/moqt_subscription.h b/quiche/quic/moqt/moqt_subscription.h index 38a314a..b058beb 100644 --- a/quiche/quic/moqt/moqt_subscription.h +++ b/quiche/quic/moqt/moqt_subscription.h
@@ -83,6 +83,7 @@ virtual quic::QuicAlarmFactory* alarm_factory() = 0; // Destroy any state associated with the subscription. It is OK destroy // SubscriptionPublisher in this method. + // TODO(martinduke): Delete once SUBSCRIBE is on the bidi stream. virtual void PublishIsDone(uint64_t request_id) = 0; // Returns nullptr if MoqtSession is closing. virtual webtransport::Session* session() = 0; @@ -101,7 +102,7 @@ SessionToPublisherInterface* absl_nonnull visitor, MoqtPublishingMonitorInterface* monitoring_interface, const quic::QuicClock* absl_nonnull clock, - MoqtTraceRecorder& trace_recorder); + MoqtTraceRecorder& trace_recorder, bool is_publish); ~SubscriptionPublisher(); SubscriptionPublisher(const SubscriptionPublisher&) = delete; @@ -210,8 +211,7 @@ webtransport::Stream* absl_nullable OpenDataStream( const NewDataStreamParameters& parameters); - void PublishIsDone(uint64_t request_id, PublishDoneCode code, - absl::string_view reason); + void PublishIsDone(PublishDoneCode code, absl::string_view reason); MoqtPriority subscriber_priority() const { return parameters_.subscriber_priority.value_or(kDefaultSubscriberPriority); @@ -228,7 +228,7 @@ SessionToPublisherInterface* absl_nonnull visitor_; uint64_t request_id_; // Subscription is in the ESTABLISHED state. - bool established_ = false; + bool established_; const uint64_t track_alias_; MoqtFramer framer_; MoqtTraceRecorder& trace_recorder_;
diff --git a/quiche/quic/moqt/moqt_subscription_test.cc b/quiche/quic/moqt/moqt_subscription_test.cc index 37e1da9..791c05f 100644 --- a/quiche/quic/moqt/moqt_subscription_test.cc +++ b/quiche/quic/moqt/moqt_subscription_test.cc
@@ -91,7 +91,9 @@ SessionErrorCallback session_error_callback) : MoqtBidiStreamBase(framer, message_parser, std::move(stream_deleted_callback), - std::move(session_error_callback)) {} + std::move(session_error_callback)) { + set_control_stream(); // TODO(martinduke): Delete + } ~TestMoqtBidiStream() override = default; void OnStreamBound() override {}; absl::Status OnRawControlMessage( @@ -136,7 +138,7 @@ publisher_ = std::make_unique<SubscriptionPublisher>( framer_, track_publisher_, &bidi_stream_, kRequestId, kTrackAlias, parameters_, &visitor_, &monitoring_interface_, &mock_clock_, - trace_recorder_); + trace_recorder_, /*is_publish=*/false); ON_CALL(visitor_, alternate_delivery_timeout).WillByDefault(Return(false)); ON_CALL(webtrans_, GetStreamById(kStreamId)) .WillByDefault(Return(&mock_uni_stream_));
diff --git a/quiche/quic/moqt/moqt_track.cc b/quiche/quic/moqt/moqt_track.cc index cea09cd..4530cd9 100644 --- a/quiche/quic/moqt/moqt_track.cc +++ b/quiche/quic/moqt/moqt_track.cc
@@ -23,8 +23,7 @@ #include "quiche/quic/moqt/moqt_key_value_pair.h" #include "quiche/quic/moqt/moqt_messages.h" #include "quiche/quic/moqt/moqt_object.h" -#include "quiche/quic/moqt/moqt_priority.h" -#include "quiche/quic/moqt/moqt_session_interface.h" +#include "quiche/quic/moqt/moqt_session_callbacks.h" #include "quiche/quic/moqt/moqt_types.h" #include "quiche/common/platform/api/quiche_bug_tracker.h" #include "quiche/common/quiche_mem_slice.h" @@ -54,10 +53,12 @@ if (publish_done_alarm_ != nullptr) { publish_done_alarm_->PermanentCancel(); } - if (register_track_alias_callback_ && track_alias_.has_value()) { - register_track_alias_callback_(*track_alias_, nullptr); + if (callbacks_.unregister) { + std::move(callbacks_.unregister)(full_track_name(), track_alias_); } - visitor_->OnPublishDone(full_track_name()); + if (visitor_ != nullptr) { + visitor_->OnPublishDone(full_track_name()); + } } void SubscribeRemoteTrack::OnObjectOrOk(const SubscribeOkData& data) {
diff --git a/quiche/quic/moqt/moqt_track.h b/quiche/quic/moqt/moqt_track.h index 9451b1f..1d85bbb 100644 --- a/quiche/quic/moqt/moqt_track.h +++ b/quiche/quic/moqt/moqt_track.h
@@ -11,6 +11,7 @@ #include <memory> #include <optional> #include <utility> +#include <variant> #include "absl/status/status.h" #include "absl/strings/string_view.h" @@ -24,6 +25,7 @@ #include "quiche/quic/moqt/moqt_names.h" #include "quiche/quic/moqt/moqt_object.h" #include "quiche/quic/moqt/moqt_priority.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/common/quiche_callbacks.h" @@ -95,21 +97,64 @@ // A track on the peer to which the session has subscribed. class SubscribeRemoteTrack : public RemoteTrack { public: - // If the second argument is null, delete the registration. Returns false if - // it fails due to a duplicate track alias, destroying the session. + struct SubscribeCallbacks { + quiche::SingleUseCallback<SubscribeRemoteTrack*(const FullTrackName&)> + query_name; + quiche::SingleUseCallback<void(const FullTrackName&, SubscribeRemoteTrack*)> + register_name; + quiche::SingleUseCallback<bool(uint64_t, SubscribeRemoteTrack*)> + register_alias; + quiche::SingleUseCallback<void(const FullTrackName&, + std::optional<uint64_t>)> + unregister; + }; + // Tells the session about changes to a track's subscription status. + // If SubscribeRemoteTrack* is null, the subscription is gone and the callback + // will always return true. + // If non-null, try to add the track. If the name is new, return true. If + // ready present but for a pending SUBSCRIBE, return the visitor for that + // SUBSCRIBE. Otherwise, return false. + using RegisterNameCallback = + quiche::MultiUseCallback<std::variant<bool, SubscribeVisitor*>( + const FullTrackName&, SubscribeRemoteTrack*)>; + // When SubscribeRemoteTrack* is non-null, this callback informs the session + // of the track alias after the receipt of SUBSCRIBE_OK or PUBLISH. + // + // If the second argument is null, it means the subscription to the track + // alias has ended, and always returns absl::OkStatus(). + // + // Returns true if the operation was successful, It can only fail on + // registration because there is already a track with that alias. using RegisterTrackAliasCallback = quiche::MultiUseCallback<bool(uint64_t, SubscribeRemoteTrack*)>; - // We're using BidiStreamDeletedCallback here because this will move to a - // bidi stream. SubscribeRemoteTrack(const MoqtSubscribe& subscribe, SubscribeVisitor* visitor, BidiStreamDeletedCallback callback, - RegisterTrackAliasCallback register_track_alias_callback) + SubscribeCallbacks callbacks) : RemoteTrack(subscribe.full_track_name, subscribe.request_id, subscribe.parameters, std::move(callback)), visitor_(visitor), - register_track_alias_callback_( - std::move(register_track_alias_callback)) {} + callbacks_(std::move(callbacks)) { + if (callbacks_.register_name) { + std::move(callbacks_.register_name)(full_track_name(), this); + } + } + + SubscribeRemoteTrack(const MoqtPublish& publish, SubscribeVisitor* visitor, + BidiStreamDeletedCallback callback, + SubscribeCallbacks callbacks) + : RemoteTrack(publish.full_track_name, publish.request_id, + publish.parameters, std::move(callback)), + visitor_(visitor), + callbacks_(std::move(callbacks)) { + OnObjectOrOk(); + visitor_->OnReply(publish.full_track_name, + SubscribeOkData(publish.parameters, publish.extensions)); + if (callbacks_.register_name) { + std::move(callbacks_.register_name)(full_track_name(), this); + callbacks_.register_name = nullptr; + } + } ~SubscribeRemoteTrack() override; void OnObjectOrOk(const SubscribeOkData& data); @@ -121,8 +166,8 @@ // destroyed. [[nodiscard]] bool set_track_alias(uint64_t track_alias) { track_alias_.emplace(track_alias); - if (register_track_alias_callback_) { - return register_track_alias_callback_(track_alias, this); + if (callbacks_.register_alias) { + return std::move(callbacks_.register_alias)(track_alias, this); } return true; } @@ -154,6 +199,11 @@ } SubscribeVisitor* visitor() const { return visitor_; } + SubscribeVisitor* ReleaseVisitor() { + SubscribeVisitor* temp = visitor_; + visitor_ = nullptr; + return temp; + } private: friend class test::MoqtSessionPeer; @@ -188,7 +238,7 @@ int currently_open_streams_ = 0; // Every stream that has received FIN or RESET_STREAM. uint64_t streams_closed_ = 0; - RegisterTrackAliasCallback register_track_alias_callback_; + SubscribeCallbacks callbacks_; // Value assigned on PUBLISH_DONE. Can destroy subscription state if // streams_closed_ == total_streams_. std::optional<uint64_t> total_streams_;
diff --git a/quiche/quic/moqt/moqt_track_test.cc b/quiche/quic/moqt/moqt_track_test.cc index 9e15a7f..2e3b52f 100644 --- a/quiche/quic/moqt/moqt_track_test.cc +++ b/quiche/quic/moqt/moqt_track_test.cc
@@ -16,13 +16,13 @@ #include "quiche/quic/moqt/moqt_messages.h" #include "quiche/quic/moqt/moqt_names.h" #include "quiche/quic/moqt/moqt_object.h" -#include "quiche/quic/moqt/moqt_priority.h" #include "quiche/quic/moqt/moqt_types.h" #include "quiche/quic/moqt/test_tools/moqt_mock_visitor.h" #include "quiche/quic/platform/api/quic_test.h" #include "quiche/quic/test_tools/mock_clock.h" #include "quiche/quic/test_tools/quic_test_utils.h" #include "quiche/common/quiche_mem_slice.h" +#include "quiche/common/test_tools/quiche_test_utils.h" #include "quiche/web_transport/web_transport.h" namespace moqt { @@ -57,13 +57,18 @@ SubscribeRemoteTrackTest() : track_( subscribe_, &visitor_, [this]() { deleted_ = true; }, - [this](uint64_t, SubscribeRemoteTrack* track) { - alias_registered_ = (track != nullptr); - if (alias_registered_) { - EXPECT_EQ(track, &track_); - } - return true; - }) {} + SubscribeRemoteTrack::SubscribeCallbacks{ + /*query_name_=*/nullptr, + /*register_name_=*/nullptr, + /*register_alias_=*/ + [this](uint64_t, SubscribeRemoteTrack* track) { + alias_registered_ = (track != nullptr); + if (alias_registered_) { + EXPECT_EQ(track, &track_); + } + return true; + }, + /*unregister_=*/nullptr}) {} MockSubscribeRemoteTrackVisitor visitor_; MoqtSubscribe subscribe_ = {/*request_id=*/1, FullTrackName("foo", "bar"),
diff --git a/quiche/quic/moqt/moqt_uni_stream_test.cc b/quiche/quic/moqt/moqt_uni_stream_test.cc index d69c198..f56c5c7 100644 --- a/quiche/quic/moqt/moqt_uni_stream_test.cc +++ b/quiche/quic/moqt/moqt_uni_stream_test.cc
@@ -493,11 +493,16 @@ .WillRepeatedly(Return(false)); track_ = std::make_unique<SubscribeRemoteTrack>( subscribe_message_, &visitor_, []() {}, - [this](uint64_t alias, SubscribeRemoteTrack* track) -> bool { - alias_ = alias; - alias_track_ = track; - return true; - }); + SubscribeRemoteTrack::SubscribeCallbacks{ + /*query_name_=*/nullptr, + /*register_name_=*/nullptr, + /*register_alias_=*/ + [this](uint64_t alias, SubscribeRemoteTrack* track) { + alias_ = alias; + alias_track_ = track; + return true; + }, + /*unregister_=*/nullptr}); EXPECT_TRUE(track_->set_track_alias(2)); CreateStream(); }
diff --git a/quiche/quic/moqt/test_tools/mock_moqt_session.h b/quiche/quic/moqt/test_tools/mock_moqt_session.h index e3e8e09..c09c8c6 100644 --- a/quiche/quic/moqt/test_tools/mock_moqt_session.h +++ b/quiche/quic/moqt/test_tools/mock_moqt_session.h
@@ -13,12 +13,12 @@ #include "quiche/quic/moqt/moqt_error.h" #include "quiche/quic/moqt/moqt_fetch_task.h" #include "quiche/quic/moqt/moqt_key_value_pair.h" -#include "quiche/quic/moqt/moqt_messages.h" #include "quiche/quic/moqt/moqt_names.h" -#include "quiche/quic/moqt/moqt_priority.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/common/platform/api/quiche_test.h" +#include "quiche/common/quiche_callbacks.h" #include "quiche/common/quiche_weak_ptr.h" namespace moqt { @@ -38,6 +38,12 @@ MoqtResponseCallback), (override)); MOCK_METHOD(void, Unsubscribe, (const FullTrackName& name), (override)); + MOCK_METHOD(bool, Publish, + (std::shared_ptr<MoqtTrackPublisher> publisher, + const MessageParameters& parameters, + const TrackExtensions& extensions, + MoqtResponseCallback response_callback), + (override)); MOCK_METHOD(bool, Fetch, (const FullTrackName& name, FetchResponseCallback callback, Location start, uint64_t end_group,
diff --git a/quiche/quic/moqt/test_tools/moqt_simulator_harness.cc b/quiche/quic/moqt/test_tools/moqt_simulator_harness.cc index 59d29d8..c282858 100644 --- a/quiche/quic/moqt/test_tools/moqt_simulator_harness.cc +++ b/quiche/quic/moqt/test_tools/moqt_simulator_harness.cc
@@ -38,10 +38,9 @@ } MoqtSessionCallbacks CreateCallbacks(quic::simulator::Simulator* simulator) { - return MoqtSessionCallbacks( - +[] {}, +[](absl::string_view) {}, +[](absl::string_view) {}, +[] {}, - DefaultIncomingPublishNamespaceCallback, - DefaultIncomingSubscribeNamespaceCallback, simulator->GetClock()); + MoqtSessionCallbacks callbacks; + callbacks.clock = simulator->GetClock(); + return callbacks; } } // namespace