Move FETCH to a bidi stream as part of draft-18 update. Remove request_id from FETCH_OK and REQUEST_ERROR so that callback data structures are aligned with wire messages. FIN always follows REQUEST_ERROR. As the messages and external API structs are now the same, this deletes a lot of code, unfortunately with a bigger blast radius than hoped. FetchResponseCallback is no longer handled in the FetchTask at all. Not in production. PiperOrigin-RevId: 982748285
diff --git a/build/source_list.bzl b/build/source_list.bzl index da09290..4ddf6e3 100644 --- a/build/source_list.bzl +++ b/build/source_list.bzl
@@ -1597,6 +1597,7 @@ "quic/moqt/moqt_bitrate_adjuster.h", "quic/moqt/moqt_control_message_queue.h", "quic/moqt/moqt_error.h", + "quic/moqt/moqt_fetch_stream.h", "quic/moqt/moqt_fetch_task.h", "quic/moqt/moqt_framer.h", "quic/moqt/moqt_key_value_pair.h", @@ -1641,6 +1642,7 @@ "quic/moqt/moqt_bitrate_adjuster.cc", "quic/moqt/moqt_control_message_queue.cc", "quic/moqt/moqt_error.cc", + "quic/moqt/moqt_fetch_stream.cc", "quic/moqt/moqt_framer.cc", "quic/moqt/moqt_key_value_pair.cc", "quic/moqt/moqt_known_track_publisher.cc", @@ -1680,6 +1682,7 @@ "quic/moqt/moqt_bidi_stream_test.cc", "quic/moqt/moqt_bitrate_adjuster_test.cc", "quic/moqt/moqt_control_message_queue_test.cc", + "quic/moqt/moqt_fetch_stream_test.cc", "quic/moqt/moqt_framer_test.cc", "quic/moqt/moqt_integration_test.cc", "quic/moqt/moqt_key_value_pair_test.cc",
diff --git a/build/source_list.gni b/build/source_list.gni index fd66f95..ab2801d 100644 --- a/build/source_list.gni +++ b/build/source_list.gni
@@ -1602,6 +1602,7 @@ "src/quiche/quic/moqt/moqt_bitrate_adjuster.h", "src/quiche/quic/moqt/moqt_control_message_queue.h", "src/quiche/quic/moqt/moqt_error.h", + "src/quiche/quic/moqt/moqt_fetch_stream.h", "src/quiche/quic/moqt/moqt_fetch_task.h", "src/quiche/quic/moqt/moqt_framer.h", "src/quiche/quic/moqt/moqt_key_value_pair.h", @@ -1646,6 +1647,7 @@ "src/quiche/quic/moqt/moqt_bitrate_adjuster.cc", "src/quiche/quic/moqt/moqt_control_message_queue.cc", "src/quiche/quic/moqt/moqt_error.cc", + "src/quiche/quic/moqt/moqt_fetch_stream.cc", "src/quiche/quic/moqt/moqt_framer.cc", "src/quiche/quic/moqt/moqt_key_value_pair.cc", "src/quiche/quic/moqt/moqt_known_track_publisher.cc", @@ -1686,6 +1688,7 @@ "src/quiche/quic/moqt/moqt_bidi_stream_test.cc", "src/quiche/quic/moqt/moqt_bitrate_adjuster_test.cc", "src/quiche/quic/moqt/moqt_control_message_queue_test.cc", + "src/quiche/quic/moqt/moqt_fetch_stream_test.cc", "src/quiche/quic/moqt/moqt_framer_test.cc", "src/quiche/quic/moqt/moqt_integration_test.cc", "src/quiche/quic/moqt/moqt_key_value_pair_test.cc",
diff --git a/build/source_list.json b/build/source_list.json index 62889c0..7404e7a 100644 --- a/build/source_list.json +++ b/build/source_list.json
@@ -1601,6 +1601,7 @@ "quiche/quic/moqt/moqt_bitrate_adjuster.h", "quiche/quic/moqt/moqt_control_message_queue.h", "quiche/quic/moqt/moqt_error.h", + "quiche/quic/moqt/moqt_fetch_stream.h", "quiche/quic/moqt/moqt_fetch_task.h", "quiche/quic/moqt/moqt_framer.h", "quiche/quic/moqt/moqt_key_value_pair.h", @@ -1645,6 +1646,7 @@ "quiche/quic/moqt/moqt_bitrate_adjuster.cc", "quiche/quic/moqt/moqt_control_message_queue.cc", "quiche/quic/moqt/moqt_error.cc", + "quiche/quic/moqt/moqt_fetch_stream.cc", "quiche/quic/moqt/moqt_framer.cc", "quiche/quic/moqt/moqt_key_value_pair.cc", "quiche/quic/moqt/moqt_known_track_publisher.cc", @@ -1685,6 +1687,7 @@ "quiche/quic/moqt/moqt_bidi_stream_test.cc", "quiche/quic/moqt/moqt_bitrate_adjuster_test.cc", "quiche/quic/moqt/moqt_control_message_queue_test.cc", + "quiche/quic/moqt/moqt_fetch_stream_test.cc", "quiche/quic/moqt/moqt_framer_test.cc", "quiche/quic/moqt/moqt_integration_test.cc", "quiche/quic/moqt/moqt_key_value_pair_test.cc",
diff --git a/quiche/quic/moqt/moqt_bidi_stream.cc b/quiche/quic/moqt/moqt_bidi_stream.cc index b5b7740..cb236c3 100644 --- a/quiche/quic/moqt/moqt_bidi_stream.cc +++ b/quiche/quic/moqt/moqt_bidi_stream.cc
@@ -65,23 +65,18 @@ } absl::Status MoqtBidiStreamBase::SendRequestError( - uint64_t request_id, RequestErrorCode error_code, + RequestErrorCode error_code, std::optional<quic::QuicTimeDelta> retry_interval, - absl::string_view reason_phrase, bool fin) { - MoqtRequestError request_error; - request_error.request_id = request_id; - request_error.error_code = error_code; - request_error.retry_interval = retry_interval; - request_error.reason_phrase = reason_phrase; - return SendOrBufferMessage(framer_->SerializeRequestError(request_error), - fin); + absl::string_view reason_phrase) { + MoqtRequestErrorInfo error_info(error_code, retry_interval, + std::string(reason_phrase)); + return SendRequestError(error_info); } -absl::Status MoqtBidiStreamBase::SendRequestError(uint64_t request_id, - MoqtRequestErrorInfo info, - bool fin) { - return SendRequestError(request_id, info.error_code, info.retry_interval, - info.reason_phrase, fin); +absl::Status MoqtBidiStreamBase::SendRequestError( + const MoqtRequestErrorInfo& info) { + return SendOrBufferMessage(framer_->SerializeRequestError(info), + /*fin=*/true); } absl::Status MoqtBidiStreamBase::SendRequestUpdate(
diff --git a/quiche/quic/moqt/moqt_bidi_stream.h b/quiche/quic/moqt/moqt_bidi_stream.h index 13fc2d2..a96b422 100644 --- a/quiche/quic/moqt/moqt_bidi_stream.h +++ b/quiche/quic/moqt/moqt_bidi_stream.h
@@ -99,11 +99,10 @@ const MessageParameters& parameters, bool fin = false); absl::Status SendRequestError( - uint64_t request_id, RequestErrorCode error_code, + RequestErrorCode error_code, std::optional<quic::QuicTimeDelta> retry_interval, - absl::string_view reason_phrase, bool fin = false); - absl::Status SendRequestError(uint64_t request_id, MoqtRequestErrorInfo info, - bool fin = false); + absl::string_view reason_phrase); + absl::Status SendRequestError(const MoqtRequestErrorInfo& info); // Can be overridden for message-specific constraints. virtual absl::Status SendRequestUpdate(uint64_t request_id, uint64_t existing_request_id, @@ -117,6 +116,7 @@ if (stream() != nullptr) { stream()->ResetWithUserCode(error); } + stream_status_ = MoqtStreamErrorToStatus(error, ""); Detach(); } @@ -127,6 +127,11 @@ } } + // Returns the status of the stream, set from any incoming RESET_STREAM/ + // STOP_SENDING message, and used any time the stream needs to send those + // frames without any explicit guidance. + absl::Status stream_status() const { return stream_status_; } + MoqtFramer* framer() const { return framer_; } // Removes any state in MoqtSession related to the stream. Overrides of this @@ -135,6 +140,10 @@ // state no longer matters. virtual void Detach() = 0; + webtransport::StreamId stream_id() const { + return stream() != nullptr ? stream()->GetStreamId() : 0; + } + protected: // Called when a WebTransport stream has been associated with the object. // Should be used to set the priority for the stream. @@ -161,6 +170,7 @@ webtransport::Stream* stream() const { return stream_parser_.has_value() ? stream_parser_->stream() : nullptr; } + void set_status(absl::Status status) { stream_status_ = status; } private: friend class test::MoqtBidiStreamTestWrapper; @@ -171,6 +181,7 @@ MoqtControlMessageQueue outgoing_message_queue_; MoqtRequestUpdateQueue request_update_queue_; SessionErrorCallback session_error_callback_; + absl::Status stream_status_ = absl::OkStatus(); }; // DispatchControlMessage is wrapped into a class so that the caller class can
diff --git a/quiche/quic/moqt/moqt_bidi_stream_test.cc b/quiche/quic/moqt/moqt_bidi_stream_test.cc index 95a8f68..99e529a 100644 --- a/quiche/quic/moqt/moqt_bidi_stream_test.cc +++ b/quiche/quic/moqt/moqt_bidi_stream_test.cc
@@ -110,19 +110,7 @@ mock_stream_, Writev(ControlMessageOfType(MoqtMessageType::kRequestError), testing::_)); QUICHE_EXPECT_OK(stream_->SendRequestError( - 1, - MoqtRequestErrorInfo{RequestErrorCode::kUnauthorized, - /*retry_interval=*/std::nullopt, ""}, - false)); - EXPECT_FALSE(stream_->detached_); - EXPECT_CALL( - mock_stream_, - Writev(ControlMessageOfType(MoqtMessageType::kRequestError), testing::_)); - QUICHE_EXPECT_OK(stream_->SendRequestError( - 1, - MoqtRequestErrorInfo{RequestErrorCode::kUnauthorized, - /*retry_interval=*/std::nullopt, ""}, - true)); + RequestErrorCode::kUnauthorized, /*retry_interval=*/std::nullopt, "")); EXPECT_TRUE(stream_->detached_); } @@ -173,9 +161,8 @@ EXPECT_CALL( mock_stream_, Writev(ControlMessageOfType(MoqtMessageType::kRequestError), testing::_)); - QUICHE_EXPECT_OK(stream_->SendRequestError(1, RequestErrorCode::kUnauthorized, - std::nullopt, "reason", - /*fin=*/true)); + QUICHE_EXPECT_OK(stream_->SendRequestError(RequestErrorCode::kUnauthorized, + std::nullopt, "reason")); EXPECT_TRUE(stream_->detached_); } @@ -223,10 +210,8 @@ QUICHE_EXPECT_OK( stream_->SendRequestUpdate(1, 0, parameters, std::move(callback))); // Simulate receiving RequestError - MoqtRequestError request_error; - request_error.request_id = 1; - request_error.error_code = RequestErrorCode::kUnauthorized; - request_error.reason_phrase = "unauthorized"; + MoqtRequestError request_error(RequestErrorCode::kUnauthorized, std::nullopt, + "unauthorized"); ExpectFin(mock_stream_); QUICHE_EXPECT_OK(stream_->OnControlMessage(request_error)); EXPECT_TRUE(callback_called);
diff --git a/quiche/quic/moqt/moqt_error.cc b/quiche/quic/moqt/moqt_error.cc index 3466337..a413f88 100644 --- a/quiche/quic/moqt/moqt_error.cc +++ b/quiche/quic/moqt/moqt_error.cc
@@ -94,6 +94,27 @@ } } +webtransport::StreamErrorCode StatusToMoqtStreamError(absl::Status status) { + switch (status.code()) { + case absl::StatusCode::kInternal: + return kResetCodeInternalError; + case absl::StatusCode::kCancelled: + return kResetCodeCancelled; + case absl::StatusCode::kDeadlineExceeded: + return kResetCodeDeliveryTimeout; + case absl::StatusCode::kAborted: + return kResetCodeSessionClosed; + case absl::StatusCode::kFailedPrecondition: + return kResetCodeUnknownObjectStatus; + case absl::StatusCode::kOutOfRange: + return kResetCodeTooFarBehind; + case absl::StatusCode::kInvalidArgument: + return kResetCodeMalformedTrack; + default: + return kResetCodeInternalError; + } +} + std::optional<MoqtError> GetMoqtErrorForStatus(const absl::Status& status) { std::optional<absl::Cord> raw_code_cord = status.GetPayload(kMoqtErrorStatusPayloadUrl);
diff --git a/quiche/quic/moqt/moqt_error.h b/quiche/quic/moqt/moqt_error.h index befb362..636ec97 100644 --- a/quiche/quic/moqt/moqt_error.h +++ b/quiche/quic/moqt/moqt_error.h
@@ -103,6 +103,7 @@ absl::Status MoqtStreamErrorToStatus(webtransport::StreamErrorCode error_code, absl::string_view reason_phrase); +webtransport::StreamErrorCode StatusToMoqtStreamError(absl::Status status); } // namespace moqt
diff --git a/quiche/quic/moqt/moqt_fetch_stream.cc b/quiche/quic/moqt/moqt_fetch_stream.cc new file mode 100644 index 0000000..2d8e43c --- /dev/null +++ b/quiche/quic/moqt/moqt_fetch_stream.cc
@@ -0,0 +1,421 @@ +// 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_fetch_stream.h" + +#include <algorithm> +#include <cstdint> +#include <memory> +#include <optional> +#include <utility> +#include <variant> + +#include "absl/base/casts.h" +#include "absl/base/nullability.h" +#include "absl/status/status.h" +#include "absl/strings/string_view.h" +#include "quiche/quic/moqt/moqt_bidi_stream.h" +#include "quiche/quic/moqt/moqt_error.h" +#include "quiche/quic/moqt/moqt_fetch_task.h" +#include "quiche/quic/moqt/moqt_framer.h" +#include "quiche/quic/moqt/moqt_key_value_pair.h" +#include "quiche/quic/moqt/moqt_live_publisher.h" +#include "quiche/quic/moqt/moqt_messages.h" +#include "quiche/quic/moqt/moqt_names.h" +#include "quiche/quic/moqt/moqt_object_subscriber.h" +#include "quiche/quic/moqt/moqt_parser.h" +#include "quiche/quic/moqt/moqt_priority.h" +#include "quiche/quic/moqt/moqt_publisher.h" +#include "quiche/quic/moqt/moqt_session_callbacks.h" +#include "quiche/quic/moqt/moqt_trace_recorder.h" +#include "quiche/quic/moqt/moqt_types.h" +#include "quiche/quic/moqt/moqt_uni_stream.h" +#include "quiche/common/platform/api/quiche_bug_tracker.h" +#include "quiche/web_transport/web_transport.h" + +namespace moqt { + +MoqtFetchRequestStream::MoqtFetchRequestStream( + MoqtFramer* framer, const MoqtControlMessageParser& message_parser, + uint64_t request_id, const FullTrackName& name, Location start, + Location end, const MessageParameters& parameters, + UpstreamFetchTask* absl_nonnull task, + SessionErrorCallback session_error_callback, + FetchResponseCallback response_callback, + RemoveFetchCallback delete_callback) + : MoqtBidiStreamBase(framer, message_parser, + std::move(session_error_callback)), + ObjectSubscriber(name, request_id, parameters, this), + start_(start), + end_(end), + task_(task), + response_callback_(std::move(response_callback)), + remove_callback_(std::move(delete_callback)) { + // This is a temporary callback until the data stream arrives. + task_->set_task_destroyed_callback([this]() { + task_ = nullptr; + Reset(kResetCodeCancelled); + }); +} + +MoqtFetchRequestStream::MoqtFetchRequestStream( + MoqtFramer* framer, const MoqtControlMessageParser& message_parser, + uint64_t request_id, const FullTrackName& name, uint64_t joining_request_id, + uint64_t joining_start, bool relative, MessageParameters parameters, + UpstreamFetchTask* absl_nonnull task, + SessionErrorCallback session_error_callback, + FetchResponseCallback response_callback, + RemoveFetchCallback delete_callback) + : MoqtBidiStreamBase(framer, message_parser, + std::move(session_error_callback)), + ObjectSubscriber(name, request_id, parameters, this), + start_(relative ? Location(0, 0) : Location(joining_start, 0)), + end_(Location(kMaxGroupId, kMaxObjectId)), + joining_request_id_(joining_request_id), + relative_groups_(relative ? std::make_optional(joining_start) + : std::nullopt), + task_(std::move(task)), + response_callback_(std::move(response_callback)), + remove_callback_(std::move(delete_callback)) { + // This is a temporary callback until the data stream arrives. + task_->set_task_destroyed_callback([this]() { + task_ = nullptr; + Reset(kResetCodeCancelled); + }); +} + +void MoqtFetchRequestStream::OnStreamOpened( + webtransport::StreamVisitor* stream) { + // Remove the stream from the request ID map, because no other stream can be + // assigned to this request ID. + if (remove_callback_ == nullptr) { + QUICHE_BUG(quiche_bug_moqt_fetch_stream_already_detached) + << "OnStreamOpened called after Detach"; + return; + } + RemoveFetchCallback delete_callback = std::move(remove_callback_); + remove_callback_ = nullptr; + std::move(delete_callback)(request_id()); + UpstreamFetchTask* task = task_; + // Interactions with the task will now be mediated through the uni stream. + // If the uni stream closes, so will the bidi stream. + task_ = nullptr; + absl::down_cast<IncomingDataStream*>(stream)->set_fetch_task(task); +} + +void MoqtFetchRequestStream::OnStreamClosed(absl::Status status, + std::optional<DataStreamIndex>) { + if (status.ok()) { + Fin(); + } else { + Reset(StatusToMoqtStreamError(status)); + } +} + +void MoqtFetchRequestStream::OnStreamBound() { + stream_parser()->set_allow_fin(true); + // TODO(martinduke): Set the priority for this stream. + MoqtFetch fetch; + fetch.request_id = request_id(); + fetch.parameters = const_parameters(); + if (!joining_request_id_.has_value()) { + fetch.fetch = StandaloneFetch(full_track_name(), start_, end_); + } else if (relative_groups_.has_value()) { + fetch.fetch = JoiningFetchRelative(*joining_request_id_, *relative_groups_); + } else { + fetch.fetch = JoiningFetchAbsolute(*joining_request_id_, start_.group); + } + SendOrBufferMessageOrFatal(framer()->SerializeFetch(fetch)); +} + +absl::Status MoqtFetchRequestStream::OnRawControlMessage( + const MoqtRawControlMessage& message) { + return ControlMessageDispatcher::DispatchControlMessage( + *this, message_parser(), message, "fetch request"); +} + +absl::Status MoqtFetchRequestStream::OnControlMessage( + const MoqtFetchOk& message) { + if (response_callback_ == nullptr) { + return absl::InvalidArgumentError("Multiple FETCH_OK on the same stream"); + } + QUIC_DLOG(INFO) << "Received the FETCH_OK for " << full_track_name(); + if (relative_groups_.has_value() && + (*relative_groups_ < message.end_location.group)) { + start_ = Location(message.end_location.group - *relative_groups_, 0); + relative_groups_.reset(); + } + end_ = std::min(end_, message.end_location); + FetchResponseCallback response_callback = std::move(response_callback_); + response_callback_ = nullptr; + std::move(response_callback)(message); + return absl::OkStatus(); +} + +absl::Status MoqtFetchRequestStream::OnControlMessage( + const MoqtRequestOk& message) { + absl::StatusOr<MessageParameters> old_parameters = + request_update_queue().NextParameters(); + if (!old_parameters.ok()) { + return old_parameters.status(); + } + Update(*old_parameters); + Update(message.parameters); + return request_update_queue().OnControlMessage(message); +} + +absl::Status MoqtFetchRequestStream::OnControlMessage( + const MoqtRequestError& message) { + if (response_callback_ != nullptr) { + FetchResponseCallback response_callback = std::move(response_callback_); + response_callback_ = nullptr; + std::move(response_callback)(message); + set_status( + RequestErrorCodeToStatus(message.error_code, message.reason_phrase)); + Fin(); + // Don't do anything to the data stream or the task, except set the status. + // Let the data stream closure control the task. Note that if the uni stream + // arrives after REQUEST_ERROR, the record of the request ID will be gone, + // and IncomingDataStream will therefore send STOP_SENDING immediately. + return absl::OkStatus(); + } + // Response to REQUEST_UPDATE. + absl::Status status = request_update_queue().OnControlMessage(message); + if (status.ok()) { + Fin(); + } + return status; +} + +void MoqtFetchRequestStream::Detach() { + if (task_ != nullptr) { + // Notify the task (which the application owns) that nothing more is coming. + // If this has already been called, UpstreamFetchTask will ignore it. + task_->OnStreamAndFetchClosed(stream_status()); + task_ = nullptr; + } + if (remove_callback_ != nullptr) { + RemoveFetchCallback callback = std::move(remove_callback_); + remove_callback_ = nullptr; + std::move(callback)(request_id()); + } + // The uni stream holds a weakptr to this class, so it doesn't need to be + // notified. +} + +MoqtFetchResponseStream::MoqtFetchResponseStream( + MoqtFramer* absl_nonnull framer, + const MoqtControlMessageParser& message_parser, + MoqtPublisher* absl_nonnull application, + SessionErrorCallback session_error_callback, + OpenStreamCallback open_stream_callback, + GetSubscriptionCallback get_subscription_callback) + : MoqtBidiStreamBase(framer, message_parser, + std::move(session_error_callback)), + data_stream_(nullptr), + application_(application), + open_stream_callback_(std::move(open_stream_callback)), + get_subscription_callback_(std::move(get_subscription_callback)), + weak_ptr_factory_(this) {} + +absl::Status MoqtFetchResponseStream::OnRawControlMessage( + const MoqtRawControlMessage& message) { + return ControlMessageDispatcher::DispatchControlMessage( + *this, message_parser(), message, "fetch response"); +} + +absl::Status MoqtFetchResponseStream::OnControlMessage( + const MoqtFetch& message) { + if (request_id_.has_value()) { + return absl::InvalidArgumentError( + "FETCH received on stream that already has a fetch"); + } + request_id_ = message.request_id; + parameters_ = message.parameters; + MoqtDeliveryOrder delivery_order = + message.parameters.group_order.value_or(MoqtDeliveryOrder::kAscending); + std::unique_ptr<MoqtFetchTask> fetch; + FetchResponseCallback response_callback = + [weak_ptr = weak_ptr_factory_.Create()]( + std::variant<FetchOkData, MoqtRequestErrorInfo> result) { + MoqtFetchResponseStream* stream = weak_ptr.GetIfAvailable(); + if (stream == nullptr) { + return; + } + if (std::holds_alternative<FetchOkData>(result)) { + const auto& ok_data = std::get<FetchOkData>(result); + stream->default_publisher_priority_ = + ok_data.extensions.default_publisher_priority(); + stream->parameters_.Update(ok_data.parameters); + stream->SendOrBufferMessageOrFatal( + stream->framer()->SerializeFetchOk(ok_data)); + return; + } + const auto& error = std::get<MoqtRequestErrorInfo>(result); + stream->CheckStatus(stream->SendRequestError(error)); + }; + if (std::holds_alternative<StandaloneFetch>(message.fetch)) { + const StandaloneFetch& standalone_fetch = + std::get<StandaloneFetch>(message.fetch); + FullTrackName track_name = standalone_fetch.full_track_name; + std::shared_ptr<MoqtTrackPublisher> track_publisher = + application_->GetTrack(track_name); + if (track_publisher == nullptr) { + QUIC_DLOG(INFO) << "FETCH for " << track_name + << " rejected by the application: not found"; + return SendRequestError(RequestErrorCode::kDoesNotExist, std::nullopt, + "not found"); + } + QUIC_DLOG(INFO) << "Received a StandaloneFETCH for " << track_name; + fetch = track_publisher->StandaloneFetch( + standalone_fetch.start_location, standalone_fetch.end_location, + delivery_order, std::move(response_callback)); + } else { + // Joining Fetch. + uint64_t joining_request_id = + std::holds_alternative<JoiningFetchRelative>(message.fetch) + ? std::get<JoiningFetchRelative>(message.fetch).joining_request_id + : std::get<JoiningFetchAbsolute>(message.fetch).joining_request_id; + if (get_subscription_callback_ == nullptr) { + QUIC_DLOG(INFO) + << "Received a JOINING_FETCH without subscription callback"; + return SendRequestError(RequestErrorCode::kInternalError, std::nullopt, + "Internal error"); + } + LivePublisher* subscription = + std::move(get_subscription_callback_)(joining_request_id); + get_subscription_callback_ = nullptr; + if (subscription == nullptr) { + QUIC_DLOG(INFO) << "Received a JOINING_FETCH for request_id " + << joining_request_id << " that does not exist"; + return SendRequestError(RequestErrorCode::kInvalidJoiningRequestId, + std::nullopt, + "Joining Fetch for non-existent request"); + } + if (!subscription->can_have_joining_fetch()) { + QUIC_DLOG(INFO) << "Received a JOINING_FETCH for joining_request_id " + << joining_request_id << " that is not forwarding"; + return absl::InvalidArgumentError( + "Joining Fetch for non-forwarding subscribe"); + } + if (subscription->established()) { + const std::optional<Location> largest_object = + subscription->parameters().largest_object; + if (!largest_object.has_value()) { + // Nothing to Fetch. + return SendRequestError(RequestErrorCode::kDoesNotExist, std::nullopt, + "not found"); + } + uint64_t start_group; + if (std::holds_alternative<JoiningFetchRelative>(message.fetch)) { + const JoiningFetchRelative& relative_fetch = + std::get<JoiningFetchRelative>(message.fetch); + start_group = + (relative_fetch.joining_start > largest_object->group) + ? 0 + : (largest_object->group - relative_fetch.joining_start); + } else { + const JoiningFetchAbsolute& absolute_fetch = + std::get<JoiningFetchAbsolute>(message.fetch); + start_group = absolute_fetch.joining_start; + if (start_group > largest_object->group) { + return SendRequestError(RequestErrorCode::kInvalidRange, std::nullopt, + "invalid range"); + } + } + fetch = subscription->publisher().StandaloneFetch( + Location{start_group, 0}, *largest_object, delivery_order, + std::move(response_callback)); + } else { + // Subscription is in PENDING state. + if (std::holds_alternative<JoiningFetchRelative>(message.fetch)) { + fetch = subscription->publisher().RelativeFetch( + std::get<JoiningFetchRelative>(message.fetch).joining_start, + delivery_order, std::move(response_callback)); + } else { + fetch = subscription->publisher().AbsoluteFetch( + std::get<JoiningFetchAbsolute>(message.fetch).joining_start, + delivery_order, std::move(response_callback)); + } + } + } + if (fetch == nullptr || !fetch->GetStatus().ok()) { + QUIC_DLOG(INFO) << "FETCH could not initialize the task"; + return absl::OkStatus(); + } + fetch_ = std::move(fetch); + // Set a temporary new-object callback that creates a data stream. When + // created, the stream visitor will replace this callback. + fetch_->SetObjectAvailableCallback([this]() { + if (open_stream_callback_ != nullptr) { + OpenStreamCallback callback = std::move(open_stream_callback_); + open_stream_callback_ = nullptr; + std::move(callback)( + stream_id(), + MoqtTrackPriority{parameters_.subscriber_priority.value_or( + kDefaultSubscriberPriority), + default_publisher_priority_}); + } + }); + return absl::OkStatus(); +} + +absl::Status MoqtFetchResponseStream::OnControlMessage( + const MoqtRequestUpdate& message) { + if (data_stream_ != nullptr && + message.parameters.subscriber_priority.has_value()) { + data_stream_->UpdatePriority(*message.parameters.subscriber_priority); + } + return SendRequestOk(message.request_id, MessageParameters()); +} + +void MoqtFetchResponseStream::OnDataStreamOpen( + webtransport::Stream* absl_nonnull stream, + MoqtTraceRecorder* trace_recorder) { + if (stream == nullptr) { + return; + } + webtransport::StreamPriority priority = { + kMoqtSendGroupId, + SendOrderForFetch(parameters_.subscriber_priority.value_or( + kDefaultSubscriberPriority))}; + // The line below will lead to updating ObjectsAvailableCallback in the + // FetchTask to call OnCanWrite() on the stream. If there is an object + // available, the callback will be invoked synchronously (i.e. before + // SetVisitor() returns). + if (!request_id_.has_value()) { + QUICHE_BUG(moqt_bug_data_stream_without_request_id) + << "OnDataStreamOpen called with no request ID"; + return; + } + auto fetch_stream = std::make_unique<OutgoingFetchStream>( + *framer(), stream, *request_id_, priority, std::move(fetch_), + [this](absl::Status status) { + // If |this| has closed, it will have called Detach() and + // data_stream_->Reset(), which deletes this callback without invoking + // it. Therefore, there can't be a use-after-free. + data_stream_ = nullptr; + if (status.ok()) { + Fin(); // Clean teardown of the data stream. + } else { + Reset(StatusToMoqtStreamError(status)); + } + }, + trace_recorder); + fetch_ = nullptr; + data_stream_ = fetch_stream.get(); + stream->SetVisitor(std::move(fetch_stream)); + data_stream_->Init(); +} + +void MoqtFetchResponseStream::Detach() { + if (data_stream_ != nullptr && !stream_status().ok()) { + // If the bidi stream FINed, let the data stream close on its own. + data_stream_->OnBidiStreamReset(StatusToMoqtStreamError(stream_status())); + } + data_stream_ = nullptr; + fetch_ = nullptr; +} + +} // namespace moqt
diff --git a/quiche/quic/moqt/moqt_fetch_stream.h b/quiche/quic/moqt/moqt_fetch_stream.h new file mode 100644 index 0000000..5024212 --- /dev/null +++ b/quiche/quic/moqt/moqt_fetch_stream.h
@@ -0,0 +1,145 @@ +// 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_FETCH_STREAM_H_ +#define QUICHE_QUIC_MOQT_MOQT_FETCH_STREAM_H_ + +#include <cstdint> +#include <memory> +#include <optional> + +#include "absl/base/nullability.h" +#include "absl/status/status.h" +#include "quiche/quic/moqt/moqt_bidi_stream.h" +#include "quiche/quic/moqt/moqt_fetch_task.h" +#include "quiche/quic/moqt/moqt_framer.h" +#include "quiche/quic/moqt/moqt_key_value_pair.h" +#include "quiche/quic/moqt/moqt_live_publisher.h" +#include "quiche/quic/moqt/moqt_messages.h" +#include "quiche/quic/moqt/moqt_names.h" +#include "quiche/quic/moqt/moqt_object_subscriber.h" +#include "quiche/quic/moqt/moqt_parser.h" +#include "quiche/quic/moqt/moqt_priority.h" +#include "quiche/quic/moqt/moqt_publisher.h" +#include "quiche/quic/moqt/moqt_session_callbacks.h" +#include "quiche/quic/moqt/moqt_trace_recorder.h" +#include "quiche/quic/moqt/moqt_types.h" +#include "quiche/quic/moqt/moqt_uni_stream.h" +#include "quiche/common/quiche_callbacks.h" +#include "quiche/common/quiche_weak_ptr.h" +#include "quiche/web_transport/web_transport.h" + +namespace moqt { + +using RemoveFetchCallback = + quiche::SingleUseCallback<void(uint64_t request_id)>; + +class MoqtFetchRequestStream : public MoqtBidiStreamBase, + public ObjectSubscriber { + public: + // Constructor for a standalone fetch. + MoqtFetchRequestStream(MoqtFramer* framer, + const MoqtControlMessageParser& message_parser, + uint64_t request_id, const FullTrackName& name, + Location start, Location end, + const MessageParameters& parameters, + UpstreamFetchTask* absl_nonnull task, + SessionErrorCallback session_error_callback, + FetchResponseCallback response_callback, + RemoveFetchCallback delete_callback); + // Constructor for a joining fetch. + MoqtFetchRequestStream(MoqtFramer* framer, + const MoqtControlMessageParser& message_parser, + uint64_t request_id, const FullTrackName& name, + uint64_t joining_request_id, uint64_t joining_start, + bool relative, MessageParameters parameters, + UpstreamFetchTask* absl_nonnull task, + SessionErrorCallback session_error_callback, + FetchResponseCallback response_callback, + RemoveFetchCallback delete_callback); + ~MoqtFetchRequestStream() { + // If stream_status_ has yet to be reported, this is not a clean close. + if (stream_status().ok()) { + set_status(absl::CancelledError("stream destroyed")); + } + Detach(); + } + + // ObjectSubscriber overrides. + bool InWindow(Location location) const override { + return (location >= start_ && location <= end_); + } + bool is_fetch() const override { return true; } + void OnStreamOpened(webtransport::StreamVisitor* stream) override; + void OnStreamClosed(absl::Status status, + std::optional<DataStreamIndex> index) override; + + // MoqtBidiStreamBase overrides. + void OnStreamBound() override; + absl::Status OnRawControlMessage( + const MoqtRawControlMessage& message) override; + absl::Status OnControlMessage(const MoqtFetchOk& message); + absl::Status OnControlMessage(const MoqtRequestOk& message); + absl::Status OnControlMessage(const MoqtRequestError& message); + void Detach() override; + + private: + Location start_, end_; + std::optional<uint64_t> joining_request_id_; + std::optional<uint64_t> relative_groups_; + + UpstreamFetchTask* task_; + FetchResponseCallback response_callback_; + RemoveFetchCallback remove_callback_; +}; + +class MoqtFetchResponseStream : public MoqtBidiStreamBase { + public: + // The stream ID is the *bidi* stream ID. The stream ID is to find the + // FETCH when the uni stream is open. + using OpenStreamCallback = quiche::SingleUseCallback<void( + webtransport::StreamId, MoqtTrackPriority)>; + // Retrieve the associated subscription for a joining FETCH. + using GetSubscriptionCallback = + quiche::SingleUseCallback<LivePublisher*(uint64_t request_id)>; + MoqtFetchResponseStream(MoqtFramer* absl_nonnull framer, + const MoqtControlMessageParser& message_parser, + MoqtPublisher* absl_nonnull application, + SessionErrorCallback session_error_callback, + OpenStreamCallback open_stream_callback, + GetSubscriptionCallback get_subscription_callback); + ~MoqtFetchResponseStream() { + if (stream_status().ok()) { + set_status(absl::CancelledError("stream destroyed")); + } + Detach(); + } + + // MoqtBidiStreamBase overrides. + void OnStreamBound() override { stream_parser()->set_allow_fin(true); } + absl::Status OnRawControlMessage( + const MoqtRawControlMessage& message) override; + absl::Status OnControlMessage(const MoqtFetch& message); + absl::Status OnControlMessage(const MoqtRequestUpdate& message); + void Detach() override; + + std::optional<uint64_t> request_id() const { return request_id_; } + void OnDataStreamOpen(webtransport::Stream* absl_nonnull stream, + MoqtTraceRecorder* trace_recorder); + + private: + OutgoingFetchStream* absl_nullable data_stream_ = nullptr; + std::optional<uint64_t> request_id_; + MessageParameters parameters_; + MoqtPriority default_publisher_priority_ = kDefaultPublisherPriority; + std::unique_ptr<MoqtFetchTask> fetch_; + MoqtPublisher* absl_nonnull application_; + OpenStreamCallback open_stream_callback_; + GetSubscriptionCallback get_subscription_callback_; + quiche::QuicheWeakPtrFactory<MoqtFetchResponseStream> weak_ptr_factory_; +}; + +} // namespace moqt + +#endif // QUICHE_QUIC_MOQT_MOQT_FETCH_STREAM_H_
diff --git a/quiche/quic/moqt/moqt_fetch_stream_test.cc b/quiche/quic/moqt/moqt_fetch_stream_test.cc new file mode 100644 index 0000000..69ecec0 --- /dev/null +++ b/quiche/quic/moqt/moqt_fetch_stream_test.cc
@@ -0,0 +1,904 @@ +// 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_fetch_stream.h" + +#include <cstdint> +#include <memory> +#include <optional> +#include <utility> +#include <variant> + +#include "absl/base/nullability.h" +#include "absl/status/status.h" +#include "absl/strings/string_view.h" +#include "quiche/quic/core/quic_types.h" +#include "quiche/quic/moqt/moqt_error.h" +#include "quiche/quic/moqt/moqt_fetch_task.h" +#include "quiche/quic/moqt/moqt_framer.h" +#include "quiche/quic/moqt/moqt_key_value_pair.h" +#include "quiche/quic/moqt/moqt_live_publisher.h" +#include "quiche/quic/moqt/moqt_messages.h" +#include "quiche/quic/moqt/moqt_names.h" +#include "quiche/quic/moqt/moqt_object.h" +#include "quiche/quic/moqt/moqt_object_subscriber.h" +#include "quiche/quic/moqt/moqt_parser.h" +#include "quiche/quic/moqt/moqt_priority.h" +#include "quiche/quic/moqt/moqt_publisher.h" +#include "quiche/quic/moqt/moqt_session_callbacks.h" +#include "quiche/quic/moqt/moqt_session_interface.h" +#include "quiche/quic/moqt/moqt_trace_recorder.h" +#include "quiche/quic/moqt/moqt_types.h" +#include "quiche/quic/moqt/moqt_uni_stream.h" +#include "quiche/quic/moqt/test_tools/mock_moqt_session.h" +#include "quiche/quic/moqt/test_tools/moqt_framer_utils.h" +#include "quiche/quic/moqt/test_tools/moqt_mock_visitor.h" +#include "quiche/quic/test_tools/mock_clock.h" +#include "quiche/common/platform/api/quiche_test.h" +#include "quiche/common/quiche_mem_slice.h" +#include "quiche/common/test_tools/quiche_test_utils.h" +#include "quiche/web_transport/test_tools/mock_web_transport.h" +#include "quiche/web_transport/web_transport.h" + +namespace moqt::test { +namespace { + +using ::testing::_; +using ::testing::Return; +using ::testing::StrictMock; + +constexpr uint64_t kRequestId = 1; +const FullTrackName kTrackName("foo", "bar"); +const Location kStart(1, 1); +const Location kEnd(3, 100); + +PublishedObject DefaultObject() { + PublishedObject object; + object.metadata = PublishedObjectMetadata{Location(0, 0), + 0, + "", + MoqtObjectStatus::kNormal, + kDefaultPublisherPriority, + true, + 3}; + object.payload.push_back(quiche::QuicheMemSlice::Copy("foo")); + object.fin_after_this = false; + return object; +} + +class MockIncomingDataStream : public IncomingDataStream { + public: + MockIncomingDataStream(MoqtStreamTypeParser& stream_type_parser, + SessionToUniStreamInterface* absl_nonnull session, + const quic::MockClock* absl_nonnull clock) + : IncomingDataStream(std::move(stream_type_parser), session, clock) {} + MOCK_METHOD(void, set_fetch_task, (UpstreamFetchTask * fetch_task), + (override)); +}; + +class MoqtFetchRequestStreamTest : public quiche::test::QuicheTest { + protected: + MoqtFetchRequestStreamTest() + : framer_(/*using_webtrans=*/true, quic::Perspective::IS_CLIENT), + message_parser_(kDefaultMoqtVersion, /*uses_web_transport=*/true, + quic::Perspective::IS_CLIENT), + stream_type_parser_(&mock_stream_), + data_stream_(stream_type_parser_, &mock_session_, &mock_clock_) { + EXPECT_CALL(mock_stream_, CanWrite).WillRepeatedly(Return(true)); + EXPECT_CALL(delete_callback_, Call).Times(testing::AnyNumber()); + EXPECT_CALL(session_error_callback_, Call).Times(testing::AnyNumber()); + } + + std::unique_ptr<MoqtFetchRequestStream> CreateAndBindStandaloneStream( + Location start = kStart, Location end = kEnd, + const MessageParameters& parameters = MessageParameters()) { + MoqtFetch expected_fetch; + expected_fetch.request_id = kRequestId; + expected_fetch.fetch = StandaloneFetch(kTrackName, start, end); + expected_fetch.parameters = parameters; + EXPECT_CALL(mock_stream_, + Writev(SerializedControlMessage(expected_fetch), _)) + .WillOnce(Return(absl::OkStatus())); + + EXPECT_CALL(task_, set_task_destroyed_callback); + auto stream = std::make_unique<MoqtFetchRequestStream>( + &framer_, message_parser_, kRequestId, kTrackName, start, end, + parameters, &task_, session_error_callback_.AsStdFunction(), + response_callback_.AsStdFunction(), delete_callback_.AsStdFunction()); + stream->BindStream(&mock_stream_); + return stream; + } + + std::unique_ptr<MoqtFetchRequestStream> CreateAndBindJoiningStream( + uint64_t joining_request_id, uint64_t joining_start, bool relative, + const MessageParameters& parameters = MessageParameters()) { + MoqtFetch expected_fetch; + expected_fetch.request_id = kRequestId; + if (relative) { + expected_fetch.fetch = + JoiningFetchRelative(joining_request_id, joining_start); + } else { + expected_fetch.fetch = + JoiningFetchAbsolute(joining_request_id, joining_start); + } + expected_fetch.parameters = parameters; + EXPECT_CALL(mock_stream_, + Writev(SerializedControlMessage(expected_fetch), _)) + .WillOnce(Return(absl::OkStatus())); + + EXPECT_CALL(task_, set_task_destroyed_callback) + .WillOnce([&](TaskDestroyedCallback callback) { + task_destroyed_callback_ = std::move(callback); + }); + auto stream = std::make_unique<MoqtFetchRequestStream>( + &framer_, message_parser_, kRequestId, kTrackName, joining_request_id, + joining_start, relative, parameters, &task_, + session_error_callback_.AsStdFunction(), + response_callback_.AsStdFunction(), delete_callback_.AsStdFunction()); + stream->BindStream(&mock_stream_); + return stream; + } + + MoqtFramer framer_; + webtransport::test::MockStream mock_stream_; + MoqtControlMessageParser message_parser_; + MoqtStreamTypeParser stream_type_parser_; + MockSessionToUniStreamInterface mock_session_; + quic::MockClock mock_clock_; + MessageParameters parameters_; + StrictMock<testing::MockFunction<void(MoqtError, absl::string_view)>> + session_error_callback_; + StrictMock<testing::MockFunction<void( + std::variant<FetchOkData, MoqtRequestErrorInfo>)>> + response_callback_; + TaskDestroyedCallback task_destroyed_callback_; + StrictMock<testing::MockFunction<void(uint64_t)>> delete_callback_; + MockUpstreamFetchTask task_; + MockIncomingDataStream data_stream_; +}; + +TEST_F(MoqtFetchRequestStreamTest, OnStreamBoundStandalone) { + std::unique_ptr<MoqtFetchRequestStream> stream = + CreateAndBindStandaloneStream(); + EXPECT_EQ(stream->request_id(), kRequestId); + EXPECT_EQ(stream->full_track_name(), kTrackName); + EXPECT_TRUE(stream->is_fetch()); +} + +TEST_F(MoqtFetchRequestStreamTest, OnStreamBoundJoiningRelative) { + std::unique_ptr<MoqtFetchRequestStream> stream = + CreateAndBindJoiningStream(/*joining_request_id=*/10, + /*joining_start=*/2, /*relative=*/true); + EXPECT_EQ(stream->request_id(), kRequestId); + EXPECT_EQ(stream->full_track_name(), kTrackName); + EXPECT_TRUE(stream->is_fetch()); +} + +TEST_F(MoqtFetchRequestStreamTest, OnStreamBoundJoiningAbsolute) { + std::unique_ptr<MoqtFetchRequestStream> stream = + CreateAndBindJoiningStream(/*joining_request_id=*/10, + /*joining_start=*/5, /*relative=*/false); + EXPECT_EQ(stream->request_id(), kRequestId); + EXPECT_EQ(stream->full_track_name(), kTrackName); + EXPECT_TRUE(stream->is_fetch()); +} + +TEST_F(MoqtFetchRequestStreamTest, InWindowStandalone) { + std::unique_ptr<MoqtFetchRequestStream> stream = + CreateAndBindStandaloneStream(Location(1, 1), Location(3, 100)); + EXPECT_FALSE(stream->InWindow(Location(1, 0))); + EXPECT_TRUE(stream->InWindow(Location(1, 1))); + EXPECT_TRUE(stream->InWindow(Location(2, 50))); + EXPECT_TRUE(stream->InWindow(Location(3, 100))); + EXPECT_FALSE(stream->InWindow(Location(3, 101))); + EXPECT_FALSE(stream->InWindow(Location(4, 0))); +} + +TEST_F(MoqtFetchRequestStreamTest, InWindowJoiningRelative) { + std::unique_ptr<MoqtFetchRequestStream> stream = + CreateAndBindJoiningStream(/*joining_request_id=*/10, + /*joining_start=*/2, /*relative=*/true); + // Before FETCH_OK, start is (0, 0), end is max. + EXPECT_TRUE(stream->InWindow(Location(0, 0))); + EXPECT_TRUE(stream->InWindow(Location(kMaxGroupId, kMaxObjectId))); + + // Deliver FETCH_OK with end_location = (10, 50). + EXPECT_CALL(response_callback_, Call); + MoqtFetchOk ok_message; + ok_message.end_location = Location(10, 50); + ok_message.end_of_track = true; + QUICHE_EXPECT_OK(stream->OnControlMessage(ok_message)); + // Window is now defined relative to end_location. + EXPECT_FALSE(stream->InWindow(Location(7, 99))); + EXPECT_TRUE(stream->InWindow(Location(8, 0))); + EXPECT_TRUE(stream->InWindow(Location(9, 10))); + EXPECT_TRUE(stream->InWindow(Location(10, 50))); + EXPECT_FALSE(stream->InWindow(Location(10, 51))); + EXPECT_FALSE(stream->InWindow(Location(11, 0))); +} + +TEST_F(MoqtFetchRequestStreamTest, InWindowJoiningRelativeUnderflow) { + std::unique_ptr<MoqtFetchRequestStream> stream = + CreateAndBindJoiningStream(/*joining_request_id=*/10, + /*joining_start=*/10, /*relative=*/true); + // Deliver FETCH_OK with end_location = (1, 50). + EXPECT_CALL(response_callback_, Call); + MoqtFetchOk ok_message; + ok_message.end_location = Location(1, 50); + ok_message.end_of_track = false; + QUICHE_EXPECT_OK(stream->OnControlMessage(ok_message)); + EXPECT_TRUE(stream->InWindow(Location(1, 50))); + EXPECT_FALSE(stream->InWindow(Location(1, 51))); +} + +TEST_F(MoqtFetchRequestStreamTest, OnControlMessageFetchOk) { + std::unique_ptr<MoqtFetchRequestStream> stream = + CreateAndBindStandaloneStream(); + + EXPECT_CALL(response_callback_, + Call(testing::VariantWith<FetchOkData>(testing::_))); + MoqtFetchOk ok_message; + ok_message.end_location = Location(3, 50); + ok_message.end_of_track = true; + + QUICHE_EXPECT_OK(stream->OnControlMessage(ok_message)); + EXPECT_TRUE(task_.GetStatus().ok()); + EXPECT_TRUE(stream->InWindow(Location(3, 50))); + EXPECT_FALSE(stream->InWindow(Location(3, 51))); +} + +TEST_F(MoqtFetchRequestStreamTest, OnControlMessageDuplicateFetchOk) { + std::unique_ptr<MoqtFetchRequestStream> stream = + CreateAndBindStandaloneStream(); + EXPECT_CALL(response_callback_, Call); + MoqtFetchOk ok_message; + ok_message.end_location = Location(3, 50); + ok_message.end_of_track = true; + QUICHE_EXPECT_OK(stream->OnControlMessage(ok_message)); + // Second FETCH_OK on the same stream should return InvalidArgumentError. + EXPECT_EQ(stream->OnControlMessage(ok_message).code(), + absl::StatusCode::kInvalidArgument); +} + +TEST_F(MoqtFetchRequestStreamTest, OnControlMessageRequestErrorInitial) { + std::unique_ptr<MoqtFetchRequestStream> stream = + CreateAndBindStandaloneStream(); + EXPECT_CALL(response_callback_, + Call(testing::VariantWith<MoqtRequestErrorInfo>(testing::_))); + ExpectFin(mock_stream_); + MoqtRequestError error_message; + error_message.error_code = RequestErrorCode::kUnauthorized; + error_message.reason_phrase = "Unauthorized"; + QUICHE_EXPECT_OK(stream->OnControlMessage(error_message)); +} + +TEST_F(MoqtFetchRequestStreamTest, SendRequestUpdateAndReceiveOk) { + std::unique_ptr<MoqtFetchRequestStream> stream = + CreateAndBindStandaloneStream(); + EXPECT_CALL(response_callback_, Call); + MoqtFetchOk ok_message; + ok_message.end_location = Location(3, 50); + ok_message.end_of_track = true; + QUICHE_EXPECT_OK(stream->OnControlMessage(ok_message)); + + // Send REQUEST_UPDATE. + EXPECT_CALL(mock_stream_, + Writev(ControlMessageOfType(MoqtMessageType::kRequestUpdate), _)) + .WillOnce(Return(absl::OkStatus())); + MessageParameters update_params; + update_params.subscriber_priority = 50; + bool update_callback_called = false; + MoqtResponseCallback update_callback = + [&](std::variant<MessageParameters, MoqtRequestErrorInfo> res) { + update_callback_called = true; + ASSERT_TRUE(std::holds_alternative<MessageParameters>(res)); + EXPECT_EQ(std::get<MessageParameters>(res).subscriber_priority, 50); + }; + QUICHE_EXPECT_OK(stream->SendRequestUpdate( + /*request_id=*/2, /*joining_start=*/0, update_params, + std::move(update_callback))); + + // Receive REQUEST_OK for the update. + MoqtRequestOk request_ok; + request_ok.request_id = 2; + request_ok.parameters.subscriber_priority = 50; + QUICHE_EXPECT_OK(stream->OnControlMessage(request_ok)); + EXPECT_TRUE(update_callback_called); + EXPECT_EQ(stream->const_parameters().subscriber_priority, 50); +} + +TEST_F(MoqtFetchRequestStreamTest, ReceiveRequestOkWithoutPendingUpdate) { + std::unique_ptr<MoqtFetchRequestStream> stream = + CreateAndBindStandaloneStream(); + EXPECT_CALL(response_callback_, Call); + MoqtFetchOk ok_message; + ok_message.end_location = Location(3, 50); + ok_message.end_of_track = true; + QUICHE_EXPECT_OK(stream->OnControlMessage(ok_message)); + MoqtRequestOk request_ok; + request_ok.request_id = 2; + EXPECT_EQ(stream->OnControlMessage(request_ok).code(), + absl::StatusCode::kFailedPrecondition); +} + +TEST_F(MoqtFetchRequestStreamTest, SendRequestUpdateAndReceiveError) { + std::unique_ptr<MoqtFetchRequestStream> stream = + CreateAndBindStandaloneStream(); + EXPECT_CALL(response_callback_, Call); + MoqtFetchOk ok_message; + ok_message.end_location = Location(3, 50); + ok_message.end_of_track = true; + QUICHE_EXPECT_OK(stream->OnControlMessage(ok_message)); + + // Send REQUEST_UPDATE. + EXPECT_CALL(mock_stream_, + Writev(ControlMessageOfType(MoqtMessageType::kRequestUpdate), _)) + .WillOnce(Return(absl::OkStatus())); + MessageParameters update_params; + update_params.subscriber_priority = 50; + bool update_callback_called = false; + MoqtRequestError request_error(RequestErrorCode::kUnauthorized, std::nullopt, + "Unauthorized"); + MoqtResponseCallback update_callback = + [&](std::variant<MessageParameters, MoqtRequestErrorInfo> res) { + update_callback_called = true; + ASSERT_TRUE(std::holds_alternative<MoqtRequestErrorInfo>(res)); + EXPECT_EQ(request_error, std::get<MoqtRequestErrorInfo>(res)); + }; + QUICHE_EXPECT_OK(stream->SendRequestUpdate( + /*request_id=*/2, /*joining_start=*/0, update_params, + std::move(update_callback))); + + ExpectFin(mock_stream_); + QUICHE_EXPECT_OK(stream->OnControlMessage(request_error)); + EXPECT_TRUE(update_callback_called); +} + +TEST_F(MoqtFetchRequestStreamTest, OnRawControlMessage) { + std::unique_ptr<MoqtFetchRequestStream> stream = + CreateAndBindStandaloneStream(); + EXPECT_CALL(response_callback_, Call); + MoqtFetchOk ok_message; + ok_message.end_location = Location(3, 50); + ok_message.end_of_track = false; + QUICHE_EXPECT_OK(stream->OnRawControlMessage( + GenericMessageToRawControlMessage(ok_message))); + EXPECT_TRUE(task_.GetStatus().ok()); +} + +TEST_F(MoqtFetchRequestStreamTest, DetachCallsRemoveCallbackAndClosesTask) { + std::unique_ptr<MoqtFetchRequestStream> stream = + CreateAndBindStandaloneStream(); + EXPECT_CALL(response_callback_, Call); + MoqtFetchOk ok_message; + ok_message.end_location = Location(3, 50); + ok_message.end_of_track = true; + QUICHE_EXPECT_OK(stream->OnControlMessage(ok_message)); + EXPECT_CALL(delete_callback_, Call(kRequestId)); + EXPECT_CALL(task_, + OnStreamAndFetchClosed(absl::CancelledError("stream destroyed"))); + stream = nullptr; // Destroys stream, triggering Detach(). + EXPECT_TRUE(task_.GetStatus().ok()); +} + +TEST_F(MoqtFetchRequestStreamTest, OnStreamClosedWithError) { + std::unique_ptr<MoqtFetchRequestStream> stream = + CreateAndBindStandaloneStream(); + EXPECT_CALL(data_stream_, set_fetch_task(&task_)); + stream->OnStreamOpened(&data_stream_); + EXPECT_CALL(mock_stream_, ResetWithUserCode(kResetCodeCancelled)); + // Data stream now has responsibility for the task. + EXPECT_CALL(task_, OnStreamAndFetchClosed).Times(0); + stream->OnStreamClosed(absl::CancelledError(), std::nullopt); +} + +TEST_F(MoqtFetchRequestStreamTest, OnStreamClosedWithFin) { + std::unique_ptr<MoqtFetchRequestStream> stream = + CreateAndBindStandaloneStream(); + EXPECT_CALL(data_stream_, set_fetch_task(&task_)); + stream->OnStreamOpened(&data_stream_); + ExpectFin(mock_stream_); + // Data stream now has responsibility for the task. + EXPECT_CALL(task_, OnStreamAndFetchClosed).Times(0); + stream->OnStreamClosed(absl::OkStatus(), std::nullopt); +} + +class MockMoqtPublisher : public MoqtPublisher { + public: + MOCK_METHOD(std::shared_ptr<MoqtTrackPublisher>, GetTrack, + (const FullTrackName& track_name), (override)); +}; + +class MoqtFetchResponseStreamTest : public quiche::test::QuicheTest { + protected: + MoqtFetchResponseStreamTest() + : framer_(/*using_webtrans=*/true, quic::Perspective::IS_SERVER), + message_parser_(kDefaultMoqtVersion, /*uses_web_transport=*/true, + quic::Perspective::IS_SERVER), + track_publisher_(std::make_shared<MockTrackPublisher>(kTrackName)) { + EXPECT_CALL(mock_stream_, CanWrite).WillRepeatedly(Return(true)); + EXPECT_CALL(mock_stream_, GetStreamId).WillRepeatedly(Return(4)); + EXPECT_CALL(session_error_callback_, Call).Times(testing::AnyNumber()); + ON_CALL(mock_data_stream_, CanWrite).WillByDefault(Return(true)); + ON_CALL(mock_stream_, CanWrite).WillByDefault(Return(true)); + } + + std::unique_ptr<MoqtFetchResponseStream> CreateAndBindStream( + bool has_subscription_callback = true) { + MoqtFetchResponseStream::GetSubscriptionCallback subscription_cb; + if (has_subscription_callback) { + subscription_cb = get_subscription_callback_.AsStdFunction(); + } + auto stream = std::make_unique<MoqtFetchResponseStream>( + &framer_, message_parser_, &mock_publisher_, + session_error_callback_.AsStdFunction(), + open_stream_callback_.AsStdFunction(), std::move(subscription_cb)); + stream->BindStream(&mock_stream_); + return stream; + } + + MoqtFramer framer_; + MoqtControlMessageParser message_parser_; + webtransport::test::MockStream mock_stream_, mock_data_stream_; + MockMoqtPublisher mock_publisher_; + std::shared_ptr<MockTrackPublisher> track_publisher_; + MockFetchTask task_; + StrictMock<testing::MockFunction<void(MoqtError, absl::string_view)>> + session_error_callback_; + StrictMock< + testing::MockFunction<void(webtransport::StreamId, MoqtTrackPriority)>> + open_stream_callback_; + StrictMock<testing::MockFunction<LivePublisher*(uint64_t)>> + get_subscription_callback_; + MoqtTraceRecorder trace_recorder_; +}; + +TEST_F(MoqtFetchResponseStreamTest, ReceiveFetchStandaloneSuccess) { + std::unique_ptr<MoqtFetchResponseStream> stream = CreateAndBindStream(); + EXPECT_FALSE(stream->request_id().has_value()); + EXPECT_CALL(mock_publisher_, GetTrack(kTrackName)) + .WillOnce(Return(track_publisher_)); + EXPECT_CALL(*track_publisher_, + StandaloneFetch(kStart, kEnd, MoqtDeliveryOrder::kAscending, _)) + .WillOnce([&](Location, Location, MoqtDeliveryOrder, + FetchResponseCallback callback) { + FetchOkData ok_data(true, kEnd); + std::move(callback)(ok_data); + return std::make_unique<MockFetchTask>(); + }); + EXPECT_CALL(mock_stream_, + Writev(ControlMessageOfType(MoqtMessageType::kFetchOk), _)) + .WillOnce(Return(absl::OkStatus())); + + MoqtFetch fetch; + fetch.request_id = kRequestId; + fetch.fetch = StandaloneFetch(kTrackName, kStart, kEnd); + QUICHE_EXPECT_OK(stream->OnControlMessage(fetch)); + EXPECT_EQ(stream->request_id(), kRequestId); +} + +TEST_F(MoqtFetchResponseStreamTest, ReceiveFetchStandaloneTrackDoesNotExist) { + std::unique_ptr<MoqtFetchResponseStream> stream = CreateAndBindStream(); + EXPECT_CALL(mock_publisher_, GetTrack(kTrackName)).WillOnce(Return(nullptr)); + MoqtRequestError expected_error = {RequestErrorCode::kDoesNotExist, + std::nullopt, "not found"}; + EXPECT_CALL(mock_stream_, Writev(SerializedControlMessage(expected_error), _)) + .WillOnce(Return(absl::OkStatus())); + + MoqtFetch fetch; + fetch.request_id = kRequestId; + fetch.fetch = StandaloneFetch(kTrackName, kStart, kEnd); + QUICHE_EXPECT_OK(stream->OnControlMessage(fetch)); +} + +TEST_F(MoqtFetchResponseStreamTest, + ReceiveFetchStandaloneErrorFromApplication) { + std::unique_ptr<MoqtFetchResponseStream> stream = CreateAndBindStream(); + EXPECT_CALL(mock_publisher_, GetTrack(kTrackName)) + .WillOnce(Return(track_publisher_)); + MoqtRequestError expected_error = {RequestErrorCode::kInternalError, + std::nullopt, "Application error"}; + EXPECT_CALL(*track_publisher_, + StandaloneFetch(kStart, kEnd, MoqtDeliveryOrder::kAscending, _)) + .WillOnce([&](Location, Location, MoqtDeliveryOrder, + FetchResponseCallback callback) { + std::move(callback)(expected_error); + return nullptr; + }); + + EXPECT_CALL(mock_stream_, Writev(SerializedControlMessage(expected_error), _)) + .WillOnce(Return(absl::OkStatus())); + + MoqtFetch fetch; + fetch.request_id = kRequestId; + fetch.fetch = StandaloneFetch(kTrackName, kStart, kEnd); + QUICHE_EXPECT_OK(stream->OnControlMessage(fetch)); +} + +TEST_F(MoqtFetchResponseStreamTest, ReceiveDuplicateFetch) { + std::unique_ptr<MoqtFetchResponseStream> stream = CreateAndBindStream(); + EXPECT_CALL(mock_publisher_, GetTrack(kTrackName)) + .WillOnce(Return(track_publisher_)); + EXPECT_CALL(*track_publisher_, + StandaloneFetch(kStart, kEnd, MoqtDeliveryOrder::kAscending, _)) + .WillOnce([&](Location, Location, MoqtDeliveryOrder, + FetchResponseCallback callback) { + FetchOkData ok_data(true, kEnd); + std::move(callback)(ok_data); + return std::make_unique<MockFetchTask>(); + }); + EXPECT_CALL(mock_stream_, + Writev(ControlMessageOfType(MoqtMessageType::kFetchOk), _)) + .WillOnce(Return(absl::OkStatus())); + + MoqtFetch fetch; + fetch.request_id = kRequestId; + fetch.fetch = StandaloneFetch(kTrackName, kStart, kEnd); + QUICHE_EXPECT_OK(stream->OnControlMessage(fetch)); + + EXPECT_THAT(stream->OnControlMessage(fetch), + quiche::test::StatusIs(absl::StatusCode::kInvalidArgument)); +} + +TEST_F(MoqtFetchResponseStreamTest, + AsyncObjectAvailableCallsOpenStreamCallback) { + std::unique_ptr<MoqtFetchResponseStream> stream = CreateAndBindStream(); + MockFetchTask* task_ptr = nullptr; + EXPECT_CALL(mock_publisher_, GetTrack(kTrackName)) + .WillOnce(Return(track_publisher_)); + EXPECT_CALL(*track_publisher_, + StandaloneFetch(kStart, kEnd, MoqtDeliveryOrder::kAscending, _)) + .WillOnce([&](Location, Location, MoqtDeliveryOrder, + FetchResponseCallback callback) { + FetchOkData ok_data(true, kEnd); + auto task = std::make_unique<MockFetchTask>(); + std::move(callback)(ok_data); + task_ptr = task.get(); + return task; + }); + EXPECT_CALL(mock_stream_, + Writev(ControlMessageOfType(MoqtMessageType::kFetchOk), _)) + .WillOnce(Return(absl::OkStatus())); + MoqtFetch fetch; + fetch.request_id = kRequestId; + fetch.parameters.subscriber_priority = 50; + fetch.fetch = StandaloneFetch(kTrackName, kStart, kEnd); + QUICHE_EXPECT_OK(stream->OnControlMessage(fetch)); + ASSERT_NE(task_ptr, nullptr); + + EXPECT_CALL(open_stream_callback_, + Call(4, MoqtTrackPriority{50, kDefaultPublisherPriority})) + .WillOnce([&](uint64_t, MoqtTrackPriority) { + stream->OnDataStreamOpen(&mock_data_stream_, &trace_recorder_); + }); + EXPECT_CALL(mock_data_stream_, SetPriority); + std::unique_ptr<webtransport::StreamVisitor> data_stream_visitor; + EXPECT_CALL(mock_data_stream_, SetVisitor) + .WillOnce([&](std::unique_ptr<webtransport::StreamVisitor> visitor) { + data_stream_visitor = std::move(visitor); + }); + EXPECT_CALL(*task_ptr, GetNextObject) + .WillOnce([&](PublishedObject& object) { + object = DefaultObject(); + return MoqtFetchTask::kSuccess; + }) + .WillOnce(Return(MoqtFetchTask::kPending)); + EXPECT_CALL(mock_data_stream_, Writev).WillOnce(Return(absl::OkStatus())); + task_ptr->CallObjectsAvailableCallback(); + EXPECT_NE(data_stream_visitor, nullptr); +} + +TEST_F(MoqtFetchResponseStreamTest, + SyncObjectAvailableCallsOpenStreamCallback) { + std::unique_ptr<MoqtFetchResponseStream> stream = CreateAndBindStream(); + MockFetchTask* task_ptr = nullptr; + EXPECT_CALL(mock_publisher_, GetTrack(kTrackName)) + .WillOnce(Return(track_publisher_)); + EXPECT_CALL(*track_publisher_, + StandaloneFetch(kStart, kEnd, MoqtDeliveryOrder::kAscending, _)) + .WillOnce([&](Location, Location, MoqtDeliveryOrder, + FetchResponseCallback callback) { + FetchOkData ok_data(true, kEnd); + auto task = std::make_unique<MockFetchTask>(true); + std::move(callback)(ok_data); + EXPECT_CALL(*task, GetNextObject) + .WillOnce([&](PublishedObject& object) { + object = DefaultObject(); + return MoqtFetchTask::kSuccess; + }) + .WillOnce(Return(MoqtFetchTask::kPending)); + EXPECT_CALL(mock_data_stream_, Writev) + .WillOnce(Return(absl::OkStatus())); + task_ptr = task.get(); + return task; + }); + EXPECT_CALL(mock_stream_, + Writev(ControlMessageOfType(MoqtMessageType::kFetchOk), _)) + .WillOnce(Return(absl::OkStatus())); + EXPECT_CALL(open_stream_callback_, + Call(4, MoqtTrackPriority{50, kDefaultPublisherPriority})) + .WillOnce([&](uint64_t, MoqtTrackPriority) { + stream->OnDataStreamOpen(&mock_data_stream_, &trace_recorder_); + }); + EXPECT_CALL(mock_data_stream_, SetPriority); + std::unique_ptr<webtransport::StreamVisitor> data_stream_visitor; + EXPECT_CALL(mock_data_stream_, SetVisitor) + .WillOnce([&](std::unique_ptr<webtransport::StreamVisitor> visitor) { + data_stream_visitor = std::move(visitor); + }); + MoqtFetch fetch; + fetch.request_id = kRequestId; + fetch.parameters.subscriber_priority = 50; + fetch.fetch = StandaloneFetch(kTrackName, kStart, kEnd); + QUICHE_EXPECT_OK(stream->OnControlMessage(fetch)); + ASSERT_NE(task_ptr, nullptr); + EXPECT_NE(data_stream_visitor, nullptr); +} + +TEST_F(MoqtFetchResponseStreamTest, OnDataStreamOpenCleanAsyncTeardown) { + std::unique_ptr<MoqtFetchResponseStream> stream = CreateAndBindStream(); + MockFetchTask* task_ptr = nullptr; + EXPECT_CALL(mock_publisher_, GetTrack(kTrackName)) + .WillOnce(Return(track_publisher_)); + EXPECT_CALL(*track_publisher_, + StandaloneFetch(kStart, kEnd, MoqtDeliveryOrder::kAscending, _)) + .WillOnce([&](Location, Location, MoqtDeliveryOrder, + FetchResponseCallback callback) { + FetchOkData ok_data(true, kEnd); + auto task = std::make_unique<MockFetchTask>(true); + std::move(callback)(ok_data); + EXPECT_CALL(*task, GetNextObject) + .WillOnce([&](PublishedObject& object) { + object = DefaultObject(); + return MoqtFetchTask::kSuccess; + }) + .WillOnce(Return(MoqtFetchTask::kPending)); + EXPECT_CALL(mock_data_stream_, Writev) + .WillOnce(Return(absl::OkStatus())); + task_ptr = task.get(); + return task; + }); + EXPECT_CALL(mock_stream_, + Writev(ControlMessageOfType(MoqtMessageType::kFetchOk), _)) + .WillOnce(Return(absl::OkStatus())); + EXPECT_CALL(open_stream_callback_, + Call(4, MoqtTrackPriority{50, kDefaultPublisherPriority})) + .WillOnce([&](uint64_t, MoqtTrackPriority) { + stream->OnDataStreamOpen(&mock_data_stream_, &trace_recorder_); + }); + EXPECT_CALL(mock_data_stream_, SetPriority); + std::unique_ptr<webtransport::StreamVisitor> data_stream_visitor; + EXPECT_CALL(mock_data_stream_, SetVisitor) + .WillOnce([&](std::unique_ptr<webtransport::StreamVisitor> visitor) { + data_stream_visitor = std::move(visitor); + }); + MoqtFetch fetch; + fetch.request_id = kRequestId; + fetch.parameters.subscriber_priority = 50; + fetch.fetch = StandaloneFetch(kTrackName, kStart, kEnd); + QUICHE_EXPECT_OK(stream->OnControlMessage(fetch)); + ASSERT_NE(task_ptr, nullptr); + EXPECT_NE(data_stream_visitor, nullptr); + EXPECT_CALL(*task_ptr, GetNextObject).WillOnce(Return(MoqtFetchTask::kEof)); + ExpectFin(mock_data_stream_); + ExpectFin(mock_stream_); + task_ptr->CallObjectsAvailableCallback(); +} + +TEST_F(MoqtFetchResponseStreamTest, OnDataStreamOpenCleanSyncTeardown) { + std::unique_ptr<MoqtFetchResponseStream> stream = CreateAndBindStream(); + MockFetchTask* task_ptr = nullptr; + EXPECT_CALL(mock_publisher_, GetTrack(kTrackName)) + .WillOnce(Return(track_publisher_)); + EXPECT_CALL(*track_publisher_, + StandaloneFetch(kStart, kEnd, MoqtDeliveryOrder::kAscending, _)) + .WillOnce([&](Location, Location, MoqtDeliveryOrder, + FetchResponseCallback callback) { + FetchOkData ok_data(true, kEnd); + auto task = std::make_unique<MockFetchTask>(true); + std::move(callback)(ok_data); + EXPECT_CALL(*task, GetNextObject).WillOnce(Return(MoqtFetchTask::kEof)); + ExpectFin(mock_data_stream_); + ExpectFin(mock_stream_); + task_ptr = task.get(); + return task; + }); + EXPECT_CALL(mock_stream_, + Writev(ControlMessageOfType(MoqtMessageType::kFetchOk), _)) + .WillOnce(Return(absl::OkStatus())); + EXPECT_CALL(open_stream_callback_, + Call(4, MoqtTrackPriority{50, kDefaultPublisherPriority})) + .WillOnce([&](uint64_t, MoqtTrackPriority) { + stream->OnDataStreamOpen(&mock_data_stream_, &trace_recorder_); + }); + EXPECT_CALL(mock_data_stream_, SetPriority); + std::unique_ptr<webtransport::StreamVisitor> data_stream_visitor; + EXPECT_CALL(mock_data_stream_, SetVisitor) + .WillOnce([&](std::unique_ptr<webtransport::StreamVisitor> visitor) { + data_stream_visitor = std::move(visitor); + }); + MoqtFetch fetch; + fetch.request_id = kRequestId; + fetch.parameters.subscriber_priority = 50; + fetch.fetch = StandaloneFetch(kTrackName, kStart, kEnd); + QUICHE_EXPECT_OK(stream->OnControlMessage(fetch)); +} + +TEST_F(MoqtFetchResponseStreamTest, OnDataStreamOpenErrorTeardown) { + std::unique_ptr<MoqtFetchResponseStream> stream = CreateAndBindStream(); + MockFetchTask* task_ptr = nullptr; + EXPECT_CALL(mock_publisher_, GetTrack(kTrackName)) + .WillOnce(Return(track_publisher_)); + EXPECT_CALL(*track_publisher_, + StandaloneFetch(kStart, kEnd, MoqtDeliveryOrder::kAscending, _)) + .WillOnce([&](Location, Location, MoqtDeliveryOrder, + FetchResponseCallback callback) { + FetchOkData ok_data(true, kEnd); + auto task = std::make_unique<MockFetchTask>(); + std::move(callback)(ok_data); + task_ptr = task.get(); + return task; + }); + EXPECT_CALL(mock_stream_, + Writev(ControlMessageOfType(MoqtMessageType::kFetchOk), _)) + .WillOnce(Return(absl::OkStatus())); + + MoqtFetch fetch; + fetch.request_id = kRequestId; + fetch.fetch = StandaloneFetch(kTrackName, kStart, kEnd); + QUICHE_EXPECT_OK(stream->OnControlMessage(fetch)); + ASSERT_NE(task_ptr, nullptr); + + webtransport::test::MockStream mock_data_stream; + std::unique_ptr<webtransport::StreamVisitor> data_stream_visitor; + EXPECT_CALL(mock_data_stream, CanWrite).WillRepeatedly(Return(true)); + EXPECT_CALL(mock_data_stream, GetStreamId).WillRepeatedly(Return(10)); + EXPECT_CALL(mock_data_stream, SetPriority).Times(testing::AtLeast(1)); + EXPECT_CALL(mock_data_stream, SetVisitor) + .WillOnce([&](std::unique_ptr<webtransport::StreamVisitor> visitor) { + data_stream_visitor = std::move(visitor); + }); + EXPECT_CALL(*task_ptr, GetNextObject).WillOnce(Return(MoqtFetchTask::kError)); + EXPECT_CALL(*task_ptr, GetStatus()) + .WillOnce(Return(absl::InternalError("Fetch failed"))); + EXPECT_CALL(mock_data_stream, ResetWithUserCode(kResetCodeInternalError)); + EXPECT_CALL(mock_stream_, ResetWithUserCode(kResetCodeInternalError)); + stream->OnDataStreamOpen(&mock_data_stream, &trace_recorder_); +} + +TEST_F(MoqtFetchResponseStreamTest, ReceiveRequestUpdate) { + std::unique_ptr<MoqtFetchResponseStream> stream = CreateAndBindStream(); + EXPECT_CALL(mock_publisher_, GetTrack(kTrackName)) + .WillOnce(Return(track_publisher_)); + EXPECT_CALL(*track_publisher_, + StandaloneFetch(kStart, kEnd, MoqtDeliveryOrder::kAscending, _)) + .WillOnce([&](Location, Location, MoqtDeliveryOrder, + FetchResponseCallback callback) { + FetchOkData ok_data(true, kEnd); + std::move(callback)(ok_data); + return std::make_unique<MockFetchTask>(); + }); + EXPECT_CALL(mock_stream_, + Writev(ControlMessageOfType(MoqtMessageType::kFetchOk), _)) + .WillOnce(Return(absl::OkStatus())); + + MoqtFetch fetch; + fetch.request_id = kRequestId; + fetch.fetch = StandaloneFetch(kTrackName, kStart, kEnd); + QUICHE_EXPECT_OK(stream->OnControlMessage(fetch)); + + EXPECT_CALL(mock_stream_, + Writev(ControlMessageOfType(MoqtMessageType::kRequestOk), _)) + .WillOnce(Return(absl::OkStatus())); + MoqtRequestUpdate update; + update.request_id = kRequestId; + update.parameters.subscriber_priority = 20; + QUICHE_EXPECT_OK(stream->OnControlMessage(update)); +} + +TEST_F(MoqtFetchResponseStreamTest, OnRawControlMessage) { + std::unique_ptr<MoqtFetchResponseStream> stream = CreateAndBindStream(); + EXPECT_CALL(mock_publisher_, GetTrack(kTrackName)) + .WillOnce(Return(track_publisher_)); + EXPECT_CALL(*track_publisher_, + StandaloneFetch(kStart, kEnd, MoqtDeliveryOrder::kAscending, _)) + .WillOnce([&](Location, Location, MoqtDeliveryOrder, + FetchResponseCallback callback) { + FetchOkData ok_data(true, kEnd); + std::move(callback)(ok_data); + return std::make_unique<MockFetchTask>(); + }); + EXPECT_CALL(mock_stream_, + Writev(ControlMessageOfType(MoqtMessageType::kFetchOk), _)) + .WillOnce(Return(absl::OkStatus())); + + MoqtFetch fetch; + fetch.request_id = kRequestId; + fetch.fetch = StandaloneFetch(kTrackName, kStart, kEnd); + QUICHE_EXPECT_OK( + stream->OnRawControlMessage(GenericMessageToRawControlMessage(fetch))); +} + +TEST_F(MoqtFetchResponseStreamTest, OnRawControlMessageUnexpectedType) { + std::unique_ptr<MoqtFetchResponseStream> stream = CreateAndBindStream(); + MoqtSubscribe subscribe; + subscribe.request_id = kRequestId; + subscribe.full_track_name = kTrackName; + EXPECT_THAT( + stream->OnRawControlMessage(GenericMessageToRawControlMessage(subscribe)), + quiche::test::StatusIs(absl::StatusCode::kInvalidArgument)); +} + +TEST_F(MoqtFetchResponseStreamTest, ReceiveJoiningFetchWithoutCallback) { + std::unique_ptr<MoqtFetchResponseStream> stream = + CreateAndBindStream(/*has_subscription_callback=*/false); + MoqtRequestError expected_error = {RequestErrorCode::kInternalError, + std::nullopt, "Internal error"}; + EXPECT_CALL(mock_stream_, Writev(SerializedControlMessage(expected_error), _)) + .WillOnce(Return(absl::OkStatus())); + + MoqtFetch fetch; + fetch.request_id = kRequestId; + fetch.fetch = JoiningFetchRelative(10, 1); + QUICHE_EXPECT_OK(stream->OnControlMessage(fetch)); +} + +TEST_F(MoqtFetchResponseStreamTest, ReceiveJoiningFetchSubscriptionNotFound) { + std::unique_ptr<MoqtFetchResponseStream> stream = CreateAndBindStream(); + EXPECT_CALL(get_subscription_callback_, Call(10)).WillOnce(Return(nullptr)); + MoqtRequestError expected_error = {RequestErrorCode::kInvalidJoiningRequestId, + std::nullopt, + "Joining Fetch for non-existent request"}; + EXPECT_CALL(mock_stream_, Writev(SerializedControlMessage(expected_error), _)) + .WillOnce(Return(absl::OkStatus())); + + MoqtFetch fetch; + fetch.request_id = kRequestId; + fetch.fetch = JoiningFetchRelative(10, 1); + QUICHE_EXPECT_OK(stream->OnControlMessage(fetch)); +} + +TEST_F(MoqtFetchResponseStreamTest, DetachResetsDataStream) { + std::unique_ptr<MoqtFetchResponseStream> stream = CreateAndBindStream(); + MockFetchTask* task_ptr = nullptr; + EXPECT_CALL(mock_publisher_, GetTrack(kTrackName)) + .WillOnce(Return(track_publisher_)); + EXPECT_CALL(*track_publisher_, + StandaloneFetch(kStart, kEnd, MoqtDeliveryOrder::kAscending, _)) + .WillOnce([&](Location, Location, MoqtDeliveryOrder, + FetchResponseCallback callback) { + FetchOkData ok_data(true, kEnd); + auto task = std::make_unique<MockFetchTask>(); + std::move(callback)(ok_data); + task_ptr = task.get(); + return task; + }); + EXPECT_CALL(mock_stream_, + Writev(ControlMessageOfType(MoqtMessageType::kFetchOk), _)) + .WillOnce(Return(absl::OkStatus())); + + MoqtFetch fetch; + fetch.request_id = kRequestId; + fetch.fetch = StandaloneFetch(kTrackName, kStart, kEnd); + QUICHE_EXPECT_OK(stream->OnControlMessage(fetch)); + ASSERT_NE(task_ptr, nullptr); + + webtransport::test::MockStream mock_data_stream; + std::unique_ptr<webtransport::StreamVisitor> data_stream_visitor; + EXPECT_CALL(mock_data_stream, CanWrite).WillRepeatedly(Return(false)); + EXPECT_CALL(mock_data_stream, GetStreamId).WillRepeatedly(Return(10)); + EXPECT_CALL(mock_data_stream, SetPriority).Times(testing::AtLeast(1)); + EXPECT_CALL(mock_data_stream, SetVisitor) + .WillOnce([&](std::unique_ptr<webtransport::StreamVisitor> visitor) { + data_stream_visitor = std::move(visitor); + }); + + stream->OnDataStreamOpen(&mock_data_stream, &trace_recorder_); + + EXPECT_CALL(mock_data_stream, ResetWithUserCode(kResetCodeCancelled)); + stream = nullptr; // Destroys stream, triggering Detach(). +} + +} // namespace +} // namespace moqt::test
diff --git a/quiche/quic/moqt/moqt_fetch_task.h b/quiche/quic/moqt/moqt_fetch_task.h index 08b8cd1..a3866d1 100644 --- a/quiche/quic/moqt/moqt_fetch_task.h +++ b/quiche/quic/moqt/moqt_fetch_task.h
@@ -7,15 +7,12 @@ #include <cstdint> #include <optional> -#include <string> -#include <utility> #include <variant> #include "absl/base/nullability.h" #include "absl/status/status.h" #include "quiche/quic/moqt/moqt_error.h" #include "quiche/quic/moqt/moqt_key_value_pair.h" -#include "quiche/quic/moqt/moqt_messages.h" #include "quiche/quic/moqt/moqt_names.h" #include "quiche/quic/moqt/moqt_object.h" #include "quiche/common/quiche_callbacks.h" @@ -23,7 +20,7 @@ namespace moqt { -// The callback we'll use for all request types going forward. Can only be used +// The callback we'll use for most request types going forward. Can only be used // once. If the argument is MessageParameters, the request was successful. // Otherwise, the request failed with the given error info. // This is not used for responses to SUBSCRIBE (SubscribeVisitor), FETCH @@ -56,10 +53,6 @@ // cancelled by deleting the object. class MoqtFetchTask { public: - // The request_id field will be ignored. - using FetchResponseCallback = quiche::SingleUseCallback<void( - std::variant<MoqtFetchOk, MoqtRequestError>)>; - virtual ~MoqtFetchTask() = default; // TODO(martinduke): Replace with GetNextResult above. @@ -89,11 +82,6 @@ // immediately. virtual void SetObjectAvailableCallback( ObjectsAvailableCallback callback) = 0; - // One of these callbacks is called as soon as the data publisher has enough - // information for either FETCH_OK or FETCH_ERROR. - // If the appropriate response is already available, the callback will be - // called immediately. - virtual void SetFetchResponseCallback(FetchResponseCallback callback) = 0; // Returns the error if fetch has completely failed, and OK otherwise. virtual absl::Status GetStatus() = 0; @@ -124,27 +112,6 @@ virtual const TrackNamespace& prefix() = 0; }; -// A fetch that starts out in the failed state. -class MoqtFailedFetch : public MoqtFetchTask { - public: - explicit MoqtFailedFetch(absl::Status status) : status_(std::move(status)) {} - - GetNextObjectResult GetNextObject(PublishedObject&) override { - return kError; - } - absl::Status GetStatus() override { return status_; } - void SetObjectAvailableCallback( - ObjectsAvailableCallback /*callback*/) override {} - void SetFetchResponseCallback(FetchResponseCallback callback) override { - MoqtRequestError error{/*request_id=*/0, StatusToRequestErrorCode(status_), - std::nullopt, std::string(status_.message())}; - std::move(callback)(error); - } - - private: - absl::Status status_; -}; - } // namespace moqt #endif // QUICHE_QUIC_MOQT_MOQT_FETCH_TASK_H_
diff --git a/quiche/quic/moqt/moqt_framer.cc b/quiche/quic/moqt/moqt_framer.cc index 7e76bcc..70cc872 100644 --- a/quiche/quic/moqt/moqt_framer.cc +++ b/quiche/quic/moqt/moqt_framer.cc
@@ -532,8 +532,7 @@ quiche::QuicheBuffer MoqtFramer::SerializeRequestError( const MoqtRequestError& message) { return SerializeControlMessage( - MoqtMessageType::kRequestError, WireMoqVarInt(message.request_id), - WireMoqVarInt(message.error_code), + MoqtMessageType::kRequestError, WireMoqVarInt(message.error_code), WireMoqVarInt(message.retry_interval.has_value() ? message.retry_interval->ToMilliseconds() + 1 : 0), @@ -650,8 +649,7 @@ quiche::QuicheBuffer MoqtFramer::SerializeFetchOk(const MoqtFetchOk& message) { return SerializeControlMessage( - MoqtMessageType::kFetchOk, WireMoqVarInt(message.request_id), - WireBoolean(message.end_of_track), + MoqtMessageType::kFetchOk, WireBoolean(message.end_of_track), WireMoqVarInt(message.end_location.group), WireMoqVarInt(message.end_location.object == kMaxObjectId ? 0 @@ -660,11 +658,6 @@ WireKeyValuePairList(message.extensions, false)); } -quiche::QuicheBuffer MoqtFramer::SerializeFetchCancel( - const MoqtFetchCancel& message) { - return SerializeControlMessage(MoqtMessageType::kFetchCancel, - WireMoqVarInt(message.request_id)); -} quiche::QuicheBuffer MoqtFramer::SerializePublish(const MoqtPublish& message) { return SerializeControlMessage(
diff --git a/quiche/quic/moqt/moqt_framer.h b/quiche/quic/moqt/moqt_framer.h index c11479b..beaa407 100644 --- a/quiche/quic/moqt/moqt_framer.h +++ b/quiche/quic/moqt/moqt_framer.h
@@ -67,7 +67,6 @@ quiche::QuicheBuffer SerializeSubscribeTracks( const MoqtSubscribeTracks& message); quiche::QuicheBuffer SerializeFetch(const MoqtFetch& message); - quiche::QuicheBuffer SerializeFetchCancel(const MoqtFetchCancel& message); quiche::QuicheBuffer SerializeFetchOk(const MoqtFetchOk& message); quiche::QuicheBuffer SerializePublish(const MoqtPublish& message); quiche::QuicheBuffer SerializeObjectAck(const MoqtObjectAck& message);
diff --git a/quiche/quic/moqt/moqt_framer_test.cc b/quiche/quic/moqt/moqt_framer_test.cc index 7082ef7..f173cc5 100644 --- a/quiche/quic/moqt/moqt_framer_test.cc +++ b/quiche/quic/moqt/moqt_framer_test.cc
@@ -55,7 +55,6 @@ MoqtMessageType::kSubscribeNamespace, MoqtMessageType::kSubscribeTracks, MoqtMessageType::kFetch, - MoqtMessageType::kFetchCancel, MoqtMessageType::kFetchOk, MoqtMessageType::kPublish, MoqtMessageType::kObjectAck, @@ -184,10 +183,6 @@ auto data = std::get<MoqtFetch>(structured_data); return framer_.SerializeFetch(data); } - case moqt::MoqtMessageType::kFetchCancel: { - auto data = std::get<MoqtFetchCancel>(structured_data); - return framer_.SerializeFetchCancel(data); - } case moqt::MoqtMessageType::kFetchOk: { auto data = std::get<MoqtFetchOk>(structured_data); return framer_.SerializeFetchOk(data); @@ -413,7 +408,6 @@ TEST_F(MoqtFramerSimpleTest, FetchOkWholeGroup) { MoqtFetchOk fetch_ok = { - /*request_id=*/1, /*end_of_track=*/false, /*end_location=*/Location{4, kMaxObjectId}, MessageParameters(), @@ -421,7 +415,7 @@ }; quiche::QuicheBuffer buffer = framer_.SerializeFetchOk(fetch_ok); // Check that object ID is zero. - EXPECT_EQ(static_cast<uint8_t>(buffer.AsSpan()[7]), 0); + EXPECT_EQ(static_cast<uint8_t>(buffer.AsSpan()[5]), 0); } TEST_F(MoqtFramerSimpleTest, RelativeJoiningFetch) {
diff --git a/quiche/quic/moqt/moqt_integration_test.cc b/quiche/quic/moqt/moqt_integration_test.cc index 8650fc8..c0c5ab4 100644 --- a/quiche/quic/moqt/moqt_integration_test.cc +++ b/quiche/quic/moqt/moqt_integration_test.cc
@@ -482,32 +482,44 @@ for (int i = 0; i < 100; ++i) { queue->AddObject(MemSliceFromString("object"), /*key=*/true); } - std::unique_ptr<MoqtFetchTask> fetch; - EXPECT_TRUE(client_->session()->Fetch( + bool fetch_ok = false; + std::unique_ptr<MoqtFetchTask> fetch = client_->session()->Fetch( full_track_name, - [&](std::unique_ptr<MoqtFetchTask> task) { fetch = std::move(task); }, - Location{0, 0}, 99, std::nullopt, MessageParameters())); - // Run until we get FETCH_OK. + [&](std::variant<FetchOkData, MoqtRequestErrorInfo> response) { + fetch_ok = std::holds_alternative<FetchOkData>(response); + }, + Location{0, 0}, 99, std::nullopt, MessageParameters()); + ASSERT_NE(fetch, nullptr); + bool eof = false; + std::vector<PublishedObject> objects; + fetch->SetObjectAvailableCallback([&]() { + PublishedObject object; + while (true) { + MoqtFetchTask::GetNextObjectResult result = fetch->GetNextObject(object); + if (result == MoqtFetchTask::GetNextObjectResult::kSuccess) { + objects.push_back(std::move(object)); + } else if (result == MoqtFetchTask::GetNextObjectResult::kEof) { + eof = true; + break; + } else { + break; + } + } + }); + // Run until we get FETCH_OK and all objects until EOF. bool success = test_harness_.RunUntilWithDefaultTimeout( - [&]() { return fetch != nullptr; }); + [&]() { return fetch_ok && eof; }); EXPECT_TRUE(success); EXPECT_TRUE(fetch->GetStatus().ok()); - MoqtFetchTask::GetNextObjectResult result; - PublishedObject object; + EXPECT_EQ(objects.size(), 3); Location expected{97, 0}; - do { - result = fetch->GetNextObject(object); - if (result == MoqtFetchTask::GetNextObjectResult::kEof) { - break; - } - EXPECT_EQ(result, MoqtFetchTask::GetNextObjectResult::kSuccess); + for (const PublishedObject& object : objects) { EXPECT_EQ(object.metadata.location, expected); EXPECT_EQ(object.metadata.status, MoqtObjectStatus::kNormal); EXPECT_EQ(object.payload[0].AsStringView(), "object"); ++expected.group; - } while (result == MoqtFetchTask::GetNextObjectResult::kSuccess); - EXPECT_EQ(result, MoqtFetchTask::GetNextObjectResult::kEof); + } EXPECT_EQ(expected, Location(100, 0)); }
diff --git a/quiche/quic/moqt/moqt_live_publisher.cc b/quiche/quic/moqt/moqt_live_publisher.cc index f69f598..fc91cd6 100644 --- a/quiche/quic/moqt/moqt_live_publisher.cc +++ b/quiche/quic/moqt/moqt_live_publisher.cc
@@ -104,7 +104,7 @@ return; } visitor()->UpdateTrackPriority( - request_id_, old_track_priority, + track_publisher_->GetTrackName(), old_track_priority, MoqtTrackPriority{new_priority, publisher_priority}); // Don't bother to update all the pending stream send orders. } @@ -154,8 +154,7 @@ } void LivePublisher::OnSubscribeRejected(MoqtRequestErrorInfo info) { - bidi_stream_->CheckStatus(bidi_stream_->SendRequestError(request_id_, info, - /*fin=*/true)); + bidi_stream_->CheckStatus(bidi_stream_->SendRequestError(info)); // Sending FIN will delete the class. } @@ -244,7 +243,7 @@ StreamRank rank = StreamRankFor(parameters); if (pending_streams_.empty() || rank > pending_streams_.rbegin()->first) { session_info->UpdateTrackPriority( - request_id_, + track_publisher_->GetTrackName(), /*old_priority=*/pending_streams_.empty() ? std::optional<MoqtTrackPriority>() : std::make_optional( @@ -450,7 +449,7 @@ pending_streams_.erase(--(it.base())); if (!pending_streams_.empty()) { session_info->UpdateTrackPriority( - request_id_, std::nullopt, + track_publisher_->GetTrackName(), std::nullopt, MoqtTrackPriority{ subscriber_priority(), pending_streams_.rbegin()->second.publisher_priority.value_or(
diff --git a/quiche/quic/moqt/moqt_live_publisher.h b/quiche/quic/moqt/moqt_live_publisher.h index 0dcd182..6bbdb00 100644 --- a/quiche/quic/moqt/moqt_live_publisher.h +++ b/quiche/quic/moqt/moqt_live_publisher.h
@@ -80,7 +80,7 @@ // streams. If it has a value, |old_priority| is the old value to be // replaced by |new_priority|. virtual void UpdateTrackPriority( - uint64_t request_id, std::optional<MoqtTrackPriority> old_priority, + const FullTrackName& name, std::optional<MoqtTrackPriority> old_priority, MoqtTrackPriority new_priority) = 0; virtual quic::QuicAlarmFactory* alarm_factory() = 0; virtual std::shared_ptr<MoqtTrackPublisher> GetTrackPublisher(
diff --git a/quiche/quic/moqt/moqt_live_publisher_test.cc b/quiche/quic/moqt/moqt_live_publisher_test.cc index dd5b0b7..df1850d 100644 --- a/quiche/quic/moqt/moqt_live_publisher_test.cc +++ b/quiche/quic/moqt/moqt_live_publisher_test.cc
@@ -203,7 +203,7 @@ EXPECT_CALL(webtrans_, CanOpenNextOutgoingUnidirectionalStream()) .WillOnce(Return(false)); EXPECT_CALL(visitor_, - UpdateTrackPriority(1, _, + UpdateTrackPriority(track_publisher_->GetTrackName(), _, MoqtTrackPriority{subscriber_priority(), publisher_priority})); publisher_->OnNewObjectAvailable(location, subgroup, publisher_priority); @@ -325,7 +325,7 @@ new_params.subscriber_priority = 20; EXPECT_CALL(*track_publisher_, extensions()) .WillRepeatedly(ReturnRef(extensions_)); - EXPECT_CALL(visitor_, UpdateTrackPriority(1, + EXPECT_CALL(visitor_, UpdateTrackPriority(track_publisher_->GetTrackName(), std::optional<MoqtTrackPriority>( {subscriber_priority(), 64}), MoqtTrackPriority{20, 64}));
diff --git a/quiche/quic/moqt/moqt_messages.cc b/quiche/quic/moqt/moqt_messages.cc index ee6e6ca..a965006 100644 --- a/quiche/quic/moqt/moqt_messages.cc +++ b/quiche/quic/moqt/moqt_messages.cc
@@ -103,8 +103,6 @@ return "PUBLISH"; case MoqtMessageType::kFetch: return "FETCH"; - case MoqtMessageType::kFetchCancel: - return "FETCH_CANCEL"; case MoqtMessageType::kFetchOk: return "FETCH_OK"; case MoqtMessageType::kObjectAck:
diff --git a/quiche/quic/moqt/moqt_messages.h b/quiche/quic/moqt/moqt_messages.h index e3ff734..2be9f7f 100644 --- a/quiche/quic/moqt/moqt_messages.h +++ b/quiche/quic/moqt/moqt_messages.h
@@ -22,6 +22,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_types.h" #include "quiche/common/platform/api/quiche_export.h" @@ -211,7 +212,6 @@ kNamespaceDone = 0x0e, kGoAway = 0x10, kFetch = 0x16, - kFetchCancel = 0x17, kFetchOk = 0x18, kPublish = 0x1d, kSubscribeNamespace = 0x50, @@ -362,12 +362,7 @@ uint64_t value_ = 0; }; -struct QUICHE_EXPORT MoqtRequestError { - uint64_t request_id; - RequestErrorCode error_code; - std::optional<quic::QuicTimeDelta> retry_interval; - std::string reason_phrase; -}; +using MoqtRequestError = MoqtRequestErrorInfo; struct QUICHE_EXPORT MoqtSubscribe { uint64_t request_id; @@ -485,17 +480,7 @@ MessageParameters parameters; }; -struct QUICHE_EXPORT MoqtFetchOk { - uint64_t request_id; - bool end_of_track; - Location end_location; - MessageParameters parameters; - TrackExtensions extensions; -}; - -struct QUICHE_EXPORT MoqtFetchCancel { - uint64_t request_id; -}; +using MoqtFetchOk = FetchOkData; struct QUICHE_EXPORT MoqtPublish { uint64_t request_id;
diff --git a/quiche/quic/moqt/moqt_namespace_stream.cc b/quiche/quic/moqt/moqt_namespace_stream.cc index c561373..c7b6cb2 100644 --- a/quiche/quic/moqt/moqt_namespace_stream.cc +++ b/quiche/quic/moqt/moqt_namespace_stream.cc
@@ -62,22 +62,14 @@ // This is irrelevant. return absl::OkStatus(); } - MoqtResponseCallback callback = task->GetResponseCallback(message.request_id); - if (callback == nullptr) { - return absl::InvalidArgumentError("Unexpected request ID in response"); - } - std::move(callback)(message.parameters); - return absl::OkStatus(); + // TODO(martinduke): update parameters. + return request_update_queue().OnControlMessage(message); } absl::Status MoqtSubscribeNamespaceRequestStream::OnControlMessage( const MoqtRequestError& message) { - if (message.request_id == request_id_) { - if (response_callback_ == nullptr) { - return absl::InvalidArgumentError("Two responses"); - } - std::move(response_callback_)(MoqtRequestErrorInfo{ - message.error_code, message.retry_interval, message.reason_phrase}); + if (response_callback_ != nullptr) { + std::move(response_callback_)(message); response_callback_ = nullptr; return absl::OkStatus(); } @@ -87,13 +79,7 @@ // This is irrelevant. return absl::OkStatus(); } - MoqtResponseCallback callback = task->GetResponseCallback(message.request_id); - if (callback == nullptr) { - return absl::InvalidArgumentError("Unexpected request ID in response"); - } - std::move(callback)(MoqtRequestErrorInfo{ - message.error_code, message.retry_interval, message.reason_phrase}); - return absl::OkStatus(); + return request_update_queue().OnControlMessage(message); } absl::Status MoqtSubscribeNamespaceRequestStream::OnControlMessage( @@ -185,7 +171,8 @@ return; } MoqtRequestUpdate message{next_request_id_, state_->request_id_, parameters}; - pending_updates_[message.request_id] = std::move(response_callback); + state_->request_update_queue().Enqueue(parameters, + std::move(response_callback)); state_->SendOrBufferMessageOrFatal( state_->framer()->SerializeRequestUpdate(message)); next_request_id_ += 2; @@ -234,18 +221,6 @@ } } -MoqtResponseCallback -MoqtSubscribeNamespaceRequestStream::NamespaceTask::GetResponseCallback( - uint64_t request_id) { - auto it = pending_updates_.find(request_id); - if (it == pending_updates_.end()) { - return nullptr; - } - MoqtResponseCallback callback = std::move(it->second); - pending_updates_.erase(it); - return callback; -} - MoqtSubscribeNamespaceResponseStream::MoqtSubscribeNamespaceResponseStream( MoqtFramer* framer, const MoqtControlMessageParser& message_parser, AddPrefixCallback add_callback, RemovePrefixCallback remove_callback, @@ -272,8 +247,7 @@ } if (!std::move(add_callback_)(message.track_namespace_prefix)) { add_callback_ = nullptr; - return SendRequestError(request_id_, RequestErrorCode::kPrefixOverlap, - std::nullopt, "", /*fin=*/true); + return SendRequestError(RequestErrorCode::kPrefixOverlap, std::nullopt, ""); } add_callback_ = nullptr; QUICHE_DCHECK(task_ == nullptr); @@ -368,18 +342,17 @@ uint64_t request_id) { return [this, request_id]( std::variant<MessageParameters, MoqtRequestErrorInfo> response) { - std::visit(absl::Overload{ - [this, request_id](const MessageParameters& parameters) { - // In draft-18, there are no useful parameters in - // SUBSCRIBE_NAMESPACE_OK, but Issue #1639 would change - // that. - CheckStatus(SendRequestOk(request_id, parameters)); - }, - [this, request_id](const MoqtRequestErrorInfo& error_info) { - CheckStatus(SendRequestError(request_id, error_info, - /*fin=*/true)); - }}, - response); + std::visit( + absl::Overload{[this, request_id](const MessageParameters& parameters) { + // In draft-18, there are no useful parameters in + // SUBSCRIBE_NAMESPACE_OK, but Issue #1639 would change + // that. + CheckStatus(SendRequestOk(request_id, parameters)); + }, + [this](const MoqtRequestErrorInfo& error_info) { + CheckStatus(SendRequestError(error_info)); + }}, + response); }; }
diff --git a/quiche/quic/moqt/moqt_namespace_stream.h b/quiche/quic/moqt/moqt_namespace_stream.h index 4d1bdbc..9820993 100644 --- a/quiche/quic/moqt/moqt_namespace_stream.h +++ b/quiche/quic/moqt/moqt_namespace_stream.h
@@ -12,7 +12,6 @@ #include <utility> #include "absl/base/nullability.h" -#include "absl/container/flat_hash_map.h" #include "absl/container/flat_hash_set.h" #include "absl/status/status.h" #include "quiche/quic/moqt/moqt_bidi_stream.h" @@ -113,7 +112,6 @@ // The stream is closed, so no more NAMESPACE messages are forthcoming. // This is an implicit NAMESPACE_DONE for all published namespaces. void DeclareEof(); - MoqtResponseCallback GetResponseCallback(uint64_t request_id); quiche::QuicheWeakPtr<NamespaceTask> GetWeakPtr() { return weak_ptr_factory_.Create(); } @@ -133,7 +131,6 @@ std::optional<webtransport::StreamErrorCode> error_; bool eof_ = false; uint64_t next_request_id_; - absl::flat_hash_map<uint64_t, MoqtResponseCallback> pending_updates_; // Must be last. quiche::QuicheWeakPtrFactory<NamespaceTask> weak_ptr_factory_; };
diff --git a/quiche/quic/moqt/moqt_namespace_stream_test.cc b/quiche/quic/moqt/moqt_namespace_stream_test.cc index fe14190..6872e85 100644 --- a/quiche/quic/moqt/moqt_namespace_stream_test.cc +++ b/quiche/quic/moqt/moqt_namespace_stream_test.cc
@@ -24,7 +24,6 @@ #include "quiche/quic/moqt/moqt_parser.h" #include "quiche/quic/moqt/moqt_session_callbacks.h" #include "quiche/quic/moqt/moqt_session_interface.h" -#include "quiche/quic/moqt/moqt_types.h" #include "quiche/quic/moqt/test_tools/moqt_framer_utils.h" #include "quiche/quic/moqt/test_tools/moqt_mock_visitor.h" #include "quiche/common/platform/api/quiche_test.h" @@ -91,25 +90,11 @@ ReceiveControlMessage(MoqtRequestOk{kRequestId}); } -TEST_F(MoqtSubscribeNamespaceRequestStreamTest, RequestOkWrongId) { - EXPECT_CALL(error_callback_, Call(MoqtError::kProtocolViolation, - "Unexpected request ID in response")); - ReceiveControlMessage(MoqtRequestOk{kRequestId + 1}); -} - TEST_F(MoqtSubscribeNamespaceRequestStreamTest, RequestError) { EXPECT_CALL(response_callback_, Call); ReceiveControlMessage( - MoqtRequestError{kRequestId, RequestErrorCode::kInternalError, - quic::QuicTimeDelta::FromMilliseconds(100), "bar"}); -} - -TEST_F(MoqtSubscribeNamespaceRequestStreamTest, RequestErrorWrongId) { - EXPECT_CALL(error_callback_, Call(MoqtError::kProtocolViolation, - "Unexpected request ID in response")); - ReceiveControlMessage( - MoqtRequestError{kRequestId + 1, RequestErrorCode::kInternalError, - quic::QuicTimeDelta::FromMilliseconds(100), "bar"}); + MoqtRequestError(RequestErrorCode::kInternalError, + quic::QuicTimeDelta::FromMilliseconds(100), "bar")); } TEST_F(MoqtSubscribeNamespaceRequestStreamTest, NamespaceBeforeResponse) { @@ -290,8 +275,8 @@ task_->Update(update_params, update_response_callback.AsStdFunction()); EXPECT_CALL(update_response_callback, Call(_)); ReceiveControlMessage( - MoqtRequestError{kRequestId + 2, RequestErrorCode::kInternalError, - quic::QuicTimeDelta::FromMilliseconds(100), "bar"}); + MoqtRequestError(RequestErrorCode::kInternalError, + quic::QuicTimeDelta::FromMilliseconds(100), "bar")); } class MoqtSubscribeNamespaceResponseStreamTest
diff --git a/quiche/quic/moqt/moqt_object_subscriber.cc b/quiche/quic/moqt/moqt_object_subscriber.cc index 478044e..1db5534 100644 --- a/quiche/quic/moqt/moqt_object_subscriber.cc +++ b/quiche/quic/moqt/moqt_object_subscriber.cc
@@ -17,7 +17,6 @@ #include "quiche/quic/core/quic_alarm_factory.h" #include "quiche/quic/core/quic_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_messages.h" @@ -59,24 +58,24 @@ default_publisher_priority_ = data.extensions.default_publisher_priority(); dynamic_groups_ = data.extensions.dynamic_groups(); visitor_->OnReply(full_track_name(), data); - OnObjectOrOk(); + error_is_allowed_ = false; } -void LiveSubscriber::OnStreamOpened() { +void LiveSubscriber::OnStreamOpened(webtransport::StreamVisitor*) { ++currently_open_streams_; if (publish_done_alarm_ != nullptr && publish_done_alarm_->IsSet()) { publish_done_alarm_->Cancel(); } } -void LiveSubscriber::OnStreamClosed(bool fin_received, +void LiveSubscriber::OnStreamClosed(absl::Status status, std::optional<DataStreamIndex> index) { ++streams_closed_; --currently_open_streams_; QUICHE_DCHECK_GE(currently_open_streams_, -1); if (index.has_value() && visitor_ != nullptr) { // If index is nullopt, there was not an object received on the stream. - if (fin_received) { + if (status.ok()) { visitor_->OnStreamFin(full_track_name(), *index); } else { visitor_->OnStreamReset(full_track_name(), *index); @@ -177,51 +176,7 @@ {group_id, object_id, delta_from_deadline})); } -UpstreamFetch::~UpstreamFetch() { - UpstreamFetchTask* task = task_.GetIfAvailable(); - if (task != nullptr) { - // Notify the task (which the application owns) that nothing more is coming. - // If this has already been called, UpstreamFetchTask will ignore it. - task->OnStreamAndFetchClosed(kResetCodeCancelled, ""); - } - task = nullptr; -} - -void UpstreamFetch::OnFetchResult(Location largest_location, - absl::Status status, - TaskDestroyedCallback callback) { - if (!status.ok()) { - std::move(ok_callback_)(std::make_unique<MoqtFailedFetch>(status)); - // This is called from OnRequestError, which will delete UpstreamFetch. So - // there is no need to call |callback|, which would inappropriately send a - // FETCH_CANCEL. - return; - } - auto task = std::make_unique<UpstreamFetchTask>(largest_location, status, - std::move(callback)); - task_ = task->weak_ptr(); - if (relative_groups_.has_value() && - (*relative_groups_ < largest_location.group)) { - start_ = Location(largest_location.group - *relative_groups_, 0); - relative_groups_.reset(); - } - end_ = std::min(end_, largest_location); - std::move(ok_callback_)(std::move(task)); - if (can_read_callback_) { - task_.GetIfAvailable()->set_can_read_callback( - std::move(can_read_callback_)); - } -} - -void UpstreamFetch::OnStreamOpened(CanReadCallback can_read_callback) { - if (task_.IsValid()) { - task_.GetIfAvailable()->set_can_read_callback(std::move(can_read_callback)); - } else { - can_read_callback_ = std::move(can_read_callback); - } -} - -UpstreamFetch::UpstreamFetchTask::~UpstreamFetchTask() { +UpstreamFetchTask::~UpstreamFetchTask() { // Set status_ so that callbacks into UpstreamFetchTask exit early. status_ = absl::CancelledError("UpstreamFetchTask destroyed"); if (task_destroyed_callback_) { @@ -229,8 +184,8 @@ } } -MoqtFetchTask::GetNextObjectResult -UpstreamFetch::UpstreamFetchTask::GetNextObject(PublishedObject& output) { +MoqtFetchTask::GetNextObjectResult UpstreamFetchTask::GetNextObject( + PublishedObject& output) { if (!next_object_.has_value()) { if (!status_.ok()) { return kError; @@ -257,21 +212,18 @@ output.metadata.publisher_priority = next_object_->publisher_priority; output.metadata.payload_length = next_object_->payload_length; output.fin_after_this = false; - // TODO(martinduke): Make sure the whole object has been delivered. - if (output.metadata.location == - largest_location_) { // This is the last object. - eof_ = true; - } if (payload_offset_ == next_object_->payload_length) { next_object_.reset(); payload_offset_ = 0; payload_length_ = 0; } - can_read_callback_(); + if (can_read_callback_) { + can_read_callback_(); + } return kSuccess; } -void UpstreamFetch::UpstreamFetchTask::NewObject(const MoqtObject& message) { +void UpstreamFetchTask::NewObject(const MoqtObject& message) { next_object_ = message; while (!payload_.empty()) { payload_.pop_front(); @@ -280,8 +232,7 @@ payload_length_ = 0; } -void UpstreamFetch::UpstreamFetchTask::AppendPayloadToObject( - absl::string_view payload) { +void UpstreamFetchTask::AppendPayloadToObject(absl::string_view payload) { QUICHE_BUG_IF(quic_bug_AppendPayloadToObjectCalledEarly, !next_object_.has_value()) << "AppendPayloadToObject called without an object"; @@ -292,28 +243,26 @@ payload_.push_back(quiche::QuicheMemSlice::Copy(payload)); } -void UpstreamFetch::UpstreamFetchTask::NotifyNewObject() { +void UpstreamFetchTask::NotifyNewObject() { if (need_object_available_callback_ && object_available_callback_) { need_object_available_callback_ = false; object_available_callback_(); } } -void UpstreamFetch::UpstreamFetchTask::OnStreamAndFetchClosed( - std::optional<webtransport::StreamErrorCode> error, - absl::string_view reason_phrase) { +void UpstreamFetchTask::OnStreamAndFetchClosed(absl::Status status) { if (eof_ || !status_.ok()) { return; } - // Delete callbacks, because IncomingDataStream and UpstreamFetch are gone. + // Delete callbacks, because IncomingDataStream and FetchRequestStream are + // gone. can_read_callback_ = nullptr; task_destroyed_callback_ = nullptr; - if (!error.has_value()) { // This was a FIN. - eof_ = true; - } else { - status_ = MoqtStreamErrorToStatus(*error, reason_phrase); - } - if (object_available_callback_) { + status_ = status; + eof_ = status.ok(); + if (need_object_available_callback_ && + object_available_callback_ != nullptr) { + need_object_available_callback_ = false; object_available_callback_(); } }
diff --git a/quiche/quic/moqt/moqt_object_subscriber.h b/quiche/quic/moqt/moqt_object_subscriber.h index a5da585..6325d3f 100644 --- a/quiche/quic/moqt/moqt_object_subscriber.h +++ b/quiche/quic/moqt/moqt_object_subscriber.h
@@ -12,7 +12,6 @@ #include <optional> #include <utility> -#include "absl/base/nullability.h" #include "absl/status/status.h" #include "absl/strings/string_view.h" #include "quiche/quic/core/quic_alarm.h" @@ -27,7 +26,6 @@ #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" #include "quiche/common/quiche_circular_deque.h" @@ -56,12 +54,12 @@ virtual ~ObjectSubscriber() {} const FullTrackName& full_track_name() const { return full_track_name_; } - // If REQUEST_ERROR arrives after OK or an object, it is a protocol violation. - virtual void OnObjectOrOk() { error_is_allowed_ = false; } - bool ErrorIsAllowed() const { return error_is_allowed_; } - uint64_t request_id() const { return request_id_; } + virtual void OnStreamOpened(webtransport::StreamVisitor* stream) = 0; + virtual void OnStreamClosed(absl::Status status, + std::optional<DataStreamIndex> index) = 0; + // Is the object one that was requested? virtual bool InWindow(Location sequence) const = 0; @@ -88,9 +86,6 @@ const uint64_t request_id_; MoqtBidiStreamBase* request_stream_; MessageParameters parameters_; - // If false, an object or OK message has been received, so any ERROR message - // is a protocol violation. - bool error_is_allowed_ = true; // Must be last. quiche::QuicheWeakPtrFactory<ObjectSubscriber> weak_ptr_factory_; @@ -117,15 +112,19 @@ ~LiveSubscriber() override; void OnObjectOrOk(const SubscribeOkData& data); - void OnObjectOrOk() override { ObjectSubscriber::OnObjectOrOk(); } + void OnObjectOrOk() { error_is_allowed_ = false; } + // If REQUEST_ERROR arrives after OK or an object, it is a protocol violation. + bool ErrorIsAllowed() const { return error_is_allowed_; } std::optional<uint64_t> track_alias() const { return track_alias_; } // Returns false if the callback returns false, meaning the session has been // destroyed. void set_track_alias(uint64_t track_alias) { track_alias_.emplace(track_alias); } - void OnStreamOpened(); - void OnStreamClosed(bool fin_received, std::optional<DataStreamIndex> index); + void OnStreamOpened(webtransport::StreamVisitor* /*unused*/) override; + // If |status.ok()|, it was a FIN. + void OnStreamClosed(absl::Status status, + std::optional<DataStreamIndex> index) override; void OnPublishDone(uint64_t stream_count, const quic::QuicClock* clock, quic::QuicAlarmFactory* alarm_factory); @@ -187,6 +186,9 @@ bool all_streams_closed() const { return total_streams_.has_value() && *total_streams_ == streams_closed_; } + // If false, an object or OK message has been received, so any ERROR message + // is a protocol violation. + bool error_is_allowed_ = true; quic::QuicTimeDelta publisher_delivery_timeout_ = kDefaultDeliveryTimeout; MoqtPriority default_publisher_priority_ = kDefaultPublisherPriority; @@ -213,197 +215,86 @@ // application reads it. using CanReadCallback = quiche::MultiUseCallback<void()>; -// If the application destroys the FetchTask, this is a signal to MoqtSession to -// cancel the FETCH and STOP_SENDING the stream. +// If the application destroys the FetchTask, this is a signal to +// the owner to cancel the FETCH and STOP_SENDING the stream. using TaskDestroyedCallback = quiche::SingleUseCallback<void()>; -// Class for upstream FETCH. It will notify the application using |callback| -// when a FETCH_OK or REQUEST_ERROR is received. -using RemoveFetchCallback = quiche::SingleUseCallback<void()>; -class UpstreamFetch : public ObjectSubscriber { +// This class is passed to the application, which views it as a MoqtFetchTask. +// A pointer to the child calls is held by the FETCH control at first, then the +// data stream when initiated, to update its state. UpstreamFetchTask is +// responsible for calling task_destroyed_callback_ so that pointers to it are +// cleared, so there is no need for a QuicheWeakPtr. +class QUICHE_EXPORT UpstreamFetchTask : public MoqtFetchTask { public: - // Standalone Fetch constructor - UpstreamFetch(const MoqtFetch& fetch, const StandaloneFetch standalone, - FetchResponseCallback callback, - RemoveFetchCallback delete_callback) - : ObjectSubscriber(standalone.full_track_name, fetch.request_id, - fetch.parameters, /*request_stream=*/nullptr), - group_order_(fetch.parameters.group_order.value_or( - MoqtDeliveryOrder::kAscending)), - start_(standalone.start_location), - end_(standalone.end_location), - subscriber_priority_(fetch.parameters.subscriber_priority.value_or( - kDefaultSubscriberPriority)), - ok_callback_(std::move(callback)), - remove_callback_(std::move(delete_callback)) {} - // Relative Joining Fetch constructor - UpstreamFetch(const MoqtFetch& fetch, FullTrackName full_track_name, - FetchResponseCallback callback, - RemoveFetchCallback delete_callback) - : ObjectSubscriber(full_track_name, fetch.request_id, fetch.parameters, - /*request_stream=*/nullptr), - group_order_(fetch.parameters.group_order.value_or( - MoqtDeliveryOrder::kAscending)), - relative_groups_( - std::get<JoiningFetchRelative>(fetch.fetch).joining_start), - subscriber_priority_(fetch.parameters.subscriber_priority.value_or( - kDefaultSubscriberPriority)), - ok_callback_(std::move(callback)), - remove_callback_(std::move(delete_callback)) {} - // Absolute Joining Fetch constructor - UpstreamFetch(const MoqtFetch& fetch, FullTrackName full_track_name, - JoiningFetchAbsolute absolute_joining, - FetchResponseCallback callback, - RemoveFetchCallback delete_callback) - : ObjectSubscriber(full_track_name, fetch.request_id, fetch.parameters, - /*request_stream=*/nullptr), - group_order_(fetch.parameters.group_order.value_or( - MoqtDeliveryOrder::kAscending)), - start_(Location(absolute_joining.joining_start, 0)), - subscriber_priority_(fetch.parameters.subscriber_priority.value_or( - kDefaultSubscriberPriority)), - ok_callback_(std::move(callback)), - remove_callback_(std::move(delete_callback)) {} - UpstreamFetch(const UpstreamFetch&) = delete; - ~UpstreamFetch(); + // If the FetchRequestStream is destroyed, it will call OnStreamAndFetchClosed + // which sets the TaskDestroyedCallback to nullptr. Thus, |callback| can + // assume that FetchRequestStream is valid. + UpstreamFetchTask() {} + ~UpstreamFetchTask() override; - bool InWindow(Location location) const override { - return (location >= start_ && location <= end_); - } + // Implementation of MoqtFetchTask. + GetNextObjectResult GetNextObject(PublishedObject& output) override; + void SetObjectAvailableCallback(ObjectsAvailableCallback callback) override { + object_available_callback_ = std::move(callback); + }; + absl::Status GetStatus() override { return status_; }; - // Called when the data stream is destroyed. - void OnStreamClosed() { Destroy(); } - - void Destroy() { - if (remove_callback_) { - RemoveFetchCallback callback = std::move(remove_callback_); - remove_callback_ = nullptr; - std::move(callback)(); - } - } - - class UpstreamFetchTask : public MoqtFetchTask { - public: - // If the UpstreamFetch is destroyed, it will call OnStreamAndFetchClosed - // which sets the TaskDestroyedCallback to nullptr. Thus, |callback| can - // assume that UpstreamFetch is valid. - UpstreamFetchTask(Location largest_location, absl::Status status, - TaskDestroyedCallback callback) - : largest_location_(largest_location), - status_(status), - task_destroyed_callback_(std::move(callback)), - weak_ptr_factory_(this) {} - ~UpstreamFetchTask() override; - - // Implementation of MoqtFetchTask. - GetNextObjectResult GetNextObject(PublishedObject& output) override; - void SetObjectAvailableCallback( - ObjectsAvailableCallback callback) override { - object_available_callback_ = std::move(callback); - }; - // TODO(martinduke): Implement the new API, but for now, only deliver the - // FetchTask on FETCH_OK. - void SetFetchResponseCallback(FetchResponseCallback callback) override {} - absl::Status GetStatus() override { return status_; }; - - quiche::QuicheWeakPtr<UpstreamFetchTask> weak_ptr() { - return weak_ptr_factory_.Create(); - } - - // MoqtSession should not use this function; use - // UpstreamFetch::OnStreamOpened() instead, in case the task does not exist - // yet. - void set_can_read_callback(CanReadCallback callback) { - can_read_callback_ = std::move(callback); + // Called by incoming data stream. + virtual void set_can_read_callback(CanReadCallback callback) { + can_read_callback_ = std::move(callback); + if (can_read_callback_) { can_read_callback_(); // Accept the first object. } + } + virtual void set_task_destroyed_callback(TaskDestroyedCallback callback) { + task_destroyed_callback_ = std::move(callback); + } - // Called when the data stream receives a new object. - void NewObject(const MoqtObject& message); - void AppendPayloadToObject(absl::string_view payload); - // MoqtSession calls this for a hint if the object has been read. - bool HasObject() const { return next_object_.has_value(); } - bool NeedsMorePayload() const { - return next_object_.has_value() && - payload_length_ < next_object_->payload_length; - } - // MoqtSession calls NotifyNewObject() after NewObject() because it has to - // exit the parser loop before the callback possibly causes another read. - // Furthermore, NewObject() may be a partial object, and so - // NotifyNewObject() is called only when the object is complete. - void NotifyNewObject(); + // Called when the data stream receives a new object. + virtual void NewObject(const MoqtObject& message); + virtual void AppendPayloadToObject(absl::string_view payload); + // The data stream calls this for a hint if the object has been read. + virtual bool HasObject() const { return next_object_.has_value(); } + virtual bool NeedsMorePayload() const { + return next_object_.has_value() && + payload_length_ < next_object_->payload_length; + } + // The data stream calls NotifyNewObject() after NewObject() because it has to + // exit the parser loop before the callback possibly causes another read. + // Furthermore, NewObject() may be a partial object, and so + // NotifyNewObject() is called only when the object is complete. + virtual void NotifyNewObject(); - // Deletes callbacks to session or stream, updates the status. If |error| - // has no value, will append an EOF to the object stream. - void OnStreamAndFetchClosed( - std::optional<webtransport::StreamErrorCode> error, - absl::string_view reason_phrase); + // Deletes callbacks to session or stream, updates the status. If |status| is + // OK, will append an EOF to the object stream. + virtual void OnStreamAndFetchClosed(absl::Status status); - uint64_t payload_offset() const { return payload_offset_; } - uint64_t payload_length() const { return payload_length_; } - - private: - Location largest_location_; - absl::Status status_ = absl::OkStatus(); - TaskDestroyedCallback task_destroyed_callback_; - - // Object delivery state. The payload_length member is used to track the - // payload bytes not yet received. The application receives a - // PublishedObject that is constructed from next_object_ and payload_. - std::optional<MoqtObject> next_object_; - quiche::QuicheCircularDeque<quiche::QuicheMemSlice> payload_; - // The starting point of payload_. Data is deleted as it is delivered. - uint64_t payload_offset_ = 0; - // Total data delivered for this object. - uint64_t payload_length_ = 0; - - // The task should only call object_available_callback_ when the last result - // was kPending. Otherwise, there can be recursive loops of - // GetNextObjectResult(). - bool need_object_available_callback_ = true; - bool eof_ = false; // The next object is EOF. - // The Fetch task signals the application when it has new objects. - ObjectsAvailableCallback object_available_callback_; - // The Fetch task signals the stream when it has dispensed of an object. - CanReadCallback can_read_callback_; - - // Must be last. - quiche::QuicheWeakPtrFactory<UpstreamFetchTask> weak_ptr_factory_; - }; - - // Arrival of FETCH_OK/REQUEST_ERROR. - void OnFetchResult(Location largest_location, absl::Status status, - TaskDestroyedCallback callback); - - UpstreamFetchTask* task() { return task_.GetIfAvailable(); } - - // Manage the relationship with the data stream. - void OnStreamOpened(CanReadCallback callback); - - bool is_fetch() const override { return true; } + virtual uint64_t payload_offset() const { return payload_offset_; } + virtual uint64_t payload_length() const { return payload_length_; } private: - MoqtDeliveryOrder group_order_; - Location start_ = Location(0, 0); - Location end_ = Location(kMaxGroupId, kMaxObjectId); - std::optional<uint64_t> relative_groups_; - MoqtPriority subscriber_priority_; - // The last object received on the stream. - std::optional<Location> last_location_; - // The highest location received on the stream. - std::optional<Location> highest_location_; - bool last_group_is_finished_ = false; // Received EndOfGroup. - std::optional<Location> end_of_track_; // Received EndOfTrack + absl::Status status_ = absl::OkStatus(); + TaskDestroyedCallback task_destroyed_callback_; - quiche::QuicheWeakPtr<UpstreamFetchTask> task_; + // Object delivery state. The payload_length member is used to track the + // payload bytes not yet received. The application receives a + // PublishedObject that is constructed from next_object_ and payload_. + std::optional<MoqtObject> next_object_; + quiche::QuicheCircularDeque<quiche::QuicheMemSlice> payload_; + // The starting point of payload_. Data is deleted as it is delivered. + uint64_t payload_offset_ = 0; + // Total data delivered for this object. + uint64_t payload_length_ = 0; - // Before FetchTask is created, an incoming stream will register the callback - // here instead. + // The task should only call object_available_callback_ when the last result + // was kPending. Otherwise, there can be recursive loops of + // GetNextObjectResult(). + bool need_object_available_callback_ = true; + bool eof_ = false; // The next object is EOF. + // The Fetch task signals the application when it has new objects. + ObjectsAvailableCallback object_available_callback_; + // The Fetch task signals the stream when it has dispensed of an object. CanReadCallback can_read_callback_; - - // Initial values from Fetch() call. - FetchResponseCallback ok_callback_; // Will be destroyed on FETCH_OK. - RemoveFetchCallback remove_callback_; }; } // namespace moqt
diff --git a/quiche/quic/moqt/moqt_object_subscriber_test.cc b/quiche/quic/moqt/moqt_object_subscriber_test.cc index 490001b..de32e4f 100644 --- a/quiche/quic/moqt/moqt_object_subscriber_test.cc +++ b/quiche/quic/moqt/moqt_object_subscriber_test.cc
@@ -24,7 +24,6 @@ #include "quiche/quic/test_tools/quic_test_utils.h" #include "quiche/common/quiche_mem_slice.h" #include "quiche/web_transport/test_tools/mock_web_transport.h" -#include "quiche/web_transport/web_transport.h" namespace moqt { @@ -91,31 +90,32 @@ } TEST_F(LiveSubscriberTest, OnPublishDoneReadyToClose) { - track_.OnStreamOpened(); - track_.OnStreamClosed(true, std::nullopt); + track_.OnStreamOpened(nullptr); + track_.OnStreamClosed(absl::OkStatus(), std::nullopt); EXPECT_CALL(visitor_, OnPublishDone); ExpectFin(wt_stream_); track_.OnPublishDone(1, &clock_, &alarm_factory_); } TEST_F(LiveSubscriberTest, OnPublishDoneAllStreamsCloseLater) { - track_.OnStreamOpened(); + track_.OnStreamOpened(nullptr); EXPECT_CALL(visitor_, OnPublishDone).Times(0); EXPECT_CALL(wt_stream_, Writev).Times(0); track_.OnPublishDone(2, &clock_, &alarm_factory_); - track_.OnStreamClosed(true, std::nullopt); - track_.OnStreamOpened(); + track_.OnStreamClosed(absl::OkStatus(), std::nullopt); + track_.OnStreamOpened(nullptr); ExpectFin(wt_stream_); EXPECT_CALL(visitor_, OnPublishDone); - track_.OnStreamClosed(true, std::nullopt); + track_.OnStreamClosed(absl::OkStatus(), std::nullopt); } TEST_F(LiveSubscriberTest, OnPublishDoneTimesOut) { - track_.OnStreamOpened(); + track_.OnStreamOpened(nullptr); EXPECT_CALL(visitor_, OnPublishDone).Times(0); EXPECT_CALL(wt_stream_, Writev).Times(0); track_.OnPublishDone(2, &clock_, &alarm_factory_); - track_.OnStreamClosed(true, std::nullopt); // No streams are open; timer set. + track_.OnStreamClosed(absl::OkStatus(), std::nullopt); + // No streams are open; timer set. quic::QuicAlarm* alarm = LiveSubscriberPeer::GetPublishDoneAlarm(&track_); EXPECT_NE(alarm, nullptr); EXPECT_TRUE(alarm->IsSet()); @@ -225,213 +225,147 @@ EXPECT_EQ(LiveSubscriberPeer::GetFetchTask(&track_), nullptr); } -class UpstreamFetchTest : public quic::test::QuicTest { +class UpstreamFetchTaskTest : public quiche::test::QuicheTest { protected: - UpstreamFetchTest() - : fetch_( - fetch_message_, std::get<StandaloneFetch>(fetch_message_.fetch), - [&](std::unique_ptr<MoqtFetchTask> task) { - fetch_task_ = std::move(task); - }, - [&]() { deleted_ = true; }) {} + UpstreamFetchTaskTest() { + EXPECT_CALL(task_destroyed_callback_, Call).Times(testing::AnyNumber()); + EXPECT_CALL(can_read_callback_, Call).Times(testing::AnyNumber()); + task_.set_task_destroyed_callback(task_destroyed_callback_.AsStdFunction()); + task_.set_can_read_callback(can_read_callback_.AsStdFunction()); + } - MoqtFetch fetch_message_ = { - /*request_id=*/1, - StandaloneFetch(FullTrackName("foo", "bar"), Location(1, 1), - Location(3, 100)), - MessageParameters(), - }; - // The pointer held by the application. - UpstreamFetch fetch_; - std::unique_ptr<MoqtFetchTask> fetch_task_; - bool deleted_ = false; + const Location kEndLocation = Location(3, 50); + testing::StrictMock<testing::MockFunction<void()>> task_destroyed_callback_; + testing::StrictMock<testing::MockFunction<void()>> can_read_callback_; + UpstreamFetchTask task_; }; -TEST_F(UpstreamFetchTest, Queries) { - EXPECT_EQ(fetch_.request_id(), 1); - EXPECT_EQ(fetch_.full_track_name(), FullTrackName("foo", "bar")); - EXPECT_TRUE(fetch_.is_fetch()); - EXPECT_FALSE(fetch_.InWindow(Location{1, 0})); - EXPECT_TRUE(fetch_.InWindow(Location{1, 1})); - EXPECT_TRUE(fetch_.InWindow(Location{3, 100})); - EXPECT_FALSE(fetch_.InWindow(Location{3, 101})); -} - -TEST_F(UpstreamFetchTest, AllowError) { - EXPECT_TRUE(fetch_.ErrorIsAllowed()); - fetch_.OnObjectOrOk(); - EXPECT_FALSE(fetch_.ErrorIsAllowed()); -} - -TEST_F(UpstreamFetchTest, FetchResponse) { - EXPECT_EQ(fetch_task_, nullptr); - fetch_.OnFetchResult(Location(3, 50), absl::OkStatus(), nullptr); - EXPECT_NE(fetch_task_, nullptr); - EXPECT_NE(fetch_.task(), nullptr); - EXPECT_TRUE(fetch_task_->GetStatus().ok()); -} - -TEST_F(UpstreamFetchTest, FetchClosedByMoqt) { - bool terminated = false; - fetch_.OnFetchResult(Location(3, 50), absl::OkStatus(), - [&]() { terminated = true; }); - bool got_eof = false; - fetch_task_->SetObjectAvailableCallback([&]() { - PublishedObject object; - EXPECT_EQ(fetch_task_->GetNextObject(object), - MoqtFetchTask::GetNextObjectResult::kEof); - got_eof = true; - }); - fetch_.task()->OnStreamAndFetchClosed(std::nullopt, ""); - EXPECT_FALSE(terminated); - EXPECT_TRUE(got_eof); -} - -TEST_F(UpstreamFetchTest, FetchClosedByApplication) { - bool terminated = false; - fetch_.OnFetchResult(Location(3, 50), absl::Status(), - [&]() { terminated = true; }); - fetch_task_.reset(); - EXPECT_TRUE(terminated); -} - -TEST_F(UpstreamFetchTest, ObjectRetrieval) { - fetch_.OnFetchResult(Location(3, 50), absl::OkStatus(), nullptr); - PublishedObject object; - EXPECT_EQ(fetch_task_->GetNextObject(object), - MoqtFetchTask::GetNextObjectResult::kPending); - MoqtObject new_object = {1, 3, 0, 128, "", MoqtObjectStatus::kNormal, - 0, true, 6}; - bool got_object = false; - fetch_task_->SetObjectAvailableCallback([&]() { - got_object = true; - EXPECT_EQ(fetch_task_->GetNextObject(object), - MoqtFetchTask::GetNextObjectResult::kSuccess); - EXPECT_EQ(object.metadata.location, Location(3, 0)); - EXPECT_EQ(object.metadata.subgroup, 0); - EXPECT_EQ(object.payload[0].AsStringView(), "foo"); - EXPECT_EQ(object.payload[1].AsStringView(), "bar"); - }); - int got_read_callback = 0; - fetch_.OnStreamOpened([&]() { ++got_read_callback; }); - EXPECT_FALSE(fetch_.task()->HasObject()); - EXPECT_FALSE(fetch_.task()->NeedsMorePayload()); - fetch_.task()->NewObject(new_object); - EXPECT_TRUE(fetch_.task()->HasObject()); - EXPECT_TRUE(fetch_.task()->NeedsMorePayload()); - fetch_.task()->AppendPayloadToObject("foo"); - EXPECT_TRUE(fetch_.task()->HasObject()); - EXPECT_TRUE(fetch_.task()->NeedsMorePayload()); - fetch_.task()->AppendPayloadToObject("bar"); - EXPECT_TRUE(fetch_.task()->HasObject()); - EXPECT_FALSE(fetch_.task()->NeedsMorePayload()); - EXPECT_FALSE(got_object); - EXPECT_EQ(got_read_callback, 1); // Call from OnStreamOpened(). - fetch_.task()->NotifyNewObject(); - EXPECT_FALSE(fetch_.task()->HasObject()); - EXPECT_FALSE(fetch_.task()->NeedsMorePayload()); - EXPECT_EQ(got_read_callback, 2); // Call from GetNextObjectResult(). - EXPECT_TRUE(got_object); -} - -TEST_F(UpstreamFetchTest, ObjectRetrievalEmptyPayload) { - fetch_.OnFetchResult(Location(3, 50), absl::OkStatus(), nullptr); - MoqtObject moqt_obj = {1, 3, 0, 128, "", MoqtObjectStatus::kEndOfGroup, - 0, true, 0}; - fetch_.task()->NewObject(moqt_obj); - fetch_.task()->NotifyNewObject(); - fetch_.OnStreamOpened([]() {}); +TEST_F(UpstreamFetchTaskTest, ObjectRetrievalMultiSlice) { + int can_read_calls = 0; + task_.set_can_read_callback([&]() { ++can_read_calls; }); + EXPECT_EQ(can_read_calls, 1); PublishedObject output; - EXPECT_EQ(fetch_task_->GetNextObject(output), + EXPECT_EQ(task_.GetNextObject(output), + MoqtFetchTask::GetNextObjectResult::kPending); + MoqtObject new_object = { + /*track_alias=*/1, + /*group_id=*/3, + /*object_id=*/0, + /*publisher_priority=*/128, + /*extension_headers=*/"", + /*object_status=*/MoqtObjectStatus::kNormal, + /*subgroup_id=*/1, + /*first_object_in_subgroup=*/true, + /*payload_length=*/6, + }; + EXPECT_FALSE(task_.HasObject()); + EXPECT_FALSE(task_.NeedsMorePayload()); + + task_.NewObject(new_object); + EXPECT_TRUE(task_.HasObject()); + EXPECT_TRUE(task_.NeedsMorePayload()); + EXPECT_EQ(task_.payload_length(), 0); + EXPECT_EQ(task_.payload_offset(), 0); + + task_.AppendPayloadToObject("foo"); + EXPECT_TRUE(task_.HasObject()); + EXPECT_TRUE(task_.NeedsMorePayload()); + EXPECT_EQ(task_.payload_length(), 3); + + task_.AppendPayloadToObject("bar"); + EXPECT_TRUE(task_.HasObject()); + EXPECT_FALSE(task_.NeedsMorePayload()); + EXPECT_EQ(task_.payload_length(), 6); + + bool object_available_called = false; + task_.SetObjectAvailableCallback([&]() { object_available_called = true; }); + task_.NotifyNewObject(); + EXPECT_TRUE(object_available_called); + + EXPECT_EQ(task_.GetNextObject(output), + MoqtFetchTask::GetNextObjectResult::kSuccess); + EXPECT_EQ(output.metadata.location, Location(3, 0)); + EXPECT_EQ(output.metadata.subgroup, 1); + EXPECT_EQ(output.metadata.status, MoqtObjectStatus::kNormal); + EXPECT_EQ(output.metadata.publisher_priority, 128); + EXPECT_EQ(output.metadata.payload_length, 6); + EXPECT_FALSE(output.fin_after_this); + ASSERT_EQ(output.payload.size(), 2); + EXPECT_EQ(output.payload[0].AsStringView(), "foo"); + EXPECT_EQ(output.payload[1].AsStringView(), "bar"); + EXPECT_EQ(can_read_calls, 2); + + EXPECT_FALSE(task_.HasObject()); + EXPECT_FALSE(task_.NeedsMorePayload()); +} + +TEST_F(UpstreamFetchTaskTest, ObjectRetrievalEmptyPayload) { + MoqtObject moqt_obj = { + /*track_alias=*/1, + /*group_id=*/3, + /*object_id=*/0, + /*publisher_priority=*/128, + /*extension_headers=*/"", + /*object_status=*/MoqtObjectStatus::kEndOfGroup, + /*subgroup_id=*/0, + /*first_object_in_subgroup=*/true, + /*payload_length=*/0, + }; + task_.NewObject(moqt_obj); + task_.NotifyNewObject(); + + PublishedObject output; + EXPECT_EQ(task_.GetNextObject(output), MoqtFetchTask::GetNextObjectResult::kSuccess); EXPECT_TRUE(output.payload.empty()); EXPECT_EQ(output.metadata.status, MoqtObjectStatus::kEndOfGroup); + EXPECT_EQ(output.metadata.location, Location(3, 0)); } -TEST_F(UpstreamFetchTest, GetNextObjectAfterEof) { - fetch_.OnFetchResult(Location(3, 50), absl::OkStatus(), nullptr); - fetch_.task()->OnStreamAndFetchClosed(std::nullopt, ""); +TEST_F(UpstreamFetchTaskTest, PartialPayloadPending) { + MoqtObject moqt_obj = { + /*track_alias=*/1, + /*group_id=*/3, + /*object_id=*/0, + /*publisher_priority=*/128, + /*extension_headers=*/"", + /*object_status=*/MoqtObjectStatus::kNormal, + /*subgroup_id=*/0, + /*first_object_in_subgroup=*/true, + /*payload_length=*/10, + }; + task_.NewObject(moqt_obj); - PublishedObject object; - EXPECT_EQ(fetch_task_->GetNextObject(object), - MoqtFetchTask::GetNextObjectResult::kEof); - // Subsequent calls should still return EOF. - EXPECT_EQ(fetch_task_->GetNextObject(object), - MoqtFetchTask::GetNextObjectResult::kEof); -} - -TEST_F(UpstreamFetchTest, GetNextObjectEofAtLargestLocation) { - Location largest(3, 50); - fetch_.OnFetchResult(largest, absl::OkStatus(), nullptr); - fetch_.OnStreamOpened([]() {}); - - MoqtObject obj1 = {1, 3, 49, 128, "", MoqtObjectStatus::kNormal, 0, false, 1}; - fetch_.task()->NewObject(obj1); - fetch_.task()->AppendPayloadToObject("a"); - fetch_.task()->NotifyNewObject(); - - PublishedObject out; - EXPECT_EQ(fetch_task_->GetNextObject(out), - MoqtFetchTask::GetNextObjectResult::kSuccess); - // Not at largest location yet. - EXPECT_EQ(fetch_task_->GetNextObject(out), + PublishedObject output; + EXPECT_EQ(task_.GetNextObject(output), MoqtFetchTask::GetNextObjectResult::kPending); - - MoqtObject obj2 = {1, 3, 50, 128, "", MoqtObjectStatus::kNormal, 0, false, 1}; - fetch_.task()->NewObject(obj2); - fetch_.task()->AppendPayloadToObject("b"); - fetch_.task()->NotifyNewObject(); - - EXPECT_EQ(fetch_task_->GetNextObject(out), - MoqtFetchTask::GetNextObjectResult::kSuccess); - // Reached largest location. EOF should be set. - EXPECT_EQ(fetch_task_->GetNextObject(out), - MoqtFetchTask::GetNextObjectResult::kEof); } -TEST_F(UpstreamFetchTest, CloseWithError) { - fetch_.OnFetchResult(Location(3, 50), absl::OkStatus(), nullptr); - fetch_.task()->OnStreamAndFetchClosed( - static_cast<webtransport::StreamErrorCode>(0x123), "reason"); +TEST_F(UpstreamFetchTaskTest, OnStreamAndFetchClosedFin) { + task_.OnStreamAndFetchClosed(absl::OkStatus()); + PublishedObject out; - EXPECT_EQ(fetch_task_->GetNextObject(out), + EXPECT_EQ(task_.GetNextObject(out), MoqtFetchTask::GetNextObjectResult::kEof); + EXPECT_EQ(task_.GetNextObject(out), MoqtFetchTask::GetNextObjectResult::kEof); + EXPECT_TRUE(task_.GetStatus().ok()); +} + +TEST_F(UpstreamFetchTaskTest, OnStreamAndFetchClosedError) { + task_.OnStreamAndFetchClosed(absl::InternalError("custom reason")); + + PublishedObject out; + EXPECT_EQ(task_.GetNextObject(out), MoqtFetchTask::GetNextObjectResult::kError); - EXPECT_FALSE(fetch_task_->GetStatus().ok()); + EXPECT_FALSE(task_.GetStatus().ok()); } -TEST_F(UpstreamFetchTest, RelativeJoiningFetch) { - MoqtFetch relative_fetch_message = { - /*request_id=*/2, - JoiningFetchRelative(1, 2), - MessageParameters(), - }; - UpstreamFetch relative_fetch( - relative_fetch_message, FullTrackName("foo", "bar"), - [&](std::unique_ptr<MoqtFetchTask> task) { - fetch_task_ = std::move(task); - }, - []() {}); - relative_fetch.OnFetchResult(Location(10, 50), absl::OkStatus(), nullptr); - EXPECT_FALSE(relative_fetch.InWindow(Location(7, 35))); - EXPECT_TRUE(relative_fetch.InWindow(Location(8, 0))); -} - -TEST_F(UpstreamFetchTest, RelativeJoiningFetchUnderflow) { - MoqtFetch relative_fetch_message = { - /*request_id=*/2, - JoiningFetchRelative(1, 10), - MessageParameters(), - }; - UpstreamFetch relative_fetch( - relative_fetch_message, FullTrackName("foo", "bar"), - [&](std::unique_ptr<MoqtFetchTask> task) { - fetch_task_ = std::move(task); - }, - []() {}); - relative_fetch.OnFetchResult(Location(1, 50), absl::OkStatus(), nullptr); - EXPECT_TRUE(relative_fetch.InWindow(Location(0, 0))); - EXPECT_TRUE(relative_fetch.InWindow(Location(1, 50))); +TEST_F(UpstreamFetchTaskTest, DestroyedByApplication) { + testing::StrictMock<testing::MockFunction<void()>> callback; + EXPECT_CALL(callback, Call); + auto dynamic_task = std::make_unique<UpstreamFetchTask>(); + dynamic_task->set_task_destroyed_callback(callback.AsStdFunction()); + dynamic_task.reset(); } } // namespace test
diff --git a/quiche/quic/moqt/moqt_outgoing_queue.cc b/quiche/quic/moqt/moqt_outgoing_queue.cc index 9472a89..77bf351 100644 --- a/quiche/quic/moqt/moqt_outgoing_queue.cc +++ b/quiche/quic/moqt/moqt_outgoing_queue.cc
@@ -13,10 +13,13 @@ #include "absl/algorithm/container.h" #include "absl/status/status.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_object.h" #include "quiche/quic/moqt/moqt_priority.h" #include "quiche/quic/moqt/moqt_publisher.h" +#include "quiche/quic/moqt/moqt_session_callbacks.h" #include "quiche/quic/moqt/moqt_types.h" #include "quiche/common/platform/api/quiche_bug_tracker.h" #include "quiche/common/quiche_mem_slice.h" @@ -132,10 +135,13 @@ } std::unique_ptr<MoqtFetchTask> MoqtOutgoingQueue::StandaloneFetch( - Location start, Location end, MoqtDeliveryOrder order) { + Location start, Location end, MoqtDeliveryOrder order, + FetchResponseCallback callback) { if (queue_.empty()) { - return std::make_unique<MoqtFailedFetch>( - absl::NotFoundError("No objects available on the track")); + std::move(callback)( + MoqtRequestErrorInfo(RequestErrorCode::kInvalidRange, std::nullopt, + "No objects available on the track")); + return nullptr; } Location first_available_object = Location(first_group_in_queue(), 0); @@ -143,39 +149,58 @@ Location(current_group_id_, queue_.back().size() - 1); if (end < first_available_object) { - return std::make_unique<MoqtFailedFetch>( - absl::NotFoundError("All of the requested objects have expired")); + std::move(callback)( + MoqtRequestErrorInfo(RequestErrorCode::kInvalidRange, std::nullopt, + "All of the requested objects have expired")); + return nullptr; } if (start > last_available_object) { - return std::make_unique<MoqtFailedFetch>( - absl::NotFoundError("All of the requested objects are in the future")); + std::move(callback)( + MoqtRequestErrorInfo(RequestErrorCode::kInvalidRange, std::nullopt, + "All of the requested objects are in the future")); + return nullptr; } Location adjusted_start = std::max(start, first_available_object); Location adjusted_end = std::min(end, last_available_object); std::vector<Location> objects = GetCachedObjectsInRange(adjusted_start, adjusted_end); + if (objects.empty()) { + std::move(callback)( + MoqtRequestErrorInfo(RequestErrorCode::kInvalidRange, std::nullopt, + "No objects in the requested range")); + return nullptr; + } // Default to ascending order. if (order == MoqtDeliveryOrder::kDescending) { ObjectsInDescendingOrder(objects); } + FetchOkData ok(closed_ && adjusted_end == largest_location(), adjusted_end, + MessageParameters(), extensions_); + std::move(callback)(ok); return std::make_unique<FetchTask>(this, std::move(objects)); } std::unique_ptr<MoqtFetchTask> MoqtOutgoingQueue::RelativeFetch( - uint64_t /*group_diff*/, MoqtDeliveryOrder /*order*/) { + uint64_t /*group_diff*/, MoqtDeliveryOrder, + FetchResponseCallback callback) { QUICHE_BUG(MoqtOutgoingQueue_RelativeFetch) << "Calling RelativeFetch() on an established subscription"; - return std::make_unique<MoqtFailedFetch>(absl::InternalError( + std::move(callback)(MoqtRequestErrorInfo( + RequestErrorCode::kNotSupported, std::nullopt, "RelativeFetch called on an established subscription")); + return nullptr; } std::unique_ptr<MoqtFetchTask> MoqtOutgoingQueue::AbsoluteFetch( - uint64_t /*group*/, MoqtDeliveryOrder /*order*/) { + uint64_t /*group*/, MoqtDeliveryOrder /*order*/, + FetchResponseCallback callback) { QUICHE_BUG(MoqtOutgoingQueue_AbsoluteFetch) << "Calling AbsoluteFetch() on an established subscription"; - return std::make_unique<MoqtFailedFetch>(absl::InternalError( - "AbsoluteFetch called on an established subscription")); + std::move(callback)(MoqtRequestErrorInfo( + RequestErrorCode::kNotSupported, std::nullopt, + "RelativeFetch called on an established subscription")); + return nullptr; } MoqtFetchTask::GetNextObjectResult MoqtOutgoingQueue::FetchTask::GetNextObject(
diff --git a/quiche/quic/moqt/moqt_outgoing_queue.h b/quiche/quic/moqt/moqt_outgoing_queue.h index 2d64c0d..f46f37e 100644 --- a/quiche/quic/moqt/moqt_outgoing_queue.h +++ b/quiche/quic/moqt/moqt_outgoing_queue.h
@@ -27,6 +27,7 @@ #include "quiche/quic/moqt/moqt_object.h" #include "quiche/quic/moqt/moqt_priority.h" #include "quiche/quic/moqt/moqt_publisher.h" +#include "quiche/quic/moqt/moqt_session_callbacks.h" #include "quiche/quic/moqt/moqt_types.h" #include "quiche/common/quiche_circular_deque.h" #include "quiche/common/quiche_mem_slice.h" @@ -75,13 +76,16 @@ const TrackExtensions& extensions() const override { return extensions_; } std::unique_ptr<MoqtFetchTask> StandaloneFetch( - Location start, Location end, MoqtDeliveryOrder order) override; + Location start, Location end, MoqtDeliveryOrder order, + FetchResponseCallback callback) override; // Joining Fetch functions should never be called because subscriptions are // never pending in MoqtOutgoingQueue. std::unique_ptr<MoqtFetchTask> RelativeFetch( - uint64_t group_diff, MoqtDeliveryOrder order) override; + uint64_t group_diff, MoqtDeliveryOrder order, + FetchResponseCallback callback) override; std::unique_ptr<MoqtFetchTask> AbsoluteFetch( - uint64_t group, MoqtDeliveryOrder order) override; + uint64_t group, MoqtDeliveryOrder order, + FetchResponseCallback callback) override; bool HasSubscribers() const { return !listeners_.empty(); } @@ -124,31 +128,6 @@ // guaranteed to resolve immediately. callback(); } - void SetFetchResponseCallback(FetchResponseCallback callback) override { - if (!status_.ok()) { - MoqtRequestError error(0, StatusToRequestErrorCode(status_), - std::nullopt, std::string(status_.message())); - std::move(callback)(error); - return; - } - if (objects_.empty()) { - MoqtRequestError error(0, StatusToRequestErrorCode(status_), - std::nullopt, "No objects in range"); - std::move(callback)(error); - return; - } - MoqtFetchOk ok; - ok.end_location = *(objects_.crbegin()); - if (objects_.size() > 1 && *(objects_.cbegin()) > ok.end_location) { - ok.extensions = TrackExtensions( - std::nullopt, std::nullopt, std::nullopt, - MoqtDeliveryOrder::kDescending, std::nullopt, std::nullopt); - ok.end_location = *(objects_.cbegin()); - } - ok.end_of_track = - queue_->closed_ && ok.end_location == queue_->largest_location(); - std::move(callback)(ok); - } private: GetNextObjectResult GetNextObjectInner(PublishedObject&);
diff --git a/quiche/quic/moqt/moqt_outgoing_queue_test.cc b/quiche/quic/moqt/moqt_outgoing_queue_test.cc index c1265f9..4b7b76e 100644 --- a/quiche/quic/moqt/moqt_outgoing_queue_test.cc +++ b/quiche/quic/moqt/moqt_outgoing_queue_test.cc
@@ -17,12 +17,13 @@ #include "absl/strings/string_view.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_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_publisher.h" +#include "quiche/quic/moqt/moqt_session_callbacks.h" #include "quiche/quic/moqt/moqt_types.h" #include "quiche/quic/moqt/test_tools/moqt_mock_visitor.h" #include "quiche/common/platform/api/quiche_expect_bug.h" @@ -39,7 +40,6 @@ using ::testing::AnyOf; using ::testing::ElementsAre; using ::testing::Field; -using ::testing::IsEmpty; using ::testing::Return; class TestMoqtOutgoingQueue : public MoqtOutgoingQueue, @@ -103,7 +103,20 @@ }; absl::StatusOr<std::vector<std::string>> FetchToVector( - std::unique_ptr<MoqtFetchTask> fetch) { + MoqtOutgoingQueue& queue, Location start, Location end, + MoqtDeliveryOrder order) { + std::optional<std::variant<FetchOkData, MoqtRequestErrorInfo>> response; + std::unique_ptr<MoqtFetchTask> fetch = queue.StandaloneFetch( + start, end, order, + [&](std::variant<FetchOkData, MoqtRequestErrorInfo> res) { + response = std::move(res); + }); + if (response.has_value() && + std::holds_alternative<MoqtRequestErrorInfo>(*response)) { + MoqtRequestErrorInfo error_info = std::get<MoqtRequestErrorInfo>(*response); + return RequestErrorCodeToStatus(error_info.error_code, + error_info.reason_phrase); + } std::vector<std::string> objects; for (;;) { PublishedObject object; @@ -125,6 +138,7 @@ return fetch->GetStatus(); } } + return objects; } TEST(MoqtOutgoingQueue, FirstObjectNotKeyframe) { @@ -304,10 +318,9 @@ TEST(MoqtOutgoingQueue, StandaloneFetch) { TestMoqtOutgoingQueue queue; - EXPECT_THAT( - FetchToVector(queue.StandaloneFetch(Location{0, 0}, Location{2, 0}, - MoqtDeliveryOrder::kAscending)), - StatusIs(absl::StatusCode::kNotFound)); + EXPECT_THAT(FetchToVector(queue, Location{0, 0}, Location{2, 0}, + MoqtDeliveryOrder::kAscending), + StatusIs(absl::StatusCode::kOutOfRange)); queue.AddObject(quiche::QuicheMemSlice::Copy("a"), true); queue.AddObject(quiche::QuicheMemSlice::Copy("b"), false); @@ -315,66 +328,59 @@ queue.AddObject(quiche::QuicheMemSlice::Copy("d"), false); queue.AddObject(quiche::QuicheMemSlice::Copy("e"), true); - EXPECT_THAT( - FetchToVector(queue.StandaloneFetch(Location{0, 0}, Location{2, 0}, - MoqtDeliveryOrder::kAscending)), - IsOkAndHolds(ElementsAre("a", "b", "c", "d", "e"))); - EXPECT_THAT( - FetchToVector(queue.StandaloneFetch(Location{0, 100}, Location{0, 1000}, - MoqtDeliveryOrder::kAscending)), - IsOkAndHolds(IsEmpty())); - EXPECT_THAT( - FetchToVector(queue.StandaloneFetch(Location{0, 0}, Location{2, 0}, - MoqtDeliveryOrder::kDescending)), - IsOkAndHolds(ElementsAre("e", "c", "d", "a", "b"))); - EXPECT_THAT( - FetchToVector(queue.StandaloneFetch(Location{0, 0}, Location{1, 0}, - MoqtDeliveryOrder::kAscending)), - IsOkAndHolds(ElementsAre("a", "b", "c"))); - EXPECT_THAT( - FetchToVector(queue.StandaloneFetch(Location{0, 0}, Location{1, 0}, - MoqtDeliveryOrder::kAscending)), - IsOkAndHolds(ElementsAre("a", "b", "c"))); - EXPECT_THAT(FetchToVector(queue.StandaloneFetch( - Location{1, 0}, Location{5, kMaxObjectId}, - MoqtDeliveryOrder::kAscending)), + EXPECT_THAT(FetchToVector(queue, Location{0, 0}, Location{2, 0}, + MoqtDeliveryOrder::kAscending), + IsOkAndHolds(ElementsAre("a", "b", "c", "d", "e"))); + EXPECT_THAT(FetchToVector(queue, Location{0, 100}, Location{0, 1000}, + MoqtDeliveryOrder::kAscending), + StatusIs(absl::StatusCode::kOutOfRange)); + EXPECT_THAT(FetchToVector(queue, Location{0, 0}, Location{2, 0}, + MoqtDeliveryOrder::kDescending), + IsOkAndHolds(ElementsAre("e", "c", "d", "a", "b"))); + EXPECT_THAT(FetchToVector(queue, Location{0, 0}, Location{1, 0}, + MoqtDeliveryOrder::kAscending), + IsOkAndHolds(ElementsAre("a", "b", "c"))); + EXPECT_THAT(FetchToVector(queue, Location{0, 0}, Location{1, 0}, + MoqtDeliveryOrder::kAscending), + IsOkAndHolds(ElementsAre("a", "b", "c"))); + EXPECT_THAT(FetchToVector(queue, Location{1, 0}, Location{5, kMaxObjectId}, + MoqtDeliveryOrder::kAscending), IsOkAndHolds(ElementsAre("c", "d", "e"))); - EXPECT_THAT(FetchToVector(queue.StandaloneFetch( - Location{3, 0}, Location{5, kMaxObjectId}, - MoqtDeliveryOrder::kAscending)), - StatusIs(absl::StatusCode::kNotFound)); + EXPECT_THAT(FetchToVector(queue, Location{3, 0}, Location{5, kMaxObjectId}, + MoqtDeliveryOrder::kAscending), + StatusIs(absl::StatusCode::kOutOfRange)); queue.AddObject(quiche::QuicheMemSlice::Copy("f"), true); queue.AddObject(quiche::QuicheMemSlice::Copy("g"), false); - EXPECT_THAT( - FetchToVector(queue.StandaloneFetch(Location{0, 0}, Location{0, 1}, - MoqtDeliveryOrder::kAscending)), - StatusIs(absl::StatusCode::kNotFound)); - EXPECT_THAT( - FetchToVector(queue.StandaloneFetch(Location{0, 0}, Location{2, 0}, - MoqtDeliveryOrder::kAscending)), - IsOkAndHolds(ElementsAre("c", "d", "e"))); - EXPECT_THAT(FetchToVector(queue.StandaloneFetch( - Location{1, 0}, Location{5, kMaxObjectId}, - MoqtDeliveryOrder::kAscending)), + EXPECT_THAT(FetchToVector(queue, Location{0, 0}, Location{0, 1}, + MoqtDeliveryOrder::kAscending), + StatusIs(absl::StatusCode::kOutOfRange)); + EXPECT_THAT(FetchToVector(queue, Location{0, 0}, Location{2, 0}, + MoqtDeliveryOrder::kAscending), + IsOkAndHolds(ElementsAre("c", "d", "e"))); + EXPECT_THAT(FetchToVector(queue, Location{1, 0}, Location{5, kMaxObjectId}, + MoqtDeliveryOrder::kAscending), IsOkAndHolds(ElementsAre("c", "d", "e", "f", "g"))); - EXPECT_THAT(FetchToVector(queue.StandaloneFetch( - Location{3, 0}, Location{5, kMaxObjectId}, - MoqtDeliveryOrder::kAscending)), + EXPECT_THAT(FetchToVector(queue, Location{3, 0}, Location{5, kMaxObjectId}, + MoqtDeliveryOrder::kAscending), IsOkAndHolds(ElementsAre("f", "g"))); } TEST(MoqtOutgoingQueue, RelativeJoiningFetch) { TestMoqtOutgoingQueue queue; EXPECT_QUICHE_BUG( - queue.RelativeFetch(1, MoqtDeliveryOrder::kAscending), + queue.RelativeFetch( + 1, MoqtDeliveryOrder::kAscending, + [](std::variant<FetchOkData, MoqtRequestErrorInfo>) {}), "Calling RelativeFetch\\(\\) on an established subscription"); } TEST(MoqtOutgoingQueue, AbsoluteJoiningFetch) { TestMoqtOutgoingQueue queue; EXPECT_QUICHE_BUG( - queue.AbsoluteFetch(1, MoqtDeliveryOrder::kAscending), + queue.AbsoluteFetch( + 1, MoqtDeliveryOrder::kAscending, + [](std::variant<FetchOkData, MoqtRequestErrorInfo>) {}), "Calling AbsoluteFetch\\(\\) on an established subscription"); } @@ -386,20 +392,21 @@ queue.AddObject(quiche::QuicheMemSlice::Copy("d"), true); queue.AddObject(quiche::QuicheMemSlice::Copy("e"), true); - EXPECT_THAT( - FetchToVector(queue.StandaloneFetch(Location{0, 0}, Location{5, 0}, - MoqtDeliveryOrder::kAscending)), - IsOkAndHolds(ElementsAre("c", "d", "e"))); + EXPECT_THAT(FetchToVector(queue, Location{0, 0}, Location{5, 0}, + MoqtDeliveryOrder::kAscending), + IsOkAndHolds(ElementsAre("c", "d", "e"))); std::unique_ptr<MoqtFetchTask> deferred_fetch = queue.StandaloneFetch( - Location{0, 0}, Location{5, 0}, MoqtDeliveryOrder::kAscending); + Location{0, 0}, Location{5, 0}, MoqtDeliveryOrder::kAscending, + [](std::variant<FetchOkData, MoqtRequestErrorInfo>) {}); queue.AddObject(quiche::QuicheMemSlice::Copy("f"), true); queue.AddObject(quiche::QuicheMemSlice::Copy("g"), true); queue.AddObject(quiche::QuicheMemSlice::Copy("h"), true); queue.AddObject(quiche::QuicheMemSlice::Copy("i"), true); - EXPECT_THAT(FetchToVector(std::move(deferred_fetch)), - IsOkAndHolds(IsEmpty())); + PublishedObject unused; + EXPECT_EQ(deferred_fetch->GetNextObject(unused), + MoqtFetchTask::GetNextObjectResult::kEof); } TEST(MoqtOutgoingQueue, ObjectIsTimestamped) { @@ -416,45 +423,34 @@ TestMoqtOutgoingQueue queue; queue.AddObject(quiche::QuicheMemSlice::Copy("a"), true); // Create (0, 0) queue.AddObject(quiche::QuicheMemSlice::Copy("b"), true); // Create (1, 0) - std::unique_ptr<MoqtFetchTask> fetch = queue.StandaloneFetch( - Location{0, 0}, Location{5, kMaxObjectId}, MoqtDeliveryOrder::kAscending); - bool end_of_track = false; - Location end_location; + int responses = 0; // end_of_track is false before Close() is called. - fetch->SetFetchResponseCallback( - [&end_of_track, - &end_location](std::variant<MoqtFetchOk, MoqtRequestError> arg) { - end_of_track = std::get<MoqtFetchOk>(arg).end_of_track; - end_location = std::get<MoqtFetchOk>(arg).end_location; + FetchOkData expected_ok(false, Location(1, 0)); + std::unique_ptr<MoqtFetchTask> fetch = queue.StandaloneFetch( + Location{0, 0}, Location{5, kMaxObjectId}, MoqtDeliveryOrder::kAscending, + [&](std::variant<FetchOkData, MoqtRequestErrorInfo> arg) { + ++responses; + EXPECT_EQ(std::get<FetchOkData>(arg), expected_ok); }); - EXPECT_FALSE(end_of_track); - EXPECT_EQ(end_location, Location(1, 0)); - queue.Close(); // Create (2, 0) EXPECT_EQ(queue.largest_location(), Location(2, 0)); - fetch = queue.StandaloneFetch(Location{0, 0}, Location{1, kMaxObjectId}, - MoqtDeliveryOrder::kAscending); // end_of_track is false if the fetch does not include the last object. - fetch->SetFetchResponseCallback( - [&end_of_track, - &end_location](std::variant<MoqtFetchOk, MoqtRequestError> arg) { - end_of_track = std::get<MoqtFetchOk>(arg).end_of_track; - end_location = std::get<MoqtFetchOk>(arg).end_location; + expected_ok.end_location = Location(1, kMaxObjectId); + fetch = queue.StandaloneFetch( + Location{0, 0}, Location{1, kMaxObjectId}, MoqtDeliveryOrder::kAscending, + [&](std::variant<FetchOkData, MoqtRequestErrorInfo> arg) { + ++responses; + EXPECT_EQ(std::get<FetchOkData>(arg), expected_ok); }); - EXPECT_FALSE(end_of_track); - EXPECT_EQ(end_location, Location(1, 1)); - - fetch = queue.StandaloneFetch(Location{0, 0}, Location{5, kMaxObjectId}, - MoqtDeliveryOrder::kAscending); // end_of_track is true if the fetch includes the last object. - fetch->SetFetchResponseCallback( - [&end_of_track, - &end_location](std::variant<MoqtFetchOk, MoqtRequestError> arg) { - end_of_track = std::get<MoqtFetchOk>(arg).end_of_track; - end_location = std::get<MoqtFetchOk>(arg).end_location; + expected_ok = FetchOkData(true, Location(2, 0)); + fetch = queue.StandaloneFetch( + Location{0, 0}, Location{5, kMaxObjectId}, MoqtDeliveryOrder::kAscending, + [&](std::variant<FetchOkData, MoqtRequestErrorInfo> arg) { + ++responses; + EXPECT_EQ(std::get<FetchOkData>(arg), expected_ok); }); - EXPECT_TRUE(end_of_track); - EXPECT_EQ(end_location, Location(2, 0)); + EXPECT_EQ(responses, 3); } // Regression test for b/459527759. `RemoveAllSubscriptions()` calls
diff --git a/quiche/quic/moqt/moqt_parser.cc b/quiche/quic/moqt/moqt_parser.cc index bf0eadf..e5db23a 100644 --- a/quiche/quic/moqt/moqt_parser.cc +++ b/quiche/quic/moqt/moqt_parser.cc
@@ -689,8 +689,7 @@ MoqtRequestError request_error; uint64_t error_code; uint64_t raw_interval; - if (!reader.ReadMoqVarInt(&request_error.request_id) || - !reader.ReadMoqVarInt(&error_code) || + if (!reader.ReadMoqVarInt(&error_code) || !reader.ReadMoqVarInt(&raw_interval) || !reader.ReadStringMoqVarInt(request_error.reason_phrase)) { return absl::InvalidArgumentError("Message missing fields"); @@ -901,8 +900,7 @@ quic::QuicDataReader reader(data); MoqtFetchOk fetch_ok; uint8_t end_of_track; - if (!reader.ReadMoqVarInt(&fetch_ok.request_id) || - !reader.ReadUInt8(&end_of_track) || + if (!reader.ReadUInt8(&end_of_track) || !reader.ReadMoqVarInt(&fetch_ok.end_location.group) || !reader.ReadMoqVarInt(&fetch_ok.end_location.object)) { return absl::InvalidArgumentError("Message missing fields"); @@ -927,17 +925,6 @@ return fetch_ok; } -absl::StatusOr<MoqtFetchCancel> MoqtControlMessageParser::ProcessFetchCancel( - absl::string_view data) const { - quic::QuicDataReader reader(data); - MoqtFetchCancel fetch_cancel; - if (!reader.ReadMoqVarInt(&fetch_cancel.request_id)) { - return absl::InvalidArgumentError("Request ID missing"); - } - QUICHE_RETURN_IF_ERROR(CheckForTrailingData(reader)); - return fetch_cancel; -} - absl::StatusOr<MoqtPublish> MoqtControlMessageParser::ProcessPublish( absl::string_view data) const { quic::QuicDataReader reader(data);
diff --git a/quiche/quic/moqt/moqt_parser.h b/quiche/quic/moqt/moqt_parser.h index b36815c..aa28ab4 100644 --- a/quiche/quic/moqt/moqt_parser.h +++ b/quiche/quic/moqt/moqt_parser.h
@@ -172,8 +172,6 @@ absl::StatusOr<MoqtSubscribeTracks> ProcessSubscribeTracks( absl::string_view data) const; absl::StatusOr<MoqtFetch> ProcessFetch(absl::string_view data) const; - absl::StatusOr<MoqtFetchCancel> ProcessFetchCancel( - absl::string_view data) const; absl::StatusOr<MoqtFetchOk> ProcessFetchOk(absl::string_view data) const; absl::StatusOr<MoqtPublish> ProcessPublish(absl::string_view data) const; absl::StatusOr<MoqtObjectAck> ProcessObjectAck(absl::string_view data) const; @@ -224,8 +222,6 @@ return parse(&MoqtControlMessageParser::ProcessSubscribeTracks); case MoqtMessageType::kFetch: return parse(&MoqtControlMessageParser::ProcessFetch); - case MoqtMessageType::kFetchCancel: - return parse(&MoqtControlMessageParser::ProcessFetchCancel); case MoqtMessageType::kFetchOk: return parse(&MoqtControlMessageParser::ProcessFetchOk); case MoqtMessageType::kPublish:
diff --git a/quiche/quic/moqt/moqt_parser_test.cc b/quiche/quic/moqt/moqt_parser_test.cc index 8b29cb6..806ce6c 100644 --- a/quiche/quic/moqt/moqt_parser_test.cc +++ b/quiche/quic/moqt/moqt_parser_test.cc
@@ -61,7 +61,6 @@ MoqtMessageType::kSubscribeNamespace, MoqtMessageType::kSubscribeTracks, MoqtMessageType::kFetch, - MoqtMessageType::kFetchCancel, MoqtMessageType::kFetchOk, MoqtMessageType::kPublish, MoqtMessageType::kObjectAck,
diff --git a/quiche/quic/moqt/moqt_publish_namespace_stream.cc b/quiche/quic/moqt/moqt_publish_namespace_stream.cc index 522fd0c..dd0932e 100644 --- a/quiche/quic/moqt/moqt_publish_namespace_stream.cc +++ b/quiche/quic/moqt/moqt_publish_namespace_stream.cc
@@ -24,6 +24,7 @@ framer()->SerializePublishNamespace( MoqtPublishNamespace{request_id_, prefix_, parameters_}), false); + stream_parser()->set_allow_fin(true); QUIC_DLOG(INFO) << "Sent PUBLISH_NAMESPACE message for " << prefix_; } @@ -59,8 +60,7 @@ auto callback = std::move(response_callback_); response_callback_ = nullptr; Fin(); - std::move(callback)(MoqtRequestErrorInfo{ - message.error_code, message.retry_interval, message.reason_phrase}); + std::move(callback)(message); return absl::OkStatus(); } // The REQUEST_ERROR is a response to the REQUEST_UPDATE message. @@ -93,8 +93,7 @@ request_id_ = message.request_id; if (!std::move(add_callback_)(message.track_namespace, this)) { add_callback_ = nullptr; - return SendRequestError(request_id_, RequestErrorCode::kInternalError, - std::nullopt, "", /*fin=*/true); + return SendRequestError(RequestErrorCode::kInternalError, std::nullopt, ""); } add_callback_ = nullptr; prefix_ = message.track_namespace; @@ -112,9 +111,8 @@ stream->SendRequestOk(id, parameters)); }, [&](const MoqtRequestErrorInfo& error) { - stream->CheckStatus(stream->SendRequestError( - id, error.error_code, error.retry_interval, - error.reason_phrase)); + stream->CheckStatus( + stream->SendRequestError(error)); }}, response); }); @@ -141,9 +139,8 @@ stream->SendRequestOk(id, parameters)); }, [&](const MoqtRequestErrorInfo& error) { - stream->CheckStatus(stream->SendRequestError( - id, error.error_code, error.retry_interval, - error.reason_phrase)); + stream->CheckStatus( + stream->SendRequestError(error)); }}, response); });
diff --git a/quiche/quic/moqt/moqt_publish_namespace_stream_test.cc b/quiche/quic/moqt/moqt_publish_namespace_stream_test.cc index f7153cd..4ee1c0e 100644 --- a/quiche/quic/moqt/moqt_publish_namespace_stream_test.cc +++ b/quiche/quic/moqt/moqt_publish_namespace_stream_test.cc
@@ -105,18 +105,15 @@ std::unique_ptr<MoqtPublishNamespaceRequestStream> request_stream = CreateAndBindStream(); bool callback_called = false; + MoqtRequestError message(RequestErrorCode::kUnauthorized, std::nullopt, + "unauthorized"); EXPECT_CALL(response_callback_, Call(_)) .WillOnce([&](std::variant<MessageParameters, MoqtRequestErrorInfo> res) { callback_called = true; ASSERT_TRUE(std::holds_alternative<MoqtRequestErrorInfo>(res)); - EXPECT_EQ(std::get<MoqtRequestErrorInfo>(res).error_code, - RequestErrorCode::kUnauthorized); + EXPECT_EQ(std::get<MoqtRequestErrorInfo>(res), message); }); ExpectFin(mock_stream_); - - MoqtRequestError message; - message.request_id = 10; - message.error_code = RequestErrorCode::kUnauthorized; QUICHE_EXPECT_OK(request_stream->OnControlMessage(message)); EXPECT_TRUE(callback_called); }
diff --git a/quiche/quic/moqt/moqt_publish_stream.cc b/quiche/quic/moqt/moqt_publish_stream.cc index a7d11f3..fcc6ebc 100644 --- a/quiche/quic/moqt/moqt_publish_stream.cc +++ b/quiche/quic/moqt/moqt_publish_stream.cc
@@ -78,12 +78,7 @@ absl::Status MoqtPublishRequestStream::OnControlMessage( const MoqtRequestError& message) { - if (message.request_id != publisher_->request_id()) { - return absl::InvalidArgumentError( - "REQUEST_OK does not match PUBLISH request ID"); - } - std::move(response_callback_)(MoqtRequestErrorInfo{ - message.error_code, message.retry_interval, message.reason_phrase}); + std::move(response_callback_)(message); return absl::OkStatus(); } @@ -133,10 +128,8 @@ subscriber_ = std::make_unique<LiveSubscriber>(message, nullptr, this); if (!std::move(add_callback_)(subscriber_.get())) { add_callback_ = nullptr; - return SendRequestError(message.request_id, - RequestErrorCode::kDuplicateSubscription, - /*retry_interval=*/std::nullopt, "", - /*fin=*/true); + return SendRequestError(RequestErrorCode::kDuplicateSubscription, + /*retry_interval=*/std::nullopt, ""); } add_callback_ = nullptr; if (subscriber_->visitor() == nullptr) { @@ -157,8 +150,8 @@ request_id, parameters, /*fin=*/false)); }, [&](const MoqtRequestErrorInfo& error_info) { - stream->CheckStatus(stream->SendRequestError( - request_id, error_info, /*fin=*/true)); + stream->CheckStatus( + stream->SendRequestError(error_info)); }}, response); })); @@ -171,10 +164,8 @@ incoming_publish_callback_ = nullptr; if (subscriber_->visitor() == nullptr) { // The application doesn't care. - CheckStatus(SendRequestError(message.request_id, - RequestErrorCode::kUninterested, - /*retry_interval=*/std::nullopt, "", - /*fin=*/true)); + CheckStatus(SendRequestError(RequestErrorCode::kUninterested, + /*retry_interval=*/std::nullopt, "")); return absl::OkStatus(); } // Notify the visitor.
diff --git a/quiche/quic/moqt/moqt_publish_stream_test.cc b/quiche/quic/moqt/moqt_publish_stream_test.cc index 36a17cb..304b471 100644 --- a/quiche/quic/moqt/moqt_publish_stream_test.cc +++ b/quiche/quic/moqt/moqt_publish_stream_test.cc
@@ -158,21 +158,15 @@ Writev(ControlMessageOfType(MoqtMessageType::kPublish), _)) .WillOnce(Return(absl::OkStatus())); stream_->BindStream(&mock_stream_); // Calls OnStreamBound - - 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"; + MoqtRequestError request_error(RequestErrorCode::kUnauthorized, + quic::QuicTimeDelta::FromSeconds(5), + "Unauthorized"); QUICHE_EXPECT_OK(stream_->OnControlMessage(request_error)); // 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); + EXPECT_EQ(std::get<MoqtRequestErrorInfo>(*response_), request_error); } TEST_F(MoqtPublishRequestStreamTest, ReceiveRequestUpdate) {
diff --git a/quiche/quic/moqt/moqt_publisher.h b/quiche/quic/moqt/moqt_publisher.h index f55ec9e..8d7ccab 100644 --- a/quiche/quic/moqt/moqt_publisher.h +++ b/quiche/quic/moqt/moqt_publisher.h
@@ -17,6 +17,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_types.h" #include "quiche/web_transport/web_transport.h" @@ -109,13 +110,16 @@ // Performs a fetch for the specified range of objects. Should also be used // for joining fetches where Largest Location is known. virtual std::unique_ptr<MoqtFetchTask> StandaloneFetch( - Location start, Location end, MoqtDeliveryOrder order) = 0; + Location start, Location end, MoqtDeliveryOrder order, + FetchResponseCallback callback) = 0; // Use only when the subscription is pending, so that Largest Location is // unknown. virtual std::unique_ptr<MoqtFetchTask> RelativeFetch( - uint64_t group_diff, MoqtDeliveryOrder order) = 0; + uint64_t group_diff, MoqtDeliveryOrder order, + FetchResponseCallback callback) = 0; virtual std::unique_ptr<MoqtFetchTask> AbsoluteFetch( - uint64_t group, MoqtDeliveryOrder order) = 0; + uint64_t group, MoqtDeliveryOrder order, + FetchResponseCallback callback) = 0; // Returns an optional monitoring interface for tracking delivery and object // ACKs for this track. Note that this only works if there is one subscriber
diff --git a/quiche/quic/moqt/moqt_relay_track_publisher.h b/quiche/quic/moqt/moqt_relay_track_publisher.h index 9622f43..44cea94 100644 --- a/quiche/quic/moqt/moqt_relay_track_publisher.h +++ b/quiche/quic/moqt/moqt_relay_track_publisher.h
@@ -16,7 +16,6 @@ #include "absl/base/nullability.h" #include "absl/container/btree_map.h" #include "absl/container/flat_hash_set.h" -#include "absl/status/status.h" #include "absl/strings/string_view.h" #include "quiche/quic/core/quic_clock.h" #include "quiche/quic/core/quic_default_clock.h" @@ -108,21 +107,29 @@ std::optional<quic::QuicTimeDelta> oack_window_size) { oack_window_size_ = oack_window_size; } - std::unique_ptr<MoqtFetchTask> StandaloneFetch(Location /*start*/, - Location /*end*/, - MoqtDeliveryOrder) override { - return std::make_unique<MoqtFailedFetch>( - absl::UnimplementedError("Fetch not implemented")); + std::unique_ptr<MoqtFetchTask> StandaloneFetch( + Location /*start*/, Location /*end*/, MoqtDeliveryOrder, + FetchResponseCallback callback) override { + std::move(callback)( + MoqtRequestErrorInfo(RequestErrorCode::kNotSupported, std::nullopt, + "StandaloneFetch not implemented")); + return nullptr; } - std::unique_ptr<MoqtFetchTask> RelativeFetch(uint64_t /*group_diff*/, - MoqtDeliveryOrder) override { - return std::make_unique<MoqtFailedFetch>( - absl::UnimplementedError("Fetch not implemented")); + std::unique_ptr<MoqtFetchTask> RelativeFetch( + uint64_t /*group_diff*/, MoqtDeliveryOrder, + FetchResponseCallback callback) override { + std::move(callback)(MoqtRequestErrorInfo(RequestErrorCode::kNotSupported, + std::nullopt, + "RelativeFetch not implemented")); + return nullptr; } - std::unique_ptr<MoqtFetchTask> AbsoluteFetch(uint64_t /*group*/, - MoqtDeliveryOrder) override { - return std::make_unique<MoqtFailedFetch>( - absl::UnimplementedError("Fetch not implemented")); + std::unique_ptr<MoqtFetchTask> AbsoluteFetch( + uint64_t /*group*/, MoqtDeliveryOrder, + FetchResponseCallback callback) override { + std::move(callback)(MoqtRequestErrorInfo(RequestErrorCode::kNotSupported, + std::nullopt, + "AbsoluteFetch not implemented")); + return nullptr; } // MoqtPublishingMonitorInterface implementation.
diff --git a/quiche/quic/moqt/moqt_session.cc b/quiche/quic/moqt/moqt_session.cc index e812048..31578a9 100644 --- a/quiche/quic/moqt/moqt_session.cc +++ b/quiche/quic/moqt/moqt_session.cc
@@ -18,7 +18,6 @@ #include "absl/container/flat_hash_map.h" #include "absl/container/flat_hash_set.h" #include "absl/container/node_hash_map.h" -#include "absl/functional/bind_front.h" #include "absl/memory/memory.h" #include "absl/status/status.h" #include "absl/status/statusor.h" @@ -29,6 +28,7 @@ #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_stream.h" #include "quiche/quic/moqt/moqt_fetch_task.h" #include "quiche/quic/moqt/moqt_framer.h" #include "quiche/quic/moqt/moqt_key_value_pair.h" @@ -179,7 +179,7 @@ options.set_send_fin(true); std::array write_vector = { quiche::QuicheMemSlice(framer_.SerializeRequestError(MoqtRequestError{ - 0, RequestErrorCode::kGoingAway, std::nullopt, ""}))}; + RequestErrorCode::kGoingAway, std::nullopt, ""}))}; if (!stream->Writev(absl::MakeSpan(write_vector), options).ok()) { stream->ResetWithUserCode(kResetCodeSessionClosed); }; @@ -610,80 +610,116 @@ return true; } -bool MoqtSession::Fetch(const FullTrackName& name, - FetchResponseCallback callback, Location start, - uint64_t end_group, std::optional<uint64_t> end_object, - MessageParameters parameters) { +std::unique_ptr<MoqtFetchTask> MoqtSession::Fetch( + const FullTrackName& name, FetchResponseCallback callback, Location start, + uint64_t end_group, std::optional<uint64_t> end_object, + const MessageParameters& parameters) { QUICHE_DCHECK(name.IsValid()); if (received_goaway_ || sent_goaway_) { QUIC_DLOG(INFO) << ENDPOINT << "Tried to send FETCH after GOAWAY"; + return nullptr; + } + webtransport::Stream* stream = session_->OpenOutgoingBidirectionalStream(); + if (stream == nullptr) { + QUIC_DLOG(INFO) << ENDPOINT << "Tried to send FETCH but no more streams"; + return nullptr; + } + uint64_t request_id = NextRequestId(); + auto task = std::make_unique<UpstreamFetchTask>(); + auto fetch = std::make_unique<MoqtFetchRequestStream>( + &framer_, ControlMessageParser(), request_id, name, start, + Location(end_group, end_object.value_or(kMaxObjectId)), parameters, + task.get(), + [weak_session = GetWeakPtr()](MoqtError code, absl::string_view reason) { + MoqtSession* session = MoqtSessionFromWeakPtr(weak_session); + if (session == nullptr) { + return; + } + session->Error(code, reason); + }, + std::move(callback), + [weak_session = GetWeakPtr()](uint64_t request_id) { + MoqtSession* session = MoqtSessionFromWeakPtr(weak_session); + if (session == nullptr) { + return; + } + session->fetch_by_id_.erase(request_id); + }); + MoqtFetchRequestStream* fetch_visitor = fetch.get(); + fetch_by_id_.emplace(request_id, fetch_visitor); + stream->SetVisitor(std::move(fetch)); + fetch_visitor->BindStream(stream); + return task; +} + +bool MoqtSession::RelativeJoiningFetch(const FullTrackName& name, + SubscribeVisitor* visitor, + uint64_t num_previous_groups, + const MessageParameters& parameters) { + QUICHE_DCHECK(name.IsValid()); + std::unique_ptr<MoqtFetchTask> fetch_task = RelativeJoiningFetch( + name, visitor, [](std::variant<FetchOkData, MoqtRequestErrorInfo>) {}, + num_previous_groups, parameters); + if (fetch_task == nullptr) { return false; } - MoqtFetch message; - Location end_location = end_object.has_value() - ? Location(end_group, *end_object) - : Location(end_group, kMaxObjectId); - message.fetch = StandaloneFetch(name, start, end_location); - message.request_id = next_request_id_; - next_request_id_ += 2; - message.parameters = parameters; - SendControlMessage(framer_.SerializeFetch(message)); - QUIC_DLOG(INFO) << ENDPOINT << "Sent FETCH message for " << name; - auto fetch = std::make_unique<UpstreamFetch>( - message, std::get<StandaloneFetch>(message.fetch), std::move(callback), - [this, id = message.request_id]() { - fetch_by_id_.erase(id); // Deletion callback - }); - fetch_by_id_.emplace(message.request_id, std::move(fetch)); + LiveSubscriber* subscribe = SubscribeByName(name); + if (subscribe == nullptr || subscribe->is_fetch()) { + // fetch_task will be released on exit. + return false; + } + subscribe->OnJoiningFetchReady(std::move(fetch_task)); return true; } -bool MoqtSession::RelativeJoiningFetch(const FullTrackName& name, - SubscribeVisitor* visitor, - uint64_t num_previous_groups, - MessageParameters parameters) { - QUICHE_DCHECK(name.IsValid()); - return RelativeJoiningFetch( - name, visitor, - [this, track_name = name](std::unique_ptr<MoqtFetchTask> fetch_task) { - // Move the fetch_task to the subscribe to plumb into its visitor. - LiveSubscriber* subscribe = SubscribeByName(track_name); - if (subscribe == nullptr || subscribe->is_fetch()) { - fetch_task.release(); - return; - } - subscribe->OnJoiningFetchReady(std::move(fetch_task)); - }, - num_previous_groups, parameters); -} - -bool MoqtSession::RelativeJoiningFetch(const FullTrackName& name, - SubscribeVisitor* visitor, - FetchResponseCallback callback, - uint64_t num_previous_groups, - MessageParameters parameters) { +std::unique_ptr<MoqtFetchTask> MoqtSession::RelativeJoiningFetch( + const FullTrackName& name, SubscribeVisitor* visitor, + FetchResponseCallback callback, uint64_t num_previous_groups, + const MessageParameters& parameters) { QUICHE_DCHECK(name.IsValid()); MessageParameters subscribe_parameters = parameters; subscribe_parameters.subscription_filter.emplace( MoqtFilterType::kLargestObject); + uint64_t subscribe_request_id = next_request_id_; if (!Subscribe(name, visitor, subscribe_parameters)) { - return false; + return nullptr; } - - MoqtFetch fetch; - fetch.request_id = next_request_id_; - next_request_id_ += 2; - fetch.fetch = JoiningFetchRelative{fetch.request_id - 2, num_previous_groups}; - fetch.parameters = parameters; - SendControlMessage(framer_.SerializeFetch(fetch)); + webtransport::Stream* stream = session_->OpenOutgoingBidirectionalStream(); + if (stream == nullptr) { + // TODO(martinduke): This is a spot where the bool return value is not all + // that helpful, but the problem will go away when the whole transaction + // occurs on one stream. + QUIC_DLOG(INFO) << ENDPOINT + << "Tried to send JOINING FETCH but no more " + "streams"; + return nullptr; + } QUIC_DLOG(INFO) << ENDPOINT << "Sent Joining FETCH message for " << name; - auto upstream_fetch = std::make_unique<UpstreamFetch>( - fetch, name, std::move(callback), - /*Deletion callback=*/[this, id = fetch.request_id]() { - fetch_by_id_.erase(id); + uint64_t request_id = NextRequestId(); + auto task = std::make_unique<UpstreamFetchTask>(); + auto fetch = std::make_unique<MoqtFetchRequestStream>( + &framer_, ControlMessageParser(), request_id, name, subscribe_request_id, + num_previous_groups, /*relative=*/true, parameters, task.get(), + [weak_session = GetWeakPtr()](MoqtError code, absl::string_view reason) { + MoqtSession* session = MoqtSessionFromWeakPtr(weak_session); + if (session == nullptr) { + return; + } + session->Error(code, reason); + }, + std::move(callback), + [weak_session = GetWeakPtr()](uint64_t request_id) { + MoqtSession* session = MoqtSessionFromWeakPtr(weak_session); + if (session == nullptr) { + return; + } + session->fetch_by_id_.erase(request_id); }); - fetch_by_id_.emplace(fetch.request_id, std::move(upstream_fetch)); - return true; + MoqtFetchRequestStream* fetch_visitor = fetch.get(); + fetch_by_id_.emplace(request_id, fetch_visitor); + stream->SetVisitor(std::move(fetch)); + fetch_visitor->BindStream(stream); + return task; } void MoqtSession::GoAway(absl::string_view new_session_uri) { @@ -712,19 +748,38 @@ } void MoqtSession::UpdateTrackPriority( - uint64_t request_id, std::optional<MoqtTrackPriority> old_priority, + const FullTrackName& name, std::optional<MoqtTrackPriority> old_priority, MoqtTrackPriority new_priority) { if (old_priority.has_value()) { auto [start, end] = - subscriptions_with_queued_streams_.equal_range(*old_priority); + requests_with_queued_streams_.equal_range(*old_priority); for (auto it = start; it != end; ++it) { - if (it->second == request_id) { - subscriptions_with_queued_streams_.erase(it); + if (std::holds_alternative<FullTrackName>(it->second) && + std::get<FullTrackName>(it->second) == name) { + requests_with_queued_streams_.erase(it); break; } } } - subscriptions_with_queued_streams_.emplace(new_priority, request_id); + requests_with_queued_streams_.emplace(new_priority, name); +} + +void MoqtSession::UpdateTrackPriority( + webtransport::StreamId stream_id, + std::optional<MoqtTrackPriority> old_priority, + MoqtTrackPriority new_priority) { + if (old_priority.has_value()) { + auto [start, end] = + requests_with_queued_streams_.equal_range(*old_priority); + for (auto it = start; it != end; ++it) { + if (std::holds_alternative<webtransport::StreamId>(it->second) && + std::get<webtransport::StreamId>(it->second) == stream_id) { + requests_with_queued_streams_.erase(it); + break; + } + } + } + requests_with_queued_streams_.emplace(new_priority, stream_id); } std::shared_ptr<MoqtTrackPublisher> MoqtSession::GetTrackPublisher( @@ -746,37 +801,6 @@ return interface; } -bool MoqtSession::OpenDataStream(PublishedFetch* fetch, - webtransport::SendOrder send_order) { - webtransport::Stream* new_stream = - session_->OpenOutgoingUnidirectionalStream(); - if (new_stream == nullptr) { - QUICHE_BUG(MoqtSession_OpenDataStream_blocked) - << "OpenDataStream called when creation of new streams is blocked."; - return false; - } - fetch->SetStreamId(new_stream->GetStreamId()); - // The line below will lead to updating ObjectsAvailableCallback in the - // FetchTask to call OnCanWrite() on the stream. If there is an object - // available, the callback will be invoked synchronously (i.e. before - // SetVisitor() returns). - new_stream->SetVisitor(std::make_unique<OutgoingFetchStream>( - framer_, new_stream, fetch->request_id(), - webtransport::StreamPriority{/*send_group_id=*/kMoqtSendGroupId, - send_order}, - fetch->release_fetch_task(), - // use weakptr to avoid use-after-free for this. - [weakptr = GetWeakPtr(), request_id = fetch->request_id()]() { - if (weakptr.IsValid()) { - auto session = - absl::down_cast<MoqtSession*>(weakptr.GetIfAvailable()); - session->incoming_fetches_.erase(request_id); - } - }, - &trace_recorder_)); - return true; -} - LiveSubscriber* MoqtSession::SubscribeByAlias(uint64_t track_alias) { auto it = subscribe_by_alias_.find(track_alias); if (it == subscribe_by_alias_.end()) { @@ -793,36 +817,44 @@ return it->second; } -UpstreamFetch* MoqtSession::FetchById(uint64_t request_id) { +MoqtFetchRequestStream* MoqtSession::FetchById(uint64_t request_id) { auto it = fetch_by_id_.find(request_id); if (it == fetch_by_id_.end()) { return nullptr; } - return it->second.get(); + return it->second; } void MoqtSession::OnCanCreateNewOutgoingUnidirectionalStream() { - while (!subscriptions_with_queued_streams_.empty() && + while (!requests_with_queued_streams_.empty() && session_->CanOpenNextOutgoingUnidirectionalStream()) { - auto next = subscriptions_with_queued_streams_.begin(); - auto subscription = published_subscriptions_.find(next->second); - if (subscription == published_subscriptions_.end()) { - auto fetch = incoming_fetches_.find(next->second); - // Create the stream if the fetch still exists. - if (fetch != incoming_fetches_.end() && - !OpenDataStream(fetch->second.get(), - SendOrderForFetch(next->first.subscriber_priority))) { - return; // A QUIC_BUG has fired because this shouldn't happen. + auto next = requests_with_queued_streams_.begin(); + if (std::holds_alternative<FullTrackName>(next->second)) { + auto it = + subscribed_track_names_.find(std::get<FullTrackName>(next->second)); + requests_with_queued_streams_.erase(next); + if (it != subscribed_track_names_.end()) { + it->second->OnCanCreateNewUniStream(); } - // FETCH needs only one stream, and can be deleted from the queue. Or, - // there is no subscribe and no fetch; the entry in the queue is invalid. - subscriptions_with_queued_streams_.erase(next); continue; } - subscriptions_with_queued_streams_.erase(next); - // Pop the item from the subscription's queue, which might update - // subscriptions_with_queued_streams_ with a second pending stream. - subscription->second->OnCanCreateNewUniStream(); + // FETCH. + webtransport::StreamId stream_id = + std::get<webtransport::StreamId>(next->second); + requests_with_queued_streams_.erase(next); + webtransport::Stream* stream = session_->GetStreamById(stream_id); + if (stream == nullptr) { + // The request is gone, so remove it from the queue and continue. + continue; + } + auto fetch = absl::down_cast<MoqtFetchResponseStream*>(stream->visitor()); + if (fetch == nullptr) { + QUICHE_BUG(queued_uni_stream_invalid_request_type) + << "Unknown stream type for request " << stream_id; + continue; + } + fetch->OnDataStreamOpen(session_->OpenOutgoingUnidirectionalStream(), + &trace_recorder_); } } @@ -835,11 +867,6 @@ } // TODO(martinduke): Write new checks for duplicate request IDs. It's // probably best to track the largest observed plus a set of holes. - if (incoming_fetches_.contains(request_id)) { - QUICHE_DLOG(INFO) << ENDPOINT << "Duplicate request ID"; - Error(MoqtError::kInvalidRequestId, "Duplicate request ID"); - return false; - } return true; } @@ -903,9 +930,8 @@ if (!queue .SendOrBufferMessage( session_->framer_.SerializeRequestError(MoqtRequestError{ - /*request_id=*/0, RequestErrorCode::kNotSupported, - std::nullopt, "SUBSCRIBE_TRACKS is not supported"}), - /*fin=*/true) + RequestErrorCode::kNotSupported, std::nullopt, + "SUBSCRIBE_TRACKS is not supported"})) .ok()) { session_->Error(MoqtError::kInternalError, "Internal write error"); return; @@ -1083,6 +1109,67 @@ temp_stream->OnCanRead(); break; } + case MoqtMessageType::kFetch: { + auto fetch_stream = std::make_unique<MoqtFetchResponseStream>( + &session_->framer_, session_->ControlMessageParser(), + session_->publisher_, + [weakptr = session_->GetWeakPtr()](MoqtError code, + absl::string_view reason) { + MoqtSession* session = MoqtSessionFromWeakPtr(weakptr); + if (session != nullptr) { + session->Error(code, reason); + } + }, + // OpenStreamCallback + [weakptr = session_->GetWeakPtr()](webtransport::StreamId stream_id, + MoqtTrackPriority priority) { + MoqtSession* session = MoqtSessionFromWeakPtr(weakptr); + if (session == nullptr) { + return; + } + if (!session->session_->CanOpenNextOutgoingUnidirectionalStream()) { + session->UpdateTrackPriority(stream_id, std::nullopt, priority); + return; + } + webtransport::Stream* wt_stream = + session->session_->GetStreamById(stream_id); + if (wt_stream == nullptr) { + QUICHE_BUG( + quiche_bug_OpenStreamCallback_called_by_nonexistent_stream) + << "OpenStreamCallback called by non-existent stream " + << stream_id; + return; + } + MoqtFetchResponseStream* response_stream = + absl::down_cast<MoqtFetchResponseStream*>(wt_stream->visitor()); + if (response_stream == nullptr) { + QUICHE_BUG(quiche_bug_fetch_response_stream_not_found) + << "Failed to get fetch response stream for id " << stream_id; + return; + } + response_stream->OnDataStreamOpen( + session->session_->OpenOutgoingUnidirectionalStream(), + &session->trace_recorder()); + }, + // GetSubscriptionCallback + [weakptr = + session_->GetWeakPtr()](uint64_t request_id) -> LivePublisher* { + MoqtSession* session = MoqtSessionFromWeakPtr(weakptr); + if (session == nullptr) { + return nullptr; + } + auto it = session->published_subscriptions_.find(request_id); + if (it == session->published_subscriptions_.end()) { + return nullptr; + } + return it->second; + }); + fetch_stream->BindStream(std::move(parser_)); + MoqtFetchResponseStream* temp_stream = fetch_stream.get(); + stream_->SetVisitor(std::move(fetch_stream)); + temp_stream->OnCanRead(); + break; + } default: session_->Error(MoqtError::kProtocolViolation, "Unexpected message type received to start bidi stream"); @@ -1241,61 +1328,6 @@ return absl::OkStatus(); } -absl::Status MoqtSession::OnControlMessage(const MoqtRequestOk& message) { - if (fetch_by_id_.contains(message.request_id)) { - return absl::InvalidArgumentError("Received REQUEST_OK for FETCH"); - } - // Response to PUBLISH/SUBSCRIBE_NAMESPACE is handled in the bidi stream.. - // TRACK_STATUS response would go here, but we don't support upstream - // TRACK_STATUS. - // If it doesn't match any state, it might be because the local application - // cancelled the request. Do nothing. - // TODO(martinduke): Do something with parameters. - return absl::OkStatus(); -} - -absl::Status MoqtSession::OnControlMessage(const MoqtRequestError& message) { - MoqtRequestErrorInfo error_info{message.error_code, message.retry_interval, - message.reason_phrase}; - // TODO(martinduke): Do something with retry_interval. - UpstreamFetch* fetch = FetchById(message.request_id); - if (fetch != nullptr) { - // It's in response to FETCH. - if (!fetch->ErrorIsAllowed()) { - return absl::InvalidArgumentError( - "Received REQUEST_ERROR after REQUEST_OK or objects"); - } - QUIC_DLOG(INFO) << ENDPOINT << "Received the REQUEST_ERROR for " - << "request_id = " << message.request_id << " (" - << fetch->full_track_name() << ")" - << ", error = " << static_cast<uint64_t>(message.error_code) - << " (" << message.reason_phrase << ")"; - absl::Status status = - RequestErrorCodeToStatus(message.error_code, message.reason_phrase); - fetch->OnFetchResult(Location(0, 0), status, nullptr); - if (!is_closing_) { - // The visitor might have closed the session. - fetch->Destroy(); - } - return absl::OkStatus(); - } - // Response to PUBLISH/SUBSCRIBE_NAMESPACE is handled in the bidi stream. - // TRACK_STATUS response would go here, but we don't support upstream - // TRACK_STATUS. - // If it doesn't match any state, it might be because the local application - // cancelled the request. Do nothing. - return absl::OkStatus(); -} - -absl::Status MoqtSession::OnControlMessage(const MoqtRequestUpdate& message) { - // TODO(martinduke): Check all the request types. - // Does not match any known request. - SendRequestErrorOnControlStream(message.request_id, - RequestErrorCode::kNotSupported, std::nullopt, - "No support for update of this type"); - return absl::OkStatus(); -} - absl::Status MoqtSession::OnControlMessage(const MoqtGoAway& message) { if (!message.new_session_uri.empty() && perspective() == quic::Perspective::IS_SERVER) { @@ -1312,206 +1344,18 @@ return absl::OkStatus(); } -absl::Status MoqtSession::OnControlMessage(const MoqtFetch& message) { - if (!ValidateRequestId(message.request_id)) { - return absl::OkStatus(); - } - if (sent_goaway_) { - QUIC_DLOG(INFO) << ENDPOINT << "Received a FETCH after GOAWAY"; - SendRequestErrorOnControlStream(message.request_id, - RequestErrorCode::kUnauthorized, - std::nullopt, "FETCH after GOAWAY"); - return absl::OkStatus(); - } - std::unique_ptr<MoqtFetchTask> fetch; - FullTrackName track_name; - if (std::holds_alternative<StandaloneFetch>(message.fetch)) { - const StandaloneFetch& standalone_fetch = - std::get<StandaloneFetch>(message.fetch); - track_name = standalone_fetch.full_track_name; - std::shared_ptr<MoqtTrackPublisher> track_publisher = - publisher_->GetTrack(track_name); - if (track_publisher == nullptr) { - QUIC_DLOG(INFO) << ENDPOINT << "FETCH for " << track_name - << " rejected by the application: not found"; - SendRequestErrorOnControlStream(message.request_id, - RequestErrorCode::kDoesNotExist, - std::nullopt, "not found"); - return absl::OkStatus(); - } - QUIC_DLOG(INFO) << ENDPOINT << "Received a StandaloneFETCH for " - << track_name; - // The check for end_object < start_object is done in - // MoqtTrackPublisher::Fetch(). - fetch = track_publisher->StandaloneFetch( - standalone_fetch.start_location, standalone_fetch.end_location, - message.parameters.group_order.value_or(MoqtDeliveryOrder::kAscending)); - } else { - // Joining Fetch processing. - uint64_t joining_request_id = - std::holds_alternative<JoiningFetchRelative>(message.fetch) - ? std::get<struct JoiningFetchRelative>(message.fetch) - .joining_request_id - : std::get<JoiningFetchAbsolute>(message.fetch).joining_request_id; - auto it = published_subscriptions_.find(joining_request_id); - if (it == published_subscriptions_.end()) { - QUIC_DLOG(INFO) << ENDPOINT << "Received a JOINING_FETCH for " - << "request_id " << joining_request_id - << " that does not exist"; - SendRequestErrorOnControlStream( - message.request_id, RequestErrorCode::kInvalidJoiningRequestId, - std::nullopt, "Joining Fetch for non-existent request"); - return absl::OkStatus(); - } - if (!it->second->can_have_joining_fetch()) { - QUIC_DLOG(INFO) << ENDPOINT << "Received a JOINING_FETCH for " - << "joining_request_id " << joining_request_id - << " that is not forwarding"; - return absl::InvalidArgumentError( - "Joining Fetch for non-forwarding subscribe"); - } - track_name = it->second->publisher().GetTrackName(); - if (it->second->established()) { - if (!it->second->parameters().largest_object.has_value()) { - // Nothing to Fetch. - SendRequestErrorOnControlStream(message.request_id, - RequestErrorCode::kDoesNotExist, - std::nullopt, "not found"); - return absl::OkStatus(); - } - const Location largest_location = - *it->second->parameters().largest_object; - uint64_t start_group; - if (std::holds_alternative<JoiningFetchRelative>(message.fetch)) { - const JoiningFetchRelative& relative_fetch = - std::get<JoiningFetchRelative>(message.fetch); - start_group = - (relative_fetch.joining_start > largest_location.group) - ? 0 - : (largest_location.group - relative_fetch.joining_start); - } else { - const JoiningFetchAbsolute& absolute_fetch = - std::get<JoiningFetchAbsolute>(message.fetch); - start_group = absolute_fetch.joining_start; - if (start_group > largest_location.group) { - SendRequestErrorOnControlStream(message.request_id, - RequestErrorCode::kInvalidRange, - std::nullopt, "invalid range"); - return absl::OkStatus(); - } - } - fetch = it->second->publisher().StandaloneFetch( - Location{start_group, 0}, largest_location, - message.parameters.group_order.value_or( - MoqtDeliveryOrder::kAscending)); - } else { - // Subscription is in PENDING state. - if (std::holds_alternative<JoiningFetchRelative>(message.fetch)) { - fetch = it->second->publisher().RelativeFetch( - std::get<JoiningFetchRelative>(message.fetch).joining_start, - message.parameters.group_order.value_or( - MoqtDeliveryOrder::kAscending)); - } else { - fetch = it->second->publisher().AbsoluteFetch( - std::get<JoiningFetchAbsolute>(message.fetch).joining_start, - message.parameters.group_order.value_or( - MoqtDeliveryOrder::kAscending)); - } - } - } - if (!fetch->GetStatus().ok()) { - QUIC_DLOG(INFO) << ENDPOINT << "FETCH for " << track_name - << " could not initialize the task"; - SendRequestErrorOnControlStream(message.request_id, - RequestErrorCode::kInvalidRange, - std::nullopt, fetch->GetStatus().message()); - return absl::OkStatus(); - } - auto published_fetch = - std::make_unique<PublishedFetch>(message.request_id, std::move(fetch)); - auto result = - incoming_fetches_.emplace(message.request_id, std::move(published_fetch)); - if (!result.second) { // Emplace failed. - QUIC_DLOG(INFO) << ENDPOINT << "FETCH for " << track_name - << " could not be added to the session"; - SendRequestErrorOnControlStream( - message.request_id, RequestErrorCode::kInternalError, std::nullopt, - "Could not initialize FETCH state"); - return absl::OkStatus(); - } - MoqtFetchTask* fetch_task = result.first->second->fetch_task_ptr(); - fetch_task->SetFetchResponseCallback( - [this, request_id = message.request_id]( - std::variant<MoqtFetchOk, MoqtRequestError> message) { - if (!incoming_fetches_.contains(request_id)) { - return; // FETCH was cancelled. - } - if (std::holds_alternative<MoqtFetchOk>(message)) { - MoqtFetchOk& fetch_ok = std::get<MoqtFetchOk>(message); - fetch_ok.request_id = request_id; - SendControlMessage(framer_.SerializeFetchOk(fetch_ok)); - return; - } - SendRequestErrorOnControlStream( - request_id, std::get<MoqtRequestError>(message).error_code, - std::get<MoqtRequestError>(message).retry_interval, - std::get<MoqtRequestError>(message).reason_phrase); - }); - // Set a temporary new-object callback that creates a data stream. When - // created, the stream visitor will replace this callback. - fetch_task->SetObjectAvailableCallback( - [this, - subscriber_priority = message.parameters.subscriber_priority.value_or( - kDefaultSubscriberPriority), - request_id = message.request_id]() { - auto it = incoming_fetches_.find(request_id); - if (it == incoming_fetches_.end()) { - return; - } - if (!session()->CanOpenNextOutgoingUnidirectionalStream() || - !OpenDataStream(it->second.get(), - SendOrderForFetch(subscriber_priority))) { - UpdateTrackPriority(request_id, std::nullopt, - MoqtTrackPriority(subscriber_priority, - kDefaultPublisherPriority)); - } - }); - return absl::OkStatus(); -} - -absl::Status MoqtSession::OnControlMessage(const MoqtFetchOk& message) { - UpstreamFetch* track = FetchById(message.request_id); - if (track == nullptr) { - QUIC_DLOG(INFO) << ENDPOINT << "Received the FETCH_OK for " - << "request_id = " << message.request_id - << " but no track exists"; - // Subscription state might have been destroyed for internal reasons. - return absl::OkStatus(); - } - QUIC_DLOG(INFO) << ENDPOINT << "Received the FETCH_OK for request_id = " - << message.request_id << " " << track->full_track_name(); - UpstreamFetch* fetch = absl::down_cast<UpstreamFetch*>(track); - fetch->OnFetchResult(message.end_location, absl::OkStatus(), - [=, this]() { CancelFetch(message.request_id); }); - return absl::OkStatus(); -} void MoqtSession::OnMalformedTrack(ObjectSubscriber* track) { - if (!track->is_fetch()) { - auto* subscribe = absl::down_cast<LiveSubscriber*>(track); - if (subscribe->visitor() != nullptr) { - subscribe->visitor()->OnMalformedTrack(track->full_track_name()); - } - Unsubscribe(track->full_track_name()); + if (track->is_fetch()) { + QUICHE_BUG(quiche_bug_malformed_fetch_track) + << "Malformed FETCH track should be handled in the data stream"; return; } - UpstreamFetch::UpstreamFetchTask* task = - absl::down_cast<UpstreamFetch*>(track)->task(); - if (task != nullptr) { - task->OnStreamAndFetchClosed(kResetCodeMalformedTrack, - "Malformed track received"); + auto* subscribe = absl::down_cast<LiveSubscriber*>(track); + if (subscribe->visitor() != nullptr) { + subscribe->visitor()->OnMalformedTrack(track->full_track_name()); } - CancelFetch(track->request_id()); + subscribe->request_stream()->Reset(kResetCodeMalformedTrack); } void MoqtSession::CleanUpState() { @@ -1537,10 +1381,6 @@ publish_namespace_requests_.erase(it); stream->Detach(); } - // WebTransport session, the incoming FETCHes are owned by this class. - while (!fetch_by_id_.empty()) { - fetch_by_id_.begin()->second->Destroy(); - } for (auto& [track_name, subscriber] : subscribe_by_name_) { // It's possible the application is going to destroy its visitor as early // as session_deleted_callback is called. So call OnPublishDone() now and @@ -1552,29 +1392,6 @@ } } -void MoqtSession::CancelFetch(uint64_t request_id) { - if (is_closing_) { - return; - } - auto it = fetch_by_id_.find(request_id); - if (it == fetch_by_id_.end()) { - return; - } - it->second->Destroy(); - // This is only called from the callback where UpstreamFetchTask has been - // destroyed, so there is no need to notify the application. - OutgoingControlStream* stream = GetOutgoingControlStream(); - if (stream == nullptr) { - return; - } - MoqtFetchCancel message; - message.request_id = request_id; - stream->SendOrBufferMessageOrFatal(framer_.SerializeFetchCancel(message)); - // The FETCH_CANCEL will cause a RESET_STREAM to return, which would be the - // same as a STOP_SENDING. However, a FETCH_CANCEL works even if the stream - // hasn't opened yet. -} - void MoqtSessionParameters::ToSetupParameters(SetupParameters& out) const { if (perspective == quic::Perspective::IS_CLIENT && !using_webtrans) { out.path = path;
diff --git a/quiche/quic/moqt/moqt_session.h b/quiche/quic/moqt/moqt_session.h index bdf0e8d..a91ae90 100644 --- a/quiche/quic/moqt/moqt_session.h +++ b/quiche/quic/moqt/moqt_session.h
@@ -10,6 +10,7 @@ #include <optional> #include <string> #include <utility> +#include <variant> #include "absl/base/casts.h" #include "absl/base/nullability.h" @@ -26,6 +27,7 @@ #include "quiche/quic/moqt/moqt_bidi_stream.h" #include "quiche/quic/moqt/moqt_control_message_queue.h" #include "quiche/quic/moqt/moqt_error.h" +#include "quiche/quic/moqt/moqt_fetch_stream.h" #include "quiche/quic/moqt/moqt_fetch_task.h" #include "quiche/quic/moqt/moqt_framer.h" #include "quiche/quic/moqt/moqt_key_value_pair.h" @@ -103,19 +105,18 @@ 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, - MessageParameters parameters) override; + std::unique_ptr<MoqtFetchTask> Fetch( + const FullTrackName& name, FetchResponseCallback callback, Location start, + uint64_t end_group, std::optional<uint64_t> end_object, + const MessageParameters& parameters) override; bool RelativeJoiningFetch(const FullTrackName& name, SubscribeVisitor* visitor, uint64_t num_previous_groups, - MessageParameters parameters) override; - bool RelativeJoiningFetch(const FullTrackName& name, - SubscribeVisitor* visitor, - FetchResponseCallback callback, - uint64_t num_previous_groups, - MessageParameters parameters) override; + const MessageParameters& parameters) override; + std::unique_ptr<MoqtFetchTask> RelativeJoiningFetch( + const FullTrackName& name, SubscribeVisitor* visitor, + FetchResponseCallback callback, uint64_t num_previous_groups, + const MessageParameters& parameters) override; bool PublishNamespace( const TrackNamespace& track_namespace, const MessageParameters& parameters, @@ -153,10 +154,14 @@ } // If |old_priority| is nullopt, the subscription does not have any pending // streams. If it has a value, |old_priority| is the old value to be replaced - // by |new_priority|. - void UpdateTrackPriority(uint64_t request_id, + // by |new_priority|. Subgroup streams send |name| as the first argument. + // Fetch streams send the request stream ID. + void UpdateTrackPriority(const FullTrackName& name, std::optional<MoqtTrackPriority> old_priority, MoqtTrackPriority new_priority) override; + void UpdateTrackPriority(webtransport::StreamId stream_id, + std::optional<MoqtTrackPriority> old_priority, + MoqtTrackPriority new_priority); quic::QuicAlarmFactory* alarm_factory() override { return alarm_factory_.get(); } @@ -343,28 +348,6 @@ quiche::QuicheWeakPtrFactory<OutgoingControlStream> weak_ptr_factory_; }; - class QUICHE_EXPORT PublishedFetch { - public: - PublishedFetch(uint64_t request_id, std::unique_ptr<MoqtFetchTask> fetch) - : request_id_(request_id), fetch_(std::move(fetch)) {} - - MoqtFetchTask* fetch_task_ptr() { return fetch_.get(); } - // Can only be called once. - std::unique_ptr<MoqtFetchTask> release_fetch_task() { - auto on_return = absl::MakeCleanup([this] { fetch_ = nullptr; }); - return std::move(fetch_); - } - uint64_t request_id() const { return request_id_; } - void SetStreamId(webtransport::StreamId id) { stream_id_ = id; } - - private: - uint64_t request_id_; - // Store the stream ID in case a FETCH_CANCEL requires a reset. - std::optional<webtransport::StreamId> stream_id_; - // Temporary storage until the stream is created. - std::unique_ptr<MoqtFetchTask> fetch_; - }; - class GoAwayTimeoutDelegate : public quic::QuicAlarm::DelegateWithoutContext { public: explicit GoAwayTimeoutDelegate(MoqtSession* session) : session_(session) {} @@ -386,19 +369,14 @@ // is present. void SendControlMessage(quiche::QuicheBuffer message); - // Returns false if creation failed. - [[nodiscard]] bool OpenDataStream(PublishedFetch* fetch, - webtransport::SendOrder send_order); LiveSubscriber* SubscribeByAlias(uint64_t track_alias); LiveSubscriber* SubscribeByName(const FullTrackName& track_name); - UpstreamFetch* FetchById(uint64_t request_id); + MoqtFetchRequestStream* FetchById(uint64_t request_id); // Checks that a subscribe ID from a SUBSCRIBE or FETCH is valid, and throws // a session error if is not. bool ValidateRequestId(uint64_t request_id); - void CancelFetch(uint64_t request_id); - // Sends an OBJECT_ACK message for a specific subscribe ID. void SendObjectAck(FullTrackName track_name, uint64_t group_id, uint64_t object_id, @@ -431,29 +409,7 @@ // TODO(martinduke): All of these should be moved to bidi streams or // deleted. - absl::Status OnControlMessage(const MoqtRequestOk& message); - absl::Status OnControlMessage(const MoqtRequestError& message); - absl::Status OnControlMessage(const MoqtRequestUpdate& message); absl::Status OnControlMessage(const MoqtGoAway& /*message*/); - absl::Status OnControlMessage(const MoqtFetch& message); - absl::Status OnControlMessage(const MoqtFetchCancel& /*message*/) { - return absl::OkStatus(); - } - absl::Status OnControlMessage(const MoqtFetchOk& message); - - // TODO(vasilvv): remove this once all requests are moved into individual - // streams. - void SendRequestErrorOnControlStream( - uint64_t request_id, RequestErrorCode error_code, - std::optional<quic::QuicTimeDelta> retry_interval, - absl::string_view reason_phrase) { - MoqtRequestError request_error; - request_error.request_id = request_id; - request_error.error_code = error_code; - request_error.retry_interval = retry_interval; - request_error.reason_phrase = reason_phrase; - SendControlMessage(framer_.SerializeRequestError(request_error)); - } uint64_t NextRequestId() { uint64_t id = next_request_id_; @@ -482,8 +438,10 @@ MoqtTraceRecorder trace_recorder_; - // Upstream FETCHes, indexed by request_id. Do not erase. - absl::flat_hash_map<uint64_t, std::unique_ptr<UpstreamFetch>> fetch_by_id_; + // Upstream FETCHes, indexed by request_id. The RemoveFetchCallback used to + // create MoqtFetchRequestStream MUST delete this entry, so the pointer is + // always valid. + absl::flat_hash_map<uint64_t, MoqtFetchRequestStream*> fetch_by_id_; // All outgoing SUBSCRIBE and incoming PUBLISH, indexed by track_alias. absl::flat_hash_map<uint64_t, LiveSubscriber*> subscribe_by_alias_; // All outgoing SUBSCRIBE and incoming PUBLISH, indexed by track name. @@ -500,18 +458,17 @@ MoqtPublisher* publisher_; // Subscriptions for local tracks by the remote peer, indexed by request ID. absl::flat_hash_map<uint64_t, LivePublisher*> published_subscriptions_; - // Keeps track of all request IDs that have queued outgoing data streams. - // The first element is the highest priority (lowest integer). - absl::btree_multimap<MoqtTrackPriority, uint64_t> - subscriptions_with_queued_streams_; + // Keeps track of all subscriptions and fetches that have queued outgoing data + // streams. For subscriptions (PUBLISH/SUBSCRIBE), stores a FullTrackName that + // can be used to retrieve LivePublisher. For FETCH, stores the response + // stream ID. The first element is the highest priority (lowest integer). + absl::btree_multimap<MoqtTrackPriority, + std::variant<FullTrackName, webtransport::StreamId>> + requests_with_queued_streams_; // This is only used to check for track_alias collisions. absl::flat_hash_set<uint64_t> used_track_aliases_; uint64_t next_local_track_alias_ = 0; - // Incoming FETCHes, indexed by fetch ID. - absl::flat_hash_map<uint64_t, std::unique_ptr<PublishedFetch>> - incoming_fetches_; - // Monitoring interfaces for expected incoming subscriptions. absl::flat_hash_map<FullTrackName, MoqtPublishingMonitorInterface*> monitoring_interfaces_for_published_tracks_;
diff --git a/quiche/quic/moqt/moqt_session_callbacks.h b/quiche/quic/moqt/moqt_session_callbacks.h index c8b8356..4221215 100644 --- a/quiche/quic/moqt/moqt_session_callbacks.h +++ b/quiche/quic/moqt/moqt_session_callbacks.h
@@ -65,6 +65,17 @@ DataStreamIndex stream) = 0; }; +struct FetchOkData { + bool end_of_track = false; + Location end_location; + MessageParameters parameters; + TrackExtensions extensions; + bool operator==(const FetchOkData& other) const = default; +}; + +using FetchResponseCallback = quiche::SingleUseCallback<void( + std::variant<FetchOkData, MoqtRequestErrorInfo>)>; + // Called when the SETUP message from the peer is received. using MoqtSessionEstablishedCallback = quiche::SingleUseCallback<void()>;
diff --git a/quiche/quic/moqt/moqt_session_interface.h b/quiche/quic/moqt/moqt_session_interface.h index 1f8d52e..8d182cd 100644 --- a/quiche/quic/moqt/moqt_session_interface.h +++ b/quiche/quic/moqt/moqt_session_interface.h
@@ -67,13 +67,6 @@ void ToSetupParameters(SetupParameters& out) const; }; - -// MoqtSession calls this when a FETCH_OK or REQUEST_ERROR is received. The -// destination of the callback owns |fetch_task| and MoqtSession will react -// safely if the owner destroys it. -using FetchResponseCallback = - quiche::SingleUseCallback<void(std::unique_ptr<MoqtFetchTask> fetch_task)>; - class MoqtSessionInterface { public: virtual ~MoqtSessionInterface() = default; @@ -119,15 +112,14 @@ 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 - // the MoqtFetchTask. - virtual bool Fetch(const FullTrackName& name, FetchResponseCallback callback, - Location start, uint64_t end_group, - std::optional<uint64_t> end_object, - MessageParameters parameters) = 0; - + // Sends a FETCH for a pre-specified object range. Once a FETCH_OK or a + // FETCH_ERROR is received, `callback` is called with the result. To cancel a + // FETCH, simply destroy the provided MoqtFetchTask. Returns nullptr if the + // FETCH cannot be sent. + virtual std::unique_ptr<MoqtFetchTask> Fetch( + const FullTrackName& name, FetchResponseCallback callback, Location start, + uint64_t end_group, std::optional<uint64_t> end_object, + const MessageParameters& parameters) = 0; // Sends both a SUBSCRIBE and a joining FETCH, beginning `num_previous_groups` // groups before the current group. The Fetch will not be flow controlled, // instead using |visitor| to deliver fetched objects when they arrive. Gaps @@ -137,16 +129,15 @@ virtual bool RelativeJoiningFetch(const FullTrackName& name, SubscribeVisitor* visitor, uint64_t num_previous_groups, - MessageParameters parameters) = 0; - + const MessageParameters& parameters) = 0; // Sends both a SUBSCRIBE and a joining FETCH, beginning `num_previous_groups` // groups before the current group. `callback` acts the same way as the // callback for the regular Fetch() call. - virtual bool RelativeJoiningFetch(const FullTrackName& name, - SubscribeVisitor* visitor, - FetchResponseCallback callback, - uint64_t num_previous_groups, - MessageParameters parameters) = 0; + virtual std::unique_ptr<MoqtFetchTask> RelativeJoiningFetch( + const FullTrackName& name, SubscribeVisitor* visitor, + FetchResponseCallback callback, uint64_t num_previous_groups, + const MessageParameters& parameters) = 0; + // Send a PUBLISH_NAMESPACE message for |track_namespace|, and call // |response_callback| when the response arrives. Will fail // immediately if there is already an unresolved PUBLISH_NAMESPACE for that
diff --git a/quiche/quic/moqt/moqt_session_test.cc b/quiche/quic/moqt/moqt_session_test.cc index 047a4e2..a4fc9c8 100644 --- a/quiche/quic/moqt/moqt_session_test.cc +++ b/quiche/quic/moqt/moqt_session_test.cc
@@ -9,7 +9,6 @@ #include <cstring> #include <memory> #include <optional> -#include <queue> #include <string> #include <utility> #include <variant> @@ -122,15 +121,6 @@ return fetch; } -std::optional<MoqtMessageType> PeekControlMessageType(absl::string_view data) { - quiche::QuicheDataReader reader(data); - uint64_t varint; - if (!reader.ReadMoqVarInt(&varint)) { - return std::nullopt; - } - return static_cast<MoqtMessageType>(varint); -} - } // namespace class MoqtSessionTest : public quic::test::QuicTest { @@ -174,6 +164,7 @@ static constexpr absl::string_view kPublishByte = "\x1d"; static constexpr absl::string_view kTrackStatusByte = "\x0d"; static constexpr absl::string_view kPublishNamespaceByte = "\x06"; + static constexpr absl::string_view kFetchByte = "\x16"; std::unique_ptr<MoqtBidiStreamBase> ResponseStream( absl::string_view first_byte, webtransport::test::MockStream* wt_stream = nullptr) { @@ -187,14 +178,14 @@ EXPECT_CALL(*stream, SetVisitor) .WillOnce([&](std::unique_ptr<webtransport::StreamVisitor> visitor) { unknown_bidi_stream = std::move(visitor); + bidi_visitor_ = unknown_bidi_stream.get(); }) .WillOnce([&](std::unique_ptr<webtransport::StreamVisitor> visitor) { final_stream = std::unique_ptr<MoqtBidiStreamBase>( absl::down_cast<MoqtBidiStreamBase*>(visitor.release())); + bidi_visitor_ = final_stream.get(); }); - EXPECT_CALL(*stream, visitor()) - .WillOnce([&]() { return unknown_bidi_stream.get(); }) - .WillRepeatedly([&]() { return final_stream.get(); }); + ON_CALL(*stream, visitor).WillByDefault([this]() { return bidi_visitor_; }); EXPECT_CALL(*stream, PeekNextReadableRegion) .WillOnce( Return(webtransport::Stream::PeekResult(first_byte, false, false))) @@ -222,13 +213,19 @@ ON_CALL(mock_session_, CanOpenNextOutgoingBidirectionalStream()) .WillByDefault(Return(true)); EXPECT_CALL(mock_session_, OpenOutgoingBidirectionalStream()) - .WillOnce(Return(stream)); + .WillOnce(Return(stream)) + // TODO(martinduke): The line below is necessary because JoiningFetch() + // opens two streams. It can be removed once these are implemented on a + // single bidi stream. + .RetiresOnSaturation(); EXPECT_CALL(*stream, SetVisitor) .WillOnce([&](std::unique_ptr<webtransport::StreamVisitor> visitor) { stream_wrapper = std::make_unique<MoqtBidiStreamTestWrapper>( std::unique_ptr<MoqtBidiStreamBase>( absl::down_cast<MoqtBidiStreamBase*>(visitor.release()))); + bidi_visitor_ = &stream_wrapper->stream(); }); + ON_CALL(*stream, visitor).WillByDefault([this]() { return bidi_visitor_; }); EXPECT_CALL(*stream, CanWrite).WillRepeatedly(Return(true)); } @@ -330,7 +327,7 @@ MoqtSession session_; std::unique_ptr<MoqtBidiStreamTestWrapper> bidi_wrapper_; webtransport::test::MockStream mock_bidi_stream_, mock_uni_stream_; - // std::shared_ptr<IncomingSubscribeInfo> last_incoming_subscribe_; + webtransport::StreamVisitor* bidi_visitor_ = nullptr; }; TEST_F(MoqtSessionTest, Queries) { @@ -652,16 +649,15 @@ publish_namespace_resolved_callback.AsStdFunction(), []() {}); - MoqtRequestError error{/*request_id=*/0, RequestErrorCode::kInternalError, - std::nullopt, "Test error"}; + MoqtRequestError error{RequestErrorCode::kInternalError, std::nullopt, + "Test error"}; EXPECT_CALL(publish_namespace_resolved_callback, Call) .WillOnce( [&](std::variant<MessageParameters, MoqtRequestErrorInfo> response) { ASSERT_TRUE(std::holds_alternative<MoqtRequestErrorInfo>(response)); - const MoqtRequestErrorInfo& error = + const MoqtRequestErrorInfo& error_info = std::get<MoqtRequestErrorInfo>(response); - EXPECT_EQ(error.error_code, RequestErrorCode::kInternalError); - EXPECT_EQ(error.reason_phrase, "Test error"); + EXPECT_EQ(error_info, error); }); ExpectFin(mock_bidi_stream_); bidi_wrapper_->ReceiveMessage(error); @@ -980,13 +976,8 @@ parameters.subscription_filter.emplace(MoqtFilterType::kLargestObject); session_.Subscribe(FullTrackName("foo", "bar"), &remote_track_visitor_, parameters); - - MoqtRequestError error = { - /*request_id=*/0, - /*error_code=*/RequestErrorCode::kInvalidRange, - /*retry_interval=*/std::nullopt, - /*reason_phrase=*/"deadbeef", - }; + MoqtRequestError error(RequestErrorCode::kInvalidRange, std::nullopt, + "deadbeef"); EXPECT_CALL(remote_track_visitor_, OnReply) .WillOnce( [&](const FullTrackName& ftn, @@ -994,8 +985,7 @@ EXPECT_EQ(ftn, FullTrackName("foo", "bar")); EXPECT_TRUE( std::holds_alternative<MoqtRequestErrorInfo>(response) && - std::get<MoqtRequestErrorInfo>(response).reason_phrase == - "deadbeef"); + std::get<MoqtRequestErrorInfo>(response) == error); }); EXPECT_CALL(mock_bidi_stream_, Writev(testing::IsEmpty(), _)); // FIN. bidi_wrapper_->ReceiveMessage(error); @@ -1111,11 +1101,7 @@ EXPECT_TRUE(params == nullptr); EXPECT_TRUE(callback == nullptr); }); - EXPECT_CALL(mock_bidi_stream_, - Writev(SerializedControlMessage(MoqtRequestError{ - kDefaultPeerRequestId, error.error_code, - error.retry_interval, error.reason_phrase}), - _)); + EXPECT_CALL(mock_bidi_stream_, Writev(SerializedControlMessage(error), _)); bidi_wrapper_->ReceiveMessage(publish_namespace); } @@ -1155,15 +1141,16 @@ EXPECT_EQ(error.error_code, RequestErrorCode::kInvalidRange); EXPECT_EQ(error.reason_phrase, "deadbeef"); }); - MoqtRequestError error = {kDefaultLocalRequestId, - RequestErrorCode::kInvalidRange, std::nullopt, - "deadbeef"}; + MoqtRequestError error(RequestErrorCode::kInvalidRange, std::nullopt, + "deadbeef"); bidi_wrapper_->ReceiveMessage(error); EXPECT_TRUE(got_callback); } TEST_F(MoqtSessionTest, SubscribeOkWithBadTrackAlias) { PrepareRequestStream(bidi_wrapper_); + EXPECT_CALL(mock_bidi_stream_, + Writev(ControlMessageOfType(MoqtMessageType::kSubscribe), _)); session_.Subscribe(FullTrackName("foo", "bar"), &remote_track_visitor_, MessageParameters()); MoqtSubscribeOk subscribe_ok = { @@ -1175,9 +1162,10 @@ bidi_wrapper_->ReceiveMessage(subscribe_ok); // Second subscribe, but OK has the same track alias. webtransport::test::MockStream bidi_stream_2; - std::unique_ptr<MoqtBidiStreamTestWrapper> bidi_wrapper_2 = - MoqtSessionPeer::CreateControlStream(&session_, &bidi_stream_2); + std::unique_ptr<MoqtBidiStreamTestWrapper> bidi_wrapper_2; PrepareRequestStream(bidi_wrapper_2, &bidi_stream_2); + EXPECT_CALL(bidi_stream_2, + Writev(ControlMessageOfType(MoqtMessageType::kSubscribe), _)); session_.Subscribe(FullTrackName("foo2", "bar2"), &remote_track_visitor_, MessageParameters()); subscribe_ok.request_id += 2; @@ -1405,38 +1393,39 @@ } TEST_F(MoqtSessionTest, UpdateTrackPriority) { - session_.UpdateTrackPriority(0, std::nullopt, MoqtTrackPriority{0x40, 0x82}); - EXPECT_EQ(MoqtSessionPeer::NextQueuedRequestIdToServer(&session_), 0); + EXPECT_CALL(mock_session_, GetStreamById).WillRepeatedly(Return(nullptr)); + session_.UpdateTrackPriority(4, std::nullopt, MoqtTrackPriority{0x40, 0x82}); + EXPECT_EQ(MoqtSessionPeer::NextQueuedStreamIdToServer(&session_), 4); // Same track, higher priority. - session_.UpdateTrackPriority(0, MoqtTrackPriority{0x40, 0x82}, + session_.UpdateTrackPriority(4, MoqtTrackPriority{0x40, 0x82}, MoqtTrackPriority{0x40, 0x80}); - EXPECT_EQ(MoqtSessionPeer::NextQueuedRequestIdToServer(&session_), 0); + EXPECT_EQ(MoqtSessionPeer::NextQueuedStreamIdToServer(&session_), 4); // New track, higher priority. - session_.UpdateTrackPriority(2, std::nullopt, MoqtTrackPriority{0x20, 0x82}); - EXPECT_EQ(MoqtSessionPeer::NextQueuedRequestIdToServer(&session_), 2); - // Pop request ID 2 from the queue. The subscription doesn't really exist, so + session_.UpdateTrackPriority(6, std::nullopt, MoqtTrackPriority{0x20, 0x82}); + EXPECT_EQ(MoqtSessionPeer::NextQueuedStreamIdToServer(&session_), 6); + // Pop request ID 6 from the queue. The subscription doesn't really exist, so // nothing else happens. EXPECT_CALL(mock_session_, CanOpenNextOutgoingUnidirectionalStream()) .WillOnce(Return(true)) .WillOnce(Return(false)); session_.OnCanCreateNewOutgoingUnidirectionalStream(); - EXPECT_EQ(MoqtSessionPeer::NextQueuedRequestIdToServer(&session_), 0); - // There's another stream for request ID 2. - session_.UpdateTrackPriority(2, std::nullopt, MoqtTrackPriority{0x20, 0x81}); - EXPECT_EQ(MoqtSessionPeer::NextQueuedRequestIdToServer(&session_), 2); - // The subscriber demotes track 2. Track 0 is first now due to higher + EXPECT_EQ(MoqtSessionPeer::NextQueuedStreamIdToServer(&session_), 4); + // There's another stream for request ID 6. + session_.UpdateTrackPriority(6, std::nullopt, MoqtTrackPriority{0x20, 0x81}); + EXPECT_EQ(MoqtSessionPeer::NextQueuedStreamIdToServer(&session_), 6); + // The subscriber demotes track 6. Track 4 is first now due to higher // publisher priority. - session_.UpdateTrackPriority(2, MoqtTrackPriority{0x20, 0x81}, + session_.UpdateTrackPriority(6, MoqtTrackPriority{0x20, 0x81}, MoqtTrackPriority{0x40, 0x81}); - EXPECT_EQ(MoqtSessionPeer::NextQueuedRequestIdToServer(&session_), 0); + EXPECT_EQ(MoqtSessionPeer::NextQueuedStreamIdToServer(&session_), 4); EXPECT_CALL(mock_session_, CanOpenNextOutgoingUnidirectionalStream()) .WillOnce(Return(true)) .WillOnce(Return(false)); session_.OnCanCreateNewOutgoingUnidirectionalStream(); // The subscription will update with the first stream. It's lower priority - // than request ID 2. - session_.UpdateTrackPriority(0, std::nullopt, MoqtTrackPriority{0x40, 0x82}); - EXPECT_EQ(MoqtSessionPeer::NextQueuedRequestIdToServer(&session_), 2); + // than request ID 6. + session_.UpdateTrackPriority(4, std::nullopt, MoqtTrackPriority{0x40, 0x82}); + EXPECT_EQ(MoqtSessionPeer::NextQueuedStreamIdToServer(&session_), 6); } // Helper functions to handle the many EXPECT_CALLs for FETCH processing and @@ -1456,7 +1445,7 @@ .WillOnce([&](std::unique_ptr<webtransport::StreamVisitor> visitor) { stream_visitor = std::move(visitor); }); - EXPECT_CALL(data_stream, SetPriority).Times(1); + EXPECT_CALL(data_stream, SetPriority); } // Sets expectations to send one object at the start of the stream, and then @@ -1521,25 +1510,29 @@ // All callbacks are called asynchronously. TEST_F(MoqtSessionTest, ProcessFetchGetEverythingFromUpstream) { bidi_wrapper_ = - MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_); + std::make_unique<MoqtBidiStreamTestWrapper>(ResponseStream(kFetchByte)); MoqtFetch fetch = DefaultFetch(); MockTrackPublisher* track = CreateTrackPublisher(); // No callbacks are synchronous. MockFetchTask will store the callbacks. auto fetch_task_ptr = std::make_unique<MockFetchTask>(); MockFetchTask* fetch_task = fetch_task_ptr.get(); + + FetchResponseCallback fetch_response_callback; EXPECT_CALL(*track, StandaloneFetch) - .WillOnce(Return(std::move(fetch_task_ptr))); + .WillOnce([&](Location, Location, MoqtDeliveryOrder, + FetchResponseCallback callback) { + fetch_response_callback = std::move(callback); + return std::move(fetch_task_ptr); + }); bidi_wrapper_->ReceiveMessage(fetch); // Compose and send the FETCH_OK. - MoqtFetchOk expected_ok; - expected_ok.request_id = fetch.request_id; - expected_ok.end_of_track = false; - expected_ok.end_location = Location(1, 4); + MoqtFetchOk expected_ok(false, Location(1, 4)); EXPECT_CALL(mock_bidi_stream_, Writev(SerializedControlMessage(expected_ok), _)); - fetch_task->CallFetchResponseCallback(expected_ok); + std::move(fetch_response_callback)(expected_ok); + // Data arrives. webtransport::test::MockStream data_stream; std::unique_ptr<webtransport::StreamVisitor> stream_visitor; @@ -1554,19 +1547,19 @@ // is the original publisher). TEST_F(MoqtSessionTest, ProcessFetchWholeRangeIsPresent) { bidi_wrapper_ = - MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_); + std::make_unique<MoqtBidiStreamTestWrapper>(ResponseStream(kFetchByte)); MoqtFetch fetch = DefaultFetch(); MockTrackPublisher* track = CreateTrackPublisher(); - MoqtFetchOk expected_ok; - expected_ok.request_id = fetch.request_id; - expected_ok.end_of_track = false; - expected_ok.end_location = Location(1, 4); - auto fetch_task_ptr = - std::make_unique<MockFetchTask>(expected_ok, std::nullopt, true); + MoqtFetchOk expected_ok(false, Location(1, 4)); + auto fetch_task_ptr = std::make_unique<MockFetchTask>(true); MockFetchTask* fetch_task = fetch_task_ptr.get(); EXPECT_CALL(*track, StandaloneFetch) - .WillOnce(Return(std::move(fetch_task_ptr))); + .WillOnce([&](Location, Location, MoqtDeliveryOrder, + FetchResponseCallback callback) { + std::move(callback)(expected_ok); + return std::move(fetch_task_ptr); + }); EXPECT_CALL(mock_bidi_stream_, Writev(SerializedControlMessage(expected_ok), _)); webtransport::test::MockStream data_stream; @@ -1583,28 +1576,29 @@ TEST_F(MoqtSessionTest, SendFragmentedFetchObject) { using ::testing::ByMove; bidi_wrapper_ = - MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_); + std::make_unique<MoqtBidiStreamTestWrapper>(ResponseStream(kFetchByte)); MoqtFetch fetch = DefaultFetch(); fetch.request_id = 3; // Use an odd ID for peer request in client session. MockTrackPublisher* track = CreateTrackPublisher(); // Disable synchronous callback to have more control. - auto fetch_task_ptr = - std::make_unique<MockFetchTask>(std::nullopt, std::nullopt, false); + auto fetch_task_ptr = std::make_unique<MockFetchTask>(); MockFetchTask* fetch_task = fetch_task_ptr.get(); + FetchResponseCallback fetch_response_callback; EXPECT_CALL(*track, StandaloneFetch) - .WillOnce(Return(ByMove(std::move(fetch_task_ptr)))); + .WillOnce([&](Location, Location, MoqtDeliveryOrder, + FetchResponseCallback callback) { + fetch_response_callback = std::move(callback); + return std::move(fetch_task_ptr); + }); // Receive FETCH, send FETCH_OK. bidi_wrapper_->ReceiveMessage(fetch); // FETCH_OK responding to the request. - MoqtFetchOk expected_ok; - expected_ok.request_id = fetch.request_id; - expected_ok.end_of_track = false; - expected_ok.end_location = Location(1, 0); + MoqtFetchOk expected_ok(false, Location(1, 0)); EXPECT_CALL(mock_bidi_stream_, Writev(SerializedControlMessage(expected_ok), _)); - fetch_task->CallFetchResponseCallback(expected_ok); + std::move(fetch_response_callback)(expected_ok); std::unique_ptr<webtransport::StreamVisitor> stream_visitor; EXPECT_CALL(mock_session_, CanOpenNextOutgoingUnidirectionalStream) @@ -1661,16 +1655,20 @@ // the rest. TEST_F(MoqtSessionTest, FetchReturnsObjectBeforeOk) { bidi_wrapper_ = - MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_); + std::make_unique<MoqtBidiStreamTestWrapper>(ResponseStream(kFetchByte)); MoqtFetch fetch = DefaultFetch(); MockTrackPublisher* track = CreateTrackPublisher(); // Object returns synchronously. - auto fetch_task_ptr = - std::make_unique<MockFetchTask>(std::nullopt, std::nullopt, true); + auto fetch_task_ptr = std::make_unique<MockFetchTask>(true); MockFetchTask* fetch_task = fetch_task_ptr.get(); + FetchResponseCallback fetch_response_callback; EXPECT_CALL(*track, StandaloneFetch) - .WillOnce(Return(std::move(fetch_task_ptr))); + .WillOnce([&](Location, Location, MoqtDeliveryOrder, + FetchResponseCallback callback) { + fetch_response_callback = std::move(callback); + return std::move(fetch_task_ptr); + }); webtransport::test::MockStream data_stream; std::unique_ptr<webtransport::StreamVisitor> stream_visitor; ExpectStreamOpen(mock_session_, fetch_task, data_stream, stream_visitor); @@ -1679,26 +1677,27 @@ MoqtFetchTask::GetNextObjectResult::kPending); bidi_wrapper_->ReceiveMessage(fetch); - MoqtFetchOk expected_ok; - expected_ok.request_id = fetch.request_id; - expected_ok.end_of_track = false; - expected_ok.end_location = Location(1, 4); + MoqtFetchOk expected_ok(false, Location(1, 4)); EXPECT_CALL(mock_bidi_stream_, Writev(SerializedControlMessage(expected_ok), _)); - fetch_task->CallFetchResponseCallback(expected_ok); + std::move(fetch_response_callback)(expected_ok); } TEST_F(MoqtSessionTest, FetchReturnsObjectBeforeError) { bidi_wrapper_ = - MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_); + std::make_unique<MoqtBidiStreamTestWrapper>(ResponseStream(kFetchByte)); MoqtFetch fetch = DefaultFetch(); MockTrackPublisher* track = CreateTrackPublisher(); - auto fetch_task_ptr = - std::make_unique<MockFetchTask>(std::nullopt, std::nullopt, true); + auto fetch_task_ptr = std::make_unique<MockFetchTask>(true); MockFetchTask* fetch_task = fetch_task_ptr.get(); + FetchResponseCallback fetch_response_callback; EXPECT_CALL(*track, StandaloneFetch) - .WillOnce(Return(std::move(fetch_task_ptr))); + .WillOnce([&](Location, Location, MoqtDeliveryOrder, + FetchResponseCallback callback) { + fetch_response_callback = std::move(callback); + return std::move(fetch_task_ptr); + }); webtransport::test::MockStream data_stream; std::unique_ptr<webtransport::StreamVisitor> stream_visitor; ExpectStreamOpen(mock_session_, fetch_task, data_stream, stream_visitor); @@ -1707,40 +1706,42 @@ MoqtFetchTask::GetNextObjectResult::kPending); bidi_wrapper_->ReceiveMessage(fetch); - MoqtRequestError expected_error{ - fetch.request_id, RequestErrorCode::kDoesNotExist, std::nullopt, "foo"}; + MoqtRequestError expected_error{RequestErrorCode::kDoesNotExist, std::nullopt, + "foo"}; EXPECT_CALL(mock_bidi_stream_, Writev(SerializedControlMessage(expected_error), _)); - fetch_task->CallFetchResponseCallback(expected_error); + std::move(fetch_response_callback)(expected_error); } TEST_F(MoqtSessionTest, InvalidFetch) { bidi_wrapper_ = - MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_); + std::make_unique<MoqtBidiStreamTestWrapper>(ResponseStream(kFetchByte)); MockTrackPublisher* track = CreateTrackPublisher(); MoqtFetch fetch = DefaultFetch(); EXPECT_CALL(*track, StandaloneFetch) .WillOnce(Return(std::make_unique<MockFetchTask>())); bidi_wrapper_->ReceiveMessage(fetch); EXPECT_CALL(mock_session_, - CloseSession(static_cast<uint64_t>(MoqtError::kInvalidRequestId), - "Duplicate request ID")) + CloseSession(static_cast<uint64_t>(MoqtError::kProtocolViolation), + "FETCH received on stream that already has a fetch")) .Times(1); bidi_wrapper_->ReceiveMessage(fetch); } TEST_F(MoqtSessionTest, FetchFails) { bidi_wrapper_ = - MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_); + std::make_unique<MoqtBidiStreamTestWrapper>(ResponseStream(kFetchByte)); MoqtFetch fetch = DefaultFetch(); MockTrackPublisher* track = CreateTrackPublisher(); auto fetch_task_ptr = std::make_unique<MockFetchTask>(); - MockFetchTask* fetch_task = fetch_task_ptr.get(); EXPECT_CALL(*track, StandaloneFetch) - .WillOnce(Return(std::move(fetch_task_ptr))); - EXPECT_CALL(*fetch_task, GetStatus()) - .WillRepeatedly(Return(absl::Status(absl::StatusCode::kInternal, "foo"))); + .WillOnce([&](Location, Location, MoqtDeliveryOrder, + FetchResponseCallback callback) { + std::move(callback)(MoqtRequestErrorInfo( + RequestErrorCode::kDoesNotExist, std::nullopt, "foo")); + return std::move(fetch_task_ptr); + }); EXPECT_CALL(mock_bidi_stream_, Writev(ControlMessageOfType(MoqtMessageType::kRequestError), _)); bidi_wrapper_->ReceiveMessage(fetch); @@ -1748,26 +1749,24 @@ TEST_F(MoqtSessionTest, FullFetchDeliveryWithFlowControl) { bidi_wrapper_ = - MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_); + std::make_unique<MoqtBidiStreamTestWrapper>(ResponseStream(kFetchByte)); MoqtFetch fetch = DefaultFetch(); MockTrackPublisher* track = CreateTrackPublisher(); - auto fetch_task_ptr = - std::make_unique<MockFetchTask>(std::nullopt, std::nullopt, true); + auto fetch_task_ptr = std::make_unique<MockFetchTask>(true); MockFetchTask* fetch_task = fetch_task_ptr.get(); EXPECT_CALL(*track, StandaloneFetch) .WillOnce(Return(std::move(fetch_task_ptr))); - bidi_wrapper_->ReceiveMessage(fetch); EXPECT_CALL(mock_session_, CanOpenNextOutgoingUnidirectionalStream()) .WillOnce(Return(false)); - fetch_task->CallObjectsAvailableCallback(); + bidi_wrapper_->ReceiveMessage(fetch); // Stream opens, but with no credit. webtransport::test::MockStream data_stream; std::unique_ptr<webtransport::StreamVisitor> stream_visitor; ExpectStreamOpen(mock_session_, fetch_task, data_stream, stream_visitor); - EXPECT_CALL(data_stream, CanWrite()).WillOnce(Return(false)); + EXPECT_CALL(data_stream, CanWrite()).WillRepeatedly(Return(false)); session_.OnCanCreateNewOutgoingUnidirectionalStream(); // Object with FIN ExpectSendObject(fetch_task, data_stream, MoqtObjectStatus::kNormal, @@ -1787,17 +1786,18 @@ SetLargestId(track, Location(4, 10)); ReceiveSubscribeSynchronousOk(track, subscribe, bidi_wrapper_.get()); - webtransport::test::MockStream control_stream; - std::unique_ptr<MoqtBidiStreamTestWrapper> control_wrapper = - MoqtSessionPeer::CreateControlStream(&session_, &control_stream); + webtransport::test::MockStream fetch_stream; + std::unique_ptr<MoqtBidiStreamTestWrapper> fetch_wrapper = + std::make_unique<MoqtBidiStreamTestWrapper>( + ResponseStream(kFetchByte, &fetch_stream)); ASSERT_TRUE(MoqtSessionPeer::RequestIdIsLivePublisher(&session_, subscribe.request_id)); MoqtFetch fetch = DefaultFetch(); fetch.request_id = 3; fetch.fetch = JoiningFetchRelative(1, 2); - EXPECT_CALL(*track, StandaloneFetch(Location(2, 0), Location(4, 10), _)) + EXPECT_CALL(*track, StandaloneFetch(Location(2, 0), Location(4, 10), _, _)) .WillOnce(Return(std::make_unique<MockFetchTask>())); - control_wrapper->ReceiveMessage(fetch); + fetch_wrapper->ReceiveMessage(fetch); } TEST_F(MoqtSessionTest, IncomingAbsoluteJoiningFetch) { @@ -1813,24 +1813,24 @@ ASSERT_TRUE(MoqtSessionPeer::RequestIdIsLivePublisher(&session_, subscribe.request_id)); - webtransport::test::MockStream control_stream; - std::unique_ptr<MoqtBidiStreamTestWrapper> control_wrapper = - MoqtSessionPeer::CreateControlStream(&session_, &control_stream); + webtransport::test::MockStream fetch_stream; + std::unique_ptr<MoqtBidiStreamTestWrapper> fetch_wrapper = + std::make_unique<MoqtBidiStreamTestWrapper>( + ResponseStream(kFetchByte, &fetch_stream)); MoqtFetch fetch = DefaultFetch(); fetch.request_id = 3; fetch.fetch = JoiningFetchAbsolute(1, 2); - EXPECT_CALL(*track, StandaloneFetch(Location(2, 0), Location(4, 10), _)) + EXPECT_CALL(*track, StandaloneFetch(Location(2, 0), Location(4, 10), _, _)) .WillOnce(Return(std::make_unique<MockFetchTask>())); - control_wrapper->ReceiveMessage(fetch); + fetch_wrapper->ReceiveMessage(fetch); } TEST_F(MoqtSessionTest, IncomingJoiningFetchBadRequestId) { bidi_wrapper_ = - MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_); + std::make_unique<MoqtBidiStreamTestWrapper>(ResponseStream(kFetchByte)); MoqtFetch fetch = DefaultFetch(); fetch.fetch = JoiningFetchRelative(1, 2); MoqtRequestError expected_error = { - /*request_id=*/1, RequestErrorCode::kInvalidJoiningRequestId, /*retry_interval=*/std::nullopt, "Joining Fetch for non-existent request", @@ -1849,9 +1849,10 @@ SetLargestId(track, Location(2, 10)); ReceiveSubscribeSynchronousOk(track, subscribe, bidi_wrapper_.get()); - webtransport::test::MockStream control_stream; - std::unique_ptr<MoqtBidiStreamTestWrapper> control_wrapper = - MoqtSessionPeer::CreateControlStream(&session_, &control_stream); + webtransport::test::MockStream fetch_stream; + std::unique_ptr<MoqtBidiStreamTestWrapper> fetch_wrapper = + std::make_unique<MoqtBidiStreamTestWrapper>( + ResponseStream(kFetchByte, &fetch_stream)); MoqtFetch fetch = DefaultFetch(); fetch.request_id = 3; fetch.fetch = JoiningFetchRelative(1, 2); @@ -1859,51 +1860,51 @@ CloseSession(static_cast<uint64_t>(MoqtError::kProtocolViolation), "Joining Fetch for non-forwarding subscribe")) .Times(1); - control_wrapper->ReceiveMessage(fetch); + fetch_wrapper->ReceiveMessage(fetch); } TEST_F(MoqtSessionTest, SendJoiningFetch) { + webtransport::test::MockStream subscribe_stream; + std::unique_ptr<MoqtBidiStreamTestWrapper> subscribe_wrapper; + // Call in reverse order of opening, so that the EXPECT_CALLs execute in the + // proper order. PrepareRequestStream(bidi_wrapper_); - webtransport::test::MockStream control_stream; - std::unique_ptr<MoqtBidiStreamTestWrapper> control_wrapper = - MoqtSessionPeer::CreateControlStream(&session_, &control_stream); + PrepareRequestStream(subscribe_wrapper, &subscribe_stream); MoqtSubscribe expected_subscribe( 0, FullTrackName("foo", "bar"), MessageParameters(MoqtFilterType::kLargestObject)); - MoqtFetch expected_fetch = { - /*request_id=*/2, - /*fetch=*/JoiningFetchRelative(0, 1), - MessageParameters(), - }; - EXPECT_CALL(mock_bidi_stream_, + MoqtFetch expected_fetch(2, JoiningFetchRelative(0, 1), MessageParameters()); + EXPECT_CALL(subscribe_stream, Writev(SerializedControlMessage(expected_subscribe), _)); - EXPECT_CALL(control_stream, + EXPECT_CALL(mock_bidi_stream_, Writev(SerializedControlMessage(expected_fetch), _)); - EXPECT_TRUE(session_.RelativeJoiningFetch(expected_subscribe.full_track_name, - &remote_track_visitor_, nullptr, 1, - MessageParameters())); + EXPECT_NE(session_.RelativeJoiningFetch(expected_subscribe.full_track_name, + &remote_track_visitor_, nullptr, 1, + MessageParameters()), + nullptr); } TEST_F(MoqtSessionTest, SendJoiningFetchNoFlowControl) { + webtransport::test::MockStream subscribe_stream; + std::unique_ptr<MoqtBidiStreamTestWrapper> subscribe_wrapper; + // Call in reverse order of opening, so that the EXPECT_CALLs execute in the + // proper order. PrepareRequestStream(bidi_wrapper_); - webtransport::test::MockStream control_stream; - std::unique_ptr<MoqtBidiStreamTestWrapper> control_wrapper = - MoqtSessionPeer::CreateControlStream(&session_, &control_stream); - EXPECT_CALL(mock_bidi_stream_, + PrepareRequestStream(subscribe_wrapper, &subscribe_stream); + EXPECT_CALL(subscribe_stream, Writev(ControlMessageOfType(MoqtMessageType::kSubscribe), _)); - EXPECT_CALL(control_stream, + EXPECT_CALL(mock_bidi_stream_, Writev(ControlMessageOfType(MoqtMessageType::kFetch), _)); EXPECT_TRUE(session_.RelativeJoiningFetch(FullTrackName("foo", "bar"), &remote_track_visitor_, 0, MessageParameters())); - EXPECT_CALL(remote_track_visitor_, OnReply).Times(1); MessageParameters parameters; parameters.largest_object = Location(2, 0); - bidi_wrapper_->ReceiveMessage( + subscribe_wrapper->ReceiveMessage( MoqtSubscribeOk(0, 2, parameters, TrackExtensions())); - control_wrapper->ReceiveMessage(MoqtFetchOk( - 2, false, Location(2, 0), MessageParameters(), TrackExtensions())); + bidi_wrapper_->ReceiveMessage(MoqtFetchOk( + false, Location(2, 0), MessageParameters(), TrackExtensions())); // Packet arrives on FETCH stream. MoqtObject object = { /*request_id=*/2, @@ -1925,9 +1926,6 @@ MoqtSessionPeer::CreateIncomingStreamVisitor(&session_, &data_stream)); data_stream.Receive(header.AsStringView(), false); EXPECT_CALL(remote_track_visitor_, OnObjectFragment).Times(1); - // Last object of the FETCH causes FETCH_CANCEL. - EXPECT_CALL(control_stream, - Writev(ControlMessageOfType(MoqtMessageType::kFetchCancel), _)); data_stream.Receive("foo", false); } @@ -2045,80 +2043,72 @@ } TEST_F(MoqtSessionTest, FetchThenOkThenCancel) { - bidi_wrapper_ = - MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_); - std::unique_ptr<MoqtFetchTask> fetch_task; - session_.Fetch( + PrepareRequestStream(bidi_wrapper_); + EXPECT_CALL(mock_bidi_stream_, + Writev(ControlMessageOfType(MoqtMessageType::kFetch), _)); + bool fetch_ok = false; + std::unique_ptr<MoqtFetchTask> fetch_task = session_.Fetch( FullTrackName("foo", "bar"), - [&](std::unique_ptr<MoqtFetchTask> task) { - fetch_task = std::move(task); + [&](std::variant<FetchOkData, MoqtRequestErrorInfo> res) { + fetch_ok = std::holds_alternative<FetchOkData>(res); }, Location(0, 0), 4, std::nullopt, MessageParameters()); - MoqtFetchOk ok = { - /*request_id=*/0, - /*end_of_track=*/false, Location(3, 25), - MessageParameters(), TrackExtensions(), - }; - bidi_wrapper_->ReceiveMessage(ok); ASSERT_NE(fetch_task, nullptr); + MoqtFetchOk ok(/*end_of_track=*/false, Location(3, 25)); + bidi_wrapper_->ReceiveMessage(ok); + EXPECT_TRUE(fetch_ok); EXPECT_TRUE(fetch_task->GetStatus().ok()); PublishedObject object; EXPECT_EQ(fetch_task->GetNextObject(object), MoqtFetchTask::GetNextObjectResult::kPending); // Cancel the fetch. - EXPECT_CALL(mock_bidi_stream_, - Writev(ControlMessageOfType(MoqtMessageType::kFetchCancel), _)); + EXPECT_CALL(mock_bidi_stream_, ResetWithUserCode(kResetCodeCancelled)); fetch_task.reset(); } TEST_F(MoqtSessionTest, FetchThenError) { - bidi_wrapper_ = - MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_); - std::unique_ptr<MoqtFetchTask> fetch_task; - session_.Fetch( + PrepareRequestStream(bidi_wrapper_); + EXPECT_CALL(mock_bidi_stream_, + Writev(ControlMessageOfType(MoqtMessageType::kFetch), _)); + MoqtRequestError error(RequestErrorCode::kUnauthorized, + /*retry_interval=*/std::nullopt, + "No username provided"); + std::unique_ptr<MoqtFetchTask> fetch_task = session_.Fetch( FullTrackName("foo", "bar"), - [&](std::unique_ptr<MoqtFetchTask> task) { - fetch_task = std::move(task); + [&](std::variant<FetchOkData, MoqtRequestErrorInfo> res) { + EXPECT_EQ(std::get<MoqtRequestErrorInfo>(res), error); }, Location(0, 0), 4, std::nullopt, MessageParameters()); - MoqtRequestError error = { - /*request_id=*/0, - RequestErrorCode::kUnauthorized, - /*retry_interval=*/std::nullopt, - "No username provided", - }; - bidi_wrapper_->ReceiveMessage(error); ASSERT_NE(fetch_task, nullptr); - EXPECT_TRUE(absl::IsPermissionDenied(fetch_task->GetStatus())); - EXPECT_EQ(fetch_task->GetStatus().message(), "No username provided"); + ExpectFin(mock_bidi_stream_); + bidi_wrapper_->ReceiveMessage(error); } // The application takes objects as they arrive. TEST_F(MoqtSessionTest, IncomingFetchObjectsGreedyApp) { - bidi_wrapper_ = - MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_); - std::unique_ptr<MoqtFetchTask> fetch_task; + PrepareRequestStream(bidi_wrapper_); + EXPECT_CALL(mock_bidi_stream_, + Writev(ControlMessageOfType(MoqtMessageType::kFetch), _)); uint64_t expected_object_id = 0; - session_.Fetch( + bool fetch_ok = false; + std::unique_ptr<MoqtFetchTask> fetch_task = session_.Fetch( FullTrackName("foo", "bar"), - [&](std::unique_ptr<MoqtFetchTask> task) { - fetch_task = std::move(task); - fetch_task->SetObjectAvailableCallback([&]() { - PublishedObject object; - MoqtFetchTask::GetNextObjectResult result; - do { - result = fetch_task->GetNextObject(object); - if (result == MoqtFetchTask::GetNextObjectResult::kSuccess) { - EXPECT_EQ(object.metadata.location.object, expected_object_id); - ++expected_object_id; - } - if (result == MoqtFetchTask::GetNextObjectResult::kError) { - break; - } - } while (result != MoqtFetchTask::GetNextObjectResult::kPending); - }); + [&](std::variant<FetchOkData, MoqtRequestErrorInfo> res) { + fetch_ok = std::holds_alternative<FetchOkData>(res); }, Location(0, 0), 4, std::nullopt, MessageParameters()); + ASSERT_NE(fetch_task, nullptr); + fetch_task->SetObjectAvailableCallback([&]() { + PublishedObject object; + MoqtFetchTask::GetNextObjectResult result; + do { + result = fetch_task->GetNextObject(object); + if (result == MoqtFetchTask::GetNextObjectResult::kSuccess) { + EXPECT_EQ(object.metadata.location.object, expected_object_id); + ++expected_object_id; + } + } while (result == MoqtFetchTask::GetNextObjectResult::kSuccess); + }); // Build queue of packets to arrive. std::queue<quiche::QuicheBuffer> headers; std::queue<std::string> payloads; @@ -2140,11 +2130,13 @@ headers.push(framer.SerializeObjectHeader( object, MoqtDataStreamType::Fetch(), metadata)); metadata = PublishedObjectMetadata(); - metadata->location.object = i; // only object ID matters. + metadata->location = Location(0, i); + metadata->subgroup = 0; + metadata->publisher_priority = 128; payloads.push("foo"); } - // Open stream, deliver two objects before FETCH_OK. Neither should be read. + // Open stream, deliver two objects before FETCH_OK. webtransport::test::InMemoryStream data_stream(kIncomingUniStreamId); data_stream.SetVisitor( MoqtSessionPeer::CreateIncomingStreamVisitor(&session_, &data_stream)); @@ -2154,19 +2146,11 @@ headers.pop(); payloads.pop(); } - EXPECT_EQ(fetch_task, nullptr); - EXPECT_GT(data_stream.ReadableBytes(), 0); // FETCH_OK arrives, objects are delivered. - MoqtFetchOk ok = { - /*request_id=*/0, - /*end_of_track=*/false, - /*end_location=*/Location(3, 25), - MessageParameters(), - TrackExtensions(), - }; + MoqtFetchOk ok(/*end_of_track=*/false, Location(3, 25)); bidi_wrapper_->ReceiveMessage(ok); - ASSERT_NE(fetch_task, nullptr); + EXPECT_TRUE(fetch_ok); EXPECT_EQ(expected_object_id, 2); // Deliver the rest of the objects. @@ -2180,19 +2164,20 @@ } TEST_F(MoqtSessionTest, IncomingFetchObjectsSlowApp) { - bidi_wrapper_ = - MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_); - std::unique_ptr<MoqtFetchTask> fetch_task; + PrepareRequestStream(bidi_wrapper_); + EXPECT_CALL(mock_bidi_stream_, + Writev(ControlMessageOfType(MoqtMessageType::kFetch), _)); uint64_t expected_object_id = 0; bool objects_available = false; - session_.Fetch( + bool fetch_ok = false; + std::unique_ptr<MoqtFetchTask> fetch_task = session_.Fetch( FullTrackName("foo", "bar"), - [&](std::unique_ptr<MoqtFetchTask> task) { - fetch_task = std::move(task); - fetch_task->SetObjectAvailableCallback( - [&]() { objects_available = true; }); + [&](std::variant<FetchOkData, MoqtRequestErrorInfo> res) { + fetch_ok = std::holds_alternative<FetchOkData>(res); }, Location(0, 0), 4, std::nullopt, MessageParameters()); + ASSERT_NE(fetch_task, nullptr); + fetch_task->SetObjectAvailableCallback([&]() { objects_available = true; }); // Build queue of packets to arrive. std::queue<quiche::QuicheBuffer> headers; std::queue<std::string> payloads; @@ -2214,11 +2199,13 @@ headers.push(framer.SerializeObjectHeader( object, MoqtDataStreamType::Fetch(), metadata)); metadata = PublishedObjectMetadata(); - metadata->location.object = i; // only object ID matters. + metadata->location = Location(0, i); + metadata->subgroup = 0; + metadata->publisher_priority = 128; payloads.push("foo"); } - // Open stream, deliver two objects before FETCH_OK. Neither should be read. + // Open stream, deliver two objects before FETCH_OK. webtransport::test::InMemoryStream data_stream(kIncomingUniStreamId); data_stream.SetVisitor( MoqtSessionPeer::CreateIncomingStreamVisitor(&session_, &data_stream)); @@ -2228,17 +2215,12 @@ headers.pop(); payloads.pop(); } - EXPECT_EQ(fetch_task, nullptr); - EXPECT_GT(data_stream.ReadableBytes(), 0); + EXPECT_TRUE(objects_available); // FETCH_OK arrives, objects are available. - MoqtFetchOk ok = { - /*request_id=*/0, - /*end_of_track=*/false, Location(3, 25), - MessageParameters(), TrackExtensions(), - }; + MoqtFetchOk ok(/*end_of_track=*/false, Location(3, 25)); bidi_wrapper_->ReceiveMessage(ok); - ASSERT_NE(fetch_task, nullptr); + EXPECT_TRUE(fetch_ok); EXPECT_TRUE(objects_available); // Get the objects @@ -2250,7 +2232,7 @@ EXPECT_EQ(new_object.metadata.location.object, expected_object_id); ++expected_object_id; } - } while (result != MoqtFetchTask::GetNextObjectResult::kPending); + } while (result == MoqtFetchTask::GetNextObjectResult::kSuccess); EXPECT_EQ(expected_object_id, 2); objects_available = false; @@ -2271,7 +2253,7 @@ EXPECT_EQ(new_object.metadata.location.object, expected_object_id); ++expected_object_id; } - } while (result != MoqtFetchTask::GetNextObjectResult::kPending); + } while (result == MoqtFetchTask::GetNextObjectResult::kSuccess); EXPECT_EQ(expected_object_id, 4); } @@ -2307,10 +2289,11 @@ EXPECT_FALSE(session_.PublishNamespace( TrackNamespace{"foo"}, MessageParameters(), +[](std::variant<MessageParameters, MoqtRequestErrorInfo>) {}, +[]() {})); - EXPECT_FALSE(session_.Fetch( - FullTrackName{TrackNamespace({"foo"}), "bar"}, - +[](std::unique_ptr<MoqtFetchTask>) {}, Location(0, 0), 5, std::nullopt, - MessageParameters())); + EXPECT_EQ(session_.Fetch( + FullTrackName{TrackNamespace({"foo"}), "bar"}, + +[](std::variant<FetchOkData, MoqtRequestErrorInfo>) {}, + Location(0, 0), 5, std::nullopt, MessageParameters()), + nullptr); // Error on additional GOAWAY. EXPECT_CALL(mock_session_, CloseSession(static_cast<uint64_t>(MoqtError::kProtocolViolation), @@ -2333,12 +2316,6 @@ Writev(ControlMessageOfType(MoqtMessageType::kGoAway), _)); session_.GoAway(""); - EXPECT_CALL(mock_bidi_stream_, - Writev(ControlMessageOfType(MoqtMessageType::kRequestError), _)); - MoqtFetch fetch = DefaultFetch(); - fetch.request_id = 5; - bidi_wrapper_->ReceiveMessage(fetch); - // All new bidi streams types are immediately rejected. webtransport::test::MockStream new_request_stream; EXPECT_CALL(mock_session_, AcceptIncomingBidirectionalStream) @@ -2376,10 +2353,11 @@ EXPECT_FALSE(session_.PublishNamespace( TrackNamespace{"foo"}, MessageParameters(), +[](std::variant<MessageParameters, MoqtRequestErrorInfo>) {}, +[]() {})); - EXPECT_FALSE(session_.Fetch( - FullTrackName(TrackNamespace({"foo"}), "bar"), - +[](std::unique_ptr<MoqtFetchTask>) {}, Location(0, 0), 5, std::nullopt, - MessageParameters())); + EXPECT_EQ(session_.Fetch( + FullTrackName(TrackNamespace({"foo"}), "bar"), + +[](std::variant<FetchOkData, MoqtRequestErrorInfo>) {}, + Location(0, 0), 5, std::nullopt, MessageParameters()), + nullptr); session_.GoAway(""); // GoAway timer fires. auto* goaway_alarm = @@ -2593,14 +2571,11 @@ }; EXPECT_CALL(mock_uni_stream_, GetStreamId()) .WillRepeatedly(Return(kIncomingUniStreamId)); - EXPECT_CALL(mock_session_, GetStreamById(kIncomingUniStreamId)) - .WillRepeatedly(Return(&mock_uni_stream_)); + EXPECT_CALL(remote_track_visitor_, + OnStreamFin(FullTrackName("foo", "bar"), DataStreamIndex(0, 0))); std::unique_ptr<webtransport::StreamVisitor> data_stream; DeliverObject(object, /*fin=*/true, mock_session_, &mock_uni_stream_, data_stream, &remote_track_visitor_); - // The data stream died and destroyed the visitor (IncomingDataStream). - EXPECT_CALL(remote_track_visitor_, - OnStreamFin(FullTrackName("foo", "bar"), DataStreamIndex(0, 0))); data_stream.reset(); } @@ -2640,10 +2615,10 @@ std::unique_ptr<webtransport::StreamVisitor> data_stream; DeliverObject(object, /*fin=*/false, mock_session_, &mock_uni_stream_, data_stream, &remote_track_visitor_); - // The data stream died and destroyed the visitor (IncomingDataStream). - data_stream->OnResetStreamReceived(kResetCodeCancelled); + // The data stream died and notified the visitor (IncomingDataStream). EXPECT_CALL(remote_track_visitor_, OnStreamReset(FullTrackName("foo", "bar"), DataStreamIndex(0, 0))); + data_stream->OnResetStreamReceived(kResetCodeCancelled); data_stream.reset(); } @@ -2735,11 +2710,13 @@ } TEST_F(MoqtSessionTest, IncomingRequestUpdateTriggersRequestError) { - bidi_wrapper_ = - MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_); + bidi_wrapper_ = std::make_unique<MoqtBidiStreamTestWrapper>( + ResponseStream(kSubscribeByte)); EXPECT_CALL(mock_bidi_stream_, Writev(ControlMessageOfType(MoqtMessageType::kRequestError), _)); - bidi_wrapper_->ReceiveMessage(MoqtRequestUpdate{3, 1, MessageParameters()}); + EXPECT_QUICHE_BUG(bidi_wrapper_->ReceiveMessage( + MoqtRequestUpdate{3, 1, MessageParameters()}), + "Received REQUEST_UPDATE, no subscription state"); } TEST_F(MoqtSessionTest, StopSendingBlocksSubgroup) { @@ -2972,14 +2949,17 @@ sub_stream.write_buffer().clear(); // 6. Fetch + webtransport::test::InMemoryStreamWithWriteBuffer fetch_stream(7); + EXPECT_CALL(mock_session_, OpenOutgoingBidirectionalStream) + .WillOnce(Return(&fetch_stream)); FullTrackName fetch_track("namespace2", "fetch_track"); - bool f1 = session_.Fetch( - fetch_track, [](std::unique_ptr<MoqtFetchTask>) {}, Location(0, 0), 1, - std::nullopt, MessageParameters()); - EXPECT_TRUE(f1); - EXPECT_EQ(get_request_id(control_stream), next_request_id); + std::unique_ptr<MoqtFetchTask> f1 = session_.Fetch( + fetch_track, [](std::variant<FetchOkData, MoqtRequestErrorInfo>) {}, + Location(0, 0), 1, std::nullopt, MessageParameters()); + EXPECT_NE(f1, nullptr); + EXPECT_EQ(get_request_id(fetch_stream), next_request_id); next_request_id += 2; - control_stream.write_buffer().clear(); + fetch_stream.write_buffer().clear(); // 7. SubscribeNamespace (duplicating the first call) webtransport::test::InMemoryStreamWithWriteBuffer sub_ns_stream_2(8);
diff --git a/quiche/quic/moqt/moqt_subscribe_stream.cc b/quiche/quic/moqt/moqt_subscribe_stream.cc index 637e370..7b9f517 100644 --- a/quiche/quic/moqt/moqt_subscribe_stream.cc +++ b/quiche/quic/moqt/moqt_subscribe_stream.cc
@@ -184,8 +184,8 @@ if (track_publisher == nullptr) { QUIC_DLOG(INFO) << "SUBSCRIBE for " << message.full_track_name << " rejected by the application: does not exist"; - return SendRequestError(message.request_id, RequestErrorCode::kDoesNotExist, - std::nullopt, "not found", /*fin=*/true); + return SendRequestError(RequestErrorCode::kDoesNotExist, std::nullopt, + "not found"); } subscription_ = std::make_unique<LivePublisher>( *framer(), track_publisher, this, message.request_id, track_alias_, @@ -194,10 +194,8 @@ bool result = std::move(add_callback_)(subscription_.get()); add_callback_ = nullptr; if (!result) { - return SendRequestError(message.request_id, - RequestErrorCode::kDuplicateSubscription, - std::nullopt, "duplicate subscription", - /*fin=*/true); + return SendRequestError(RequestErrorCode::kDuplicateSubscription, + std::nullopt, "duplicate subscription"); } } // Don't add the publisher until we know it's successful. @@ -209,9 +207,8 @@ const MoqtRequestUpdate& message) { if (subscription_ == nullptr) { QUICHE_BUG(INFO) << "Received REQUEST_UPDATE, no subscription state"; - return SendRequestError(message.request_id, - RequestErrorCode::kInternalError, std::nullopt, - "no subscription", /*fin=*/true); + return SendRequestError(RequestErrorCode::kInternalError, std::nullopt, + "no subscription"); } subscription_->Update(message.parameters); return SendRequestOk(message.request_id, MessageParameters());
diff --git a/quiche/quic/moqt/moqt_subscribe_stream_test.cc b/quiche/quic/moqt/moqt_subscribe_stream_test.cc index 02cb398..a380b38 100644 --- a/quiche/quic/moqt/moqt_subscribe_stream_test.cc +++ b/quiche/quic/moqt/moqt_subscribe_stream_test.cc
@@ -179,21 +179,19 @@ Writev(ControlMessageOfType(MoqtMessageType::kSubscribe), _)) .WillOnce(Return(absl::OkStatus())); stream_->BindStream(&mock_stream_); + MoqtRequestError request_error(RequestErrorCode::kUnauthorized, std::nullopt, + "unauthorized"); EXPECT_CALL(mock_subscribe_visitor_, OnReply(track_name_, _)) - .WillOnce( - [](const FullTrackName&, - const std::variant<SubscribeOkData, MoqtRequestErrorInfo>& reply) { - ASSERT_TRUE(std::holds_alternative<MoqtRequestErrorInfo>(reply)); - EXPECT_EQ(std::get<MoqtRequestErrorInfo>(reply).error_code, - RequestErrorCode::kUnauthorized); - }); + .WillOnce([err = request_error]( + const FullTrackName&, + const std::variant<SubscribeOkData, MoqtRequestErrorInfo>& + reply) { + ASSERT_TRUE(std::holds_alternative<MoqtRequestErrorInfo>(reply)); + EXPECT_EQ(std::get<MoqtRequestErrorInfo>(reply), err); + }); EXPECT_CALL(mock_stream_, Writev(testing::IsEmpty(), _)) .WillOnce(Return(absl::OkStatus())); EXPECT_CALL(mock_remove_callback_, Call); - MoqtRequestError request_error; - request_error.request_id = kRequestId; - request_error.error_code = RequestErrorCode::kUnauthorized; - request_error.reason_phrase = "unauthorized"; QUICHE_EXPECT_OK(stream_->OnControlMessage(request_error)); }
diff --git a/quiche/quic/moqt/moqt_track_status_stream.cc b/quiche/quic/moqt/moqt_track_status_stream.cc index 282cb27..62c3ea1 100644 --- a/quiche/quic/moqt/moqt_track_status_stream.cc +++ b/quiche/quic/moqt/moqt_track_status_stream.cc
@@ -117,9 +117,8 @@ } publisher_ = session()->GetTrackPublisher(message.full_track_name); if (publisher_ == nullptr) { - return SendRequestError(message.request_id, RequestErrorCode::kDoesNotExist, - std::nullopt, "Track does not exist", - /*fin=*/true); + return SendRequestError(RequestErrorCode::kDoesNotExist, std::nullopt, + "Track does not exist"); } // If the upstream subscription is already established, the code below will // invoke `OnSubscribeAccepted` immediately. @@ -146,9 +145,7 @@ return; } // Since `fin` is true, this will also reset `publisher_` if present. - CheckStatus(SendRequestError(*request_id_, info.error_code, - info.retry_interval, info.reason_phrase, - /*fin=*/true)); + CheckStatus(SendRequestError(info)); } void MoqtTrackStatusResponseStream::OnTrackPublisherGone() {
diff --git a/quiche/quic/moqt/moqt_track_status_stream_test.cc b/quiche/quic/moqt/moqt_track_status_stream_test.cc index cfb6b4b..274fc94 100644 --- a/quiche/quic/moqt/moqt_track_status_stream_test.cc +++ b/quiche/quic/moqt/moqt_track_status_stream_test.cc
@@ -142,24 +142,17 @@ EXPECT_CALL(mock_stream_, Writev(ControlMessageOfType(MoqtMessageType::kTrackStatus), _)); stream.BindStream(&mock_stream_); - + MoqtRequestError error(RequestErrorCode::kDoesNotExist, std::nullopt, + "Track does not exist"); EXPECT_CALL(response_callback_, Call) .WillOnce([&](std::variant<MessageParameters, MoqtRequestErrorInfo> v) { ASSERT_TRUE(std::holds_alternative<MoqtRequestErrorInfo>(v)); - auto info = std::get<MoqtRequestErrorInfo>(v); - EXPECT_EQ(info.error_code, RequestErrorCode::kDoesNotExist); - EXPECT_EQ(info.reason_phrase, "Track does not exist"); + EXPECT_EQ(std::get<MoqtRequestErrorInfo>(v), error); }); EXPECT_CALL( mock_stream_, Writev(IsEmpty(), Property(&webtransport::StreamWriteOptions::send_fin, true))); - - MoqtRequestError error; - error.request_id = kRequestId; - error.error_code = RequestErrorCode::kDoesNotExist; - error.reason_phrase = "Track does not exist"; - QUICHE_EXPECT_OK( stream.OnRawControlMessage(GenericMessageToRawControlMessage(error))); } @@ -191,21 +184,15 @@ EXPECT_CALL(mock_stream_, Writev(ControlMessageOfType(MoqtMessageType::kTrackStatus), _)); stream.BindStream(&mock_stream_); - + MoqtRequestError error(RequestErrorCode::kDoesNotExist, std::nullopt, + "Track does not exist"); EXPECT_CALL(response_callback_, Call); EXPECT_CALL( mock_stream_, Writev(IsEmpty(), Property(&webtransport::StreamWriteOptions::send_fin, true))); - - MoqtRequestError error; - error.request_id = kRequestId; - error.error_code = RequestErrorCode::kDoesNotExist; - error.reason_phrase = "Track does not exist"; - QUICHE_EXPECT_OK( stream.OnRawControlMessage(GenericMessageToRawControlMessage(error))); - EXPECT_THAT( stream.OnRawControlMessage(GenericMessageToRawControlMessage(error)), StatusIs(absl::StatusCode::kInvalidArgument, "Duplicate REQUEST_ERROR"));
diff --git a/quiche/quic/moqt/moqt_uni_stream.cc b/quiche/quic/moqt/moqt_uni_stream.cc index 6320e27..924a002 100644 --- a/quiche/quic/moqt/moqt_uni_stream.cc +++ b/quiche/quic/moqt/moqt_uni_stream.cc
@@ -91,7 +91,6 @@ : OutgoingUniStream(framer, stream, priority, track_alias), index_(index), visitor_(std::move(visitor)), - track_alias_(track_alias), publisher_(track_publisher), next_object_(first_object) { @@ -273,21 +272,28 @@ : OutgoingUniStream(framer, stream, priority, request_id), incoming_objects_(std::move(incoming_objects)), close_callback_(std::move(close_callback)) { - incoming_objects_->SetObjectAvailableCallback( - [this]() { this->OnCanWrite(); }); trace_recorder->RecordFetchStreamCreated(stream->GetStreamId()); } -OutgoingFetchStream::~OutgoingFetchStream() { - if (close_callback_ != nullptr) { - std::move(close_callback_)(); +void OutgoingFetchStream::Init() { + if (incoming_objects_ != nullptr) { + incoming_objects_->SetObjectAvailableCallback( + [this]() { this->OnCanWrite(); }); } - close_callback_ = nullptr; +} + +OutgoingFetchStream::~OutgoingFetchStream() { + // If Detach() has not been called, this was not FINed and there should be a + // non-OK status. + if (status_.ok()) { + status_ = absl::CancelledError("stream destroyed"); + } + Detach(); } void OutgoingFetchStream::OnCanWrite() { PublishedObject object; - while (stream().CanWrite()) { + while (stream().CanWrite() && incoming_objects_ != nullptr) { MoqtFetchTask::GetNextObjectResult result = incoming_objects_->GetNextObject(object); switch (result) { @@ -323,22 +329,36 @@ case MoqtFetchTask::GetNextObjectResult::kEof: // TODO(martinduke): Either prefetch the next object, or alter the API // so that we're not sending FIN in a separate frame. - if (!webtransport::SendFinOnStream(stream()).ok()) { + status_ = webtransport::SendFinOnStream(stream()); + if (!status_.ok()) { QUICHE_DVLOG(1) << "Sending FIN onStream " << stream().GetStreamId() << " failed"; } + Detach(); return; - case MoqtFetchTask::GetNextObjectResult::kError: - stream().ResetWithUserCode(static_cast<webtransport::StreamErrorCode>( - incoming_objects_->GetStatus().code())); + case MoqtFetchTask::GetNextObjectResult::kError: { + status_ = incoming_objects_->GetStatus(); + stream().ResetWithUserCode(StatusToMoqtStreamError(status_)); + Detach(); return; + } } } } void OutgoingFetchStream::OnStopSendingReceived( webtransport::StreamErrorCode error_code) { + status_ = MoqtStreamErrorToStatus(error_code, "stop sending"); stream().ResetWithUserCode(error_code); + Detach(); +} + +void OutgoingFetchStream::Detach() { + if (close_callback_ != nullptr) { + std::move(close_callback_)(status_); + close_callback_ = nullptr; + } + incoming_objects_.reset(); } IncomingDataStream::~IncomingDataStream() { @@ -349,22 +369,11 @@ "learning track alias"; return; } - if (!track_.IsValid()) { - return; + if (status_.ok()) { + // Notify counterparts that the stream was not cleanly closed. + status_ = absl::CancelledError("stream destroyed"); } - if (IsFetch()) { - auto fetch = absl::down_cast<UpstreamFetch*>(track_.GetIfAvailable()); - if (fetch != nullptr) { - fetch->OnStreamClosed(); - } - return; - } - // It's a subscribe. - auto subscribe = absl::down_cast<LiveSubscriber*>(track_.GetIfAvailable()); - if (subscribe == nullptr) { - return; - } - subscribe->OnStreamClosed(fin_received_, index_); + Detach(); } void IncomingDataStream::OnObjectMessage(const MoqtObject& message, @@ -457,19 +466,14 @@ bytes_received_this_object_); } } else { // FETCH - track->OnObjectOrOk(); - UpstreamFetch* fetch = absl::down_cast<UpstreamFetch*>(track); - UpstreamFetch::UpstreamFetchTask* task = fetch->task(); - if (task == nullptr) { - // The application killed the FETCH. - stream_->SendStopSending(kResetCodeCancelled); + if (fetch_task_ == nullptr) { return; } - if (!task->HasObject()) { - task->NewObject(message); + if (!fetch_task_->HasObject()) { + fetch_task_->NewObject(message); } - if (task->NeedsMorePayload() && !payload.empty()) { - task->AppendPayloadToObject(payload); + if (fetch_task_->NeedsMorePayload() && !payload.empty()) { + fetch_task_->AppendPayloadToObject(payload); } } if (end_of_message) { @@ -486,24 +490,50 @@ QUICHE_BUG(quic_bug_read_one_object_parser_unexpected_state) << "Requesting object, parser in unexpected state"; } - if (!track_.IsValid()) { + if (fetch_task_ == nullptr) { return; } - UpstreamFetch* fetch = - absl::down_cast<UpstreamFetch*>(track_.GetIfAvailable()); - UpstreamFetch::UpstreamFetchTask* task = fetch->task(); - if (task == nullptr) { - return; - } - if (task->HasObject() && !task->NeedsMorePayload()) { + if (fetch_task_->HasObject() && !fetch_task_->NeedsMorePayload()) { return; // The message is complete. Do not read more. } - uint64_t start_length = task->payload_length(); + uint64_t start_length = fetch_task_->payload_length(); parser_.ReadAtMostOneObject(); // If it read an object, it called OnObjectMessage and may have altered the // task's object state. - if (task->payload_length() > start_length) { - task->NotifyNewObject(); + if (fetch_task_ != nullptr && fetch_task_->payload_length() > start_length) { + fetch_task_->NotifyNewObject(); + } +} + +void IncomingDataStream::set_fetch_task(UpstreamFetchTask* fetch_task) { + fetch_task_ = fetch_task; + // Replace the bidi stream's callback with a new one. + fetch_task_->set_task_destroyed_callback([this]() { + fetch_task_ = nullptr; + stream_->SendStopSending(kResetCodeCancelled); + ObjectSubscriber* track = track_.GetIfAvailable(); + if (track != nullptr) { + // Inform the bidi stream, since its callback was overwritten. + track->OnStreamClosed( + MoqtStreamErrorToStatus(kResetCodeCancelled, "canceled"), + std::nullopt); + } + }); + fetch_task_->set_can_read_callback([this]() { MaybeReadOneObject(); }); +} + +void IncomingDataStream::Detach() { + if (fetch_task_ != nullptr) { + UpstreamFetchTask* fetch_task = fetch_task_; + fetch_task_ = nullptr; + fetch_task->set_can_read_callback(nullptr); + fetch_task->set_task_destroyed_callback(nullptr); + fetch_task->OnStreamAndFetchClosed(status_); + } + ObjectSubscriber* track = track_.GetIfAvailable(); + if (track != nullptr) { + track->OnStreamClosed(status_, index_); + track_ = quiche::QuicheWeakPtr<ObjectSubscriber>(); } } @@ -537,7 +567,7 @@ stream_->SendStopSending(kResetCodeCancelled); return; } - subscribe->OnStreamOpened(); + subscribe->OnStreamOpened(this); parser_.set_default_publisher_priority( subscribe->default_publisher_priority()); visitor_ = subscribe->visitor(); @@ -553,13 +583,19 @@ stream_->SendStopSending(kResetCodeCancelled); return; } - UpstreamFetch* fetch = - absl::down_cast<UpstreamFetch*>(track_.GetIfAvailable()); if (!knew_track_alias) { // If the task already exists (FETCH_OK has arrived), the callback will // immediately execute to read the first object. Otherwise, it will only - // execute when the task is created or a cached object is read. - fetch->OnStreamOpened([this]() { MaybeReadOneObject(); }); + // execute when the task is created or a cached object is read. Note that + // if the entire stream is delivered and closed synchronously during this + // call, fetch_task_ may already be null here. + auto track = track_.GetIfAvailable(); + if (track == nullptr) { + QUICHE_BUG(quic_bug_fetch_stream_destroyed_before_track_alias_read) + << "Fetch stream destroyed before track alias was read"; + return; + } + track->OnStreamOpened(this); return; } MaybeReadOneObject();
diff --git a/quiche/quic/moqt/moqt_uni_stream.h b/quiche/quic/moqt/moqt_uni_stream.h index e1234fb..28a3e4f 100644 --- a/quiche/quic/moqt/moqt_uni_stream.h +++ b/quiche/quic/moqt/moqt_uni_stream.h
@@ -12,6 +12,7 @@ #include <utility> #include "absl/base/nullability.h" +#include "absl/status/status.h" #include "absl/strings/string_view.h" #include "quiche/quic/core/quic_alarm.h" #include "quiche/quic/core/quic_alarm_factory.h" @@ -25,12 +26,13 @@ #include "quiche/quic/moqt/moqt_parser.h" #include "quiche/quic/moqt/moqt_priority.h" #include "quiche/quic/moqt/moqt_publisher.h" -#include "quiche/quic/moqt/moqt_session_interface.h" +#include "quiche/quic/moqt/moqt_session_callbacks.h" #include "quiche/quic/moqt/moqt_trace_recorder.h" #include "quiche/quic/moqt/moqt_types.h" #include "quiche/common/platform/api/quiche_export.h" #include "quiche/common/quiche_callbacks.h" #include "quiche/common/quiche_weak_ptr.h" +#include "quiche/web_transport/stream_helpers.h" #include "quiche/web_transport/web_transport.h" namespace moqt { @@ -173,7 +175,7 @@ std::unique_ptr<quic::QuicAlarm> delivery_timeout_alarm_; }; -using FetchStreamCloseCallback = quiche::SingleUseCallback<void()>; +using FetchStreamCloseCallback = quiche::SingleUseCallback<void(absl::Status)>; class QUICHE_EXPORT OutgoingFetchStream : public OutgoingUniStream { public: @@ -190,7 +192,19 @@ void OnCanWrite() override; void OnStopSendingReceived(webtransport::StreamErrorCode error_code) override; + void Init(); + + void OnBidiStreamReset(webtransport::StreamErrorCode error_code) { + status_ = MoqtStreamErrorToStatus(error_code, "reset"); + close_callback_ = nullptr; // No need to inform the bidi stream. + stream().ResetWithUserCode(error_code); + Detach(); + } + private: + void Detach(); + + absl::Status status_ = absl::OkStatus(); std::unique_ptr<MoqtFetchTask> incoming_objects_; FetchStreamCloseCallback close_callback_; }; @@ -226,33 +240,39 @@ // webtransport::StreamVisitor implementation. void OnCanRead() override; void OnCanWrite() override {} - void OnResetStreamReceived(webtransport::StreamErrorCode) override {} - void OnStopSendingReceived(webtransport::StreamErrorCode /*error*/) override { + void OnResetStreamReceived(webtransport::StreamErrorCode error) override { + status_ = MoqtStreamErrorToStatus(error, "reset"); + Detach(); } + void OnStopSendingReceived(webtransport::StreamErrorCode) override {} void OnWriteSideInDataRecvdState() override {} // MoqtParserVisitor implementation. - // TODO: Handle a stream FIN. void OnObjectMessage(const MoqtObject& message, absl::string_view payload, bool end_of_message) override; - void OnFin() override { fin_received_ = true; } + void OnFin() override { Detach(); } void OnParsingError(MoqtError error_code, absl::string_view reason) override; webtransport::Stream* stream() const { return stream_; } void MaybeReadOneObject(); + virtual void set_fetch_task(UpstreamFetchTask* fetch_task); + private: friend class test::MoqtSessionPeer; bool IsFetch() const { return parser_.stream_type().has_value() && parser_.stream_type()->IsFetch(); } + // Notifies ObjectSubscriber and UpstreamFetchTask that the stream is being + // destroyed. + void Detach(); uint64_t next_object_id_ = 0; bool no_more_objects_ = false; // EndOfGroup or EndOfTrack was received. std::optional<DataStreamIndex> index_; // Only set for subscribe. - bool fin_received_ = false; + absl::Status status_ = absl::OkStatus(); webtransport::Stream* stream_; SubscribeVisitor* visitor_ = nullptr; // Once the subscribe ID is identified, set it here. @@ -260,6 +280,7 @@ MoqtDataParser parser_; std::string partial_object_; uint64_t bytes_received_this_object_ = 0; + UpstreamFetchTask* fetch_task_ = nullptr; // FETCH only. SessionToUniStreamInterface* session_; const quic::QuicClock* absl_nonnull clock_; };
diff --git a/quiche/quic/moqt/moqt_uni_stream_test.cc b/quiche/quic/moqt/moqt_uni_stream_test.cc index 4265330..38a938b 100644 --- a/quiche/quic/moqt/moqt_uni_stream_test.cc +++ b/quiche/quic/moqt/moqt_uni_stream_test.cc
@@ -15,6 +15,7 @@ #include "absl/types/span.h" #include "quiche/quic/core/quic_alarm_factory.h" #include "quiche/quic/core/quic_time.h" +#include "quiche/quic/core/quic_types.h" #include "quiche/quic/moqt/moqt_error.h" #include "quiche/quic/moqt/moqt_fetch_task.h" #include "quiche/quic/moqt/moqt_framer.h" @@ -339,8 +340,10 @@ EXPECT_CALL(mock_stream_, SetPriority); stream_ = std::make_unique<OutgoingFetchStream>( framer_, &mock_stream_, 10, webtransport::StreamPriority(), - std::move(task_), [this]() { close_callback_called_ = true; }, + std::move(task_), + [this](absl::Status) { close_callback_called_ = true; }, &trace_recorder_); + stream_->Init(); } ~OutgoingFetchStreamTest() override { stream_.reset(); @@ -433,9 +436,7 @@ .WillOnce(Return(MoqtFetchTask::kError)); EXPECT_CALL(*task_ptr_, GetStatus()) .WillOnce(Return(absl::InternalError("error"))); - EXPECT_CALL( - mock_stream_, - ResetWithUserCode(static_cast<uint64_t>(absl::StatusCode::kInternal))); + EXPECT_CALL(mock_stream_, ResetWithUserCode(kResetCodeInternalError)); stream_->OnCanWrite(); } @@ -467,18 +468,21 @@ 0, // payload_length }; -class MockSessionToUniStreamInterface : public SessionToUniStreamInterface { +class MockObjectSubscriber : public ObjectSubscriber { public: - MockSessionToUniStreamInterface() = default; - ~MockSessionToUniStreamInterface() override = default; + MockObjectSubscriber() + : ObjectSubscriber(FullTrackName(), 0, MessageParameters(), nullptr) { + ON_CALL(*this, is_fetch).WillByDefault(Return(true)); + ON_CALL(*this, InWindow).WillByDefault(Return(true)); + } - MOCK_METHOD(bool, deliver_partial_objects, (), (const, override)); - MOCK_METHOD(void, OnMalformedTrack, (ObjectSubscriber*), (override)); - MOCK_METHOD(quiche::QuicheWeakPtr<ObjectSubscriber>, GetSubscribe, (uint64_t), + MOCK_METHOD(void, OnStreamOpened, (webtransport::StreamVisitor * stream), (override)); - MOCK_METHOD(quiche::QuicheWeakPtr<ObjectSubscriber>, GetFetch, (uint64_t), + MOCK_METHOD(void, OnStreamClosed, + (absl::Status status, std::optional<DataStreamIndex> index), (override)); - MOCK_METHOD(void, Error, (MoqtError, absl::string_view), (override)); + MOCK_METHOD(bool, InWindow, (Location sequence), (const, override)); + MOCK_METHOD(bool, is_fetch, (), (const, override)); }; class IncomingDataStreamTest : public quic::test::QuicTest { @@ -515,6 +519,8 @@ webtransport::test::InMemoryStream mock_stream_; testing::NiceMock<MockSessionToUniStreamInterface> session_; + MockObjectSubscriber mock_control_stream_; + MockUpstreamFetchTask mock_fetch_task_; quic::MockClock mock_clock_; FullTrackName ftn_; MoqtSubscribe subscribe_message_; @@ -535,7 +541,7 @@ EXPECT_CALL(visitor_, OnObjectFragment); stream_->OnObjectMessage(kDefaultObject, "", true); EXPECT_CALL(visitor_, OnStreamReset); - stream_.reset(); + stream_->OnResetStreamReceived(kResetCodeCancelled); } TEST_F(IncomingDataStreamTest, DestructorAfterFin) { @@ -543,9 +549,8 @@ ProcessAlias(2); EXPECT_CALL(visitor_, OnObjectFragment); stream_->OnObjectMessage(kDefaultObject, "", true); - stream_->OnFin(); EXPECT_CALL(visitor_, OnStreamFin); - stream_.reset(); + stream_->OnFin(); } TEST_F(IncomingDataStreamTest, OnParsingError) { @@ -640,58 +645,70 @@ TEST_F(IncomingDataStreamTest, PartialObjectFetch) { EXPECT_CALL(session_, deliver_partial_objects()).WillRepeatedly(Return(true)); - MoqtFetch fetch; - fetch.request_id = 3; - StandaloneFetch standalone(ftn_, Location(0, 0), Location(0, 9)); - int objects_available_callbacks = 0; - std::unique_ptr<MoqtFetchTask> fetch_task; - auto upstream_fetch = std::make_unique<UpstreamFetch>( - fetch, standalone, - [&](std::unique_ptr<MoqtFetchTask> t) { fetch_task = std::move(t); }, - []() {}); - upstream_fetch->OnFetchResult(Location(0, 9), absl::OkStatus(), []() {}); - UpstreamFetch::UpstreamFetchTask* task = upstream_fetch->task(); - task->SetObjectAvailableCallback([&]() { ++objects_available_callbacks; }); uint8_t stream_header[] = {0x05, 0x03}; mock_stream_.Receive( absl::string_view(reinterpret_cast<const char*>(stream_header), 2), false); EXPECT_CALL(session_, GetFetch(3)) - .WillOnce(Return(upstream_fetch->weak_ptr())); + .WillOnce(Return(mock_control_stream_.weak_ptr())); + EXPECT_CALL(mock_control_stream_, OnStreamOpened(stream_.get())) + .WillOnce([&](webtransport::StreamVisitor* visitor) { + EXPECT_EQ(visitor, stream_.get()); + stream_->set_fetch_task(&mock_fetch_task_); + }); + TaskDestroyedCallback task_destroyed_callback; + CanReadCallback can_read_callback; + EXPECT_CALL(mock_fetch_task_, set_task_destroyed_callback) + .WillOnce([&](TaskDestroyedCallback callback) { + task_destroyed_callback = std::move(callback); + }); + EXPECT_CALL(mock_fetch_task_, set_can_read_callback) + .WillOnce([&](CanReadCallback callback) { + can_read_callback = std::move(callback); + }); stream_->OnCanRead(); - MoqtObject sent_object = MoqtObject( + const MoqtObject sent_object = MoqtObject( /*request_id=*/0, /*group_id=*/0, /*object_id=*/0, /*publisher_priority=*/0x80, /*extension_headers=*/"", MoqtObjectStatus::kNormal, /*subgroup_id=*/0, /*first_object_in_subgroup=*/true, /*payload_length=*/12); + EXPECT_CALL(mock_fetch_task_, HasObject).WillOnce(Return(false)); + EXPECT_CALL(mock_fetch_task_, NewObject) + .WillOnce([&](const MoqtObject& message) { + EXPECT_EQ(message.group_id, sent_object.group_id); + EXPECT_EQ(message.object_id, sent_object.object_id); + EXPECT_EQ(message.publisher_priority, sent_object.publisher_priority); + EXPECT_EQ(message.extension_headers, sent_object.extension_headers); + EXPECT_EQ(message.object_status, sent_object.object_status); + EXPECT_EQ(message.subgroup_id, sent_object.subgroup_id); + EXPECT_EQ(message.first_object_in_subgroup, + sent_object.first_object_in_subgroup); + EXPECT_EQ(message.payload_length, sent_object.payload_length); + }); + EXPECT_CALL(mock_fetch_task_, NeedsMorePayload).WillOnce(Return(true)); + EXPECT_CALL(mock_fetch_task_, AppendPayloadToObject("foo")); stream_->OnObjectMessage(sent_object, "foo", false); - task->NotifyNewObject(); - EXPECT_EQ(objects_available_callbacks, 1); - PublishedObject received_object; - EXPECT_EQ(task->GetNextObject(received_object), - MoqtFetchTask::GetNextObjectResult::kSuccess); - EXPECT_EQ(task->GetNextObject(received_object), - MoqtFetchTask::GetNextObjectResult::kPending); - EXPECT_EQ(sent_object.object_id, received_object.metadata.location.object); - EXPECT_EQ("foo", received_object.payload[0].AsStringView()); + // Second and third fragments. + EXPECT_CALL(mock_fetch_task_, HasObject).WillOnce(Return(true)); + EXPECT_CALL(mock_fetch_task_, NeedsMorePayload).WillOnce(Return(true)); + EXPECT_CALL(mock_fetch_task_, AppendPayloadToObject("bar")); stream_->OnObjectMessage(sent_object, "bar", false); - task->NotifyNewObject(); - EXPECT_EQ(objects_available_callbacks, 2); + EXPECT_CALL(mock_fetch_task_, HasObject).WillOnce(Return(true)); + EXPECT_CALL(mock_fetch_task_, NeedsMorePayload).WillOnce(Return(true)); + EXPECT_CALL(mock_fetch_task_, AppendPayloadToObject("baz")); stream_->OnObjectMessage(sent_object, "baz", false); - task->NotifyNewObject(); - EXPECT_EQ(objects_available_callbacks, 2); - received_object.payload.clear(); - EXPECT_EQ(task->GetNextObject(received_object), - MoqtFetchTask::GetNextObjectResult::kSuccess); - EXPECT_EQ(task->GetNextObject(received_object), - MoqtFetchTask::GetNextObjectResult::kPending); - EXPECT_EQ(sent_object.object_id, received_object.metadata.location.object); - ASSERT_EQ(received_object.payload.size(), 2); - EXPECT_EQ("bar", received_object.payload[0].AsStringView()); - EXPECT_EQ("baz", received_object.payload[1].AsStringView()); + + // Cleanup + EXPECT_CALL(mock_fetch_task_, + OnStreamAndFetchClosed(absl::CancelledError("stream destroyed"))); + EXPECT_CALL(mock_fetch_task_, set_task_destroyed_callback(nullptr)); + EXPECT_CALL(mock_fetch_task_, set_can_read_callback(nullptr)); + EXPECT_CALL(mock_control_stream_, + OnStreamClosed(absl::CancelledError("stream destroyed"), + std::optional<DataStreamIndex>())); } TEST_F(IncomingDataStreamTest, OnObjectMessageInvalidTrack) { @@ -769,14 +786,9 @@ } TEST_F(IncomingDataStreamTest, OnCanReadFetchNewTrackAliasSuccess) { - MoqtFetch fetch; - fetch.request_id = 3; - StandaloneFetch standalone(ftn_, Location(0, 0), Location(0, 9)); - auto upstream_fetch = std::make_unique<UpstreamFetch>( - fetch, standalone, [](std::unique_ptr<MoqtFetchTask>) {}, []() {}); - upstream_fetch->OnFetchResult(Location(0, 0), absl::OkStatus(), []() {}); EXPECT_CALL(session_, GetFetch(3)) - .WillOnce(Return(upstream_fetch->weak_ptr())); + .WillOnce(Return(mock_control_stream_.weak_ptr())); + EXPECT_CALL(mock_control_stream_, OnStreamOpened(stream_.get())); char fetch_bytes[] = {0x05, 0x03}; mock_stream_.Receive(absl::string_view(fetch_bytes, 2), false); stream_->OnCanRead();
diff --git a/quiche/quic/moqt/test_tools/mock_moqt_session.h b/quiche/quic/moqt/test_tools/mock_moqt_session.h index de07a25..6d5bc15 100644 --- a/quiche/quic/moqt/test_tools/mock_moqt_session.h +++ b/quiche/quic/moqt/test_tools/mock_moqt_session.h
@@ -23,7 +23,6 @@ #include "quiche/quic/moqt/moqt_framer.h" #include "quiche/quic/moqt/moqt_key_value_pair.h" #include "quiche/quic/moqt/moqt_live_publisher.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" @@ -33,6 +32,7 @@ #include "quiche/quic/moqt/moqt_trace_recorder.h" #include "quiche/quic/moqt/moqt_types.h" #include "quiche/common/platform/api/quiche_test.h" +#include "quiche/common/quiche_callbacks.h" #include "quiche/common/quiche_mem_slice.h" #include "quiche/common/quiche_weak_ptr.h" #include "quiche/web_transport/test_tools/mock_web_transport.h" @@ -47,7 +47,8 @@ ~MockSessionToPublisherInterface() override = default; MOCK_METHOD(bool, alternate_delivery_timeout, (), (const, override)); MOCK_METHOD(void, UpdateTrackPriority, - (uint64_t, std::optional<MoqtTrackPriority>, MoqtTrackPriority), + (const FullTrackName&, std::optional<MoqtTrackPriority>, + MoqtTrackPriority), (override)); MOCK_METHOD(quic::QuicAlarmFactory*, alarm_factory, (), (override)); MOCK_METHOD(std::shared_ptr<MoqtTrackPublisher>, GetTrackPublisher, @@ -85,20 +86,21 @@ const TrackExtensions& extensions, MoqtResponseCallback response_callback), (override)); - MOCK_METHOD(bool, Fetch, + MOCK_METHOD(std::unique_ptr<MoqtFetchTask>, Fetch, (const FullTrackName& name, FetchResponseCallback callback, Location start, uint64_t end_group, std::optional<uint64_t> end_object, - MessageParameters parameters), + const MessageParameters& parameters), (override)); MOCK_METHOD(bool, RelativeJoiningFetch, (const FullTrackName& name, SubscribeVisitor* visitor, - uint64_t num_previous_groups, MessageParameters parameters), + uint64_t num_previous_groups, + const MessageParameters& parameters), (override)); - MOCK_METHOD(bool, RelativeJoiningFetch, + MOCK_METHOD(std::unique_ptr<MoqtFetchTask>, RelativeJoiningFetch, (const FullTrackName& name, SubscribeVisitor* visitor, FetchResponseCallback callback, uint64_t num_previous_groups, - MessageParameters parameters), + const MessageParameters& parameters), (override)); MOCK_METHOD(bool, PublishNamespace, (const TrackNamespace& track_namespace,
diff --git a/quiche/quic/moqt/test_tools/moqt_framer_utils.cc b/quiche/quic/moqt/test_tools/moqt_framer_utils.cc index f134c23..146fa38 100644 --- a/quiche/quic/moqt/test_tools/moqt_framer_utils.cc +++ b/quiche/quic/moqt/test_tools/moqt_framer_utils.cc
@@ -67,9 +67,6 @@ quiche::QuicheBuffer operator()(const MoqtFetch& message) { return framer.SerializeFetch(message); } - quiche::QuicheBuffer operator()(const MoqtFetchCancel& message) { - return framer.SerializeFetchCancel(message); - } quiche::QuicheBuffer operator()(const MoqtFetchOk& message) { return framer.SerializeFetchOk(message); }
diff --git a/quiche/quic/moqt/test_tools/moqt_framer_utils.h b/quiche/quic/moqt/test_tools/moqt_framer_utils.h index ac05e16..7fd5a9d 100644 --- a/quiche/quic/moqt/test_tools/moqt_framer_utils.h +++ b/quiche/quic/moqt/test_tools/moqt_framer_utils.h
@@ -20,13 +20,11 @@ namespace moqt::test { -using AnyMoqtControlMessage = - std::variant<MoqtSetup, MoqtRequestOk, MoqtRequestError, MoqtSubscribe, - MoqtSubscribeOk, MoqtPublishDone, MoqtRequestUpdate, - MoqtPublishNamespace, MoqtTrackStatus, MoqtGoAway, - MoqtSubscribeNamespace, MoqtSubscribeTracks, MoqtFetch, - MoqtFetchCancel, MoqtFetchOk, MoqtPublish, MoqtNamespace, - MoqtNamespaceDone, MoqtObjectAck>; +using AnyMoqtControlMessage = std::variant< + MoqtSetup, MoqtRequestOk, MoqtRequestError, MoqtSubscribe, MoqtSubscribeOk, + MoqtPublishDone, MoqtRequestUpdate, MoqtPublishNamespace, MoqtTrackStatus, + MoqtGoAway, MoqtSubscribeNamespace, MoqtSubscribeTracks, MoqtFetch, + MoqtFetchOk, MoqtPublish, MoqtNamespace, MoqtNamespaceDone, MoqtObjectAck>; std::string SerializeGenericMessage(const AnyMoqtControlMessage& frame, bool use_webtrans = false);
diff --git a/quiche/quic/moqt/test_tools/moqt_mock_visitor.h b/quiche/quic/moqt/test_tools/moqt_mock_visitor.h index 92a9e4c..238235c 100644 --- a/quiche/quic/moqt/test_tools/moqt_mock_visitor.h +++ b/quiche/quic/moqt/test_tools/moqt_mock_visitor.h
@@ -21,15 +21,16 @@ #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_live_publisher.h" #include "quiche/quic/moqt/moqt_messages.h" #include "quiche/quic/moqt/moqt_names.h" #include "quiche/quic/moqt/moqt_object.h" +#include "quiche/quic/moqt/moqt_object_subscriber.h" #include "quiche/quic/moqt/moqt_priority.h" #include "quiche/quic/moqt/moqt_publisher.h" -#include "quiche/quic/moqt/moqt_session.h" #include "quiche/quic/moqt/moqt_session_callbacks.h" -#include "quiche/quic/moqt/moqt_session_interface.h" #include "quiche/quic/moqt/moqt_types.h" +#include "quiche/quic/moqt/moqt_uni_stream.h" #include "quiche/common/platform/api/quiche_test.h" #include "quiche/common/quiche_mem_slice.h" #include "quiche/common/quiche_weak_ptr.h" @@ -91,11 +92,12 @@ MOCK_METHOD(std::optional<quic::QuicTimeDelta>, expiration, (), (const, override)); MOCK_METHOD(std::unique_ptr<MoqtFetchTask>, StandaloneFetch, - (Location, Location, MoqtDeliveryOrder), (override)); + (Location, Location, MoqtDeliveryOrder, FetchResponseCallback), + (override)); MOCK_METHOD(std::unique_ptr<MoqtFetchTask>, RelativeFetch, - (uint64_t, MoqtDeliveryOrder), (override)); + (uint64_t, MoqtDeliveryOrder, FetchResponseCallback), (override)); MOCK_METHOD(std::unique_ptr<MoqtFetchTask>, AbsoluteFetch, - (uint64_t, MoqtDeliveryOrder), (override)); + (uint64_t, MoqtDeliveryOrder, FetchResponseCallback), (override)); private: FullTrackName track_name_; @@ -134,19 +136,31 @@ } // TODO(martinduke): Support Fetch std::unique_ptr<MoqtFetchTask> StandaloneFetch( - Location start, Location end, MoqtDeliveryOrder delivery_order) override { - return std::make_unique<MoqtFailedFetch>( - absl::UnimplementedError("Fetch not implemented")); + Location start, Location end, MoqtDeliveryOrder delivery_order, + FetchResponseCallback callback) override { + std::move(callback)(MoqtRequestErrorInfo{ + .error_code = RequestErrorCode::kDoesNotExist, + .reason_phrase = "Fetch not implemented", + }); + return nullptr; } std::unique_ptr<MoqtFetchTask> RelativeFetch( - uint64_t offset, MoqtDeliveryOrder delivery_order) override { - return std::make_unique<MoqtFailedFetch>( - absl::UnimplementedError("Fetch not implemented")); + uint64_t offset, MoqtDeliveryOrder delivery_order, + FetchResponseCallback callback) override { + std::move(callback)(MoqtRequestErrorInfo{ + .error_code = RequestErrorCode::kDoesNotExist, + .reason_phrase = "Fetch not implemented", + }); + return nullptr; } std::unique_ptr<MoqtFetchTask> AbsoluteFetch( - uint64_t offset, MoqtDeliveryOrder delivery_order) override { - return std::make_unique<MoqtFailedFetch>( - absl::UnimplementedError("Fetch not implemented")); + uint64_t offset, MoqtDeliveryOrder delivery_order, + FetchResponseCallback callback) override { + std::move(callback)(MoqtRequestErrorInfo{ + .error_code = RequestErrorCode::kDoesNotExist, + .reason_phrase = "Fetch not implemented", + }); + return nullptr; } void AddObject(Location location, uint64_t subgroup, absl::string_view payload, bool fin, @@ -226,15 +240,8 @@ class MockFetchTask : public MoqtFetchTask { public: MockFetchTask() {}; // No synchronous callbacks. - MockFetchTask(std::optional<MoqtFetchOk> fetch_ok, - std::optional<MoqtRequestError> fetch_error, - bool synchronous_object_available) - : synchronous_fetch_ok_(fetch_ok), - synchronous_fetch_error_(fetch_error), - synchronous_object_available_(synchronous_object_available) { - QUICHE_DCHECK(!synchronous_fetch_ok_.has_value() || - !synchronous_fetch_error_.has_value()); - } + explicit MockFetchTask(bool synchronous_object_available) + : synchronous_object_available_(synchronous_object_available) {} MOCK_METHOD(MoqtFetchTask::GetNextObjectResult, GetNextObject, (PublishedObject & output), (override)); @@ -246,37 +253,60 @@ // The first call is installed by the session to trigger stream creation. // An object might not exist yet. objects_available_callback_(); + // This class could be destroyed by the line above. + return; } // The second call is a result of the stream replacing the callback, which // means there is an object available. synchronous_object_available_ = true; } - void SetFetchResponseCallback(FetchResponseCallback callback) override { - if (synchronous_fetch_ok_.has_value()) { - std::move(callback)(*synchronous_fetch_ok_); - return; - } - if (synchronous_fetch_error_.has_value()) { - std::move(callback)(*synchronous_fetch_error_); - return; - } - fetch_response_callback_ = std::move(callback); - } void CallObjectsAvailableCallback() { objects_available_callback_(); }; - void CallFetchResponseCallback( - std::variant<MoqtFetchOk, MoqtRequestError> response) { - std::move(fetch_response_callback_)(response); - } private: - FetchResponseCallback fetch_response_callback_; ObjectsAvailableCallback objects_available_callback_; - std::optional<MoqtFetchOk> synchronous_fetch_ok_; - std::optional<MoqtRequestError> synchronous_fetch_error_; bool synchronous_object_available_ = false; }; +class MockUpstreamFetchTask : public UpstreamFetchTask { + public: + MockUpstreamFetchTask() { + ON_CALL(*this, HasObject).WillByDefault([this]() { + return UpstreamFetchTask::HasObject(); + }); + ON_CALL(*this, NeedsMorePayload).WillByDefault([this]() { + return UpstreamFetchTask::NeedsMorePayload(); + }); + } + ~MockUpstreamFetchTask() override = default; + + MOCK_METHOD(void, set_can_read_callback, (CanReadCallback callback), + (override)); + MOCK_METHOD(void, set_task_destroyed_callback, + (TaskDestroyedCallback callback), (override)); + MOCK_METHOD(void, NewObject, (const MoqtObject& message), (override)); + MOCK_METHOD(void, AppendPayloadToObject, (absl::string_view payload), + (override)); + MOCK_METHOD(bool, HasObject, (), (const, override)); + MOCK_METHOD(bool, NeedsMorePayload, (), (const, override)); + MOCK_METHOD(void, NotifyNewObject, (), (override)); + MOCK_METHOD(void, OnStreamAndFetchClosed, (absl::Status status), (override)); +}; + +class MockSessionToUniStreamInterface : public SessionToUniStreamInterface { + public: + MockSessionToUniStreamInterface() = default; + ~MockSessionToUniStreamInterface() override = default; + + MOCK_METHOD(bool, deliver_partial_objects, (), (const, override)); + MOCK_METHOD(void, OnMalformedTrack, (ObjectSubscriber*), (override)); + MOCK_METHOD(quiche::QuicheWeakPtr<ObjectSubscriber>, GetSubscribe, (uint64_t), + (override)); + MOCK_METHOD(quiche::QuicheWeakPtr<ObjectSubscriber>, GetFetch, (uint64_t), + (override)); + MOCK_METHOD(void, Error, (MoqtError, absl::string_view), (override)); +}; + class MockNamespaceTask : public MoqtNamespaceTask { public: explicit MockNamespaceTask(const TrackNamespace& prefix)
diff --git a/quiche/quic/moqt/test_tools/moqt_session_peer.h b/quiche/quic/moqt/test_tools/moqt_session_peer.h index 024a6cb..8f07ab2 100644 --- a/quiche/quic/moqt/test_tools/moqt_session_peer.h +++ b/quiche/quic/moqt/test_tools/moqt_session_peer.h
@@ -10,11 +10,10 @@ #include <optional> #include <string> #include <utility> +#include <variant> -#include "absl/base/casts.h" #include "absl/base/nullability.h" #include "absl/container/flat_hash_set.h" -#include "absl/memory/memory.h" #include "absl/status/status.h" #include "absl/strings/string_view.h" #include "quiche/quic/core/quic_alarm.h" @@ -22,16 +21,12 @@ #include "quiche/quic/core/quic_time.h" #include "quiche/quic/moqt/moqt_bidi_stream.h" #include "quiche/quic/moqt/moqt_error.h" -#include "quiche/quic/moqt/moqt_fetch_task.h" -#include "quiche/quic/moqt/moqt_key_value_pair.h" #include "quiche/quic/moqt/moqt_live_publisher.h" #include "quiche/quic/moqt/moqt_messages.h" -#include "quiche/quic/moqt/moqt_names.h" #include "quiche/quic/moqt/moqt_object_subscriber.h" #include "quiche/quic/moqt/moqt_parser.h" #include "quiche/quic/moqt/moqt_session.h" #include "quiche/quic/moqt/moqt_session_interface.h" -#include "quiche/quic/moqt/moqt_types.h" #include "quiche/quic/moqt/moqt_uni_stream.h" #include "quiche/quic/moqt/test_tools/moqt_framer_utils.h" #include "quiche/common/platform/api/quiche_logging.h" @@ -165,15 +160,6 @@ session->peer_setup_received_ = value; } - static MoqtSession::PublishedFetch* GetFetch(MoqtSession* session, - uint64_t fetch_id) { - auto it = session->incoming_fetches_.find(fetch_id); - if (it == session->incoming_fetches_.end()) { - return nullptr; - } - return it->second.get(); - } - static void ValidateRequestId(MoqtSession* session, uint64_t id) { session->ValidateRequestId(id); } @@ -224,11 +210,18 @@ return session->parameters_; } - static std::optional<uint64_t> NextQueuedRequestIdToServer( + static std::optional<uint64_t> NextQueuedStreamIdToServer( MoqtSession* session) { - return session->subscriptions_with_queued_streams_.empty() - ? std::optional<uint64_t>() - : session->subscriptions_with_queued_streams_.begin()->second; + if (session->requests_with_queued_streams_.empty()) { + return std::nullopt; + } + for (const auto& [priority, target] : + session->requests_with_queued_streams_) { + if (std::holds_alternative<webtransport::StreamId>(target)) { + return std::get<webtransport::StreamId>(target); + } + } + return std::nullopt; } static uint64_t GetLastTrackAlias(MoqtSession* session) {
diff --git a/quiche/quic/moqt/test_tools/moqt_test_message.h b/quiche/quic/moqt/test_tools/moqt_test_message.h index 45106fe..79c6c30 100644 --- a/quiche/quic/moqt/test_tools/moqt_test_message.h +++ b/quiche/quic/moqt/test_tools/moqt_test_message.h
@@ -117,8 +117,8 @@ MoqtSubscribe, MoqtSubscribeOk, MoqtPublishDone, MoqtRequestUpdate, MoqtPublishNamespace, MoqtTrackStatus, MoqtGoAway, MoqtSubscribeNamespace, MoqtSubscribeTracks, - MoqtFetch, MoqtFetchCancel, MoqtFetchOk, MoqtPublish, - MoqtNamespace, MoqtNamespaceDone, MoqtObjectAck>; + MoqtFetch, MoqtFetchOk, MoqtPublish, MoqtNamespace, + MoqtNamespaceDone, MoqtObjectAck>; // The total actual size of the message. size_t total_message_size() const { return wire_image_size_; } @@ -822,10 +822,6 @@ bool EqualFieldValues(const MessageStructuredData& values) const override { auto cast = std::get<MoqtRequestError>(values); - if (cast.request_id != request_error_.request_id) { - QUIC_LOG(INFO) << "REQUEST_ERROR request_id mismatch"; - return false; - } if (cast.error_code != request_error_.error_code) { QUIC_LOG(INFO) << "REQUEST_ERROR error code mismatch"; return false; @@ -841,7 +837,7 @@ return true; } - void ExpandVarints() override { ExpandVarintsImpl("vvv---"); } + void ExpandVarints() override { ExpandVarintsImpl("vv---"); } MessageStructuredData structured_data() const override { return TestMessageBase::MessageStructuredData(request_error_); @@ -849,16 +845,14 @@ protected: MoqtRequestError request_error_ = { - /*request_id=*/2, /*error_code=*/RequestErrorCode::kInvalidRange, /*retry_interval=*/quic::QuicTimeDelta::FromSeconds(10), /*reason_phrase=*/"bar", }; private: - uint8_t raw_packet_[11] = { - 0x05, 0x00, 0x08, - 0x02, // request_id = 2 + uint8_t raw_packet_[10] = { + 0x05, 0x00, 0x07, 0x11, // error_code = 17 0xa7, 0x11, // retry_interval = 10000 ms 0x03, 0x62, 0x61, 0x72, // reason_phrase = "bar" @@ -1463,10 +1457,6 @@ } bool EqualFieldValues(const MessageStructuredData& values) const override { auto cast = std::get<MoqtFetchOk>(values); - if (cast.request_id != fetch_ok_.request_id) { - QUIC_LOG(INFO) << "FETCH_OK request_id mismatch"; - return false; - } if (cast.end_of_track != fetch_ok_.end_of_track) { QUIC_LOG(INFO) << "FETCH_OK end_of_track mismatch"; return false; @@ -1486,16 +1476,15 @@ return true; } - void ExpandVarints() override { ExpandVarintsImpl("v-vvvv--vv"); } + void ExpandVarints() override { ExpandVarintsImpl("-vvvv--vv"); } MessageStructuredData structured_data() const override { return TestMessageBase::MessageStructuredData(fetch_ok_); } private: - uint8_t raw_packet_[13] = { - 0x18, 0x00, 0x0a, - 0x01, // request_id = 1 + uint8_t raw_packet_[12] = { + 0x18, 0x00, 0x09, 0x00, // end_of_track = false 0x05, 0x04, // end_location = 5, 3 0x00, // no parameters @@ -1504,7 +1493,6 @@ }; MoqtFetchOk fetch_ok_ = { - /*request_id =*/1, /*end_of_track=*/false, /*end_location=*/Location{5, 3}, MessageParameters(), @@ -1515,39 +1503,6 @@ }; }; -class QUICHE_NO_EXPORT FetchCancelMessage : public TestMessageBase { - public: - FetchCancelMessage() : TestMessageBase() { - SetWireImage(raw_packet_, sizeof(raw_packet_)); - } - bool EqualFieldValues(const MessageStructuredData& values) const override { - auto cast = std::get<MoqtFetchCancel>(values); - if (cast.request_id != fetch_cancel_.request_id) { - QUIC_LOG(INFO) << "FETCH_CANCEL subscribe_id mismatch"; - return false; - } - return true; - } - - void ExpandVarints() override { ExpandVarintsImpl("v"); } - - MessageStructuredData structured_data() const override { - return TestMessageBase::MessageStructuredData(fetch_cancel_); - } - - private: - uint8_t raw_packet_[4] = { - 0x17, - 0x00, - 0x01, - 0x01, // request_id = 1 - }; - - MoqtFetchCancel fetch_cancel_ = { - /*request_id =*/1, - }; -}; - class QUICHE_NO_EXPORT PublishMessage : public TestMessageBase { public: PublishMessage() : TestMessageBase() { @@ -1691,8 +1646,6 @@ return std::make_unique<SubscribeTracksMessage>(); case MoqtMessageType::kFetch: return std::make_unique<FetchMessage>(); - case MoqtMessageType::kFetchCancel: - return std::make_unique<FetchCancelMessage>(); case MoqtMessageType::kFetchOk: return std::make_unique<FetchOkMessage>(); case MoqtMessageType::kPublish: