Move TRACK_STATUS to a seperate stream. PiperOrigin-RevId: 963021196
diff --git a/build/source_list.bzl b/build/source_list.bzl index 066b955..3aa9288 100644 --- a/build/source_list.bzl +++ b/build/source_list.bzl
@@ -1620,6 +1620,7 @@ "quic/moqt/moqt_stream_map.h", "quic/moqt/moqt_subscribe_stream.h", "quic/moqt/moqt_trace_recorder.h", + "quic/moqt/moqt_track_status_stream.h", "quic/moqt/moqt_types.h", "quic/moqt/moqt_uni_stream.h", "quic/moqt/relay_namespace_tree.h", @@ -1658,6 +1659,7 @@ "quic/moqt/moqt_stream_map.cc", "quic/moqt/moqt_subscribe_stream.cc", "quic/moqt/moqt_trace_recorder.cc", + "quic/moqt/moqt_track_status_stream.cc", "quic/moqt/moqt_uni_stream.cc", "quic/moqt/relay_namespace_tree.cc", "quic/moqt/tools/chat_client.cc", @@ -1694,6 +1696,7 @@ "quic/moqt/moqt_session_test.cc", "quic/moqt/moqt_stream_map_test.cc", "quic/moqt/moqt_subscribe_stream_test.cc", + "quic/moqt/moqt_track_status_stream_test.cc", "quic/moqt/moqt_uni_stream_test.cc", "quic/moqt/relay_namespace_tree_test.cc", "quic/moqt/session_namespace_tree_test.cc",
diff --git a/build/source_list.gni b/build/source_list.gni index 87af639..553b0a4 100644 --- a/build/source_list.gni +++ b/build/source_list.gni
@@ -1625,6 +1625,7 @@ "src/quiche/quic/moqt/moqt_stream_map.h", "src/quiche/quic/moqt/moqt_subscribe_stream.h", "src/quiche/quic/moqt/moqt_trace_recorder.h", + "src/quiche/quic/moqt/moqt_track_status_stream.h", "src/quiche/quic/moqt/moqt_types.h", "src/quiche/quic/moqt/moqt_uni_stream.h", "src/quiche/quic/moqt/relay_namespace_tree.h", @@ -1663,6 +1664,7 @@ "src/quiche/quic/moqt/moqt_stream_map.cc", "src/quiche/quic/moqt/moqt_subscribe_stream.cc", "src/quiche/quic/moqt/moqt_trace_recorder.cc", + "src/quiche/quic/moqt/moqt_track_status_stream.cc", "src/quiche/quic/moqt/moqt_uni_stream.cc", "src/quiche/quic/moqt/relay_namespace_tree.cc", "src/quiche/quic/moqt/tools/chat_client.cc", @@ -1700,6 +1702,7 @@ "src/quiche/quic/moqt/moqt_session_test.cc", "src/quiche/quic/moqt/moqt_stream_map_test.cc", "src/quiche/quic/moqt/moqt_subscribe_stream_test.cc", + "src/quiche/quic/moqt/moqt_track_status_stream_test.cc", "src/quiche/quic/moqt/moqt_uni_stream_test.cc", "src/quiche/quic/moqt/relay_namespace_tree_test.cc", "src/quiche/quic/moqt/session_namespace_tree_test.cc",
diff --git a/build/source_list.json b/build/source_list.json index 231fdb0..8e9046a 100644 --- a/build/source_list.json +++ b/build/source_list.json
@@ -1624,6 +1624,7 @@ "quiche/quic/moqt/moqt_stream_map.h", "quiche/quic/moqt/moqt_subscribe_stream.h", "quiche/quic/moqt/moqt_trace_recorder.h", + "quiche/quic/moqt/moqt_track_status_stream.h", "quiche/quic/moqt/moqt_types.h", "quiche/quic/moqt/moqt_uni_stream.h", "quiche/quic/moqt/relay_namespace_tree.h", @@ -1662,6 +1663,7 @@ "quiche/quic/moqt/moqt_stream_map.cc", "quiche/quic/moqt/moqt_subscribe_stream.cc", "quiche/quic/moqt/moqt_trace_recorder.cc", + "quiche/quic/moqt/moqt_track_status_stream.cc", "quiche/quic/moqt/moqt_uni_stream.cc", "quiche/quic/moqt/relay_namespace_tree.cc", "quiche/quic/moqt/tools/chat_client.cc", @@ -1699,6 +1701,7 @@ "quiche/quic/moqt/moqt_session_test.cc", "quiche/quic/moqt/moqt_stream_map_test.cc", "quiche/quic/moqt/moqt_subscribe_stream_test.cc", + "quiche/quic/moqt/moqt_track_status_stream_test.cc", "quiche/quic/moqt/moqt_uni_stream_test.cc", "quiche/quic/moqt/relay_namespace_tree_test.cc", "quiche/quic/moqt/session_namespace_tree_test.cc",
diff --git a/quiche/quic/moqt/moqt_integration_test.cc b/quiche/quic/moqt/moqt_integration_test.cc index 2c7e36b..4cfa3b7 100644 --- a/quiche/quic/moqt/moqt_integration_test.cc +++ b/quiche/quic/moqt/moqt_integration_test.cc
@@ -1223,6 +1223,56 @@ ASSERT_TRUE(success); } +TEST_F(MoqtIntegrationTest, TrackStatusSuccess) { + EstablishSession(); + FullTrackName track_name("test", "data"); + auto queue = std::make_shared<MoqtOutgoingQueue>(track_name); + queue->AddObject(quiche::QuicheMemSlice::Copy("object 1"), /*key=*/true); + queue->AddObject(quiche::QuicheMemSlice::Copy("object 2"), /*key=*/true); + MoqtKnownTrackPublisher known_track_publisher; + known_track_publisher.Add(queue); + server_->session()->set_publisher(&known_track_publisher); + + bool received_response = false; + MessageParameters received_parameters; + client_->session()->TrackStatus( + track_name, MessageParameters(), + [&](std::variant<MessageParameters, MoqtRequestErrorInfo> response) { + received_response = true; + ASSERT_TRUE(std::holds_alternative<MessageParameters>(response)); + received_parameters = std::get<MessageParameters>(response); + }); + + bool success = test_harness_.RunUntilWithDefaultTimeout( + [&]() { return received_response; }); + EXPECT_TRUE(success); + EXPECT_TRUE(received_parameters.largest_object.has_value()); + EXPECT_EQ(received_parameters.largest_object->group, 1u); + EXPECT_EQ(received_parameters.largest_object->object, 0u); +} + +TEST_F(MoqtIntegrationTest, TrackStatusDoesNotExist) { + EstablishSession(); + FullTrackName track_name("test", "nonexistent"); + MoqtKnownTrackPublisher known_track_publisher; + server_->session()->set_publisher(&known_track_publisher); + + bool received_response = false; + MoqtRequestErrorInfo received_error; + client_->session()->TrackStatus( + track_name, MessageParameters(), + [&](std::variant<MessageParameters, MoqtRequestErrorInfo> response) { + received_response = true; + ASSERT_TRUE(std::holds_alternative<MoqtRequestErrorInfo>(response)); + received_error = std::get<MoqtRequestErrorInfo>(response); + }); + + bool success = test_harness_.RunUntilWithDefaultTimeout( + [&]() { return received_response; }); + EXPECT_TRUE(success); + EXPECT_EQ(received_error.error_code, RequestErrorCode::kDoesNotExist); +} + } // namespace } // namespace moqt::test
diff --git a/quiche/quic/moqt/moqt_session.cc b/quiche/quic/moqt/moqt_session.cc index 909d49d..753e5ba 100644 --- a/quiche/quic/moqt/moqt_session.cc +++ b/quiche/quic/moqt/moqt_session.cc
@@ -46,6 +46,7 @@ #include "quiche/quic/moqt/moqt_session_callbacks.h" #include "quiche/quic/moqt/moqt_session_interface.h" #include "quiche/quic/moqt/moqt_subscribe_stream.h" +#include "quiche/quic/moqt/moqt_track_status_stream.h" #include "quiche/quic/moqt/moqt_types.h" #include "quiche/quic/moqt/moqt_uni_stream.h" #include "quiche/quic/platform/api/quic_logging.h" @@ -346,6 +347,41 @@ // Do nothing. } +bool MoqtSession::TrackStatus(const FullTrackName& name, + const MessageParameters& parameters, + MoqtResponseCallback response_callback) { + QUICHE_DCHECK(name.IsValid()); + if (received_goaway_ || sent_goaway_) { + QUIC_DLOG(INFO) << ENDPOINT << "Tried to send TRACK_STATUS after GOAWAY"; + std::move(response_callback)(MoqtRequestErrorInfo{ + RequestErrorCode::kGoingAway, std::nullopt, "GOAWAY received"}); + return false; + } + + webtransport::Stream* stream = session_->OpenOutgoingBidirectionalStream(); + if (stream == nullptr) { + std::move(response_callback)( + MoqtRequestErrorInfo{RequestErrorCode::kInternalError, std::nullopt, + "Flow control blocked"}); + return false; + } + + uint64_t request_id = NextRequestId(); + auto stream_visitor = std::make_unique<MoqtTrackStatusRequestStream>( + &framer_, ControlMessageParser(), request_id, name, parameters, + [session_weak = GetWeakPtr()](MoqtError code, absl::string_view reason) { + MoqtSession* session = MoqtSessionFromWeakPtr(session_weak); + if (session != nullptr) { + session->Error(code, reason); + } + }, + std::move(response_callback)); + MoqtTrackStatusRequestStream* stream_visitor_ptr = stream_visitor.get(); + stream->SetVisitor(std::move(stream_visitor)); + stream_visitor_ptr->BindStream(stream); + return true; +} + bool MoqtSession::PublishNamespace( const TrackNamespace& track_namespace, const MessageParameters& parameters, MoqtResponseCallback response_callback, @@ -876,7 +912,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) || - incoming_track_status_.contains(request_id) || incoming_publish_namespaces_by_id_.contains(request_id)) { QUICHE_DLOG(INFO) << ENDPOINT << "Duplicate request ID"; Error(MoqtError::kInvalidRequestId, "Duplicate request ID"); @@ -1077,6 +1112,24 @@ temp_stream->OnCanRead(); break; } + case MoqtMessageType::kTrackStatus: { + auto track_status_stream = + std::make_unique<MoqtTrackStatusResponseStream>( + &session_->framer_, session_->ControlMessageParser(), + [weakptr = session_->GetWeakPtr()](MoqtError code, + absl::string_view reason) { + MoqtSession* session = MoqtSessionFromWeakPtr(weakptr); + if (session != nullptr) { + session->Error(code, reason); + } + }, + session_->weak_ptr_factory_for_publishers_.Create()); + track_status_stream->BindStream(std::move(parser_)); + MoqtTrackStatusResponseStream* temp_stream = track_status_stream.get(); + stream_->SetVisitor(std::move(track_status_stream)); + temp_stream->OnCanRead(); + break; + } default: session_->Error(MoqtError::kProtocolViolation, "Unexpected message type received to start bidi stream"); @@ -1316,34 +1369,6 @@ return absl::OkStatus(); } -absl::Status MoqtSession::OnControlMessage(const MoqtTrackStatus& message) { - if (!ValidateRequestId(message.request_id)) { - return absl::OkStatus(); - } - if (sent_goaway_) { - QUIC_DLOG(INFO) << ENDPOINT - << "Received a TRACK_STATUS_REQUEST after GOAWAY"; - SendRequestErrorOnControlStream( - message.request_id, RequestErrorCode::kUnauthorized, std::nullopt, - "TRACK_STATUS_REQUEST after GOAWAY"); - return absl::OkStatus(); - } - // TODO(martinduke): Handle authentication. - std::shared_ptr<MoqtTrackPublisher> track = - publisher_->GetTrack(message.full_track_name); - if (track == nullptr) { - SendRequestErrorOnControlStream(message.request_id, - RequestErrorCode::kDoesNotExist, - std::nullopt, "Track does not exist"); - return absl::OkStatus(); - } - auto [it, inserted] = incoming_track_status_.emplace( - message.request_id, std::make_unique<DownstreamTrackStatus>( - message.request_id, this, track.get())); - track->AddObjectListener(it->second.get()); - return absl::OkStatus(); -} - absl::Status MoqtSession::OnControlMessage(const MoqtGoAway& message) { if (!message.new_session_uri.empty() && perspective() == quic::Perspective::IS_SERVER) {
diff --git a/quiche/quic/moqt/moqt_session.h b/quiche/quic/moqt/moqt_session.h index 21ff0a6..7f88f09 100644 --- a/quiche/quic/moqt/moqt_session.h +++ b/quiche/quic/moqt/moqt_session.h
@@ -38,6 +38,7 @@ #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_track_status_stream.h" #include "quiche/quic/moqt/moqt_types.h" #include "quiche/quic/moqt/moqt_uni_stream.h" #include "quiche/quic/moqt/session_namespace_tree.h" @@ -138,6 +139,9 @@ const MessageParameters& parameters, MoqtResponseCallback response_callback) override; void UnsubscribeTracks(TrackNamespace& prefix) override; + bool TrackStatus(const FullTrackName& name, + const MessageParameters& parameters, + MoqtResponseCallback response_callback) override; quiche::QuicheWeakPtr<MoqtSessionInterface> GetWeakPtr() override { return weak_ptr_factory_.Create(); } @@ -319,69 +323,6 @@ std::unique_ptr<MoqtFetchTask> fetch_; }; - class QUICHE_EXPORT DownstreamTrackStatus : public MoqtObjectListener { - public: - DownstreamTrackStatus(uint64_t request_id, - MoqtSession* absl_nonnull session, - MoqtTrackPublisher* absl_nonnull publisher) - : request_id_(request_id), session_(session), publisher_(publisher) {} - ~DownstreamTrackStatus() { - if (publisher_ != nullptr) { - publisher_->RemoveObjectListener(this); - } - } - DownstreamTrackStatus(const DownstreamTrackStatus&) = delete; - DownstreamTrackStatus(DownstreamTrackStatus&&) = delete; - - void OnSubscribeAccepted() override { - if (publisher_ == nullptr) { - QUICHE_NOTREACHED(); - return; - } - MessageParameters parameters; - parameters.expires = publisher_->expiration(); - parameters.largest_object = publisher_->largest_location(); - ControlStream* control_stream = session_->GetControlStream(); - if (control_stream != nullptr) { - control_stream->CheckStatus( - control_stream->SendRequestOk(request_id_, parameters)); - } - session_->incoming_track_status_.erase(request_id_); - // No class access below this line! - } - - void OnSubscribeRejected(MoqtRequestErrorInfo info) override { - ControlStream* control_stream = session_->GetControlStream(); - if (control_stream != nullptr) { - control_stream->CheckStatus(control_stream->SendRequestError( - request_id_, info.error_code, info.retry_interval, - info.reason_phrase)); - } - session_->incoming_track_status_.erase(request_id_); - // No class access below this line! - } - - void OnNewObjectAvailable(Location, std::optional<uint64_t> /*subgroup*/, - MoqtPriority) override {} - void OnNewFinAvailable(Location /*location*/, - uint64_t /*subgroup*/) override {} - void OnSubgroupAbandoned( - uint64_t /*group*/, uint64_t /*subgroup*/, - webtransport::StreamErrorCode /*error_code*/) override {} - void OnGroupAbandoned(uint64_t /*group_id*/) override {} - void OnTrackPublisherGone() override { - publisher_ = nullptr; - OnSubscribeRejected(MoqtRequestErrorInfo(RequestErrorCode::kDoesNotExist, - std::nullopt, - "Track publisher gone")); - } - - private: - uint64_t request_id_; - MoqtSession* session_; - MoqtTrackPublisher* publisher_; - }; - class GoAwayTimeoutDelegate : public quic::QuicAlarm::DelegateWithoutContext { public: explicit GoAwayTimeoutDelegate(MoqtSession* session) : session_(session) {} @@ -448,7 +389,6 @@ absl::Status OnControlMessage(const MoqtPublishNamespace& message); absl::Status OnControlMessage(const MoqtPublishNamespaceDone& /*message*/); absl::Status OnControlMessage(const MoqtPublishNamespaceCancel& message); - absl::Status OnControlMessage(const MoqtTrackStatus& message); absl::Status OnControlMessage(const MoqtGoAway& /*message*/); absl::Status OnControlMessage(const MoqtMaxRequestId& message); absl::Status OnControlMessage(const MoqtFetch& message); @@ -529,9 +469,6 @@ absl::flat_hash_map<uint64_t, std::unique_ptr<PublishedFetch>> incoming_fetches_; - absl::flat_hash_map<uint64_t, std::unique_ptr<DownstreamTrackStatus>> - incoming_track_status_; - // Monitoring interfaces for expected incoming subscriptions. absl::flat_hash_map<FullTrackName, MoqtPublishingMonitorInterface*> monitoring_interfaces_for_published_tracks_; @@ -578,7 +515,7 @@ std::shared_ptr<Empty> liveness_token_; }; -static MoqtSession* absl_nullable MoqtSessionFromWeakPtr( +inline MoqtSession* absl_nullable MoqtSessionFromWeakPtr( const quiche::QuicheWeakPtr<MoqtSessionInterface>& weak_ptr) { return absl::down_cast<MoqtSession*>(weak_ptr.GetIfAvailable()); }
diff --git a/quiche/quic/moqt/moqt_session_interface.h b/quiche/quic/moqt/moqt_session_interface.h index c08cf50..4d0801e 100644 --- a/quiche/quic/moqt/moqt_session_interface.h +++ b/quiche/quic/moqt/moqt_session_interface.h
@@ -186,9 +186,12 @@ // TODO(martinduke): Add an API for absolute joining fetch. - // TODO: Add SubscribeNamespace, UnsubscribeNamespace method. - // TODO: Add PublishNamespaceCancel method. - // TODO: Add TrackStatusRequest method. + // Sends TRACK_STATUS request to the peer. Returns `false` if the request + // immediately fails (usually due to flow control), and `true` otherwise; + // `response_callback` will be eventually invoked in either case. + virtual bool TrackStatus(const FullTrackName& name, + const MessageParameters& parameters, + MoqtResponseCallback response_callback) = 0; // TODO: Add RequestUpdate, PublishDone method. virtual quiche::QuicheWeakPtr<MoqtSessionInterface> GetWeakPtr() = 0; };
diff --git a/quiche/quic/moqt/moqt_session_test.cc b/quiche/quic/moqt/moqt_session_test.cc index 0581f28..436fc68 100644 --- a/quiche/quic/moqt/moqt_session_test.cc +++ b/quiche/quic/moqt/moqt_session_test.cc
@@ -170,6 +170,7 @@ static constexpr absl::string_view kSubscribeByte = "\x03"; static constexpr absl::string_view kSubscribeNamespaceByte = "\x50"; static constexpr absl::string_view kPublishByte = "\x1d"; + static constexpr absl::string_view kTrackStatusByte = "\x0d"; std::unique_ptr<MoqtBidiStreamBase> ResponseStream( absl::string_view first_byte, webtransport::test::MockStream* wt_stream = nullptr) { @@ -2437,8 +2438,8 @@ } TEST_F(MoqtSessionTest, IncomingTrackStatusThenSynchronousOk) { - bidi_wrapper_ = - MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_); + bidi_wrapper_ = std::make_unique<MoqtBidiStreamTestWrapper>( + ResponseStream(kTrackStatusByte)); auto* track = CreateTrackPublisher(); MoqtTrackStatus track_status = DefaultSubscribe(); @@ -2463,8 +2464,8 @@ } TEST_F(MoqtSessionTest, IncomingTrackStatusThenAsynchronousOk) { - bidi_wrapper_ = - MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_); + bidi_wrapper_ = std::make_unique<MoqtBidiStreamTestWrapper>( + ResponseStream(kTrackStatusByte)); auto* track = CreateTrackPublisher(); MoqtTrackStatus track_status = DefaultSubscribe(); @@ -2487,8 +2488,8 @@ } TEST_F(MoqtSessionTest, IncomingTrackStatusThenSynchronousError) { - bidi_wrapper_ = - MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_); + bidi_wrapper_ = std::make_unique<MoqtBidiStreamTestWrapper>( + ResponseStream(kTrackStatusByte)); auto* track = CreateTrackPublisher(); MoqtTrackStatus track_status = DefaultSubscribe(); @@ -2508,8 +2509,8 @@ } TEST_F(MoqtSessionTest, IncomingTrackStatusThenAsynchronousError) { - bidi_wrapper_ = - MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_); + bidi_wrapper_ = std::make_unique<MoqtBidiStreamTestWrapper>( + ResponseStream(kTrackStatusByte)); auto* track = CreateTrackPublisher(); MoqtTrackStatus track_status = DefaultSubscribe();
diff --git a/quiche/quic/moqt/moqt_track_status_stream.cc b/quiche/quic/moqt/moqt_track_status_stream.cc new file mode 100644 index 0000000..282cb27 --- /dev/null +++ b/quiche/quic/moqt/moqt_track_status_stream.cc
@@ -0,0 +1,168 @@ +// Copyright 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_track_status_stream.h" + +#include <cstdint> +#include <memory> +#include <optional> +#include <utility> + +#include "absl/base/nullability.h" +#include "absl/status/status.h" +#include "quiche/quic/moqt/moqt_bidi_stream.h" +#include "quiche/quic/moqt/moqt_error.h" +#include "quiche/quic/moqt/moqt_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_parser.h" +#include "quiche/common/platform/api/quiche_logging.h" +#include "quiche/common/quiche_weak_ptr.h" + +namespace moqt { + +MoqtTrackStatusRequestStream::MoqtTrackStatusRequestStream( + MoqtFramer* absl_nonnull framer, + const MoqtControlMessageParser& message_parser, uint64_t request_id, + const FullTrackName& full_track_name, const MessageParameters& parameters, + SessionErrorCallback session_error_callback, + MoqtResponseCallback response_callback) + : MoqtBidiStreamBase(framer, message_parser, + std::move(session_error_callback)), + request_id_(request_id), + full_track_name_(full_track_name), + parameters_(parameters), + response_callback_(std::move(response_callback)) {} + +void MoqtTrackStatusRequestStream::OnStreamBound() { + stream_parser()->set_allow_fin(true); + MoqtTrackStatus message; + message.request_id = request_id_; + message.full_track_name = full_track_name_; + message.parameters = parameters_; + SendOrBufferMessageOrFatal(framer()->SerializeTrackStatus(message)); +} + +absl::Status MoqtTrackStatusRequestStream::OnRawControlMessage( + const MoqtRawControlMessage& message) { + return ControlMessageDispatcher::DispatchControlMessage( + *this, message_parser(), message, "track status request"); +} + +absl::Status MoqtTrackStatusRequestStream::OnControlMessage( + const MoqtRequestOk& message) { + if (response_callback_ == nullptr) { + return absl::InvalidArgumentError("Duplicate REQUEST_OK"); + } + MoqtResponseCallback callback = std::move(response_callback_); + response_callback_ = nullptr; + Fin(); + // `message.request_id` is ignored, since request IDs in REQUEST_OK are + // deprecated and not present in draft-18. + std::move(callback)(message.parameters); + return absl::OkStatus(); +} + +absl::Status MoqtTrackStatusRequestStream::OnControlMessage( + const MoqtRequestError& message) { + if (response_callback_ == nullptr) { + return absl::InvalidArgumentError("Duplicate REQUEST_ERROR"); + } + MoqtResponseCallback callback = std::move(response_callback_); + response_callback_ = nullptr; + Fin(); + // `message.request_id` is ignored, since request IDs in REQUEST_ERROR are + // deprecated and not present in draft-18. + std::move(callback)(MoqtRequestErrorInfo{ + message.error_code, message.retry_interval, message.reason_phrase}); + return absl::OkStatus(); +} + +void MoqtTrackStatusRequestStream::Detach() { + if (response_callback_ != nullptr) { + MoqtResponseCallback callback = std::move(response_callback_); + response_callback_ = nullptr; + std::move(callback)(MoqtRequestErrorInfo{RequestErrorCode::kInternalError, + std::nullopt, "Stream closed"}); + } +} + +MoqtTrackStatusResponseStream::MoqtTrackStatusResponseStream( + MoqtFramer* absl_nonnull framer, + const MoqtControlMessageParser& message_parser, + SessionErrorCallback session_error_callback, + quiche::QuicheWeakPtr<SessionToPublisherInterface> session) + : MoqtBidiStreamBase(framer, message_parser, + std::move(session_error_callback)), + session_(session) {} + +absl::Status MoqtTrackStatusResponseStream::OnRawControlMessage( + const MoqtRawControlMessage& message) { + return ControlMessageDispatcher::DispatchControlMessage( + *this, message_parser(), message, "track status response"); +} + +absl::Status MoqtTrackStatusResponseStream::OnControlMessage( + const MoqtTrackStatus& message) { + if (request_id_.has_value()) { + return absl::InvalidArgumentError("Duplicate TRACK_STATUS received"); + } + request_id_ = message.request_id; + if (session() == nullptr) { + return absl::InternalError("Session unavailable"); + } + 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); + } + // If the upstream subscription is already established, the code below will + // invoke `OnSubscribeAccepted` immediately. + publisher_->AddObjectListener(this); + return absl::OkStatus(); +} + +void MoqtTrackStatusResponseStream::OnSubscribeAccepted() { + if (publisher_ == nullptr || !request_id_.has_value()) { + QUICHE_NOTREACHED(); + return; + } + MessageParameters parameters; + parameters.expires = publisher_->expiration(); + parameters.largest_object = publisher_->largest_location(); + // Since `fin` is true, this will also reset `publisher_`. + CheckStatus(SendRequestOk(*request_id_, parameters, /*fin=*/true)); +} + +void MoqtTrackStatusResponseStream::OnSubscribeRejected( + MoqtRequestErrorInfo info) { + if (!request_id_.has_value()) { + QUICHE_NOTREACHED(); + 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)); +} + +void MoqtTrackStatusResponseStream::OnTrackPublisherGone() { + publisher_ = nullptr; + OnSubscribeRejected(MoqtRequestErrorInfo(RequestErrorCode::kDoesNotExist, + std::nullopt, + "Track publisher destroyed")); +} + +void MoqtTrackStatusResponseStream::Detach() { + if (publisher_ != nullptr) { + publisher_->RemoveObjectListener(this); + publisher_ = nullptr; + } +} + +} // namespace moqt
diff --git a/quiche/quic/moqt/moqt_track_status_stream.h b/quiche/quic/moqt/moqt_track_status_stream.h new file mode 100644 index 0000000..16f9886 --- /dev/null +++ b/quiche/quic/moqt/moqt_track_status_stream.h
@@ -0,0 +1,101 @@ +// Copyright 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_TRACK_STATUS_STREAM_H_ +#define QUICHE_QUIC_MOQT_MOQT_TRACK_STATUS_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_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_parser.h" +#include "quiche/quic/moqt/moqt_priority.h" +#include "quiche/quic/moqt/moqt_publisher.h" +#include "quiche/quic/moqt/moqt_types.h" +#include "quiche/common/quiche_weak_ptr.h" +#include "quiche/web_transport/web_transport.h" + +namespace moqt { + +// MoqtTrackStatusRequestStream represents an outgoing TRACK_STATUS request. +class MoqtTrackStatusRequestStream : public MoqtBidiStreamBase { + public: + MoqtTrackStatusRequestStream(MoqtFramer* absl_nonnull framer, + const MoqtControlMessageParser& message_parser, + uint64_t request_id, + const FullTrackName& full_track_name, + const MessageParameters& parameters, + SessionErrorCallback session_error_callback, + MoqtResponseCallback response_callback); + ~MoqtTrackStatusRequestStream() { Detach(); } + + // MoqtBidiStreamBase overrides. + void OnStreamBound() override; + absl::Status OnRawControlMessage( + const MoqtRawControlMessage& message) override; + absl::Status OnControlMessage(const MoqtRequestOk& message); + absl::Status OnControlMessage(const MoqtRequestError& message); + + void Detach() override; + + private: + const uint64_t request_id_; + const FullTrackName full_track_name_; + const MessageParameters parameters_; + MoqtResponseCallback response_callback_; +}; + +// MoqtTrackStatusResponseStream represents an incoming TRACK_STATUS request. +class MoqtTrackStatusResponseStream : public MoqtBidiStreamBase, + public MoqtObjectListener { + public: + MoqtTrackStatusResponseStream( + MoqtFramer* absl_nonnull framer, + const MoqtControlMessageParser& message_parser, + SessionErrorCallback session_error_callback, + quiche::QuicheWeakPtr<SessionToPublisherInterface> session); + ~MoqtTrackStatusResponseStream() { Detach(); } + + // MoqtBidiStreamBase overrides. + void OnStreamBound() override { stream_parser()->set_allow_fin(true); } + absl::Status OnRawControlMessage( + const MoqtRawControlMessage& message) override; + absl::Status OnControlMessage(const MoqtTrackStatus& message); + + // MoqtObjectListener overrides. + void OnSubscribeAccepted() override; + void OnSubscribeRejected(MoqtRequestErrorInfo info) override; + void OnNewObjectAvailable(Location, std::optional<uint64_t>, + MoqtPriority) override {} + void OnNewFinAvailable(Location, uint64_t) override {} + void OnSubgroupAbandoned(uint64_t, uint64_t, + webtransport::StreamErrorCode) override {} + void OnGroupAbandoned(uint64_t) override {} + void OnTrackPublisherGone() override; + + void Detach() override; + + private: + SessionToPublisherInterface* absl_nullable session() const { + return session_.GetIfAvailable(); + } + + std::optional<uint64_t> request_id_; + const quiche::QuicheWeakPtr<SessionToPublisherInterface> session_; + std::shared_ptr<MoqtTrackPublisher> publisher_ = nullptr; +}; + +} // namespace moqt + +#endif // QUICHE_QUIC_MOQT_MOQT_TRACK_STATUS_STREAM_H_
diff --git a/quiche/quic/moqt/moqt_track_status_stream_test.cc b/quiche/quic/moqt/moqt_track_status_stream_test.cc new file mode 100644 index 0000000..cfb6b4b --- /dev/null +++ b/quiche/quic/moqt/moqt_track_status_stream_test.cc
@@ -0,0 +1,349 @@ +// Copyright 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_track_status_stream.h" + +#include <cstdint> +#include <memory> +#include <optional> +#include <string> +#include <variant> + +#include "absl/status/status.h" +#include "absl/strings/string_view.h" +#include "quiche/quic/core/quic_time.h" +#include "quiche/quic/core/quic_types.h" +#include "quiche/quic/moqt/moqt_error.h" +#include "quiche/quic/moqt/moqt_fetch_task.h" +#include "quiche/quic/moqt/moqt_framer.h" +#include "quiche/quic/moqt/moqt_key_value_pair.h" +#include "quiche/quic/moqt/moqt_messages.h" +#include "quiche/quic/moqt/moqt_names.h" +#include "quiche/quic/moqt/moqt_parser.h" +#include "quiche/quic/moqt/moqt_publisher.h" +#include "quiche/quic/moqt/moqt_session_interface.h" +#include "quiche/quic/moqt/moqt_types.h" +#include "quiche/quic/moqt/test_tools/mock_moqt_session.h" +#include "quiche/quic/moqt/test_tools/moqt_framer_utils.h" +#include "quiche/quic/moqt/test_tools/moqt_mock_visitor.h" +#include "quiche/common/platform/api/quiche_test.h" +#include "quiche/common/test_tools/quiche_test_utils.h" +#include "quiche/web_transport/test_tools/mock_web_transport.h" +#include "quiche/web_transport/web_transport.h" + +namespace moqt::test { +namespace { + +using ::quiche::test::StatusIs; +using ::testing::_; +using ::testing::IsEmpty; +using ::testing::Property; +using ::testing::Return; +using ::testing::StrictMock; + +class MoqtTrackStatusRequestStreamTest : public quiche::test::QuicheTest { + public: + MoqtTrackStatusRequestStreamTest() + : framer_(/*using_webtrans=*/true, quic::Perspective::IS_CLIENT), + message_parser_(kDefaultMoqtVersion, /*uses_web_transport=*/true, + quic::Perspective::IS_CLIENT), + track_name_("foo", "bar") {} + + MoqtTrackStatusRequestStream CreateStream( + const MessageParameters& parameters = MessageParameters()) { + return MoqtTrackStatusRequestStream(&framer_, message_parser_, kRequestId, + track_name_, parameters, + session_error_callback_.AsStdFunction(), + response_callback_.AsStdFunction()); + } + + protected: + static constexpr uint64_t kRequestId = 2; + + MoqtFramer framer_; + MoqtControlMessageParser message_parser_; + FullTrackName track_name_; + StrictMock<testing::MockFunction<void(MoqtError, absl::string_view)>> + session_error_callback_; + StrictMock<testing::MockFunction<void( + std::variant<MessageParameters, MoqtRequestErrorInfo>)>> + response_callback_; + StrictMock<webtransport::test::MockStream> mock_stream_; +}; + +TEST_F(MoqtTrackStatusRequestStreamTest, SendRequestOnStreamBound) { + MoqtTrackStatusRequestStream stream = CreateStream(); + EXPECT_CALL(mock_stream_, CanWrite).WillRepeatedly(Return(true)); + EXPECT_CALL(mock_stream_, + Writev(ControlMessageOfType(MoqtMessageType::kTrackStatus), _)); + 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::kInternalError); + EXPECT_EQ(info.reason_phrase, "Stream closed"); + }); + stream.BindStream(&mock_stream_); +} + +TEST_F(MoqtTrackStatusRequestStreamTest, SendRequestWithParameters) { + MessageParameters parameters; + parameters.delivery_timeout = quic::QuicTimeDelta::FromSeconds(5); + MoqtTrackStatusRequestStream stream = CreateStream(parameters); + EXPECT_CALL(mock_stream_, CanWrite).WillRepeatedly(Return(true)); + MoqtTrackStatus expected_message; + expected_message.request_id = kRequestId; + expected_message.full_track_name = track_name_; + expected_message.parameters = parameters; + EXPECT_CALL(mock_stream_, + Writev(SerializedControlMessage(expected_message), _)); + 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::kInternalError); + EXPECT_EQ(info.reason_phrase, "Stream closed"); + }); + stream.BindStream(&mock_stream_); +} + +TEST_F(MoqtTrackStatusRequestStreamTest, ReceiveOkResponse) { + MoqtTrackStatusRequestStream stream = CreateStream(); + EXPECT_CALL(mock_stream_, CanWrite).WillRepeatedly(Return(true)); + EXPECT_CALL(mock_stream_, + Writev(ControlMessageOfType(MoqtMessageType::kTrackStatus), _)); + stream.BindStream(&mock_stream_); + + MessageParameters parameters; + parameters.expires = quic::QuicTimeDelta::FromSeconds(10); + parameters.largest_object = Location(1, 2); + + EXPECT_CALL(response_callback_, Call) + .WillOnce([&](std::variant<MessageParameters, MoqtRequestErrorInfo> v) { + ASSERT_TRUE(std::holds_alternative<MessageParameters>(v)); + auto params = std::get<MessageParameters>(v); + EXPECT_EQ(params.expires, parameters.expires); + EXPECT_EQ(params.largest_object, parameters.largest_object); + }); + EXPECT_CALL(mock_stream_, Writev(testing::IsEmpty(), _)); + + MoqtRequestOk ok; + ok.request_id = kRequestId; + ok.parameters = parameters; + + QUICHE_EXPECT_OK( + stream.OnRawControlMessage(GenericMessageToRawControlMessage(ok))); +} + +TEST_F(MoqtTrackStatusRequestStreamTest, ReceiveErrorResponse) { + MoqtTrackStatusRequestStream stream = CreateStream(); + EXPECT_CALL(mock_stream_, CanWrite).WillRepeatedly(Return(true)); + EXPECT_CALL(mock_stream_, + Writev(ControlMessageOfType(MoqtMessageType::kTrackStatus), _)); + stream.BindStream(&mock_stream_); + + 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_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))); +} + +TEST_F(MoqtTrackStatusRequestStreamTest, DuplicateRequestOk) { + MoqtTrackStatusRequestStream stream = CreateStream(); + EXPECT_CALL(mock_stream_, CanWrite).WillRepeatedly(Return(true)); + EXPECT_CALL(mock_stream_, + Writev(ControlMessageOfType(MoqtMessageType::kTrackStatus), _)); + stream.BindStream(&mock_stream_); + + EXPECT_CALL(response_callback_, Call); + EXPECT_CALL(mock_stream_, Writev(testing::IsEmpty(), _)); + + MoqtRequestOk ok; + ok.request_id = kRequestId; + + QUICHE_EXPECT_OK( + stream.OnRawControlMessage(GenericMessageToRawControlMessage(ok))); + + EXPECT_THAT( + stream.OnRawControlMessage(GenericMessageToRawControlMessage(ok)), + StatusIs(absl::StatusCode::kInvalidArgument, "Duplicate REQUEST_OK")); +} + +TEST_F(MoqtTrackStatusRequestStreamTest, DuplicateRequestError) { + MoqtTrackStatusRequestStream stream = CreateStream(); + EXPECT_CALL(mock_stream_, CanWrite).WillRepeatedly(Return(true)); + EXPECT_CALL(mock_stream_, + Writev(ControlMessageOfType(MoqtMessageType::kTrackStatus), _)); + stream.BindStream(&mock_stream_); + + 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")); +} + +class MoqtTrackStatusResponseStreamTest : public quiche::test::QuicheTest { + public: + MoqtTrackStatusResponseStreamTest() + : framer_(/*using_webtrans=*/true, quic::Perspective::IS_SERVER), + message_parser_(kDefaultMoqtVersion, /*uses_web_transport=*/true, + quic::Perspective::IS_SERVER), + track_name_("foo", "bar"), + mock_publisher_(track_name_) {} + + MoqtTrackStatusResponseStream CreateStream() { + return MoqtTrackStatusResponseStream( + &framer_, message_parser_, session_error_callback_.AsStdFunction(), + session_.weak_ptr_factory_.Create()); + } + + protected: + static constexpr uint64_t kRequestId = 2; + MoqtFramer framer_; + MoqtControlMessageParser message_parser_; + FullTrackName track_name_; + StrictMock<testing::MockFunction<void(MoqtError, absl::string_view)>> + session_error_callback_; + MockSessionToPublisherInterface session_; + MockTrackPublisher mock_publisher_; + StrictMock<webtransport::test::MockStream> mock_stream_; +}; + +TEST_F(MoqtTrackStatusResponseStreamTest, ProcessTrackStatusSuccess) { + MoqtTrackStatusResponseStream stream = CreateStream(); + EXPECT_CALL(mock_stream_, CanWrite).WillRepeatedly(Return(true)); + stream.BindStream(&mock_stream_); + + EXPECT_CALL(session_, GetTrackPublisher(track_name_)) + .WillOnce(Return( + std::shared_ptr<MoqtTrackPublisher>(&mock_publisher_, [](auto*) {}))); + + MoqtObjectListener* listener = nullptr; + EXPECT_CALL(mock_publisher_, AddObjectListener) + .WillOnce(testing::SaveArg<0>(&listener)); + + MoqtTrackStatus track_status; + track_status.request_id = kRequestId; + track_status.full_track_name = track_name_; + + QUICHE_EXPECT_OK(stream.OnRawControlMessage( + GenericMessageToRawControlMessage(track_status))); + ASSERT_NE(listener, nullptr); + + EXPECT_CALL(mock_publisher_, expiration) + .WillRepeatedly(Return(quic::QuicTimeDelta::FromSeconds(5))); + EXPECT_CALL(mock_publisher_, largest_location) + .WillRepeatedly(Return(Location(10, 20))); + EXPECT_CALL(mock_stream_, + Writev(ControlMessageOfType(MoqtMessageType::kRequestOk), _)); + EXPECT_CALL(mock_publisher_, RemoveObjectListener(listener)); + + listener->OnSubscribeAccepted(); +} + +TEST_F(MoqtTrackStatusResponseStreamTest, TrackDoesNotExist) { + MoqtTrackStatusResponseStream stream = CreateStream(); + EXPECT_CALL(mock_stream_, CanWrite).WillRepeatedly(Return(true)); + stream.BindStream(&mock_stream_); + + EXPECT_CALL(session_, GetTrackPublisher(track_name_)) + .WillOnce(Return(nullptr)); + + EXPECT_CALL(mock_stream_, + Writev(ControlMessageOfType(MoqtMessageType::kRequestError), _)); + + MoqtTrackStatus track_status; + track_status.request_id = kRequestId; + track_status.full_track_name = track_name_; + + QUICHE_EXPECT_OK(stream.OnRawControlMessage( + GenericMessageToRawControlMessage(track_status))); +} + +TEST_F(MoqtTrackStatusResponseStreamTest, DuplicateTrackStatus) { + MoqtTrackStatusResponseStream stream = CreateStream(); + EXPECT_CALL(mock_stream_, CanWrite).WillRepeatedly(Return(true)); + stream.BindStream(&mock_stream_); + + EXPECT_CALL(session_, GetTrackPublisher(track_name_)) + .WillOnce(Return( + std::shared_ptr<MoqtTrackPublisher>(&mock_publisher_, [](auto*) {}))); + + MoqtObjectListener* listener = nullptr; + EXPECT_CALL(mock_publisher_, AddObjectListener) + .WillOnce(testing::SaveArg<0>(&listener)); + + MoqtTrackStatus track_status; + track_status.request_id = kRequestId; + track_status.full_track_name = track_name_; + + QUICHE_EXPECT_OK(stream.OnRawControlMessage( + GenericMessageToRawControlMessage(track_status))); + ASSERT_NE(listener, nullptr); + + EXPECT_THAT(stream.OnRawControlMessage( + GenericMessageToRawControlMessage(track_status)), + StatusIs(absl::StatusCode::kInvalidArgument, + "Duplicate TRACK_STATUS received")); + + EXPECT_CALL(mock_publisher_, RemoveObjectListener(listener)); +} + +TEST_F(MoqtTrackStatusResponseStreamTest, TrackPublisherDestroyed) { + MoqtTrackStatusResponseStream stream = CreateStream(); + EXPECT_CALL(mock_stream_, CanWrite).WillRepeatedly(Return(true)); + stream.BindStream(&mock_stream_); + + EXPECT_CALL(session_, GetTrackPublisher(track_name_)) + .WillOnce(Return( + std::shared_ptr<MoqtTrackPublisher>(&mock_publisher_, [](auto*) {}))); + + MoqtObjectListener* listener = nullptr; + EXPECT_CALL(mock_publisher_, AddObjectListener) + .WillOnce(testing::SaveArg<0>(&listener)); + + MoqtTrackStatus track_status; + track_status.request_id = kRequestId; + track_status.full_track_name = track_name_; + + QUICHE_EXPECT_OK(stream.OnRawControlMessage( + GenericMessageToRawControlMessage(track_status))); + ASSERT_NE(listener, nullptr); + + EXPECT_CALL(mock_stream_, + Writev(ControlMessageOfType(MoqtMessageType::kRequestError), _)); + + listener->OnTrackPublisherGone(); +} + +} // namespace +} // namespace moqt::test
diff --git a/quiche/quic/moqt/test_tools/mock_moqt_session.h b/quiche/quic/moqt/test_tools/mock_moqt_session.h index f4b55c0..e6e719d 100644 --- a/quiche/quic/moqt/test_tools/mock_moqt_session.h +++ b/quiche/quic/moqt/test_tools/mock_moqt_session.h
@@ -125,6 +125,10 @@ (TrackNamespace&, const MessageParameters&, MoqtResponseCallback), (override)); MOCK_METHOD(void, UnsubscribeTracks, (TrackNamespace&), (override)); + MOCK_METHOD(bool, TrackStatus, + (const FullTrackName&, const MessageParameters&, + MoqtResponseCallback), + (override)); quiche::QuicheWeakPtr<MoqtSessionInterface> GetWeakPtr() override { return weak_factory_.Create();
diff --git a/quiche/quic/moqt/test_tools/moqt_framer_utils.cc b/quiche/quic/moqt/test_tools/moqt_framer_utils.cc index f35d15b..f9c3052 100644 --- a/quiche/quic/moqt/test_tools/moqt_framer_utils.cc +++ b/quiche/quic/moqt/test_tools/moqt_framer_utils.cc
@@ -4,12 +4,18 @@ #include "quiche/quic/moqt/test_tools/moqt_framer_utils.h" +#include <cstdint> #include <string> #include <variant> +#include "absl/strings/string_view.h" +#include "quiche/quic/core/quic_types.h" #include "quiche/quic/moqt/moqt_framer.h" #include "quiche/quic/moqt/moqt_messages.h" +#include "quiche/quic/moqt/moqt_parser.h" +#include "quiche/common/platform/api/quiche_test.h" #include "quiche/common/quiche_buffer_allocator.h" +#include "quiche/common/quiche_data_reader.h" namespace moqt::test { @@ -105,4 +111,20 @@ return std::string(std::visit(FramingVisitor{framer}, frame).AsStringView()); } +MoqtRawControlMessage UnframeRawControlMessage(absl::string_view message) { + quiche::QuicheDataReader reader(message); + uint64_t raw_type; + uint16_t message_size; + bool parse_success = reader.ReadMoqVarInt(&raw_type) && + reader.ReadUInt16(&message_size) && + reader.BytesRemaining() == message_size; + if (!parse_success) { + ADD_FAILURE() << "Failed to unframe the control message"; + return MoqtRawControlMessage(); + } + return MoqtRawControlMessage{ + .type = static_cast<MoqtMessageType>(raw_type), + .payload = std::string(reader.ReadRemainingPayload())}; +} + } // namespace moqt::test
diff --git a/quiche/quic/moqt/test_tools/moqt_framer_utils.h b/quiche/quic/moqt/test_tools/moqt_framer_utils.h index 43a95f7..abea3ce 100644 --- a/quiche/quic/moqt/test_tools/moqt_framer_utils.h +++ b/quiche/quic/moqt/test_tools/moqt_framer_utils.h
@@ -13,6 +13,7 @@ #include "absl/strings/str_join.h" #include "absl/strings/string_view.h" #include "quiche/quic/moqt/moqt_messages.h" +#include "quiche/quic/moqt/moqt_parser.h" #include "quiche/common/platform/api/quiche_test.h" #include "quiche/common/quiche_data_reader.h" #include "quiche/common/quiche_mem_slice.h" @@ -66,6 +67,15 @@ return true; } +// Parses `message` into a MoqtRawControlMessage. Assumes `message` is exactly +// one message; reports a test failure if not. +MoqtRawControlMessage UnframeRawControlMessage(absl::string_view message); + +inline MoqtRawControlMessage GenericMessageToRawControlMessage( + const AnyMoqtControlMessage& message) { + return UnframeRawControlMessage(SerializeGenericMessage(message)); +} + } // namespace moqt::test #endif // QUICHE_QUIC_MOQT_TEST_TOOLS_MOQT_FRAMER_UTILS_H_