Write the MoQT Relay application. In this CL, no requests are actually forwarded. This just sets up the MoqtServer and (optional) MoQT client to handle requests. PiperOrigin-RevId: 805344884
diff --git a/build/source_list.bzl b/build/source_list.bzl index 84c347e..7666fae 100644 --- a/build/source_list.bzl +++ b/build/source_list.bzl
@@ -1578,6 +1578,7 @@ "quic/moqt/tools/moq_chat.h", "quic/moqt/tools/moqt_client.h", "quic/moqt/tools/moqt_mock_visitor.h", + "quic/moqt/tools/moqt_relay.h", "quic/moqt/tools/moqt_server.h", ] moqt_srcs = [ @@ -1602,6 +1603,7 @@ "quic/moqt/moqt_probe_manager.cc", "quic/moqt/moqt_probe_manager_test.cc", "quic/moqt/moqt_relay_publisher.cc", + "quic/moqt/moqt_relay_publisher_test.cc", "quic/moqt/moqt_relay_track_publisher.cc", "quic/moqt/moqt_relay_track_publisher_test.cc", "quic/moqt/moqt_session.cc", @@ -1623,6 +1625,9 @@ "quic/moqt/tools/moqt_client.cc", "quic/moqt/tools/moqt_end_to_end_test.cc", "quic/moqt/tools/moqt_ingestion_server_bin.cc", + "quic/moqt/tools/moqt_relay.cc", + "quic/moqt/tools/moqt_relay_bin.cc", + "quic/moqt/tools/moqt_relay_test.cc", "quic/moqt/tools/moqt_server.cc", "quic/moqt/tools/moqt_server_test.cc", "quic/moqt/tools/moqt_simulator_bin.cc",
diff --git a/build/source_list.gni b/build/source_list.gni index 09c6eca..958532c 100644 --- a/build/source_list.gni +++ b/build/source_list.gni
@@ -1582,6 +1582,7 @@ "src/quiche/quic/moqt/tools/moq_chat.h", "src/quiche/quic/moqt/tools/moqt_client.h", "src/quiche/quic/moqt/tools/moqt_mock_visitor.h", + "src/quiche/quic/moqt/tools/moqt_relay.h", "src/quiche/quic/moqt/tools/moqt_server.h", ] moqt_srcs = [ @@ -1606,6 +1607,7 @@ "src/quiche/quic/moqt/moqt_probe_manager.cc", "src/quiche/quic/moqt/moqt_probe_manager_test.cc", "src/quiche/quic/moqt/moqt_relay_publisher.cc", + "src/quiche/quic/moqt/moqt_relay_publisher_test.cc", "src/quiche/quic/moqt/moqt_relay_track_publisher.cc", "src/quiche/quic/moqt/moqt_relay_track_publisher_test.cc", "src/quiche/quic/moqt/moqt_session.cc", @@ -1627,6 +1629,9 @@ "src/quiche/quic/moqt/tools/moqt_client.cc", "src/quiche/quic/moqt/tools/moqt_end_to_end_test.cc", "src/quiche/quic/moqt/tools/moqt_ingestion_server_bin.cc", + "src/quiche/quic/moqt/tools/moqt_relay.cc", + "src/quiche/quic/moqt/tools/moqt_relay_bin.cc", + "src/quiche/quic/moqt/tools/moqt_relay_test.cc", "src/quiche/quic/moqt/tools/moqt_server.cc", "src/quiche/quic/moqt/tools/moqt_server_test.cc", "src/quiche/quic/moqt/tools/moqt_simulator_bin.cc",
diff --git a/build/source_list.json b/build/source_list.json index af36bae..12af36a 100644 --- a/build/source_list.json +++ b/build/source_list.json
@@ -1581,6 +1581,7 @@ "quiche/quic/moqt/tools/moq_chat.h", "quiche/quic/moqt/tools/moqt_client.h", "quiche/quic/moqt/tools/moqt_mock_visitor.h", + "quiche/quic/moqt/tools/moqt_relay.h", "quiche/quic/moqt/tools/moqt_server.h" ], "moqt_srcs": [ @@ -1605,6 +1606,7 @@ "quiche/quic/moqt/moqt_probe_manager.cc", "quiche/quic/moqt/moqt_probe_manager_test.cc", "quiche/quic/moqt/moqt_relay_publisher.cc", + "quiche/quic/moqt/moqt_relay_publisher_test.cc", "quiche/quic/moqt/moqt_relay_track_publisher.cc", "quiche/quic/moqt/moqt_relay_track_publisher_test.cc", "quiche/quic/moqt/moqt_session.cc", @@ -1626,6 +1628,9 @@ "quiche/quic/moqt/tools/moqt_client.cc", "quiche/quic/moqt/tools/moqt_end_to_end_test.cc", "quiche/quic/moqt/tools/moqt_ingestion_server_bin.cc", + "quiche/quic/moqt/tools/moqt_relay.cc", + "quiche/quic/moqt/tools/moqt_relay_bin.cc", + "quiche/quic/moqt/tools/moqt_relay_test.cc", "quiche/quic/moqt/tools/moqt_server.cc", "quiche/quic/moqt/tools/moqt_server_test.cc", "quiche/quic/moqt/tools/moqt_simulator_bin.cc"
diff --git a/quiche/quic/moqt/moqt_publisher.h b/quiche/quic/moqt/moqt_publisher.h index 4f723a2..0ad158b 100644 --- a/quiche/quic/moqt/moqt_publisher.h +++ b/quiche/quic/moqt/moqt_publisher.h
@@ -15,8 +15,6 @@ #include "quiche/quic/moqt/moqt_messages.h" #include "quiche/quic/moqt/moqt_object.h" #include "quiche/quic/moqt/moqt_priority.h" -#include "quiche/quic/moqt/moqt_session_interface.h" -#include "quiche/quic/moqt/moqt_track.h" #include "quiche/web_transport/web_transport.h" namespace moqt {
diff --git a/quiche/quic/moqt/moqt_relay_publisher.cc b/quiche/quic/moqt/moqt_relay_publisher.cc index 432413f..601e605 100644 --- a/quiche/quic/moqt/moqt_relay_publisher.cc +++ b/quiche/quic/moqt/moqt_relay_publisher.cc
@@ -7,13 +7,19 @@ #include <memory> #include "absl/base/nullability.h" +#include "absl/strings/string_view.h" #include "quiche/quic/moqt/moqt_messages.h" #include "quiche/quic/moqt/moqt_publisher.h" #include "quiche/quic/moqt/moqt_relay_track_publisher.h" +#include "quiche/quic/moqt/moqt_session_callbacks.h" +#include "quiche/quic/moqt/moqt_session_interface.h" #include "quiche/common/platform/api/quiche_bug_tracker.h" +#include "quiche/common/quiche_weak_ptr.h" namespace moqt { +using quiche::QuicheWeakPtr; + absl_nullable std::shared_ptr<MoqtTrackPublisher> MoqtRelayPublisher::GetTrack( const FullTrackName& track_name) { auto it = tracks_.find(track_name); @@ -23,7 +29,35 @@ return it->second; } -void MoqtRelayPublisher::Add( +void MoqtRelayPublisher::SetDefaultUpstreamSession( + MoqtSessionInterface* default_upstream_session) { + MoqtSessionInterface* old_session = + default_upstream_session_.GetIfAvailable(); + if (old_session != nullptr) { + // The Publisher no longer cares if the old session is terminated. + old_session->callbacks().session_terminated_callback = + [](absl::string_view) {}; + } + // Update callbacks. + // goaway_received_callback has already been set by MoqtClient. It will + // handle connecting to new URI and calling AddDefaultUpstreamSession() again + // when that session is ready. + default_upstream_session->callbacks().session_terminated_callback = + [this](absl::string_view error_message) { + QUICHE_LOG(INFO) << "Default upstream session terminated, error = " + << error_message; + default_upstream_session_ = QuicheWeakPtr<MoqtSessionInterface>(); + }; + AddNamespaceCallbacks(default_upstream_session); + default_upstream_session_ = default_upstream_session->GetWeakPtr(); +} + +void MoqtRelayPublisher::AddNamespaceCallbacks( + MoqtSessionInterface* /*session*/) { + // TODO(martinduke): Implement this. +} + +void MoqtRelayPublisher::AddTrack( std::shared_ptr<MoqtRelayTrackPublisher> track_publisher) { const FullTrackName& track_name = track_publisher->GetTrackName(); auto [it, success] = tracks_.emplace(track_name, track_publisher); @@ -31,7 +65,7 @@ << "Trying to add a duplicate track into a RelayPublisher"; } -void MoqtRelayPublisher::Delete(const FullTrackName& track_name) { +void MoqtRelayPublisher::DeleteTrack(const FullTrackName& track_name) { tracks_.erase(track_name); }
diff --git a/quiche/quic/moqt/moqt_relay_publisher.h b/quiche/quic/moqt/moqt_relay_publisher.h index 8d8e0f2..e1eda85 100644 --- a/quiche/quic/moqt/moqt_relay_publisher.h +++ b/quiche/quic/moqt/moqt_relay_publisher.h
@@ -12,6 +12,8 @@ #include "quiche/quic/moqt/moqt_messages.h" #include "quiche/quic/moqt/moqt_publisher.h" #include "quiche/quic/moqt/moqt_relay_track_publisher.h" +#include "quiche/quic/moqt/moqt_session_interface.h" +#include "quiche/common/quiche_weak_ptr.h" namespace moqt { @@ -19,7 +21,8 @@ // and namespaces with upstream sessions that can deliver those things. class MoqtRelayPublisher : public MoqtPublisher { public: - MoqtRelayPublisher() = default; + explicit MoqtRelayPublisher(bool broadcast_mode) + : broadcast_mode_(broadcast_mode) {} MoqtRelayPublisher(const MoqtRelayPublisher&) = delete; MoqtRelayPublisher(MoqtRelayPublisher&&) = delete; MoqtRelayPublisher& operator=(const MoqtRelayPublisher&) = delete; @@ -32,14 +35,32 @@ void AddNamespaceListener(NamespaceListener* /*listener*/) override {} void RemoveNamespaceListener(NamespaceListener* /*listener*/) override {} - void Add(std::shared_ptr<MoqtRelayTrackPublisher> track_publisher); - void Delete(const FullTrackName& track_name); + // There is a new default upstream session. When there is no other namespace + // information, requests will route here. + void SetDefaultUpstreamSession( + MoqtSessionInterface* default_upstream_session); + // There is a new incoming session. MoqtRelayPublisher will set the callbacks + // for this session, but need not keep any state at this time. + virtual void AddNamespaceCallbacks(MoqtSessionInterface* session); + + // Returns the default upstream session. + quiche::QuicheWeakPtr<MoqtSessionInterface>& GetDefaultUpstreamSession() { + return default_upstream_session_; + } private: + void AddTrack(std::shared_ptr<MoqtRelayTrackPublisher> track_publisher); + void DeleteTrack(const FullTrackName& track_name); + absl::flat_hash_map<FullTrackName, std::shared_ptr<MoqtRelayTrackPublisher>> tracks_; // TODO(martinduke): Add a map of Namespaces to source sessions and // namespace listeners. + + quiche::QuicheWeakPtr<MoqtSessionInterface> default_upstream_session_; + // If true, PUBLISH_NAMESPACE messages will be forwarded to all sessions, + // whether or not they are subscribed. + bool broadcast_mode_; }; } // namespace moqt
diff --git a/quiche/quic/moqt/moqt_relay_publisher_test.cc b/quiche/quic/moqt/moqt_relay_publisher_test.cc new file mode 100644 index 0000000..abe7af0 --- /dev/null +++ b/quiche/quic/moqt/moqt_relay_publisher_test.cc
@@ -0,0 +1,130 @@ +// Copyright 2025 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_relay_publisher.h" + +#include <cstdint> +#include <optional> +#include <utility> + +#include "absl/strings/string_view.h" +#include "quiche/quic/moqt/moqt_messages.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_track.h" +#include "quiche/common/platform/api/quiche_test.h" +#include "quiche/common/quiche_weak_ptr.h" + +namespace moqt { +namespace test { + +class MockMoqtSession : public MoqtSessionInterface { + public: + MOCK_METHOD(MoqtSessionCallbacks&, callbacks, (), (override)); + MOCK_METHOD(void, Error, (MoqtError code, absl::string_view error), + (override)); + MOCK_METHOD(bool, SubscribeAbsolute, + (const FullTrackName& name, uint64_t start_group, + uint64_t start_object, SubscribeRemoteTrack::Visitor* visitor, + VersionSpecificParameters parameters), + (override)); + MOCK_METHOD(bool, SubscribeAbsolute, + (const FullTrackName& name, uint64_t start_group, + uint64_t start_object, uint64_t end_group, + SubscribeRemoteTrack::Visitor* visitor, + VersionSpecificParameters parameters), + (override)); + MOCK_METHOD(bool, SubscribeCurrentObject, + (const FullTrackName& name, + SubscribeRemoteTrack::Visitor* visitor, + VersionSpecificParameters parameters), + (override)); + MOCK_METHOD(bool, SubscribeNextGroup, + (const FullTrackName& name, + SubscribeRemoteTrack::Visitor* visitor, + VersionSpecificParameters parameters), + (override)); + MOCK_METHOD(bool, SubscribeUpdate, + (const FullTrackName& name, std::optional<Location> start, + std::optional<uint64_t> end_group, + std::optional<MoqtPriority> subscriber_priority, + std::optional<bool> forward, + VersionSpecificParameters parameters), + (override)); + MOCK_METHOD(void, Unsubscribe, (const FullTrackName& name), (override)); + MOCK_METHOD(bool, Fetch, + (const FullTrackName& name, FetchResponseCallback callback, + Location start, uint64_t end_group, + std::optional<uint64_t> end_object, MoqtPriority priority, + std::optional<MoqtDeliveryOrder> delivery_order, + VersionSpecificParameters parameters), + (override)); + MOCK_METHOD(bool, RelativeJoiningFetch, + (const FullTrackName& name, + SubscribeRemoteTrack::Visitor* visitor, + uint64_t num_previous_groups, + VersionSpecificParameters parameters), + (override)); + MOCK_METHOD(bool, RelativeJoiningFetch, + (const FullTrackName& name, + SubscribeRemoteTrack::Visitor* visitor, + FetchResponseCallback callback, uint64_t num_previous_groups, + MoqtPriority priority, + std::optional<MoqtDeliveryOrder> delivery_order, + VersionSpecificParameters parameters), + (override)); + + quiche::QuicheWeakPtr<MoqtSessionInterface> GetWeakPtr() override { + return weak_factory_.Create(); + } + quiche::QuicheWeakPtrFactory<MoqtSessionInterface> weak_factory_{this}; +}; + +class MoqtRelayPublisherTest : public quiche::test::QuicheTest { + public: + MoqtRelayPublisherTest() : publisher_(false) {} + + MoqtSessionCallbacks callbacks_; + MockMoqtSession session_; + MoqtRelayPublisher publisher_; +}; + +TEST_F(MoqtRelayPublisherTest, SetDefaultUpstreamSession) { + EXPECT_FALSE(publisher_.GetDefaultUpstreamSession().IsValid()); + EXPECT_CALL(session_, callbacks).WillOnce(testing::ReturnRef(callbacks_)); + publisher_.SetDefaultUpstreamSession(&session_); + EXPECT_TRUE(publisher_.GetDefaultUpstreamSession().IsValid()); + EXPECT_EQ(publisher_.GetDefaultUpstreamSession().GetIfAvailable(), &session_); + // Destroy the session. + std::move(callbacks_.session_terminated_callback)("test"); + EXPECT_FALSE(publisher_.GetDefaultUpstreamSession().IsValid()); +} + +TEST_F(MoqtRelayPublisherTest, SetDefaultUpstreamSessionTwice) { + EXPECT_FALSE(publisher_.GetDefaultUpstreamSession().IsValid()); + EXPECT_CALL(session_, callbacks()).WillOnce(testing::ReturnRef(callbacks_)); + publisher_.SetDefaultUpstreamSession(&session_); + EXPECT_TRUE(publisher_.GetDefaultUpstreamSession().IsValid()); + EXPECT_EQ(publisher_.GetDefaultUpstreamSession().GetIfAvailable(), &session_); + + MockMoqtSession session2; + MoqtSessionCallbacks callbacks2; + EXPECT_CALL(session_, callbacks).WillOnce(testing::ReturnRef(callbacks_)); + EXPECT_CALL(session2, callbacks).WillOnce(testing::ReturnRef(callbacks2)); + publisher_.SetDefaultUpstreamSession(&session2); + EXPECT_TRUE(publisher_.GetDefaultUpstreamSession().IsValid()); + EXPECT_EQ(publisher_.GetDefaultUpstreamSession().GetIfAvailable(), &session2); + + // Destroying the old session doesn't affect the publisher. + std::move(callbacks_.session_terminated_callback)("test"); + EXPECT_TRUE(publisher_.GetDefaultUpstreamSession().IsValid()); + + // Destroying the new session does. + std::move(callbacks2.session_terminated_callback)("test"); + EXPECT_FALSE(publisher_.GetDefaultUpstreamSession().IsValid()); +} + +} // namespace test +} // namespace moqt
diff --git a/quiche/quic/moqt/moqt_session.cc b/quiche/quic/moqt/moqt_session.cc index 1e73ae9..1dc806b 100644 --- a/quiche/quic/moqt/moqt_session.cc +++ b/quiche/quic/moqt/moqt_session.cc
@@ -106,6 +106,7 @@ publisher_(DefaultPublisher::GetInstance()), local_max_request_id_(parameters.max_request_id), alarm_factory_(std::move(alarm_factory)), + weak_ptr_factory_(this), liveness_token_(std::make_shared<Empty>()) { if (parameters_.using_webtrans) { session_->SetOnDraining([this]() {
diff --git a/quiche/quic/moqt/moqt_session.h b/quiche/quic/moqt/moqt_session.h index 140de3f..c274748 100644 --- a/quiche/quic/moqt/moqt_session.h +++ b/quiche/quic/moqt/moqt_session.h
@@ -18,7 +18,6 @@ #include "absl/container/btree_set.h" #include "absl/container/flat_hash_map.h" #include "absl/container/flat_hash_set.h" -#include "absl/status/statusor.h" #include "absl/strings/string_view.h" #include "quiche/quic/core/quic_alarm.h" #include "quiche/quic/core/quic_alarm_factory.h" @@ -94,8 +93,6 @@ void OnCanCreateNewOutgoingBidirectionalStream() override {} void OnCanCreateNewOutgoingUnidirectionalStream() override; - void Error(MoqtError code, absl::string_view error) override; - quic::Perspective perspective() const { return parameters_.perspective; } // Returns true if message was sent. @@ -117,15 +114,13 @@ void CancelAnnounce(TrackNamespace track_namespace, RequestErrorCode code, absl::string_view reason_phrase); - // Returns true if SUBSCRIBE was sent. If there is already a subscription to - // the track, the message will still be sent. However, the visitor will be - // ignored. If |visitor| is nullptr, forward will be set to false. - // Subscribe from (start_group, start_object) to the end of the track. + // MoqtSessionInterface implementation. + MoqtSessionCallbacks& callbacks() override { return callbacks_; } + void Error(MoqtError code, absl::string_view error) override; bool SubscribeAbsolute(const FullTrackName& name, uint64_t start_group, uint64_t start_object, SubscribeRemoteTrack::Visitor* visitor, VersionSpecificParameters parameters) override; - // Subscribe from (start_group, start_object) to the end of end_group. bool SubscribeAbsolute(const FullTrackName& name, uint64_t start_group, uint64_t start_object, uint64_t end_group, SubscribeRemoteTrack::Visitor* visitor, @@ -141,44 +136,32 @@ std::optional<MoqtPriority> subscriber_priority, std::optional<bool> forward, VersionSpecificParameters parameters) override; - // Returns false if the subscription is not found. The session immediately - // destroys all subscription state. - void Unsubscribe(const FullTrackName& name); - // |callback| will be called when FETCH_OK or FETCH_ERROR is received, and - // delivers a pointer to MoqtFetchTask for application use. The callback - // transfers ownership of MoqtFetchTask to the application. - // To cancel a FETCH, simply destroy the FetchTask. + void Unsubscribe(const FullTrackName& name) override; bool Fetch(const FullTrackName& name, FetchResponseCallback callback, Location start, uint64_t end_group, std::optional<uint64_t> end_object, MoqtPriority priority, std::optional<MoqtDeliveryOrder> delivery_order, VersionSpecificParameters parameters) override; - // 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 - // in the FETCH will not be filled by with ObjectDoesNotExist. If the FETCH - // fails for any reason, the application will not receive a notification; it - // will just appear to be missing objects. bool RelativeJoiningFetch(const FullTrackName& name, SubscribeRemoteTrack::Visitor* visitor, uint64_t num_previous_groups, VersionSpecificParameters parameters) override; - // Sends both a SUBSCRIBE and a joining FETCH, beginning |num_previous_groups| - // groups before the current group. The application provides |callback| to - // fully control acceptance of Fetched objects. bool RelativeJoiningFetch(const FullTrackName& name, SubscribeRemoteTrack::Visitor* visitor, FetchResponseCallback callback, uint64_t num_previous_groups, MoqtPriority priority, std::optional<MoqtDeliveryOrder> delivery_order, VersionSpecificParameters parameters) override; + quiche::QuicheWeakPtr<MoqtSessionInterface> GetWeakPtr() override { + return weak_ptr_factory_.Create(); + } // Send a GOAWAY message to the peer. |new_session_uri| must be empty if // called by the client. void GoAway(absl::string_view new_session_uri); webtransport::Session* session() { return session_; } - MoqtSessionCallbacks& callbacks() override { return callbacks_; } + MoqtPublisher* publisher() { return publisher_; } void set_publisher(MoqtPublisher* publisher) { publisher_ = publisher; } bool support_object_acks() const { return parameters_.support_object_acks; } @@ -249,17 +232,18 @@ void OnSubscribeErrorMessage(const MoqtSubscribeError& message) override; void OnUnsubscribeMessage(const MoqtUnsubscribe& message) override; // There is no state to update for SUBSCRIBE_DONE. - void OnSubscribeDoneMessage(const MoqtSubscribeDone& /*message*/) override; + void OnSubscribeDoneMessage(const MoqtSubscribeDone& message) override; void OnSubscribeUpdateMessage(const MoqtSubscribeUpdate& message) override; void OnAnnounceMessage(const MoqtAnnounce& message) override; void OnAnnounceOkMessage(const MoqtAnnounceOk& message) override; void OnAnnounceErrorMessage(const MoqtAnnounceError& message) override; - void OnUnannounceMessage(const MoqtUnannounce& /*message*/) override; + void OnUnannounceMessage(const MoqtUnannounce& message) override; void OnAnnounceCancelMessage(const MoqtAnnounceCancel& message) override; void OnTrackStatusMessage(const MoqtTrackStatus& message) override; - void OnTrackStatusOkMessage(const MoqtTrackStatusOk& message) override {} + void OnTrackStatusOkMessage(const MoqtTrackStatusOk& /*message*/) override { + } void OnTrackStatusErrorMessage( - const MoqtTrackStatusError& message) override {} + const MoqtTrackStatusError& /*message*/) override {} void OnGoAwayMessage(const MoqtGoAway& /*message*/) override; void OnSubscribeNamespaceMessage( const MoqtSubscribeNamespace& message) override; @@ -271,13 +255,13 @@ const MoqtUnsubscribeNamespace& message) override; void OnMaxRequestIdMessage(const MoqtMaxRequestId& message) override; void OnFetchMessage(const MoqtFetch& message) override; - void OnFetchCancelMessage(const MoqtFetchCancel& message) override {} + void OnFetchCancelMessage(const MoqtFetchCancel& /*message*/) override {} void OnFetchOkMessage(const MoqtFetchOk& message) override; void OnFetchErrorMessage(const MoqtFetchError& message) override; void OnRequestsBlockedMessage(const MoqtRequestsBlocked& message) override; void OnPublishMessage(const MoqtPublish& message) override; - void OnPublishOkMessage(const MoqtPublishOk& message) override {}; - void OnPublishErrorMessage(const MoqtPublishError& message) override {}; + void OnPublishOkMessage(const MoqtPublishOk& /*message*/) override {} + void OnPublishErrorMessage(const MoqtPublishError& /*message*/) override {} void OnObjectAckMessage(const MoqtObjectAck& message) override { auto subscription_it = session_->published_subscriptions_.find(message.subscribe_id); @@ -321,8 +305,10 @@ // webtransport::StreamVisitor implementation. void OnCanRead() override; void OnCanWrite() override {} - void OnResetStreamReceived(webtransport::StreamErrorCode error) override {} - void OnStopSendingReceived(webtransport::StreamErrorCode error) override {} + void OnResetStreamReceived( + webtransport::StreamErrorCode /*error*/) override {} + void OnStopSendingReceived( + webtransport::StreamErrorCode /*error*/) override {} void OnWriteSideInDataRecvdState() override {} // MoqtParserVisitor implementation. @@ -413,7 +399,7 @@ QUICHE_CHECK(window_.has_value()); return window_->start(); } - MoqtFilterType filter_type() const { return filter_type_; }; + MoqtFilterType filter_type() const { return filter_type_; } void OnDataStreamCreated(webtransport::StreamId id, DataStreamIndex start_sequence); @@ -511,8 +497,10 @@ // webtransport::StreamVisitor implementation. void OnCanRead() override {} void OnCanWrite() override; - void OnResetStreamReceived(webtransport::StreamErrorCode error) override {} - void OnStopSendingReceived(webtransport::StreamErrorCode error) override {} + void OnResetStreamReceived( + webtransport::StreamErrorCode /*error*/) override {} + void OnStopSendingReceived( + webtransport::StreamErrorCode /*error*/) override {} void OnWriteSideInDataRecvdState() override {} class DeliveryTimeoutDelegate @@ -599,10 +587,11 @@ // webtransport::StreamVisitor implementation. void OnCanRead() override {} // Write-only stream. void OnCanWrite() override; - void OnResetStreamReceived(webtransport::StreamErrorCode error) override { + void OnResetStreamReceived( + webtransport::StreamErrorCode /*error*/) override { } // Write-only stream - void OnStopSendingReceived(webtransport::StreamErrorCode error) override { - } + void OnStopSendingReceived( + webtransport::StreamErrorCode /*error*/) override {} void OnWriteSideInDataRecvdState() override {} private: @@ -674,13 +663,14 @@ // No class access below this line! } - void OnNewObjectAvailable(Location sequence, uint64_t subgroup, - MoqtPriority publisher_priority) override {} - void OnNewFinAvailable(Location location, uint64_t subgroup) override {} + void OnNewObjectAvailable(Location /*sequence*/, uint64_t /*subgroup*/, + MoqtPriority /*publisher_priority*/) 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 {} + uint64_t /*group*/, uint64_t /*subgroup*/, + webtransport::StreamErrorCode /*error_code*/) override {} + void OnGroupAbandoned(uint64_t /*group_id*/) override {} void OnTrackPublisherGone() override { publisher_ = nullptr; OnSubscribeRejected(MoqtSubscribeErrorReason( @@ -788,7 +778,6 @@ void OnMalformedTrack(RemoteTrack* track); bool is_closing_ = false; - webtransport::Session* session_; MoqtSessionParameters parameters_; MoqtSessionCallbacks callbacks_; @@ -879,6 +868,8 @@ // the first object of group n+1 arrives. bool alternate_delivery_timeout_ = false; + quiche::QuicheWeakPtrFactory<MoqtSessionInterface> weak_ptr_factory_; + // Must be last. Token used to make sure that the streams do not call into // the session when the session has already been destroyed.
diff --git a/quiche/quic/moqt/moqt_session_interface.h b/quiche/quic/moqt/moqt_session_interface.h index a6c0c46..a703a3d 100644 --- a/quiche/quic/moqt/moqt_session_interface.h +++ b/quiche/quic/moqt/moqt_session_interface.h
@@ -14,6 +14,7 @@ #include "quiche/quic/moqt/moqt_session_callbacks.h" #include "quiche/quic/moqt/moqt_track.h" #include "quiche/common/quiche_callbacks.h" +#include "quiche/common/quiche_weak_ptr.h" namespace moqt { @@ -115,6 +116,8 @@ // TODO: Add AnnounceCancel method. // TODO: Add TrackStatusRequest method. // TODO: Add SubscribeUpdate, SubscribeDone method. + + virtual quiche::QuicheWeakPtr<MoqtSessionInterface> GetWeakPtr() = 0; }; } // namespace moqt
diff --git a/quiche/quic/moqt/tools/moqt_relay.cc b/quiche/quic/moqt/tools/moqt_relay.cc new file mode 100644 index 0000000..65ad937 --- /dev/null +++ b/quiche/quic/moqt/tools/moqt_relay.cc
@@ -0,0 +1,114 @@ +// Copyright (c) 2025 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/tools/moqt_relay.h" + +#include <cstdint> +#include <memory> +#include <string> +#include <utility> + +#include "absl/status/statusor.h" +#include "absl/strings/string_view.h" +#include "quiche/quic/core/crypto/proof_source.h" +#include "quiche/quic/core/crypto/proof_verifier.h" +#include "quiche/quic/core/io/quic_event_loop.h" +#include "quiche/quic/core/quic_server_id.h" +#include "quiche/quic/moqt/moqt_session.h" +#include "quiche/quic/moqt/moqt_session_callbacks.h" +#include "quiche/quic/moqt/tools/moqt_client.h" +#include "quiche/quic/moqt/tools/moqt_server.h" +#include "quiche/quic/platform/api/quic_default_proof_providers.h" +#include "quiche/quic/platform/api/quic_socket_address.h" +#include "quiche/quic/tools/fake_proof_verifier.h" +#include "quiche/quic/tools/quic_name_lookup.h" +#include "quiche/quic/tools/quic_url.h" +#include "quiche/common/platform/api/quiche_logging.h" +#include "quiche/common/quiche_ip_address.h" + +namespace moqt { + +MoqtRelay::MoqtRelay(std::unique_ptr<quic::ProofSource> proof_source, + std::string bind_address, uint16_t bind_port, + absl::string_view default_upstream, + bool ignore_certificate, bool broadcast_mode) + : MoqtRelay(std::move(proof_source), bind_address, bind_port, + default_upstream, ignore_certificate, broadcast_mode, nullptr) { +} + +// protected members. +MoqtRelay::MoqtRelay(std::unique_ptr<quic::ProofSource> proof_source, + std::string bind_address, uint16_t bind_port, + absl::string_view default_upstream, + bool ignore_certificate, bool broadcast_mode, + quic::QuicEventLoop* client_event_loop) + : ignore_certificate_(ignore_certificate), + client_event_loop_(client_event_loop), + server_(std::make_unique<MoqtServer>(std::move(proof_source), + [this](absl::string_view path) { + return IncomingSessionHandler( + path); + })), + publisher_(broadcast_mode) { + quiche::QuicheIpAddress bind_ip_address; + QUICHE_CHECK(bind_ip_address.FromString(bind_address)); + // CreateUDPSocketAndListen() creates the event loop that we will pass to + // MoqtClient. + server_->quic_server().CreateUDPSocketAndListen( + quic::QuicSocketAddress(bind_ip_address, bind_port)); + if (!default_upstream.empty()) { + quic::QuicUrl url(default_upstream, "https"); + if (client_event_loop == nullptr) { + client_event_loop = server_->quic_server().event_loop(); + } + default_upstream_client_ = + CreateClient(url, ignore_certificate, client_event_loop_); + default_upstream_client_->Connect(url.PathParamsQuery(), + CreateClientCallbacks()); + } +} + +// private members. +std::unique_ptr<moqt::MoqtClient> MoqtRelay::CreateClient( + quic::QuicUrl url, bool ignore_certificate, + quic::QuicEventLoop* event_loop) { + quic::QuicServerId server_id(url.host(), url.port()); + quic::QuicSocketAddress peer_address = + quic::tools::LookupAddress(AF_UNSPEC, server_id); + std::unique_ptr<quic::ProofVerifier> verifier; + if (ignore_certificate) { + verifier = std::make_unique<quic::FakeProofVerifier>(); + } else { + verifier = quic::CreateDefaultProofVerifier(server_id.host()); + } + return std::make_unique<moqt::MoqtClient>(peer_address, server_id, + std::move(verifier), event_loop); +} + +MoqtSessionCallbacks MoqtRelay::CreateClientCallbacks() { + MoqtSessionCallbacks callbacks; + callbacks.session_established_callback = [this]() { + default_upstream_client_->session()->set_publisher(&publisher_); + publisher_.SetDefaultUpstreamSession(default_upstream_client_->session()); + }; + callbacks.goaway_received_callback = [](absl::string_view new_session_uri) { + QUICHE_LOG(INFO) << "GoAway received, new session uri = " + << new_session_uri; + // There's no asynchronous means today to connect to a new URL. + // Therefore, just ignore GOAWAY. + }; + return callbacks; +} + +absl::StatusOr<MoqtConfigureSessionCallback> MoqtRelay::IncomingSessionHandler( + absl::string_view /*path*/) { + return [this](MoqtSession* session) { + session->set_publisher(&publisher_); + session->callbacks().session_established_callback = [this, session]() { + publisher_.AddNamespaceCallbacks(session); + }; + }; +} + +} // namespace moqt
diff --git a/quiche/quic/moqt/tools/moqt_relay.h b/quiche/quic/moqt/tools/moqt_relay.h new file mode 100644 index 0000000..3398685 --- /dev/null +++ b/quiche/quic/moqt/tools/moqt_relay.h
@@ -0,0 +1,77 @@ +// Copyright (c) 2025 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_TOOLS_MOQT_RELAY_H_ +#define QUICHE_QUIC_MOQT_TOOLS_MOQT_RELAY_H_ + +#include <cstdint> +#include <memory> +#include <string> + +#include "absl/status/statusor.h" +#include "absl/strings/string_view.h" +#include "quiche/quic/core/crypto/proof_source.h" +#include "quiche/quic/core/io/quic_event_loop.h" +#include "quiche/quic/moqt/moqt_relay_publisher.h" +#include "quiche/quic/moqt/moqt_session_callbacks.h" +#include "quiche/quic/moqt/tools/moqt_client.h" +#include "quiche/quic/moqt/tools/moqt_server.h" +#include "quiche/quic/tools/quic_url.h" + +namespace moqt { + +// Implements a pure MoqtRelay. It binds to |bind_address| and |bind_port| to +// listen for sessions, and optionally connects to |default_upstream| on +// startup that serves as a default route for requests. +// Requests for a track are forwarded to whatever session has published the +// relevant namespace, or the default route if not published. +// Incoming namespace subscriptions are stored locally. +// Incoming PUBLISH_NAMESPACE are forwarded to all adjacent sessions if +// broadcast_mode is true, otherwise only to sessions that have subscribed. +class MoqtRelay { + public: + // If |default_upstream| is empty, no default upstream session is created. + MoqtRelay(std::unique_ptr<quic::ProofSource> proof_source, + std::string bind_address, uint16_t bind_port, + absl::string_view default_upstream, bool ignore_certificate, + bool broadcast_mode); + + void HandleEventsForever() { server_->quic_server().HandleEventsForever(); } + + protected: // Constructor for MoqtTestRelay. + // If |client_event_loop| is null, the event loop from |server_| is used. For + // test relays, it is not null, and the provided event loop is used for the + // client. It will be the same event loop as the remote server, rather than + // the local server. + MoqtRelay(std::unique_ptr<quic::ProofSource> proof_source, + std::string bind_address, uint16_t bind_port, + absl::string_view default_upstream, bool ignore_certificate, + bool broadcast_mode, quic::QuicEventLoop* client_event_loop); + // Other functions for MoqtTestRelay. + MoqtServer* server() { return server_.get(); } + MoqtClient* client() { return default_upstream_client_.get(); } + MoqtRelayPublisher* publisher() { return &publisher_; } + + private: + std::unique_ptr<moqt::MoqtClient> CreateClient( + quic::QuicUrl url, bool ignore_certificate, + quic::QuicEventLoop* event_loop); + + MoqtSessionCallbacks CreateClientCallbacks(); + + absl::StatusOr<MoqtConfigureSessionCallback> IncomingSessionHandler( + absl::string_view path); + + const bool ignore_certificate_; + quic::QuicEventLoop* client_event_loop_; + + // Pointer to a client that has received GOAWAY. + std::unique_ptr<MoqtClient> default_upstream_client_; + std::unique_ptr<MoqtServer> server_; + MoqtRelayPublisher publisher_; +}; + +} // namespace moqt + +#endif // QUICHE_QUIC_MOQT_TOOLS_MOQT_RELAY_H_
diff --git a/quiche/quic/moqt/tools/moqt_relay_bin.cc b/quiche/quic/moqt/tools/moqt_relay_bin.cc new file mode 100644 index 0000000..e281370 --- /dev/null +++ b/quiche/quic/moqt/tools/moqt_relay_bin.cc
@@ -0,0 +1,59 @@ +// Copyright (c) 2025 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 <poll.h> +#include <unistd.h> + +#include <cstdint> +#include <string> +#include <vector> + +#include "absl/strings/string_view.h" +#include "quiche/quic/moqt/tools/moqt_relay.h" +#include "quiche/common/platform/api/quiche_command_line_flags.h" +#include "quiche/common/platform/api/quiche_default_proof_providers.h" + +DEFINE_QUICHE_COMMAND_LINE_FLAG( + bool, disable_certificate_verification, false, + "If true, don't verify the server certificate."); + +DEFINE_QUICHE_COMMAND_LINE_FLAG(std::string, bind_address, "127.0.0.1", + "Local IP address to bind to"); + +DEFINE_QUICHE_COMMAND_LINE_FLAG(uint16_t, port, 9667, + "Port for the server to listen on"); + +DEFINE_QUICHE_COMMAND_LINE_FLAG( + std::string, default_upstream, "", + "If set, connect to the upstream URL and forward all requests there if " + "there is no explicitly advertised source."); + +DEFINE_QUICHE_COMMAND_LINE_FLAG( + bool, broadcast_mode, false, + "If set, PUBLISH_NAMESPACE messages will be forwarded to all sessions, " + "whether or not they are subscribed."); + +// A pure MoQT relay. Accepts connections. Will try to route requests from a +// session to a different appropriate upstream session. If the namespace for the +// request has not been advertised, it will reject the request. If +// |default_upstream| is set, it connects on startup to that hosts, and forwards +// such requests there instead. +int main(int argc, char* argv[]) { + const char* usage = "Usage: moqt_relay [options]"; + std::vector<std::string> args = + quiche::QuicheParseCommandLineFlags(usage, argc, argv); + if (!args.empty()) { + quiche::QuichePrintCommandLineFlagHelp(usage); + return 1; + } + moqt::MoqtRelay relay( + quiche::CreateDefaultProofSource(), + quiche::GetQuicheCommandLineFlag(FLAGS_bind_address), + quiche::GetQuicheCommandLineFlag(FLAGS_port), + quiche::GetQuicheCommandLineFlag(FLAGS_default_upstream), + quiche::GetQuicheCommandLineFlag(FLAGS_disable_certificate_verification), + quiche::GetQuicheCommandLineFlag(FLAGS_broadcast_mode)); + relay.HandleEventsForever(); + return 0; +}
diff --git a/quiche/quic/moqt/tools/moqt_relay_test.cc b/quiche/quic/moqt/tools/moqt_relay_test.cc new file mode 100644 index 0000000..159eb72 --- /dev/null +++ b/quiche/quic/moqt/tools/moqt_relay_test.cc
@@ -0,0 +1,184 @@ +// Copyright (c) 2025 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/tools/moqt_relay.h" + +#include <cstdint> +#include <string> +#include <utility> + +#include "absl/strings/string_view.h" +#include "quiche/quic/core/io/quic_event_loop.h" +#include "quiche/quic/core/quic_time.h" +#include "quiche/quic/moqt/moqt_relay_publisher.h" +#include "quiche/quic/moqt/moqt_session.h" +#include "quiche/quic/moqt/moqt_session_interface.h" +#include "quiche/quic/moqt/tools/moqt_client.h" +#include "quiche/quic/moqt/tools/moqt_server.h" +#include "quiche/quic/test_tools/crypto_test_utils.h" +#include "quiche/common/platform/api/quiche_test.h" +#include "quiche/common/quiche_weak_ptr.h" + +namespace moqt { +namespace test { + +constexpr quic::QuicTime::Delta kEventLoopDuration = + quic::QuicTime::Delta::FromMilliseconds(50); + +class TestMoqtRelay : public MoqtRelay { + public: + TestMoqtRelay(std::string bind_address, uint16_t bind_port, + absl::string_view default_upstream, bool ignore_certificate, + bool promiscuous_mode, quic::QuicEventLoop* event_loop) + : MoqtRelay(quic::test::crypto_test_utils::ProofSourceForTesting(), + bind_address, bind_port, default_upstream, ignore_certificate, + promiscuous_mode, event_loop) {} + + quic::QuicEventLoop* server_event_loop() { + return server()->quic_server().event_loop(); + } + + void RunOneEvent() { + server_event_loop()->RunEventLoopOnce(kEventLoopDuration); + } + + MoqtSession* client_session() { + return (client() == nullptr) ? nullptr : client()->session(); + } + + MoqtRelayPublisher* publisher() { return MoqtRelay::publisher(); } +}; + +class MoqtRelayTest : public quiche::test::QuicheTest { + public: + MoqtRelayTest() + : upstream_("127.0.0.1", 9991, "", true, false, nullptr), // no client. + relay_("127.0.0.1", 9992, "https://127.0.0.1:9991", true, false, + upstream_.server_event_loop()), + downstream_("127.0.0.1", 9993, "https://127.0.0.1:9992", true, false, + relay_.server_event_loop()) { + RunUntilConnected(relay_, upstream_); + RunUntilConnected(downstream_, relay_); + } + + inline bool ClientFullyConnected(TestMoqtRelay& client) { + return client.publisher()->GetDefaultUpstreamSession().IsValid() && + client.publisher()->GetDefaultUpstreamSession().GetIfAvailable() == + client.client_session(); + } + + void RunUntilConnected(TestMoqtRelay& client, TestMoqtRelay& server) { + int iterations_remaining = 20; + while (!ClientFullyConnected(client) && iterations_remaining-- > 0) { + server.RunOneEvent(); + } + ASSERT_GT(iterations_remaining, 0); + } + + TestMoqtRelay upstream_, relay_, downstream_; +}; + +TEST_F(MoqtRelayTest, NodeChainEstablished) { + // relay_ and downstream_ have a default session. + ASSERT_NE(downstream_.client_session(), nullptr); + EXPECT_EQ(downstream_.client_session()->publisher(), downstream_.publisher()); + ASSERT_NE(downstream_.publisher(), nullptr); + EXPECT_EQ( + downstream_.publisher()->GetDefaultUpstreamSession().GetIfAvailable(), + downstream_.client_session()->GetWeakPtr().GetIfAvailable()); + + ASSERT_NE(relay_.client_session(), nullptr); + EXPECT_EQ(relay_.client_session()->publisher(), relay_.publisher()); + ASSERT_NE(relay_.publisher(), nullptr); + EXPECT_EQ(relay_.publisher()->GetDefaultUpstreamSession().GetIfAvailable(), + relay_.client_session()->GetWeakPtr().GetIfAvailable()); + + EXPECT_EQ(upstream_.client_session(), nullptr); + ASSERT_NE(upstream_.publisher(), nullptr); + EXPECT_EQ(upstream_.publisher()->GetDefaultUpstreamSession().GetIfAvailable(), + nullptr); +} + +TEST_F(MoqtRelayTest, CloseSession) { + ASSERT_NE(relay_.client_session(), nullptr); + std::move(relay_.client_session()->callbacks().session_terminated_callback)( + ""); + EXPECT_FALSE(relay_.publisher()->GetDefaultUpstreamSession().IsValid()); +} + +#if 0 // TODO(martinduke): Re-enable these tests when GOAWAY support exists. +TEST_F(MoqtRelayTest, GoAwayAtClient) { + ASSERT_NE(relay_.client_session(), nullptr); + // Provide the same URI again. + MoqtSessionInterface* original_default_session = + relay_.publisher()->GetDefaultUpstreamSession().GetIfAvailable(); + EXPECT_NE(original_default_session, nullptr); + std::move(relay_.client_session()->callbacks().goaway_received_callback)( + "https://127.0.0.1:9991"); + RunUntilConnected(relay_, upstream_); + EXPECT_TRUE(relay_.publisher()->GetDefaultUpstreamSession().IsValid()); + EXPECT_EQ(relay_.publisher()->GetDefaultUpstreamSession().GetIfAvailable(), + relay_.client_session()); + EXPECT_NE(relay_.publisher()->GetDefaultUpstreamSession().GetIfAvailable(), + original_default_session); + + // Terminating the original session does nothing. + std::move(original_default_session->callbacks().session_terminated_callback)( + "test"); + EXPECT_TRUE(relay_.publisher()->GetDefaultUpstreamSession().IsValid()); + EXPECT_EQ(relay_.publisher()->GetDefaultUpstreamSession().GetIfAvailable(), + relay_.client_session()); +} + +TEST_F(MoqtRelayTest, TwoGoAwaysAtClient) { + ASSERT_NE(relay_.client_session(), nullptr); + // Provide the same URI again. + MoqtSessionInterface* original_default_session = + relay_.publisher()->GetDefaultUpstreamSession().GetIfAvailable(); + EXPECT_NE(original_default_session, nullptr); + std::move(relay_.client_session()->callbacks().goaway_received_callback)( + "https://127.0.0.1:9991"); + RunUntilConnected(relay_, upstream_); + EXPECT_TRUE(relay_.publisher()->GetDefaultUpstreamSession().IsValid()); + EXPECT_EQ(relay_.publisher()->GetDefaultUpstreamSession().GetIfAvailable(), + relay_.client_session()); + EXPECT_NE(relay_.publisher()->GetDefaultUpstreamSession().GetIfAvailable(), + original_default_session); + + // The original session still exists, but the second session also receives + // a GOAWAY. + MoqtSessionInterface* second_default_session = + relay_.publisher()->GetDefaultUpstreamSession().GetIfAvailable(); + EXPECT_NE(second_default_session, nullptr); + EXPECT_NE(second_default_session, original_default_session); + std::move(second_default_session->callbacks().goaway_received_callback)( + "https://127.0.0.1:9991"); + RunUntilConnected(relay_, upstream_); + EXPECT_TRUE(relay_.publisher()->GetDefaultUpstreamSession().IsValid()); + EXPECT_EQ(relay_.publisher()->GetDefaultUpstreamSession().GetIfAvailable(), + relay_.client_session()); + EXPECT_NE(relay_.publisher()->GetDefaultUpstreamSession().GetIfAvailable(), + second_default_session); + + // The original session is now been destroyed, along with its client. The test + // might reuse original_session's address for the third session, + // unfortunately, so comparing a pointer to original_default_session is + // dangerous. + // second_default_session still exists, so this call doesn't segfault. + std::move(second_default_session->callbacks().session_terminated_callback)( + "test"); + // Third session is still connected. + EXPECT_TRUE(relay_.publisher()->GetDefaultUpstreamSession().IsValid()); + EXPECT_EQ(relay_.publisher()->GetDefaultUpstreamSession().GetIfAvailable(), + relay_.client_session()); + EXPECT_NE(relay_.publisher()->GetDefaultUpstreamSession().GetIfAvailable(), + second_default_session); +} +#endif + +// TODO(martinduke): Write tests for server sessions once there is related state +// that we can access. + +} // namespace test +} // namespace moqt