Merge CLIENT_SETUP and SERVER_SETUP into a single control message. Note that in reality, this is not supposed to be a control message, but rather a pair of unidirectional streams; this will be fixed in a subsequent CL. PiperOrigin-RevId: 925922509
diff --git a/quiche/quic/core/quic_types.h b/quiche/quic/core/quic_types.h index b1ce3ff..6b19dc0 100644 --- a/quiche/quic/core/quic_types.h +++ b/quiche/quic/core/quic_types.h
@@ -217,6 +217,11 @@ enum class Perspective : uint8_t { IS_SERVER, IS_CLIENT }; +constexpr Perspective FlipPerspective(Perspective perspective) { + return perspective == Perspective::IS_CLIENT ? Perspective::IS_SERVER + : Perspective::IS_CLIENT; +} + QUICHE_EXPORT std::string PerspectiveToString(Perspective perspective); QUICHE_EXPORT std::ostream& operator<<(std::ostream& os, const Perspective& perspective);
diff --git a/quiche/quic/moqt/moqt_bidi_stream_test.cc b/quiche/quic/moqt/moqt_bidi_stream_test.cc index a2cee5b..b2dc6db 100644 --- a/quiche/quic/moqt/moqt_bidi_stream_test.cc +++ b/quiche/quic/moqt/moqt_bidi_stream_test.cc
@@ -10,6 +10,7 @@ #include "absl/status/status.h" #include "absl/strings/string_view.h" #include "absl/types/span.h" +#include "quiche/quic/core/quic_types.h" #include "quiche/quic/moqt/moqt_error.h" #include "quiche/quic/moqt/moqt_framer.h" #include "quiche/quic/moqt/moqt_key_value_pair.h" @@ -52,11 +53,12 @@ class MoqtBidiStreamTest : public quiche::test::QuicheTest { public: MoqtBidiStreamTest() - : framer_(true), + : framer_(true, quic::Perspective::IS_CLIENT), stream_(std::make_unique<TestMoqtBidiStream>( &framer_, MoqtControlMessageParser(kDefaultMoqtVersion, - /*webtransport=*/true), + /*webtransport=*/true, + quic::Perspective::IS_CLIENT), deleted_callback_.AsStdFunction(), error_callback_.AsStdFunction())) {} @@ -136,7 +138,7 @@ TEST_F(MoqtBidiStreamTest, DispatchControlMessage) { webtransport::test::InMemoryStream stream(0); stream_->BindStream(&stream); - MoqtFramer framer(/*using_webtrans=*/true); + MoqtFramer framer(/*using_webtrans=*/true, quic::Perspective::IS_SERVER); stream.Receive(framer.SerializeRequestOk(MoqtRequestOk()).AsStringView()); stream_->OnCanRead(); EXPECT_EQ(stream_->ok_received(), 1u);
diff --git a/quiche/quic/moqt/moqt_framer.cc b/quiche/quic/moqt/moqt_framer.cc index 9b72e9f..17fd0e0 100644 --- a/quiche/quic/moqt/moqt_framer.cc +++ b/quiche/quic/moqt/moqt_framer.cc
@@ -495,25 +495,12 @@ WireOptional<WireBytes>(raw_payload)); } -quiche::QuicheBuffer MoqtFramer::SerializeClientSetup( - const MoqtClientSetup& message) { +quiche::QuicheBuffer MoqtFramer::SerializeSetup(const MoqtSetup& message) { KeyValuePairList parameters; - if (!FillAndValidateSetupParameters(MoqtMessageType::kClientSetup, - message.parameters, parameters)) { + if (!FillAndValidateSetupParameters(message.parameters, parameters)) { return quiche::QuicheBuffer(); } - return SerializeControlMessage(MoqtMessageType::kClientSetup, - WireKeyValuePairList(parameters)); -} - -quiche::QuicheBuffer MoqtFramer::SerializeServerSetup( - const MoqtServerSetup& message) { - KeyValuePairList parameters; - if (!FillAndValidateSetupParameters(MoqtMessageType::kServerSetup, - message.parameters, parameters)) { - return quiche::QuicheBuffer(); - } - return SerializeControlMessage(MoqtMessageType::kServerSetup, + return SerializeControlMessage(MoqtMessageType::kSetup, WireKeyValuePairList(parameters)); } @@ -727,13 +714,12 @@ } bool MoqtFramer::FillAndValidateSetupParameters( - MoqtMessageType message_type, const SetupParameters& parameters, - KeyValuePairList& out) { - if (SetupParametersAllowedByMessage(parameters, message_type, + const SetupParameters& parameters, KeyValuePairList& out) { + if (SetupParametersAllowedByMessage(parameters, perspective_, using_webtrans_) != MoqtError::kNoError) { QUICHE_BUG(QUICHE_BUG_invalid_setup_parameters) << "Invalid setup parameters for " - << MoqtMessageTypeToString(message_type); + << MoqtMessageTypeToString(MoqtMessageType::kSetup); return false; } out = parameters.ToKeyValuePairList();
diff --git a/quiche/quic/moqt/moqt_framer.h b/quiche/quic/moqt/moqt_framer.h index 499ba34..96a7251 100644 --- a/quiche/quic/moqt/moqt_framer.h +++ b/quiche/quic/moqt/moqt_framer.h
@@ -8,6 +8,7 @@ #include <optional> #include "absl/strings/string_view.h" +#include "quiche/quic/core/quic_types.h" #include "quiche/quic/moqt/moqt_key_value_pair.h" #include "quiche/quic/moqt/moqt_messages.h" #include "quiche/quic/moqt/moqt_object.h" @@ -27,7 +28,8 @@ // different streams. class QUICHE_EXPORT MoqtFramer { public: - MoqtFramer(bool using_webtrans) : using_webtrans_(using_webtrans) {} + MoqtFramer(bool using_webtrans, quic::Perspective perspective) + : using_webtrans_(using_webtrans), perspective_(perspective) {} // Serialize functions. Takes structured data and serializes it into a // QuicheBuffer for delivery to the stream. @@ -42,8 +44,7 @@ quiche::QuicheBuffer SerializeObjectDatagram(const MoqtObject& message, absl::string_view payload, MoqtPriority default_priority); - quiche::QuicheBuffer SerializeClientSetup(const MoqtClientSetup& message); - quiche::QuicheBuffer SerializeServerSetup(const MoqtServerSetup& message); + quiche::QuicheBuffer SerializeSetup(const MoqtSetup& message); quiche::QuicheBuffer SerializeRequestOk(const MoqtRequestOk& message); quiche::QuicheBuffer SerializeRequestError(const MoqtRequestError& message); // Returns an empty buffer if there is an illegal combination of locations. @@ -81,12 +82,12 @@ private: // Returns true if the parameters are valid for the message type. - bool FillAndValidateSetupParameters(MoqtMessageType message_type, - const SetupParameters& parameters, + bool FillAndValidateSetupParameters(const SetupParameters& parameters, KeyValuePairList& out); // Returns true if the metadata is internally consistent. static bool ValidateObjectMetadata(const MoqtObject& object); const bool using_webtrans_; + const quic::Perspective perspective_; }; } // namespace moqt
diff --git a/quiche/quic/moqt/moqt_framer_test.cc b/quiche/quic/moqt/moqt_framer_test.cc index 45403c7..79f74fb 100644 --- a/quiche/quic/moqt/moqt_framer_test.cc +++ b/quiche/quic/moqt/moqt_framer_test.cc
@@ -13,6 +13,7 @@ #include "absl/strings/str_cat.h" #include "absl/strings/string_view.h" +#include "quiche/quic/core/quic_types.h" #include "quiche/quic/moqt/moqt_key_value_pair.h" #include "quiche/quic/moqt/moqt_messages.h" #include "quiche/quic/moqt/moqt_names.h" @@ -28,10 +29,14 @@ namespace moqt::test { struct MoqtFramerTestParams { - MoqtFramerTestParams(MoqtMessageType message_type, bool uses_web_transport) - : message_type(message_type), uses_web_transport(uses_web_transport) {} + MoqtFramerTestParams(MoqtMessageType message_type, bool uses_web_transport, + quic::Perspective perspective) + : message_type(message_type), + uses_web_transport(uses_web_transport), + perspective(perspective) {} MoqtMessageType message_type; bool uses_web_transport; + quic::Perspective perspective; }; std::vector<MoqtFramerTestParams> GetMoqtFramerTestParams() { @@ -58,19 +63,22 @@ MoqtMessageType::kRequestsBlocked, MoqtMessageType::kPublish, MoqtMessageType::kObjectAck, - MoqtMessageType::kClientSetup, - MoqtMessageType::kServerSetup, + MoqtMessageType::kSetup, }; for (const MoqtMessageType message_type : message_types) { - if (message_type == MoqtMessageType::kClientSetup) { + if (message_type == MoqtMessageType::kSetup) { for (const bool uses_web_transport : {false, true}) { - params.push_back( - MoqtFramerTestParams(message_type, uses_web_transport)); + for (const quic::Perspective perspective : + {quic::Perspective::IS_CLIENT, quic::Perspective::IS_SERVER}) { + params.push_back(MoqtFramerTestParams( + message_type, uses_web_transport, perspective)); + } } } else { // All other types are processed the same for either perspective or // transport. - params.push_back(MoqtFramerTestParams(message_type, true)); + params.push_back(MoqtFramerTestParams(message_type, true, + quic::Perspective::IS_CLIENT)); } } return params; @@ -79,7 +87,8 @@ std::string ParamNameFormatter( const testing::TestParamInfo<MoqtFramerTestParams>& info) { return MoqtMessageTypeToString(info.param.message_type) + "_" + - (info.param.uses_web_transport ? "WebTransport" : "QUIC"); + (info.param.uses_web_transport ? "WebTransport" : "QUIC") + "_" + + quic::PerspectiveToString(info.param.perspective); } // If |change_in_object_id| is 0, it's the first object in the stream. @@ -117,10 +126,11 @@ MoqtFramerTest() : message_type_(GetParam().message_type), webtrans_(GetParam().uses_web_transport), - framer_(GetParam().uses_web_transport) {} + perspective_(GetParam().perspective), + framer_(GetParam().uses_web_transport, GetParam().perspective) {} std::unique_ptr<TestMessageBase> MakeMessage(MoqtMessageType message_type) { - return CreateTestMessage(message_type, webtrans_); + return CreateTestMessage(message_type, webtrans_, perspective_); } quiche::QuicheBuffer SerializeMessage( @@ -210,13 +220,9 @@ auto data = std::get<MoqtObjectAck>(structured_data); return framer_.SerializeObjectAck(data); } - case MoqtMessageType::kClientSetup: { - auto data = std::get<MoqtClientSetup>(structured_data); - return framer_.SerializeClientSetup(data); - } - case MoqtMessageType::kServerSetup: { - auto data = std::get<MoqtServerSetup>(structured_data); - return framer_.SerializeServerSetup(data); + case MoqtMessageType::kSetup: { + auto data = std::get<MoqtSetup>(structured_data); + return framer_.SerializeSetup(data); } default: // kObjectDatagram is a totally different code path. @@ -226,6 +232,7 @@ MoqtMessageType message_type_; bool webtrans_; + quic::Perspective perspective_; MoqtFramer framer_; }; @@ -245,7 +252,9 @@ class MoqtFramerSimpleTest : public quic::test::QuicTest { public: - MoqtFramerSimpleTest() : framer_(/*web_transport=*/true) {} + MoqtFramerSimpleTest() + : framer_(/*web_transport=*/true, + /*perspective=*/quic::Perspective::IS_SERVER) {} MoqtFramer framer_; // Obtain a pointer to an arbitrary offset in a serialized buffer.
diff --git a/quiche/quic/moqt/moqt_messages.cc b/quiche/quic/moqt/moqt_messages.cc index ed48af6..3512ecb 100644 --- a/quiche/quic/moqt/moqt_messages.cc +++ b/quiche/quic/moqt/moqt_messages.cc
@@ -8,6 +8,7 @@ #include <string> #include "absl/strings/str_cat.h" +#include "quiche/quic/core/quic_types.h" #include "quiche/quic/moqt/moqt_error.h" #include "quiche/quic/moqt/moqt_key_value_pair.h" #include "quiche/quic/moqt/moqt_types.h" @@ -24,10 +25,10 @@ } MoqtError SetupParametersAllowedByMessage(const SetupParameters& parameters, - MoqtMessageType message_type, + quic::Perspective sender_perspective, bool webtrans) { bool should_have_path_and_authority = - !webtrans && message_type == MoqtMessageType::kClientSetup; + !webtrans && sender_perspective == quic::Perspective::IS_CLIENT; if (should_have_path_and_authority != parameters.path.has_value()) { return MoqtError::kInvalidPath; } @@ -70,10 +71,8 @@ std::string MoqtMessageTypeToString(const MoqtMessageType message_type) { switch (message_type) { - case MoqtMessageType::kClientSetup: - return "CLIENT_SETUP"; - case MoqtMessageType::kServerSetup: - return "SERVER_SETUP"; + case MoqtMessageType::kSetup: + return "SETUP"; case MoqtMessageType::kSubscribe: return "SUBSCRIBE"; case MoqtMessageType::kSubscribeOk:
diff --git a/quiche/quic/moqt/moqt_messages.h b/quiche/quic/moqt/moqt_messages.h index ed13591..d8c0ac8 100644 --- a/quiche/quic/moqt/moqt_messages.h +++ b/quiche/quic/moqt/moqt_messages.h
@@ -15,6 +15,7 @@ #include <variant> #include "quiche/quic/core/quic_time.h" +#include "quiche/quic/core/quic_types.h" #include "quiche/quic/core/quic_versions.h" #include "quiche/quic/moqt/moqt_error.h" #include "quiche/quic/moqt/moqt_key_value_pair.h" @@ -214,8 +215,7 @@ kFetchOk = 0x18, kRequestsBlocked = 0x1a, kPublish = 0x1d, - kClientSetup = 0x20, - kServerSetup = 0x21, + kSetup = 0x2f00, // QUICHE-specific extensions. @@ -231,12 +231,7 @@ auto operator<=>(const SubgroupPriority&) const = default; }; -// TODO(martinduke): Collapse both Setup messages into SetupParameters. -struct QUICHE_EXPORT MoqtClientSetup { - SetupParameters parameters; -}; - -struct QUICHE_EXPORT MoqtServerSetup { +struct QUICHE_EXPORT MoqtSetup { SetupParameters parameters; }; @@ -538,7 +533,7 @@ // Returns false if the parameters cannot be in |message type|. MoqtError SetupParametersAllowedByMessage(const SetupParameters& parameters, - MoqtMessageType message_type, + quic::Perspective sender_perspective, bool webtrans); std::string MoqtMessageTypeToString(MoqtMessageType message_type);
diff --git a/quiche/quic/moqt/moqt_namespace_stream_test.cc b/quiche/quic/moqt/moqt_namespace_stream_test.cc index 21d3866..63b51dc 100644 --- a/quiche/quic/moqt/moqt_namespace_stream_test.cc +++ b/quiche/quic/moqt/moqt_namespace_stream_test.cc
@@ -44,13 +44,14 @@ const TrackNamespace kPrefix({"foo"}); MoqtControlMessageParser ControlMessageParser() { - return MoqtControlMessageParser(kDefaultMoqtVersion, true); + return MoqtControlMessageParser(kDefaultMoqtVersion, true, + quic::Perspective::IS_CLIENT); } class MoqtNamespaceSubscriberStreamTest : public quiche::test::QuicheTest { public: MoqtNamespaceSubscriberStreamTest() - : framer_(true), + : framer_(true, quic::Perspective::IS_CLIENT), stream_(&framer_, ControlMessageParser(), kRequestId, deleted_callback_.AsStdFunction(), error_callback_.AsStdFunction(), @@ -295,7 +296,7 @@ class MoqtNamespacePublisherStreamTest : public quiche::test::QuicheTest { public: MoqtNamespacePublisherStreamTest() - : framer_(false), + : framer_(false, quic::Perspective::IS_CLIENT), tree_(), application_callback_(mock_application_.AsStdFunction()), stream_(&framer_, ControlMessageParser(),
diff --git a/quiche/quic/moqt/moqt_parser.cc b/quiche/quic/moqt/moqt_parser.cc index 2ea2cbe..19e3377 100644 --- a/quiche/quic/moqt/moqt_parser.cc +++ b/quiche/quic/moqt/moqt_parser.cc
@@ -26,6 +26,7 @@ #include "quiche/http2/adapter/header_validator.h" #include "quiche/quic/core/quic_data_reader.h" #include "quiche/quic/core/quic_time.h" +#include "quiche/quic/core/quic_types.h" #include "quiche/quic/moqt/moqt_error.h" #include "quiche/quic/moqt/moqt_key_value_pair.h" #include "quiche/quic/moqt/moqt_messages.h" @@ -613,31 +614,19 @@ return message; } -absl::StatusOr<MoqtClientSetup> MoqtControlMessageParser::ProcessClientSetup( +absl::StatusOr<MoqtSetup> MoqtControlMessageParser::ProcessSetup( absl::string_view data) const { quic::QuicDataReader reader(data); - MoqtClientSetup setup; + MoqtSetup setup; KeyValuePairList parameters; QUICHE_RETURN_IF_ERROR(ParseKeyValuePairList(reader, parameters)); - QUICHE_RETURN_IF_ERROR(FillAndValidateSetupParameters( - parameters, setup.parameters, MoqtMessageType::kClientSetup)); + QUICHE_RETURN_IF_ERROR( + FillAndValidateSetupParameters(parameters, setup.parameters)); // TODO(martinduke): Validate construction of the PATH (Sec 8.3.2.1) QUICHE_RETURN_IF_ERROR(CheckForTrailingData(reader)); return setup; } -absl::StatusOr<MoqtServerSetup> MoqtControlMessageParser::ProcessServerSetup( - absl::string_view data) const { - quic::QuicDataReader reader(data); - MoqtServerSetup setup; - KeyValuePairList parameters; - QUICHE_RETURN_IF_ERROR(ParseKeyValuePairList(reader, parameters)); - QUICHE_RETURN_IF_ERROR(FillAndValidateSetupParameters( - parameters, setup.parameters, MoqtMessageType::kServerSetup)); - QUICHE_RETURN_IF_ERROR(CheckForTrailingData(reader)); - return setup; -} - absl::StatusOr<MoqtSubscribe> MoqtControlMessageParser::ProcessSubscribe( absl::string_view data) const { quic::QuicDataReader reader(data); @@ -1066,11 +1055,10 @@ } absl::Status MoqtControlMessageParser::FillAndValidateSetupParameters( - const KeyValuePairList& in, SetupParameters& out, - MoqtMessageType message_type) const { + const KeyValuePairList& in, SetupParameters& out) const { QUICHE_RETURN_IF_ERROR(out.FromKeyValuePairList(in)); - MoqtError error = - SetupParametersAllowedByMessage(out, message_type, uses_web_transport_); + MoqtError error = SetupParametersAllowedByMessage( + out, FlipPerspective(perspective_), uses_web_transport_); if (error != MoqtError::kNoError) { return MoqtErrorStatusWithCode("Setup parameter parsing error", error); }
diff --git a/quiche/quic/moqt/moqt_parser.h b/quiche/quic/moqt/moqt_parser.h index 65ec8db..f36415c 100644 --- a/quiche/quic/moqt/moqt_parser.h +++ b/quiche/quic/moqt/moqt_parser.h
@@ -21,6 +21,7 @@ #include "absl/strings/string_view.h" #include "absl/types/span.h" #include "quiche/quic/core/quic_data_reader.h" +#include "quiche/quic/core/quic_types.h" #include "quiche/quic/moqt/moqt_error.h" #include "quiche/quic/moqt/moqt_key_value_pair.h" #include "quiche/quic/moqt/moqt_messages.h" @@ -114,14 +115,12 @@ public: // `moqt_version` is not currently used, as we only support one version. MoqtControlMessageParser(absl::string_view /*moqt_version*/, - bool uses_web_transport) - : uses_web_transport_(uses_web_transport) {} + bool uses_web_transport, + quic::Perspective perspective) + : uses_web_transport_(uses_web_transport), perspective_(perspective) {} // Parsers for individual messages. - absl::StatusOr<MoqtClientSetup> ProcessClientSetup( - absl::string_view data) const; - absl::StatusOr<MoqtServerSetup> ProcessServerSetup( - absl::string_view data) const; + absl::StatusOr<MoqtSetup> ProcessSetup(absl::string_view data) const; absl::StatusOr<MoqtRequestOk> ProcessRequestOk(absl::string_view data) const; absl::StatusOr<MoqtRequestError> ProcessRequestError( absl::string_view data) const; @@ -175,10 +174,8 @@ return callback(*std::move(parsed_message)); }; switch (message.type) { - case MoqtMessageType::kClientSetup: - return parse(&MoqtControlMessageParser::ProcessClientSetup); - case MoqtMessageType::kServerSetup: - return parse(&MoqtControlMessageParser::ProcessServerSetup); + case MoqtMessageType::kSetup: + return parse(&MoqtControlMessageParser::ProcessSetup); case MoqtMessageType::kRequestOk: return parse(&MoqtControlMessageParser::ProcessRequestOk); case MoqtMessageType::kRequestError: @@ -239,9 +236,8 @@ // large. Sets a ParseError if the name is malformed. absl::Status ReadFullTrackName(quic::QuicDataReader& reader, FullTrackName& full_track_name) const; - absl::Status FillAndValidateSetupParameters( - const KeyValuePairList& in, SetupParameters& out, - MoqtMessageType message_type) const; + absl::Status FillAndValidateSetupParameters(const KeyValuePairList& in, + SetupParameters& out) const; // |reader| points to the beginning of a KeyValuePairList. Returns false if // there is any sort of error. (The function calls ParseError(), so the // caller has no need to do so.) @@ -249,6 +245,7 @@ MessageParameters& out) const; bool uses_web_transport_; + const quic::Perspective perspective_; }; // Parses an MoQT datagram. Returns the payload bytes, or std::nullopt on error.
diff --git a/quiche/quic/moqt/moqt_parser_fuzz_test.cc b/quiche/quic/moqt/moqt_parser_fuzz_test.cc index ee73052..5da2857 100644 --- a/quiche/quic/moqt/moqt_parser_fuzz_test.cc +++ b/quiche/quic/moqt/moqt_parser_fuzz_test.cc
@@ -19,13 +19,14 @@ namespace { void MoqtControlParserNeverCrashes(bool is_data_stream, bool uses_web_transport, + quic::Perspective perspective, absl::string_view stream_data, bool fin) { webtransport::test::InMemoryStream stream(/*stream_id=*/0); MoqtParserTestVisitor visitor(/*enable_logging=*/false); MoqtControlStreamParser control_stream_parser(&stream); - MoqtControlMessageParser control_message_parser(kDefaultMoqtVersion, - uses_web_transport); + MoqtControlMessageParser control_message_parser( + kDefaultMoqtVersion, uses_web_transport, perspective); MoqtDataParser data_parser(&stream, &visitor); if (is_data_stream) { @@ -47,6 +48,8 @@ FUZZ_TEST(MoqtParserTest, MoqtControlParserNeverCrashes) .WithDomains(fuzztest::Arbitrary<bool>(), fuzztest::Arbitrary<bool>(), + fuzztest::ElementOf({quic::Perspective::IS_CLIENT, + quic::Perspective::IS_SERVER}), fuzztest::Arbitrary<std::string>(), fuzztest::Arbitrary<bool>()); @@ -62,6 +65,7 @@ MoqtControlParserNeverCrashes( /*is_data_stream=*/false, /*uses_web_transport=*/false, + /*perspective=*/quic::Perspective::IS_SERVER, /*stream_data=*/std::string(kStreamData.begin(), kStreamData.end()), /*fin=*/true); }
diff --git a/quiche/quic/moqt/moqt_parser_test.cc b/quiche/quic/moqt/moqt_parser_test.cc index d73a1dd..e3c09f0 100644 --- a/quiche/quic/moqt/moqt_parser_test.cc +++ b/quiche/quic/moqt/moqt_parser_test.cc
@@ -21,6 +21,7 @@ #include "absl/strings/string_view.h" #include "quiche/quic/core/quic_data_writer.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_key_value_pair.h" #include "quiche/quic/moqt/moqt_messages.h" @@ -59,8 +60,6 @@ MoqtMessageType::kNamespace, MoqtMessageType::kNamespaceDone, MoqtMessageType::kPublishNamespaceCancel, - MoqtMessageType::kClientSetup, - MoqtMessageType::kServerSetup, MoqtMessageType::kGoAway, MoqtMessageType::kSubscribeNamespace, MoqtMessageType::kMaxRequestId, @@ -70,6 +69,7 @@ MoqtMessageType::kRequestsBlocked, MoqtMessageType::kPublish, MoqtMessageType::kObjectAck, + MoqtMessageType::kSetup, }; using GeneralizedMessageType = @@ -77,27 +77,39 @@ } // namespace struct MoqtParserTestParams { - MoqtParserTestParams(MoqtMessageType message_type, bool uses_web_transport) - : message_type(message_type), uses_web_transport(uses_web_transport) {} + MoqtParserTestParams( + MoqtMessageType message_type, bool uses_web_transport, + quic::Perspective perspective = quic::Perspective::IS_SERVER) + : message_type(message_type), + uses_web_transport(uses_web_transport), + perspective(perspective) {} explicit MoqtParserTestParams(MoqtDataStreamType message_type) - : message_type(message_type), uses_web_transport(true) {} + : message_type(message_type), + uses_web_transport(true), + perspective(quic::Perspective::IS_SERVER) {} + GeneralizedMessageType message_type; bool uses_web_transport; + quic::Perspective perspective; }; std::vector<MoqtParserTestParams> GetMoqtParserTestParams() { std::vector<MoqtParserTestParams> params; for (MoqtMessageType message_type : kMessageTypes) { - if (message_type == MoqtMessageType::kClientSetup) { + if (message_type == MoqtMessageType::kSetup) { for (const bool uses_web_transport : {false, true}) { - params.push_back( - MoqtParserTestParams(message_type, uses_web_transport)); + for (const quic::Perspective perspective : + {quic::Perspective::IS_CLIENT, quic::Perspective::IS_SERVER}) { + params.push_back(MoqtParserTestParams( + message_type, uses_web_transport, perspective)); + } } } else { // All other types are processed the same for either perspective or // transport. - params.push_back(MoqtParserTestParams(message_type, true)); + params.push_back(MoqtParserTestParams(message_type, true, + quic::Perspective::IS_SERVER)); } } for (MoqtDataStreamType type : AllMoqtDataStreamTypes()) { @@ -116,7 +128,8 @@ const testing::TestParamInfo<MoqtParserTestParams>& info) { return std::visit([](auto x) { return TypeFormatter(x); }, info.param.message_type) + - "_" + (info.param.uses_web_transport ? "WebTransport" : "QUIC"); + "_" + (info.param.uses_web_transport ? "WebTransport" : "QUIC") + "_" + + quic::PerspectiveToString(info.param.perspective); } std::optional<MoqtError> ExtractMoqtErrorForStatus(const absl::Status& status) { @@ -132,9 +145,10 @@ MoqtParserTest() : message_type_(GetParam().message_type), webtrans_(GetParam().uses_web_transport), + perspective_(GetParam().perspective), control_stream_(/*stream_id=*/0), control_parser_(&control_stream_), - message_parser_(kDefaultMoqtVersion, webtrans_), + message_parser_(kDefaultMoqtVersion, webtrans_, perspective_), data_stream_(/*stream_id=*/0), data_parser_(&data_stream_, &data_visitor_) { // The default object has priority 0x07, so setting this will let the @@ -151,7 +165,7 @@ return CreateTestDataStream(std::get<MoqtDataStreamType>(message_type_)); } return CreateTestMessage(std::get<MoqtMessageType>(message_type_), - webtrans_); + webtrans_, FlipPerspective(perspective_)); } void ProcessData(absl::string_view data, bool fin) { @@ -212,6 +226,7 @@ GeneralizedMessageType message_type_; bool webtrans_; + quic::Perspective perspective_; webtransport::test::InMemoryStream control_stream_; MoqtControlStreamParser control_parser_; MoqtControlMessageParser message_parser_; @@ -412,12 +427,14 @@ absl::StatusOr<std::vector<AnyMoqtControlMessage>> ParseAllMessages( absl::string_view data, absl::string_view moqt_version = kDefaultMoqtVersion, - bool uses_web_transport = true) { + bool uses_web_transport = true, + quic::Perspective perspective = quic::Perspective::IS_SERVER) { webtransport::test::InMemoryStream stream(/*stream_id=*/0); stream.Receive(data, /*fin=*/true); MoqtControlStreamParser stream_parser(&stream); stream_parser.set_allow_fin(true); - MoqtControlMessageParser message_parser(moqt_version, uses_web_transport); + MoqtControlMessageParser message_parser(moqt_version, uses_web_transport, + perspective); std::vector<AnyMoqtControlMessage> result; while (!stream_parser.fin_read()) { absl::StatusOr<MoqtRawControlMessage> raw_message = @@ -591,7 +608,7 @@ TEST_F(MoqtMessageSpecificTest, ClientSetupMaxRequestIdAppearsTwice) { char setup[] = { - 0x20, 0x00, 0x0a, + 0xaf, 0x00, 0x00, 0x0a, 0x03, // 3 params 0x01, 0x03, 0x66, 0x6f, 0x6f, // path = "foo" 0x01, 0x32, // max_request_id = 50 @@ -605,25 +622,27 @@ TEST_F(MoqtMessageSpecificTest, ServerSetupAuthorizationTokenTagRegister) { char setup[] = { - 0x21, 0x00, 0x0b, + 0xaf, 0x00, 0x00, 0x0b, 0x02, // 2 params 0x02, 0x32, // max_request_id = 50 0x01, 0x06, 0x01, 0x10, 0x00, 0x62, 0x61, 0x72, // REGISTER 0x01 }; absl::StatusOr<std::vector<AnyMoqtControlMessage>> parsed = - ParseAllMessages(absl::string_view(setup, sizeof(setup))); + ParseAllMessages(absl::string_view(setup, sizeof(setup)), + kDefaultMoqtVersion, true, quic::Perspective::IS_CLIENT); // No error even though the registration exceeds the max cache size of 0. QUICHE_EXPECT_OK(parsed.status()); } TEST_F(MoqtMessageSpecificTest, SetupPathFromServer) { char setup[] = { - 0x21, 0x00, 0x06, + 0xaf, 0x00, 0x00, 0x06, 0x01, // 1 param 0x01, 0x03, 0x66, 0x6f, 0x6f, // path = "foo" }; absl::StatusOr<std::vector<AnyMoqtControlMessage>> parsed = - ParseAllMessages(absl::string_view(setup, sizeof(setup))); + ParseAllMessages(absl::string_view(setup, sizeof(setup)), + kDefaultMoqtVersion, true, quic::Perspective::IS_CLIENT); ASSERT_THAT(parsed.status(), StatusIs(absl::StatusCode::kInvalidArgument, HasSubstr("Setup parameter parsing error"))); @@ -633,19 +652,20 @@ TEST_F(MoqtMessageSpecificTest, SetupAuthorityFromServer) { char setup[] = { - 0x21, 0x00, 0x06, + 0xaf, 0x00, 0x00, 0x06, 0x01, // 1 param 0x05, 0x03, 0x66, 0x6f, 0x6f, // authority = "foo" }; absl::StatusOr<std::vector<AnyMoqtControlMessage>> parsed = - ParseAllMessages(absl::string_view(setup, sizeof(setup))); + ParseAllMessages(absl::string_view(setup, sizeof(setup)), + kDefaultMoqtVersion, true, quic::Perspective::IS_CLIENT); EXPECT_EQ(ExtractMoqtErrorForStatus(parsed.status()), MoqtError::kInvalidAuthority); } TEST_F(MoqtMessageSpecificTest, SetupPathAppearsTwice) { char setup[] = { - 0x20, 0x00, 0x0b, + 0xaf, 0x00, 0x00, 0x0b, 0x02, // 2 params 0x01, 0x03, 0x66, 0x6f, 0x6f, // path = "foo" 0x00, 0x03, 0x66, 0x6f, 0x6f, // path = "foo" @@ -658,7 +678,7 @@ TEST_F(MoqtMessageSpecificTest, SetupPathOverWebtrans) { char setup[] = { - 0x20, 0x00, 0x06, + 0xaf, 0x00, 0x00, 0x06, 0x01, // 1 param 0x01, 0x03, 0x66, 0x6f, 0x6f, // path = "foo" }; @@ -670,7 +690,7 @@ TEST_F(MoqtMessageSpecificTest, SetupAuthorityOverWebtrans) { char setup[] = { - 0x20, 0x00, 0x06, + 0xaf, 0x00, 0x00, 0x06, 0x01, // 1 param 0x05, 0x03, 0x66, 0x6f, 0x6f, // authority = "foo" }; @@ -682,9 +702,7 @@ TEST_F(MoqtMessageSpecificTest, SetupPathMissing) { char setup[] = { - 0x20, - 0x00, - 0x01, + 0xaf, 0x00, 0x00, 0x01, 0x00, // no param }; absl::StatusOr<std::vector<AnyMoqtControlMessage>> parsed = ParseAllMessages( @@ -695,19 +713,20 @@ TEST_F(MoqtMessageSpecificTest, ServerSetupMaxRequestIdAppearsTwice) { char setup[] = { - 0x21, 0x00, 0x05, 0x02, // 2 params - 0x02, 0x32, // max_request_id = 50 - 0x00, 0x32, // max_request_id = 50 + 0xaf, 0x00, 0x00, 0x05, 0x02, // 2 params + 0x02, 0x32, // max_request_id = 50 + 0x00, 0x32, // max_request_id = 50 }; absl::StatusOr<std::vector<AnyMoqtControlMessage>> parsed = ParseAllMessages( - absl::string_view(setup, sizeof(setup)), kDefaultMoqtVersion, kRawQuic); + absl::string_view(setup, sizeof(setup)), kDefaultMoqtVersion, kRawQuic, + quic::Perspective::IS_CLIENT); EXPECT_EQ(ExtractMoqtErrorForStatus(parsed.status()), MoqtError::kProtocolViolation); } TEST_F(MoqtMessageSpecificTest, ClientSetupMalformedPath) { char setup[] = { - 0x20, 0x00, 0x06, + 0xaf, 0x00, 0x00, 0x06, 0x01, // 1 param 0x01, 0x03, 0x66, 0x5c, 0x6f, // path = "f\o" }; @@ -719,7 +738,7 @@ TEST_F(MoqtMessageSpecificTest, ClientSetupMalformedAuthority) { char setup[] = { - 0x20, 0x00, 0x0b, + 0xaf, 0x00, 0x00, 0x0b, 0x02, // 2 params 0x01, 0x03, 0x66, 0x6f, 0x6f, // path = "foo" 0x04, 0x03, 0x66, 0x5c, 0x6f, // authority = "f\o" @@ -732,16 +751,17 @@ TEST_F(MoqtMessageSpecificTest, ServerSetupUnknownParameterIsOk) { char setup[] = { - 0x21, 0x00, 0x0b, + 0xaf, 0x00, 0x00, 0x0b, 0x02, // 2 params 0x1f, 0x03, 0x62, 0x61, 0x72, // 0x1f = "bar" 0x00, 0x03, 0x62, 0x61, 0x72, // 0x1f = "bar" }; absl::StatusOr<std::vector<AnyMoqtControlMessage>> parsed = ParseAllMessages( - absl::string_view(setup, sizeof(setup)), kDefaultMoqtVersion, kRawQuic); + absl::string_view(setup, sizeof(setup)), kDefaultMoqtVersion, kRawQuic, + quic::Perspective::IS_CLIENT); ASSERT_TRUE(parsed.ok()); ASSERT_EQ(parsed->size(), 1); - MoqtServerSetup message = std::get<MoqtServerSetup>((*parsed)[0]); + MoqtSetup message = std::get<MoqtSetup>((*parsed)[0]); EXPECT_EQ(message.parameters, SetupParameters()); } @@ -1088,7 +1108,7 @@ TEST_F(MoqtMessageSpecificTest, Setup2KB) { char big_message[2 * kMaxMessageHeaderSize]; quic::QuicDataWriter writer(sizeof(big_message), big_message); - writer.WriteMoqVarInt(static_cast<uint64_t>(MoqtMessageType::kServerSetup)); + writer.WriteMoqVarInt(static_cast<uint64_t>(MoqtMessageType::kSetup)); writer.WriteUInt16(8 + kMaxMessageHeaderSize); writer.WriteMoqVarInt(0x1); // version writer.WriteMoqVarInt(0x1); // num_params @@ -1097,7 +1117,8 @@ writer.WriteRepeatedByte(0x04, kMaxMessageHeaderSize); // Send incomplete message absl::StatusOr<std::vector<AnyMoqtControlMessage>> parsed = - ParseAllMessages(absl::string_view(big_message, writer.length())); + ParseAllMessages(absl::string_view(big_message, writer.length()), + kDefaultMoqtVersion, true, quic::Perspective::IS_CLIENT); EXPECT_THAT( parsed.status(), StatusIs(absl::StatusCode::kInvalidArgument, @@ -1553,8 +1574,8 @@ TEST_F(MoqtMessageSpecificTest, ParseKeyValuePairListIntegerOverflow) { char setup[] = { - 0x20, 0x00, 0x0c, // kClientSetup, length = 12 - 0x02, // num_params + 0xaf, 0x00, 0x00, 0x0c, // kSetup, length = 12 + 0x02, // num_params 0xff, 0xff, 0xff, 0xff, 0xff, 0xff, 0xff, 0xff, 0xff, // type_diff = max 0x00, // string length = 0 0x01, // type_diff = 1 (overflows)
diff --git a/quiche/quic/moqt/moqt_session.cc b/quiche/quic/moqt/moqt_session.cc index 076d528..3e34f70 100644 --- a/quiche/quic/moqt/moqt_session.cc +++ b/quiche/quic/moqt/moqt_session.cc
@@ -85,7 +85,7 @@ : session_(session), parameters_(parameters), callbacks_(std::move(callbacks)), - framer_(parameters.using_webtrans), + framer_(parameters.using_webtrans, parameters.perspective), publisher_(DefaultPublisher::GetInstance()), local_max_request_id_(parameters.max_request_id), alarm_factory_(std::move(alarm_factory)), @@ -142,9 +142,9 @@ control_stream->BindStream(stream); trace_recorder_.RecordControlStreamCreated(stream->GetStreamId()); stream->SetVisitor(std::move(control_stream)); - MoqtClientSetup setup; + MoqtSetup setup; parameters_.ToSetupParameters(setup.parameters); - SendControlMessage(framer_.SerializeClientSetup(setup)); + SendControlMessage(framer_.SerializeSetup(setup)); QUIC_DLOG(INFO) << ENDPOINT << "Send CLIENT_SETUP"; } @@ -836,7 +836,7 @@ return; } switch (*message_type) { - case MoqtMessageType::kClientSetup: { + case MoqtMessageType::kSetup: { if (session_->control_stream_.GetIfAvailable() != nullptr) { session_->Error(MoqtError::kProtocolViolation, "Multiple control streams"); @@ -889,40 +889,33 @@ } absl::Status MoqtSession::ControlStream::OnControlMessage( - const MoqtClientSetup& message) { - if (session_->perspective() == Perspective::IS_CLIENT) { - return absl::InvalidArgumentError("Received CLIENT_SETUP from server"); + const MoqtSetup& message) { + if (session_->parameters_.perspective == Perspective::IS_SERVER) { + session_->peer_supports_object_ack_ = + message.parameters.support_object_acks.value_or( + kDefaultSupportObjectAcks); + session_->peer_max_request_id_ = + message.parameters.max_request_id.value_or(kDefaultMaxRequestId); + QUICHE_DLOG(INFO) << "Received CLIENT_SETUP"; + MoqtSetup response; + session_->parameters_.ToSetupParameters(response.parameters); + QUICHE_RETURN_IF_ERROR( + SendOrBufferMessage(session_->framer_.SerializeSetup(response))); + QUICHE_DLOG(INFO) << "Sent SERVER_SETUP"; + // TODO: handle path. + std::move(session_->callbacks_.session_established_callback)(); + return absl::OkStatus(); + } else { + session_->peer_supports_object_ack_ = + message.parameters.support_object_acks.value_or( + kDefaultSupportObjectAcks); + QUIC_DLOG(INFO) << ENDPOINT << "Received the SETUP message"; + // TODO: handle path. + session_->peer_max_request_id_ = + message.parameters.max_request_id.value_or(kDefaultMaxRequestId); + std::move(session_->callbacks_.session_established_callback)(); + return absl::OkStatus(); } - session_->peer_supports_object_ack_ = - message.parameters.support_object_acks.value_or( - kDefaultSupportObjectAcks); - session_->peer_max_request_id_ = - message.parameters.max_request_id.value_or(kDefaultMaxRequestId); - QUICHE_DLOG(INFO) << "Received CLIENT_SETUP"; - MoqtServerSetup response; - session_->parameters_.ToSetupParameters(response.parameters); - QUICHE_RETURN_IF_ERROR( - SendOrBufferMessage(session_->framer_.SerializeServerSetup(response))); - QUICHE_DLOG(INFO) << "Sent SERVER_SETUP"; - // TODO: handle path. - std::move(session_->callbacks_.session_established_callback)(); - return absl::OkStatus(); -} - -absl::Status MoqtSession::ControlStream::OnControlMessage( - const MoqtServerSetup& message) { - if (perspective() == Perspective::IS_SERVER) { - return absl::InvalidArgumentError("Received SERVER_SETUP from client"); - } - session_->peer_supports_object_ack_ = - message.parameters.support_object_acks.value_or( - kDefaultSupportObjectAcks); - QUIC_DLOG(INFO) << ENDPOINT << "Received the SETUP message"; - // TODO: handle path. - session_->peer_max_request_id_ = - message.parameters.max_request_id.value_or(kDefaultMaxRequestId); - std::move(session_->callbacks_.session_established_callback)(); - return absl::OkStatus(); } absl::Status MoqtSession::ControlStream::OnControlMessage(
diff --git a/quiche/quic/moqt/moqt_session.h b/quiche/quic/moqt/moqt_session.h index fe3da70..0e42ce2 100644 --- a/quiche/quic/moqt/moqt_session.h +++ b/quiche/quic/moqt/moqt_session.h
@@ -258,8 +258,7 @@ const MoqtRawControlMessage& message) override; // MoqtControlParserVisitor implementation. - absl::Status OnControlMessage(const MoqtClientSetup& message); - absl::Status OnControlMessage(const MoqtServerSetup& message); + absl::Status OnControlMessage(const MoqtSetup& message); absl::Status OnControlMessage(const MoqtRequestOk& message); absl::Status OnControlMessage(const MoqtRequestError& message); absl::Status OnControlMessage(const MoqtSubscribe& message); @@ -455,7 +454,8 @@ MoqtControlMessageParser ControlMessageParser() const { return MoqtControlMessageParser(parameters_.version, - parameters_.using_webtrans); + parameters_.using_webtrans, + parameters_.perspective); } bool is_closing_ = false;
diff --git a/quiche/quic/moqt/moqt_session_test.cc b/quiche/quic/moqt/moqt_session_test.cc index 8d85d8f..6ddd7f8 100644 --- a/quiche/quic/moqt/moqt_session_test.cc +++ b/quiche/quic/moqt/moqt_session_test.cc
@@ -200,7 +200,7 @@ webtransport::test::MockStream* stream, std::unique_ptr<webtransport::StreamVisitor>& visitor, MockSubscribeRemoteTrackVisitor* track_visitor) { - MoqtFramer framer(true); + MoqtFramer framer(true, quic::Perspective::IS_SERVER); std::optional<PublishedObjectMetadata> previous_object; if (visitor != nullptr) { previous_object = PublishedObjectMetadata(); @@ -285,7 +285,7 @@ EXPECT_CALL(mock_stream_, GetStreamId()) .WillRepeatedly(Return(webtransport::StreamId(4))); EXPECT_CALL(mock_stream_, - Writev(ControlMessageOfType(MoqtMessageType::kClientSetup), _)); + Writev(ControlMessageOfType(MoqtMessageType::kSetup), _)); session_.OnSessionReady(); // Receive SERVER_SETUP @@ -293,7 +293,7 @@ MoqtSessionPeer::FetchParserVisitorFromWebtransportStreamVisitor( std::move(visitor)); // Handle the server setup - MoqtServerSetup setup; // No fields are set. + MoqtSetup setup; // No fields are set. EXPECT_CALL(session_callbacks_.session_established_callback, Call()).Times(1); stream_input->ReceiveMessage(setup); } @@ -335,10 +335,11 @@ session_callbacks_.AsSessionCallbacks()); // Load a CLIENT_SETUP message into an in-memory stream. webtransport::test::InMemoryStreamWithWriteBuffer in_memory_stream(0); - MoqtFramer framer(session_parameters.using_webtrans); - MoqtClientSetup setup; + MoqtFramer framer(session_parameters.using_webtrans, + quic::Perspective::IS_CLIENT); + MoqtSetup setup; session_parameters.ToSetupParameters(setup.parameters); - quiche::QuicheBuffer buffer = framer.SerializeClientSetup(setup); + quiche::QuicheBuffer buffer = framer.SerializeSetup(setup); in_memory_stream.Receive(absl::string_view(buffer.data(), buffer.size()), /*fin=*/false); @@ -348,7 +349,7 @@ EXPECT_CALL(session_callbacks_.session_established_callback, Call()); server_session.OnIncomingBidirectionalStreamAvailable(); EXPECT_EQ(PeekControlMessageType(in_memory_stream.write_buffer()), - MoqtMessageType::kServerSetup); + MoqtMessageType::kSetup); EXPECT_NE(MoqtSessionPeer::GetControlStream(&server_session), nullptr); } @@ -1895,7 +1896,7 @@ /*subgroup=*/0, /*payload_length=*/3, }; - MoqtFramer framer(true); + MoqtFramer framer(true, quic::Perspective::IS_SERVER); std::optional<PublishedObjectMetadata> metadata; quiche::QuicheBuffer header = framer.SerializeObjectHeader( object, MoqtDataStreamType::Fetch(), metadata); @@ -1917,7 +1918,7 @@ "foo"); auto bidi_stream = std::make_unique<webtransport::test::InMemoryStreamWithWriteBuffer>(4); - MoqtFramer framer(true); + MoqtFramer framer(true, quic::Perspective::IS_SERVER); MoqtSubscribeNamespace subscribe_namespace = { /*request_id=*/1, prefix, SubscribeNamespaceOption::kBoth, parameters}; bidi_stream->Receive( @@ -1970,7 +1971,7 @@ parameters.authorization_tokens.emplace_back(AuthTokenType::kOutOfBand, "foo"); webtransport::test::InMemoryStreamWithWriteBuffer bidi_stream(4); - MoqtFramer framer(true); + MoqtFramer framer(true, quic::Perspective::IS_SERVER); MoqtSubscribeNamespace subscribe_namespace = { /*request_id=*/1, prefix, SubscribeNamespaceOption::kBoth, parameters}; bidi_stream.Receive( @@ -2001,7 +2002,7 @@ "foo"); webtransport::test::InMemoryStreamWithWriteBuffer bidi_stream1(4), bidi_stream2(8); - MoqtFramer framer(true); + MoqtFramer framer(true, quic::Perspective::IS_SERVER); MoqtSubscribeNamespace subscribe_namespace = { /*request_id=*/1, foo, SubscribeNamespaceOption::kBoth, parameters}; bidi_stream1.Receive( @@ -2125,11 +2126,11 @@ /*subgroup=*/0, /*payload_length=*/3, }; - MoqtFramer framer_(true); + MoqtFramer framer(true, quic::Perspective::IS_SERVER); std::optional<PublishedObjectMetadata> metadata; for (int i = 0; i < 4; ++i) { object.object_id = i; - headers.push(framer_.SerializeObjectHeader( + headers.push(framer.SerializeObjectHeader( object, MoqtDataStreamType::Fetch(), metadata)); metadata = PublishedObjectMetadata(); metadata->location.object = i; // only object ID matters. @@ -2198,11 +2199,11 @@ /*subgroup=*/0, /*payload_length=*/3, }; - MoqtFramer framer_(true); + MoqtFramer framer(true, quic::Perspective::IS_SERVER); std::optional<PublishedObjectMetadata> metadata; for (int i = 0; i < 4; ++i) { object.object_id = i; - headers.push(framer_.SerializeObjectHeader( + headers.push(framer.SerializeObjectHeader( object, MoqtDataStreamType::Fetch(), metadata)); metadata = PublishedObjectMetadata(); metadata->location.object = i; // only object ID matters. @@ -2338,13 +2339,15 @@ fetch.request_id = 5; stream_input->ReceiveMessage(fetch); - MoqtFramer framer(true); + MoqtFramer framer(true, quic::Perspective::IS_CLIENT); SessionNamespaceTree tree; MoqtIncomingSubscribeNamespaceCallback callback = DefaultIncomingSubscribeNamespaceCallback; MoqtNamespacePublisherStream namespace_stream( - &framer, MoqtControlMessageParser(kDefaultMoqtVersion, true), nullptr, - &tree, callback); + &framer, + MoqtControlMessageParser(kDefaultMoqtVersion, true, + quic::Perspective::IS_CLIENT), + nullptr, &tree, callback); namespace_stream.BindStream(&mock_stream_); EXPECT_CALL(mock_stream_, Writev(ControlMessageOfType(MoqtMessageType::kRequestError), _)); @@ -2691,7 +2694,7 @@ EXPECT_CALL(mock_session_, CloseSession); EXPECT_CALL(session_callbacks_.session_terminated_callback, Call); control_stream->ReceiveMessage( - MoqtClientSetup(SetupParameters("/", "example.com", 0))); + MoqtSetup(SetupParameters("/", "example.com", 0))); } TEST_F(MoqtSessionTest, NamespaceNotAllowedOnControlStream) {
diff --git a/quiche/quic/moqt/moqt_subscription_test.cc b/quiche/quic/moqt/moqt_subscription_test.cc index 89045f3..227b8cb 100644 --- a/quiche/quic/moqt/moqt_subscription_test.cc +++ b/quiche/quic/moqt/moqt_subscription_test.cc
@@ -221,8 +221,9 @@ static constexpr uint64_t kTrackAlias = 10; static constexpr uint64_t kRequestId = 1; - MoqtFramer framer_{true}; - MoqtControlMessageParser message_parser_{kDefaultMoqtVersion, true}; + MoqtFramer framer_{true, quic::Perspective::IS_CLIENT}; + MoqtControlMessageParser message_parser_{kDefaultMoqtVersion, true, + quic::Perspective::IS_CLIENT}; webtransport::test::MockSession webtrans_; StrictMock<webtransport::test::MockStream> mock_bidi_stream_; webtransport::test::MockStream mock_uni_stream_;
diff --git a/quiche/quic/moqt/moqt_uni_stream_test.cc b/quiche/quic/moqt/moqt_uni_stream_test.cc index fcfea13..67c08df 100644 --- a/quiche/quic/moqt/moqt_uni_stream_test.cc +++ b/quiche/quic/moqt/moqt_uni_stream_test.cc
@@ -120,7 +120,7 @@ EXPECT_CALL(visitor_, alarm_factory()).WillOnce(Return(&alarm_factory_)); } - MoqtFramer framer_{true}; + MoqtFramer framer_{true, quic::Perspective::IS_CLIENT}; StrictMock<webtransport::test::MockStream> mock_stream_; DataStreamIndex index_; std::shared_ptr<StrictMock<MockTrackPublisher>> track_publisher_; @@ -349,7 +349,7 @@ } protected: - MoqtFramer framer_{true}; + MoqtFramer framer_{true, quic::Perspective::IS_CLIENT}; StrictMock<webtransport::test::MockStream> mock_stream_; std::unique_ptr<StrictMock<MockFetchTask>> task_; MockFetchTask* task_ptr_;
diff --git a/quiche/quic/moqt/test_tools/moqt_framer_utils.cc b/quiche/quic/moqt/test_tools/moqt_framer_utils.cc index 87dc14b..7214426 100644 --- a/quiche/quic/moqt/test_tools/moqt_framer_utils.cc +++ b/quiche/quic/moqt/test_tools/moqt_framer_utils.cc
@@ -16,11 +16,8 @@ namespace { struct FramingVisitor { - quiche::QuicheBuffer operator()(const MoqtClientSetup& message) { - return framer.SerializeClientSetup(message); - } - quiche::QuicheBuffer operator()(const MoqtServerSetup& message) { - return framer.SerializeServerSetup(message); + quiche::QuicheBuffer operator()(const MoqtSetup& message) { + return framer.SerializeSetup(message); } quiche::QuicheBuffer operator()(const MoqtRequestOk& message) { return framer.SerializeRequestOk(message); @@ -97,7 +94,14 @@ std::string SerializeGenericMessage(const AnyMoqtControlMessage& frame, bool use_webtrans) { - MoqtFramer framer(use_webtrans); + quic::Perspective perspective = quic::Perspective::IS_CLIENT; + if (std::holds_alternative<MoqtSetup>(frame)) { + const MoqtSetup& setup = std::get<MoqtSetup>(frame); + if (!use_webtrans && !setup.parameters.path.has_value()) { + perspective = quic::Perspective::IS_SERVER; + } + } + MoqtFramer framer(use_webtrans, perspective); return std::string(std::visit(FramingVisitor{framer}, frame).AsStringView()); }
diff --git a/quiche/quic/moqt/test_tools/moqt_framer_utils.h b/quiche/quic/moqt/test_tools/moqt_framer_utils.h index f5b693b..be2f36b 100644 --- a/quiche/quic/moqt/test_tools/moqt_framer_utils.h +++ b/quiche/quic/moqt/test_tools/moqt_framer_utils.h
@@ -19,15 +19,13 @@ namespace moqt::test { -using AnyMoqtControlMessage = - std::variant<MoqtClientSetup, MoqtServerSetup, MoqtRequestOk, - MoqtRequestError, MoqtSubscribe, MoqtSubscribeOk, - MoqtUnsubscribe, MoqtPublishDone, MoqtRequestUpdate, - MoqtPublishNamespace, MoqtPublishNamespaceDone, - MoqtPublishNamespaceCancel, MoqtTrackStatus, MoqtGoAway, - MoqtSubscribeNamespace, MoqtMaxRequestId, MoqtFetch, - MoqtFetchCancel, MoqtFetchOk, MoqtRequestsBlocked, MoqtPublish, - MoqtNamespace, MoqtNamespaceDone, MoqtObjectAck>; +using AnyMoqtControlMessage = std::variant< + MoqtSetup, MoqtRequestOk, MoqtRequestError, MoqtSubscribe, MoqtSubscribeOk, + MoqtUnsubscribe, MoqtPublishDone, MoqtRequestUpdate, MoqtPublishNamespace, + MoqtPublishNamespaceDone, MoqtPublishNamespaceCancel, MoqtTrackStatus, + MoqtGoAway, MoqtSubscribeNamespace, MoqtMaxRequestId, MoqtFetch, + MoqtFetchCancel, MoqtFetchOk, MoqtRequestsBlocked, MoqtPublish, + MoqtNamespace, MoqtNamespaceDone, MoqtObjectAck>; std::string SerializeGenericMessage(const AnyMoqtControlMessage& frame, bool use_webtrans = false);
diff --git a/quiche/quic/moqt/test_tools/moqt_test_message.h b/quiche/quic/moqt/test_tools/moqt_test_message.h index d5f0fa0..a19cc8a 100644 --- a/quiche/quic/moqt/test_tools/moqt_test_message.h +++ b/quiche/quic/moqt/test_tools/moqt_test_message.h
@@ -19,6 +19,7 @@ #include "quiche/quic/core/quic_data_reader.h" #include "quiche/quic/core/quic_data_writer.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_key_value_pair.h" #include "quiche/quic/moqt/moqt_messages.h" @@ -108,14 +109,15 @@ public: virtual ~TestMessageBase() = default; - using MessageStructuredData = std::variant< - MoqtClientSetup, MoqtServerSetup, MoqtObject, MoqtRequestOk, - MoqtRequestError, MoqtSubscribe, MoqtSubscribeOk, MoqtUnsubscribe, - MoqtPublishDone, MoqtRequestUpdate, MoqtPublishNamespace, - MoqtPublishNamespaceDone, MoqtPublishNamespaceCancel, MoqtTrackStatus, - MoqtGoAway, MoqtSubscribeNamespace, MoqtMaxRequestId, MoqtFetch, - MoqtFetchCancel, MoqtFetchOk, MoqtRequestsBlocked, MoqtPublish, - MoqtNamespace, MoqtNamespaceDone, MoqtObjectAck>; + using MessageStructuredData = + std::variant<MoqtSetup, MoqtObject, MoqtRequestOk, MoqtRequestError, + MoqtSubscribe, MoqtSubscribeOk, MoqtUnsubscribe, + MoqtPublishDone, MoqtRequestUpdate, MoqtPublishNamespace, + MoqtPublishNamespaceDone, MoqtPublishNamespaceCancel, + MoqtTrackStatus, MoqtGoAway, MoqtSubscribeNamespace, + MoqtMaxRequestId, MoqtFetch, MoqtFetchCancel, MoqtFetchOk, + MoqtRequestsBlocked, MoqtPublish, MoqtNamespace, + MoqtNamespaceDone, MoqtObjectAck>; // The total actual size of the message. size_t total_message_size() const { return wire_image_size_; } @@ -603,15 +605,15 @@ // Should not send PATH or AUTHORITY. client_setup_.parameters.path = std::nullopt; client_setup_.parameters.authority = std::nullopt; - raw_packet_[2] -= 17; // adjust payload length - raw_packet_[3] = 0x02; // only two parameters + raw_packet_[3] -= 17; // adjust payload length + raw_packet_[4] = 0x02; // only two parameters // Move MaxRequestId up in the packet. - memmove(raw_packet_ + 4, raw_packet_ + 10, 2); + memmove(raw_packet_ + 5, raw_packet_ + 11, 2); // Move MoqtImplementation up in the packet. - memmove(raw_packet_ + 6, raw_packet_ + 23, + memmove(raw_packet_ + 7, raw_packet_ + 24, kTestImplementationString.length() + 2); - raw_packet_[4] = 0x02; // Diff from 0. - raw_packet_[6] = 0x05; // Diff from 2. + raw_packet_[5] = 0x02; // Diff from 0. + raw_packet_[7] = 0x05; // Diff from 2. SetWireImage(raw_packet_, sizeof(raw_packet_) - 17); } else { SetWireImage(raw_packet_, sizeof(raw_packet_)); @@ -619,7 +621,7 @@ } bool EqualFieldValues(const MessageStructuredData& values) const override { - auto cast = std::get<MoqtClientSetup>(values); + auto cast = std::get<MoqtSetup>(values); if (cast.parameters != client_setup_.parameters) { QUIC_LOG(INFO) << "CLIENT_SETUP parameter mismatch"; return false; @@ -644,8 +646,8 @@ // string parameters in order. Unfortunately, this means that // kMoqtImplementation goes last even though it is always present, while // kPath and KAuthority aren't. - uint8_t raw_packet_[53] = { - 0x20, 0x00, 0x32, // type, length + uint8_t raw_packet_[54] = { + 0xaf, 0x00, 0x00, 0x32, // type, length 0x04, // 4 parameters 0x01, 0x04, 0x70, 0x61, 0x74, 0x68, // path = "path" 0x01, 0x32, // max_request_id = 50 @@ -655,7 +657,7 @@ 0x02, 0x1c, 0x4d, 0x6f, 0x71, 0x20, 0x54, 0x65, 0x73, 0x74, 0x20, 0x49, 0x6d, 0x70, 0x6c, 0x65, 0x6d, 0x65, 0x6e, 0x74, 0x61, 0x74, 0x69, 0x6f, 0x6e, 0x20, 0x54, 0x79, 0x70, 0x65}; - MoqtClientSetup client_setup_ = { + MoqtSetup client_setup_ = { SetupParameters("path", "authority", 50), }; }; @@ -668,7 +670,7 @@ } bool EqualFieldValues(const MessageStructuredData& values) const override { - auto cast = std::get<MoqtServerSetup>(values); + auto cast = std::get<MoqtSetup>(values); if (cast.parameters != server_setup_.parameters) { QUIC_LOG(INFO) << "SERVER_SETUP parameter mismatch"; return false; @@ -683,15 +685,15 @@ } private: - uint8_t raw_packet_[36] = {0x21, 0x00, 0x21, // type, length - 0x02, // two parameters - 0x02, 0x32, // max_subscribe_id = 50 + uint8_t raw_packet_[37] = {0xaf, 0x00, 0x00, 0x21, // type, length + 0x02, // two parameters + 0x02, 0x32, // max_subscribe_id = 50 // moqt_implementation: 0x05, 0x1c, 0x4d, 0x6f, 0x71, 0x20, 0x54, 0x65, 0x73, 0x74, 0x20, 0x49, 0x6d, 0x70, 0x6c, 0x65, 0x6d, 0x65, 0x6e, 0x74, 0x61, 0x74, 0x69, 0x6f, 0x6e, 0x20, 0x54, 0x79, 0x70, 0x65}; - MoqtServerSetup server_setup_ = { + MoqtSetup server_setup_ = { SetupParameters(50), }; }; @@ -1793,7 +1795,8 @@ // Factory function for test messages. static inline std::unique_ptr<TestMessageBase> CreateTestMessage( - MoqtMessageType message_type, bool is_webtrans) { + MoqtMessageType message_type, bool is_webtrans = true, + quic::Perspective perspective = quic::Perspective::IS_CLIENT) { switch (message_type) { case MoqtMessageType::kRequestOk: return std::make_unique<RequestOkMessage>(); @@ -1839,10 +1842,12 @@ return std::make_unique<PublishMessage>(); case MoqtMessageType::kObjectAck: return std::make_unique<ObjectAckMessage>(); - case MoqtMessageType::kClientSetup: - return std::make_unique<ClientSetupMessage>(is_webtrans); - case MoqtMessageType::kServerSetup: - return std::make_unique<ServerSetupMessage>(); + case MoqtMessageType::kSetup: + if (perspective == quic::Perspective::IS_CLIENT) { + return std::make_unique<ClientSetupMessage>(is_webtrans); + } else { + return std::make_unique<ServerSetupMessage>(); + } default: return nullptr; }