diff --git a/api/uma_metrics.h b/api/uma_metrics.h index 981736d78e..0ea31c7532 100644 --- a/api/uma_metrics.h +++ b/api/uma_metrics.h @@ -252,6 +252,8 @@ enum SdpMungingType { kDataChannelSctpInit = 100, kDataChannelMaxMessageSize = 101, kDataChannelSctpPort = 102, + // Overflow area. + kIceOptionsSped = 110, kMaxValue, }; diff --git a/experiments/field_trials.py b/experiments/field_trials.py index 8cf948d1f2..b99b2c7b18 100755 --- a/experiments/field_trials.py +++ b/experiments/field_trials.py @@ -98,6 +98,9 @@ def bug_url(self) -> str: FieldTrial('WebRTC-DisableSslGroupIds', 404763475, date(2025,9,1)), + FieldTrial('WebRTC-DtlsStunPiggybackControllerSped', + 367395350, + date(2027, 1, 1)), FieldTrial('WebRTC-ElasticBitrateAllocation', 350555527, date(2025, 3, 1)), @@ -130,7 +133,7 @@ def bug_url(self) -> str: date(2024, 4, 1)), FieldTrial('WebRTC-IceHandshakeDtls', 367395350, - date(2026, 1, 1)), + date(2027, 1, 1)), FieldTrial('WebRTC-IncomingTimestampOnMarkerBitOnly', 42224805, date(2024, 4, 1)), diff --git a/p2p/BUILD.gn b/p2p/BUILD.gn index 3b76b26a8f..9293676af2 100644 --- a/p2p/BUILD.gn +++ b/p2p/BUILD.gn @@ -666,6 +666,9 @@ rtc_library("dtls_stun_piggyback_controller") { "dtls/dtls_stun_piggyback_callbacks.h", "dtls/dtls_stun_piggyback_controller.cc", "dtls/dtls_stun_piggyback_controller.h", + "dtls/dtls_stun_piggyback_controller_interface.h", + "dtls/dtls_stun_piggyback_controller_sped.cc", + "dtls/dtls_stun_piggyback_controller_sped.h", ] deps = [ ":dtls_utils", diff --git a/p2p/base/connection.cc b/p2p/base/connection.cc index ed91bad568..c3ea687782 100644 --- a/p2p/base/connection.cc +++ b/p2p/base/connection.cc @@ -652,20 +652,9 @@ void Connection::MaybeHandleDtlsPiggybackingAttributes( if (dtls_piggyback_ack != nullptr) { piggyback_acks = dtls_piggyback_ack->GetUInt32Vector(); } - // A response implicitly acknowledges the original embedded packet - // when the ack attribute is included. - if (dtls_piggyback_ack != nullptr && original_request != nullptr) { - const StunByteStringAttribute* request_dtls_piggyback = - original_request->msg()->GetByteString(STUN_ATTR_META_DTLS_IN_STUN); - if (request_dtls_piggyback) { - uint32_t sent_hash = - ComputeDtlsPacketHash(request_dtls_piggyback->array_view()); - if (!piggyback_acks) { - piggyback_acks = {}; - } - piggyback_acks->push_back(sent_hash); - } - } + // TODO: bugs.webrtc.org/367395350 - a binding response could + // implicitly acknowledge data sent in its associated binding + // request. dtls_stun_piggyback_callbacks_.recv_data(piggyback_data, piggyback_acks); } diff --git a/p2p/base/connection_unittest.cc b/p2p/base/connection_unittest.cc index 02ce1f736a..92d25d5911 100644 --- a/p2p/base/connection_unittest.cc +++ b/p2p/base/connection_unittest.cc @@ -92,7 +92,7 @@ class ConnectionTest : public ::testing::Test { void SendPingAndCaptureReply(Connection* lconn, Connection* rconn, int64_t ms, - BufferT* reply) { + Buffer* reply) { TestPort* lport = lconn->PortForTest() == lport_.get() ? lport_.get() : rport_.get(); TestPort* rport = @@ -116,7 +116,7 @@ class ConnectionTest : public ::testing::Test { void SendPingAndReceiveResponse(Connection* lconn, Connection* rconn, int64_t ms) { - BufferT reply; + Buffer reply; SendPingAndCaptureReply(lconn, rconn, ms, &reply); lconn->OnReadPacket(ReceivedIpPacket(reply, SocketAddress(), std::nullopt)); @@ -206,7 +206,7 @@ TEST_F(ConnectionTest, ConnectionForgetLearnedStateDiscardsPendingPings) { EXPECT_TRUE(lconn->writable()); EXPECT_TRUE(lconn->receiving()); - BufferT reply; + Buffer reply; SendPingAndCaptureReply(lconn, rconn, 10, &reply); lconn->ForgetLearnedState(); @@ -405,5 +405,76 @@ TEST_F(ConnectionTest, TooBigDeltaIsNotSent) { EXPECT_FALSE(received_goog_delta_ack); } +class DtlsStunPiggybackConnectionTest : public ConnectionTest {}; + +TEST_F(DtlsStunPiggybackConnectionTest, Callbacks) { + std::optional request_data_size; + std::optional request_ack_size; + std::optional response_data_size; + std::optional response_ack_size; + + Connection* lconn = CreateConnection(ICEROLE_CONTROLLING); + lconn->RegisterDtlsPiggyback(DtlsStunPiggybackCallbacks( + [&](auto type) { + std::optional data = "request"; + std::optional> ack = {{0}}; + return std::make_pair(data, ack); + }, + [&](auto data, auto ack) { + if (data) + response_data_size = data->size(); + if (ack) + response_ack_size = ack->size(); + })); + Connection* rconn = CreateConnection(ICEROLE_CONTROLLED); + rconn->RegisterDtlsPiggyback(DtlsStunPiggybackCallbacks( + [&](auto type) { return std::make_pair(std::nullopt, std::nullopt); }, + [&](auto data, auto ack) { + if (data) + request_data_size = data->size(); + if (ack) + request_ack_size = ack->size(); + })); + Buffer reply; + SendPingAndCaptureReply(lconn, rconn, env().clock().CurrentTime().ms(), + &reply); + lconn->OnReadPacket(ReceivedIpPacket(reply, SocketAddress(), std::nullopt)); + + EXPECT_EQ(request_data_size, 7); + EXPECT_EQ(request_ack_size, 1); + EXPECT_EQ(response_data_size, std::nullopt); + EXPECT_EQ(response_ack_size, std::nullopt); +} + +TEST_F(DtlsStunPiggybackConnectionTest, NoImplicitDtlsInStunAck) { + std::optional ack_size; + + Connection* lconn = CreateConnection(ICEROLE_CONTROLLING); + lconn->RegisterDtlsPiggyback(DtlsStunPiggybackCallbacks( + [&](auto type) { + std::optional data = "test"; + std::optional> ack; + return std::make_pair(data, ack); + }, + [&](auto data, auto ack) { + if (ack) + ack_size = ack->size(); + })); + Connection* rconn = CreateConnection(ICEROLE_CONTROLLED); + rconn->RegisterDtlsPiggyback(DtlsStunPiggybackCallbacks( + [&](auto type) { + std::vector empty; + std::optional data; + std::optional> ack = empty; + return std::make_pair(data, ack); + }, + [&](auto data, auto ack) {})); + Buffer reply; + SendPingAndCaptureReply(lconn, rconn, env().clock().CurrentTime().ms(), + &reply); + lconn->OnReadPacket(ReceivedIpPacket(reply, SocketAddress(), std::nullopt)); + EXPECT_EQ(ack_size, 0); +} + } // namespace } // namespace webrtc diff --git a/p2p/base/transport_description.h b/p2p/base/transport_description.h index b3c2fbf1c0..9d7d840aff 100644 --- a/p2p/base/transport_description.h +++ b/p2p/base/transport_description.h @@ -97,6 +97,8 @@ struct IceParameters { constexpr auto* ICE_OPTION_TRICKLE = "trickle"; constexpr auto* ICE_OPTION_RENOMINATION = "renomination"; constexpr auto* ICE_OPTION_GOOG_SPED_V1 = "goog-sped-v1"; +// STUN Protocol for Embedding DTLS. +constexpr auto* ICE_OPTION_SPED = "sped"; std::optional StringToConnectionRole( absl::string_view role_str); diff --git a/p2p/base/transport_description_factory.cc b/p2p/base/transport_description_factory.cc index 14105ae41e..93eb92295d 100644 --- a/p2p/base/transport_description_factory.cc +++ b/p2p/base/transport_description_factory.cc @@ -48,6 +48,9 @@ std::unique_ptr TransportDescriptionFactory::CreateOffer( if (options.enable_ice_renomination) { desc->AddOption(ICE_OPTION_RENOMINATION); } + if (options.dtls_handshake_in_stun) { + desc->AddOption(ICE_OPTION_SPED); + } if (SSLStreamAdapter::IsBoringSsl() && field_trials_.IsEnabled("WebRTC-IceHandshakeDtls") && @@ -105,7 +108,9 @@ std::unique_ptr TransportDescriptionFactory::CreateAnswer( current_description->HasOption(ICE_OPTION_GOOG_SPED_V1))) { desc->AddOption(ICE_OPTION_GOOG_SPED_V1); } - + if (options.dtls_handshake_in_stun) { + desc->AddOption(ICE_OPTION_SPED); + } // Special affordance for testing: Answer without DTLS params // if we are insecure without a certificate, or if we are // insecure with a non-DTLS offer. diff --git a/p2p/base/transport_description_factory.h b/p2p/base/transport_description_factory.h index cdd55347f8..106144e1af 100644 --- a/p2p/base/transport_description_factory.h +++ b/p2p/base/transport_description_factory.h @@ -28,6 +28,9 @@ struct TransportOptions { // If true, ICE renomination is supported and will be used if it is also // supported by the remote side. bool enable_ice_renomination = false; + // If true, SPED (STUN Protocol for Embedding DTLS) is supported and will + // be used if it is also supported by the remote side. + bool dtls_handshake_in_stun = false; }; // Creates transport descriptions according to the supplied configuration. diff --git a/p2p/base/transport_description_factory_unittest.cc b/p2p/base/transport_description_factory_unittest.cc index a46d737d59..e0740c9fa7 100644 --- a/p2p/base/transport_description_factory_unittest.cc +++ b/p2p/base/transport_description_factory_unittest.cc @@ -454,6 +454,28 @@ TEST_F( EXPECT_FALSE(new_answer->HasOption(ICE_OPTION_GOOG_SPED_V1)); } +TEST_F(TransportDescriptionFactoryTest, AddsDtlsInStunIceOption) { + webrtc::TransportOptions options; + options.dtls_handshake_in_stun = true; + std::unique_ptr offer = + f1_.CreateOffer(options, nullptr, &ice_credentials_); + ASSERT_THAT(offer, NotNull()); + EXPECT_TRUE(offer->HasOption("sped")); + std::unique_ptr answer = + f2_.CreateAnswer(offer.get(), options, true, nullptr, &ice_credentials_); + EXPECT_TRUE(answer->HasOption("sped")); + + options.dtls_handshake_in_stun = false; + std::unique_ptr offer2 = + f1_.CreateOffer(options, nullptr, &ice_credentials_); + ASSERT_THAT(offer2, NotNull()); + EXPECT_FALSE(offer2->HasOption("sped")); + options.dtls_handshake_in_stun = true; + std::unique_ptr answer2 = + f2_.CreateAnswer(offer2.get(), options, true, nullptr, &ice_credentials_); + EXPECT_TRUE(answer2->HasOption("sped")); +} + // Test CreateOffer with IceCredentialsIterator. TEST_F(TransportDescriptionFactoryTest, CreateOfferIceCredentialsIterator) { std::vector credentials = { diff --git a/p2p/dtls/dtls_ice_integration_fixture.h b/p2p/dtls/dtls_ice_integration_fixture.h index edc64bfb35..d682a15e62 100644 --- a/p2p/dtls/dtls_ice_integration_fixture.h +++ b/p2p/dtls/dtls_ice_integration_fixture.h @@ -559,6 +559,11 @@ class Base { ep.env, ep.ice_transport, crypto_options, ep.config.max_protocol_version); + // No SDP exchange in this fixture; stand in for JsepTransport's call. + if (ep.config.dtls_in_stun) { + ep.dtls->MaybeStartDtlsInStun(); + } + if (ice_lite_agent) { ep.dtls->SetFakeIceLite(); } diff --git a/p2p/dtls/dtls_stun_piggyback_controller.cc b/p2p/dtls/dtls_stun_piggyback_controller.cc index cb7d0c4cfe..ae6ae97625 100644 --- a/p2p/dtls/dtls_stun_piggyback_controller.cc +++ b/p2p/dtls/dtls_stun_piggyback_controller.cc @@ -140,6 +140,12 @@ DtlsStunPiggybackController::GetDataToPiggyback( RTC_DCHECK(!writing_packets_); if (pending_packets_.empty()) { + // In confirmed state include an empty data attribute. Can happen e.g. + // with PQC after receiving a partial flight. + // In unconfirmed and pending states do not include the attribute. + if (state_ == State::CONFIRMED) { + return ""; + } return std::nullopt; } diff --git a/p2p/dtls/dtls_stun_piggyback_controller.h b/p2p/dtls/dtls_stun_piggyback_controller.h index 18af1bdef3..3d47a3a59b 100644 --- a/p2p/dtls/dtls_stun_piggyback_controller.h +++ b/p2p/dtls/dtls_stun_piggyback_controller.h @@ -20,6 +20,7 @@ #include "absl/strings/string_view.h" #include "api/sequence_checker.h" #include "api/transport/stun.h" +#include "p2p/dtls/dtls_stun_piggyback_controller_interface.h" #include "p2p/dtls/dtls_utils.h" #include "rtc_base/network/received_packet.h" #include "rtc_base/system/no_unique_address.h" @@ -29,7 +30,8 @@ namespace webrtc { // This class is not thread safe; all methods must be called on the same thread // as the constructor. -class DtlsStunPiggybackController { +class DtlsStunPiggybackController + : public DtlsStunPiggybackControllerInterface { public: // Never ack more than 4 packets. static constexpr unsigned kMaxAckSize = 4; @@ -41,25 +43,9 @@ class DtlsStunPiggybackController { // NOLINTNEXTLINE(readability/casting) - not a cast; false positive! absl::AnyInvocable piggyback_complete_callback); - ~DtlsStunPiggybackController(); - - enum class State { - // We don't know if peer support DTLS piggybacked in STUN. - // We will piggyback DTLS until we get a piggybacked response - // or a STUN response with piggyback support. - TENTATIVE = 0, - // The peer supports DTLS in STUN and we continue the handshake. - CONFIRMED = 1, - // We are waiting for the final ack. Semantic differs depending - // on DTLS role. - PENDING = 2, - // We successfully completed the DTLS handshake in STUN. - COMPLETE = 3, - // The peer does not support piggybacking DTLS in STUN. - OFF = 4, - }; - - State state() const { + ~DtlsStunPiggybackController() override; + + State state() const override { RTC_DCHECK_RUN_ON(&sequence_checker_); return state_; } @@ -67,43 +53,44 @@ class DtlsStunPiggybackController { // Called by DtlsTransport when the handshake is complete "locally", // i.e. we can send encrypted packets to peer (but we don't strictly know // that peer can decode them). - void SetDtlsHandshakeComplete(bool is_dtls_client, bool is_dtls13); + void SetDtlsHandshakeComplete(bool is_dtls_client, bool is_dtls13) override; // Called by DtlsTransport when a packet has been received and passed // to layers above us. This means that dtls is writable for the peer, // and maybe we are complete. - void ApplicationPacketReceived(const ReceivedIpPacket& packet); + void ApplicationPacketReceived(const ReceivedIpPacket& packet) override; // Called by DtlsTransport when DTLS failed. - void SetDtlsFailed(); + void SetDtlsFailed() override; // Intercepts DTLS packets which should go into the STUN packets during the // handshake. - void CapturePacket(std::span data); - void ClearCachedPacketForTesting(); + void CapturePacket(std::span data) override; + void ClearCachedPacketForTesting() override; // Inform piggybackcontroller that a flight is complete. - void Flush(); + void Flush() override; // Called by Connection, when sending a STUN BINDING { REQUEST / RESPONSE } // to obtain optional DTLS data or ACKs. std::optional GetDataToPiggyback( - StunMessageType stun_message_type); + StunMessageType stun_message_type) override; std::optional> GetAckToPiggyback( - StunMessageType stun_message_type); - std::vector> GetPending(); + StunMessageType stun_message_type) override; + std::vector> GetPending() override; // Called by Connection when receiving a STUN BINDING { REQUEST / RESPONSE }. - void ReportDataPiggybacked(std::optional> data, - std::optional> acks); + void ReportDataPiggybacked( + std::optional> data, + std::optional> acks) override; // Called by // * DTLSTransport when receiving a DTLS packet (possibly after the packet // was emitted by this class). // * This class when processing a DTLS packet. - void ReportDtlsPacket(std::span data); + void ReportDtlsPacket(std::span data) override; - int GetCountOfReceivedData() const { return data_recv_count_; } + int GetCountOfReceivedData() const override { return data_recv_count_; } private: State state_ RTC_GUARDED_BY(sequence_checker_) = State::TENTATIVE; diff --git a/p2p/dtls/dtls_stun_piggyback_controller_interface.h b/p2p/dtls/dtls_stun_piggyback_controller_interface.h new file mode 100644 index 0000000000..f916fa10d4 --- /dev/null +++ b/p2p/dtls/dtls_stun_piggyback_controller_interface.h @@ -0,0 +1,77 @@ +/* + * Copyright 2026 The WebRTC Project Authors. All rights reserved. + * + * Use of this source code is governed by a BSD-style license + * that can be found in the LICENSE file in the root of the source + * tree. An additional intellectual property rights grant can be found + * in the file PATENTS. All contributing project authors may + * be found in the AUTHORS file in the root of the source tree. + */ + +#ifndef P2P_DTLS_DTLS_STUN_PIGGYBACK_CONTROLLER_INTERFACE_H_ +#define P2P_DTLS_DTLS_STUN_PIGGYBACK_CONTROLLER_INTERFACE_H_ + +#include +#include +#include +#include + +#include "absl/strings/string_view.h" +#include "api/transport/stun.h" +#include "rtc_base/network/received_packet.h" + +namespace webrtc { + +// Abstract interface for piggybacking DTLS handshake packets in STUN +// connectivity checks. Implementations are not thread safe; all methods must +// be called on the same thread as the constructor. +class DtlsStunPiggybackControllerInterface { + public: + enum class State { + // We don't know if peer support DTLS piggybacked in STUN. + // We will piggyback DTLS until we get a piggybacked response + // or a STUN response with piggyback support. + TENTATIVE = 0, + // The peer supports DTLS in STUN and we continue the handshake. + CONFIRMED = 1, + // We are waiting for the final ack. Semantic differs depending + // on DTLS role. + PENDING = 2, + // We successfully completed the DTLS handshake in STUN. + COMPLETE = 3, + // The peer does not support piggybacking DTLS in STUN. + OFF = 4, + }; + + virtual ~DtlsStunPiggybackControllerInterface() = default; + + virtual State state() const = 0; + + virtual void SetDtlsHandshakeComplete(bool is_dtls_client, + bool is_dtls13) = 0; + virtual void ApplicationPacketReceived(const ReceivedIpPacket& packet) = 0; + virtual void SetDtlsFailed() = 0; + + virtual void CapturePacket(std::span data) = 0; + virtual void ClearCachedPacketForTesting() = 0; + + virtual void Flush() = 0; + + virtual std::optional GetDataToPiggyback( + StunMessageType stun_message_type) = 0; + virtual std::optional> GetAckToPiggyback( + StunMessageType stun_message_type) = 0; + virtual std::vector> GetPending() = 0; + + virtual void ReportDataPiggybacked( + std::optional> data, + std::optional> acks) = 0; + + virtual void ReportDtlsPacket(std::span data) = 0; + + virtual int GetCountOfReceivedData() const = 0; +}; + +} // namespace webrtc + +#endif // P2P_DTLS_DTLS_STUN_PIGGYBACK_CONTROLLER_INTERFACE_H_ diff --git a/p2p/dtls/dtls_stun_piggyback_controller_sped.cc b/p2p/dtls/dtls_stun_piggyback_controller_sped.cc new file mode 100644 index 0000000000..d3e1563458 --- /dev/null +++ b/p2p/dtls/dtls_stun_piggyback_controller_sped.cc @@ -0,0 +1,277 @@ +/* + * Copyright 2024 The WebRTC Project Authors. All rights reserved. + * + * Use of this source code is governed by a BSD-style license + * that can be found in the LICENSE file in the root of the source + * tree. An additional intellectual property rights grant can be found + * in the file PATENTS. All contributing project authors may + * be found in the AUTHORS file in the root of the source tree. + */ + +#include "p2p/dtls/dtls_stun_piggyback_controller_sped.h" + +#include +#include +#include +#include +#include +#include + +#include "absl/container/flat_hash_set.h" +#include "absl/functional/any_invocable.h" +#include "absl/strings/string_view.h" +#include "api/sequence_checker.h" +#include "api/transport/stun.h" +#include "p2p/dtls/dtls_utils.h" +#include "rtc_base/checks.h" +#include "rtc_base/logging.h" +#include "rtc_base/network/received_packet.h" +#include "rtc_base/strings/str_join.h" + +namespace webrtc { + +DtlsStunPiggybackControllerSped::DtlsStunPiggybackControllerSped( + absl::AnyInvocable)> dtls_data_callback, + // NOLINTNEXTLINE(readability/casting) - not a cast; false positive! + absl::AnyInvocable piggyback_complete_callback) + : dtls_data_callback_(std::move(dtls_data_callback)), + piggyback_complete_callback_(std::move(piggyback_complete_callback)) {} + +DtlsStunPiggybackControllerSped::~DtlsStunPiggybackControllerSped() { + RTC_DCHECK(dtls_data_callback_); + RTC_DCHECK(piggyback_complete_callback_); +} + +void DtlsStunPiggybackControllerSped::SetDtlsHandshakeComplete( + bool is_dtls_client, + bool is_dtls13) { + RTC_DCHECK_RUN_ON(&sequence_checker_); + + // Peer does not support this so fallback to a normal DTLS handshake + // happened. + if (state_ == State::OFF) { + return; + } + + // As DTLS 1.2 client we have nothing more to send at this point + // but will continue to send ACK attributes until receiving + // the last flight from the server. + if (is_dtls_client && !is_dtls13) { + pending_packets_.clear(); + } + state_ = State::PENDING; +} + +void DtlsStunPiggybackControllerSped::ApplicationPacketReceived( + const ReceivedIpPacket& packet) { + RTC_DCHECK_RUN_ON(&sequence_checker_); + // TODO: bugs.webrtc.org/367395350 - remove this. +} + +void DtlsStunPiggybackControllerSped::SetDtlsFailed() { + RTC_DCHECK_RUN_ON(&sequence_checker_); + + if (state_ == State::TENTATIVE || state_ == State::CONFIRMED || + state_ == State::PENDING) { + RTC_LOG(LS_INFO) + << "DTLS-STUN piggybacking DTLS failed during negotiation."; + } + state_ = State::OFF; + CallCompleteCallback(/*success=*/false); +} + +void DtlsStunPiggybackControllerSped::CapturePacket( + std::span data) { + RTC_DCHECK_RUN_ON(&sequence_checker_); + if (!IsDtlsPacket(data)) { + return; + } + + // BoringSSL writes burst of packets...but the interface + // is made for 1-packet at a time. Use the writing_packets_ variable to keep + // track of a full flight. The writing_packets_ is reset in Flush. + if (!writing_packets_) { + pending_packets_.clear(); + writing_packets_ = true; + } + + pending_packets_.Add(data); +} + +void DtlsStunPiggybackControllerSped::ClearCachedPacketForTesting() { + RTC_DCHECK_RUN_ON(&sequence_checker_); + pending_packets_.clear(); +} + +void DtlsStunPiggybackControllerSped::Flush() { + // Flush is called by the StreamInterface (and the underlying SSL BIO) + // after a flight of packets has been sent. + RTC_DCHECK_RUN_ON(&sequence_checker_); + writing_packets_ = false; +} + +std::optional +DtlsStunPiggybackControllerSped::GetDataToPiggyback( + StunMessageType stun_message_type) { + RTC_DCHECK_RUN_ON(&sequence_checker_); + RTC_DCHECK(stun_message_type == STUN_BINDING_REQUEST || + stun_message_type == STUN_BINDING_RESPONSE); + + if (state_ == State::COMPLETE) { + return std::nullopt; + } + + if (state_ == State::OFF) { + return std::nullopt; + } + + // No longer writing packets...since we're now about to send them. + RTC_DCHECK(!writing_packets_); + + if (pending_packets_.empty()) { + // In confirmed state include an empty data attribute. Can happen e.g. + // with PQC after receiving a partial flight. + // In unconfirmed and pending states do not include the attribute. + if (state_ == State::CONFIRMED) { + return ""; + } + return std::nullopt; + } + + const auto packet = pending_packets_.GetNext(); + return absl::string_view(reinterpret_cast(packet.data()), + packet.size()); +} + +std::optional> +DtlsStunPiggybackControllerSped::GetAckToPiggyback( + StunMessageType stun_message_type) { + RTC_DCHECK_RUN_ON(&sequence_checker_); + + if (state_ == State::OFF || state_ == State::COMPLETE) { + return std::nullopt; + } + return handshake_messages_received_; +} + +std::vector> +DtlsStunPiggybackControllerSped::GetPending() { + RTC_DCHECK_RUN_ON(&sequence_checker_); + return pending_packets_.GetAll(); +} + +void DtlsStunPiggybackControllerSped::ReportDataPiggybacked( + std::optional> data, + std::optional> acks) { + RTC_DCHECK_RUN_ON(&sequence_checker_); + + // Drop silently when receiving acked data when the peer previously did not + // support or we already moved to the complete state. + if (state_ == State::OFF || state_ == State::COMPLETE) { + return; + } + + // We sent dtls piggybacked but got nothing in return or + // we received a stun request with neither attribute set + // => peer does not support. + if (state_ == State::TENTATIVE && !data.has_value() && !acks.has_value()) { + RTC_LOG(LS_INFO) << "DTLS-STUN piggybacking not supported by peer."; + state_ = State::OFF; + // TODO: bugs.webrtc.org/367395350 - we cached a client hello + // which needs to be sent by the DTLS transport now. + // CallCompleteCallback(/*success=*/false); + return; + } + + // In PENDING state the peer may have stopped sending the ack + // when it moved to the COMPLETE state. Move to the same state. + if (state_ == State::PENDING && !data.has_value() && !acks.has_value()) { + RTC_LOG(LS_INFO) << "DTLS-STUN piggybacking complete."; + state_ = State::COMPLETE; + CallCompleteCallback(/*success=*/true); + return; + } + + // We sent dtls piggybacked and got something in return => peer does support. + if (state_ == State::TENTATIVE) { + state_ = State::CONFIRMED; + } + + if (acks.has_value()) { + if (!pending_packets_.empty()) { + // Unpack the ACK attribute (a list of uint32_t) + absl::flat_hash_set acked_packets; + for (const auto& ack : *acks) { + acked_packets.insert(ack); + } + RTC_LOG(LS_VERBOSE) << "DTLS-STUN piggybacking ACK: " + << StrJoin(acked_packets, ","); + + // Remove all acked packets from pending_packets_. + pending_packets_.Prune(acked_packets); + } + } + + // The response to the final flight of the handshake will not contain + // the DTLS data but will contain an ack. + // Must not happen on the initial server to client packet which + // has no DTLS data yet. + if (state_ == State::PENDING && !data.has_value() && acks.has_value()) { + RTC_LOG(LS_INFO) << "DTLS-STUN piggybacking complete."; + state_ = State::COMPLETE; + CallCompleteCallback(/*success=*/true); + return; + } + + if (!data.has_value() || data->empty()) { + return; + } + // Drop non-DTLS packets. + if (!IsDtlsPacket(*data)) { + RTC_LOG(LS_WARNING) << "Dropping non-DTLS data."; + return; + } + data_recv_count_++; + ReportDtlsPacket(*data); + + // Forwards the data to the DTLS layer. Note that this will call + // ProcessDtlsPacket() again which does not change the state. + dtls_data_callback_(*data); +} + +void DtlsStunPiggybackControllerSped::ReportDtlsPacket( + std::span data) { + RTC_DCHECK_RUN_ON(&sequence_checker_); + + if (state_ == State::OFF || state_ == State::COMPLETE) { + return; + } + + // Extract the received message id of the handshake + // from the packet and prepare the ack to be sent. + uint32_t hash = ComputeDtlsPacketHash(data); + + // Check if we already received this packet. + if (std::find(handshake_messages_received_.begin(), + handshake_messages_received_.end(), + hash) == handshake_messages_received_.end()) { + // If needed, limit size of ack attribute by removing oldest ack. + while (handshake_messages_received_.size() >= kMaxAckSize) { + handshake_messages_received_.erase(handshake_messages_received_.begin()); + } + handshake_messages_received_.push_back(hash); + } +} + +void DtlsStunPiggybackControllerSped::CallCompleteCallback(bool success) { + RTC_DCHECK_RUN_ON(&sequence_checker_); + pending_packets_.clear(); + handshake_messages_received_.clear(); + if (!piggyback_complete_callback_) { + RTC_DCHECK_NOTREACHED() << "CompleteCallback called twice!"; + return; + } + std::move(piggyback_complete_callback_)(success); +} + +} // namespace webrtc diff --git a/p2p/dtls/dtls_stun_piggyback_controller_sped.h b/p2p/dtls/dtls_stun_piggyback_controller_sped.h new file mode 100644 index 0000000000..014ea260da --- /dev/null +++ b/p2p/dtls/dtls_stun_piggyback_controller_sped.h @@ -0,0 +1,119 @@ +/* + * Copyright 2024 The WebRTC Project Authors. All rights reserved. + * + * Use of this source code is governed by a BSD-style license + * that can be found in the LICENSE file in the root of the source + * tree. An additional intellectual property rights grant can be found + * in the file PATENTS. All contributing project authors may + * be found in the AUTHORS file in the root of the source tree. + */ + +#ifndef P2P_DTLS_DTLS_STUN_PIGGYBACK_CONTROLLER_SPED_H_ +#define P2P_DTLS_DTLS_STUN_PIGGYBACK_CONTROLLER_SPED_H_ + +#include +#include +#include +#include + +#include "absl/functional/any_invocable.h" +#include "absl/strings/string_view.h" +#include "api/sequence_checker.h" +#include "api/transport/stun.h" +#include "p2p/dtls/dtls_stun_piggyback_controller_interface.h" +#include "p2p/dtls/dtls_utils.h" +#include "rtc_base/network/received_packet.h" +#include "rtc_base/system/no_unique_address.h" +#include "rtc_base/thread_annotations.h" + +namespace webrtc { + +// This class is not thread safe; all methods must be called on the same thread +// as the constructor. +class DtlsStunPiggybackControllerSped + : public DtlsStunPiggybackControllerInterface { + public: + // Never ack more than 4 packets. + static constexpr unsigned kMaxAckSize = 4; + + // dtls_data_callback will be called with any DTLS packets received + // piggybacked. + DtlsStunPiggybackControllerSped( + absl::AnyInvocable)> dtls_data_callback, + // NOLINTNEXTLINE(readability/casting) - not a cast; false positive! + absl::AnyInvocable piggyback_complete_callback); + + ~DtlsStunPiggybackControllerSped() override; + + State state() const override { + RTC_DCHECK_RUN_ON(&sequence_checker_); + return state_; + } + + // Called by DtlsTransport when the handshake is complete "locally", + // i.e. we can send encrypted packets to peer (but we don't strictly know + // that peer can decode them). + void SetDtlsHandshakeComplete(bool is_dtls_client, bool is_dtls13) override; + + // Called by DtlsTransport when a packet has been received and passed + // to layers above us. This means that dtls is writable for the peer, + // and maybe we are complete. + void ApplicationPacketReceived(const ReceivedIpPacket& packet) override; + + // Called by DtlsTransport when DTLS failed. + void SetDtlsFailed() override; + + // Intercepts DTLS packets which should go into the STUN packets during the + // handshake. + void CapturePacket(std::span data) override; + void ClearCachedPacketForTesting() override; + + // Inform piggybackcontroller that a flight is complete. + void Flush() override; + + // Called by Connection, when sending a STUN BINDING { REQUEST / RESPONSE } + // to obtain optional DTLS data or ACKs. + std::optional GetDataToPiggyback( + StunMessageType stun_message_type) override; + std::optional> GetAckToPiggyback( + StunMessageType stun_message_type) override; + std::vector> GetPending() override; + + // Called by Connection when receiving a STUN BINDING { REQUEST / RESPONSE }. + void ReportDataPiggybacked( + std::optional> data, + std::optional> acks) override; + + // Called by + // * DTLSTransport when receiving a DTLS packet (possibly after the packet + // was emitted by this class). + // * This class when processing a DTLS packet. + void ReportDtlsPacket(std::span data) override; + + int GetCountOfReceivedData() const override { return data_recv_count_; } + + private: + State state_ RTC_GUARDED_BY(sequence_checker_) = State::TENTATIVE; + bool writing_packets_ RTC_GUARDED_BY(sequence_checker_) = false; + PacketStash pending_packets_ RTC_GUARDED_BY(sequence_checker_); + absl::AnyInvocable)> dtls_data_callback_ + RTC_GUARDED_BY(sequence_checker_); + // NOLINTNEXTLINE(readability/casting) - not a cast; false positive! + absl::AnyInvocable piggyback_complete_callback_ + RTC_GUARDED_BY(sequence_checker_); + + std::vector handshake_messages_received_ + RTC_GUARDED_BY(sequence_checker_); + + // Count of embedded data attributes received. + int data_recv_count_ = 0; + + void CallCompleteCallback(bool success); + + // In practice this will be the network thread. + RTC_NO_UNIQUE_ADDRESS SequenceChecker sequence_checker_; +}; + +} // namespace webrtc + +#endif // P2P_DTLS_DTLS_STUN_PIGGYBACK_CONTROLLER_SPED_H_ diff --git a/p2p/dtls/dtls_stun_piggyback_controller_unittest.cc b/p2p/dtls/dtls_stun_piggyback_controller_unittest.cc index cadc0ada54..df09ea2341 100644 --- a/p2p/dtls/dtls_stun_piggyback_controller_unittest.cc +++ b/p2p/dtls/dtls_stun_piggyback_controller_unittest.cc @@ -15,10 +15,14 @@ #include #include #include +#include #include +#include "absl/functional/any_invocable.h" #include "absl/strings/string_view.h" #include "api/transport/stun.h" +#include "p2p/dtls/dtls_stun_piggyback_controller_interface.h" +#include "p2p/dtls/dtls_stun_piggyback_controller_sped.h" #include "p2p/dtls/dtls_utils.h" #include "rtc_base/byte_buffer.h" #include "rtc_base/checks.h" @@ -63,6 +67,27 @@ const std::vector empty = {}; const std::vector kPayload = {0x1, 0x2, 0x3}; +// The two implementations of the piggybacking protocol. They agree on which +// handshake flights are piggybacked and only differ in how the piggybacking +// session terminates. +enum class Variant { kGoogSped, kSped }; + +static_assert(DtlsStunPiggybackController::kMaxAckSize == + DtlsStunPiggybackControllerSped::kMaxAckSize); + +std::unique_ptr CreateController( + Variant variant, + absl::AnyInvocable)> dtls_data_callback, + // NOLINTNEXTLINE(readability/casting) - not a cast; false positive! + absl::AnyInvocable piggyback_complete_callback) { + if (variant == Variant::kSped) { + return std::make_unique( + std::move(dtls_data_callback), std::move(piggyback_complete_callback)); + } + return std::make_unique( + std::move(dtls_data_callback), std::move(piggyback_complete_callback)); +} + std::vector FromAckAttribute(std::span attr) { ByteBufferReader ack_reader(attr); std::vector values; @@ -101,98 +126,101 @@ std::unique_ptr WrapInStun( } // namespace using ::testing::ElementsAreArray; -using ::testing::MockFunction; +using ::testing::IsEmpty; using ::testing::NotNull; -using State = DtlsStunPiggybackController::State; +using ::testing::SizeIs; +using State = DtlsStunPiggybackControllerInterface::State; -class DtlsStunPiggybackControllerTest : public ::testing::Test { +class DtlsStunPiggybackControllerTestBase : public ::testing::Test { protected: - DtlsStunPiggybackControllerTest() - : client_( + explicit DtlsStunPiggybackControllerTestBase(Variant variant) + : client_(CreateController( + variant, [this](std::span data) { ClientPacketSink(data); }, - [this](bool success) { ClientCompleteCallback(success); }), - server_( + [this](bool success) { ClientCompleteCallback(success); })), + server_(CreateController( + variant, [this](std::span data) { ServerPacketSink(data); }, - [this](bool success) { ServerCompleteCallback(success); }), + [this](bool success) { ServerCompleteCallback(success); })), packet_(kPayload, SocketAddress(), std::nullopt) {} // Send from client to server embedded in STUN. void SendClientToServerEmbedded(const std::vector& packet, StunMessageType type) { if (!packet.empty()) { - client_.CapturePacket(packet); - client_.Flush(); + client_->CapturePacket(packet); + client_->Flush(); } else { - client_.ClearCachedPacketForTesting(); + client_->ClearCachedPacketForTesting(); } std::unique_ptr attr_data; std::optional> view_data; - if (auto data = client_.GetDataToPiggyback(type)) { + if (auto data = client_->GetDataToPiggyback(type)) { attr_data = WrapInStun(STUN_ATTR_META_DTLS_IN_STUN, *data); view_data = attr_data->array_view(); } std::unique_ptr attr_ack; std::optional> view_acks; - if (auto ack = client_.GetAckToPiggyback(type)) { + if (auto ack = client_->GetAckToPiggyback(type)) { attr_ack = WrapInStun(STUN_ATTR_META_DTLS_IN_STUN_ACK, *ack); view_acks = FromAckAttribute(attr_ack->array_view()); } - server_.ReportDataPiggybacked(view_data, view_acks); + server_->ReportDataPiggybacked(view_data, view_acks); } // Send from client to server as plain DTLS. void SendClientToServerDtls(const std::vector packet) { if (!packet.empty()) { - client_.CapturePacket(packet); - client_.Flush(); + client_->CapturePacket(packet); + client_->Flush(); } else { - client_.ClearCachedPacketForTesting(); + client_->ClearCachedPacketForTesting(); } - server_.ReportDtlsPacket(packet); + server_->ReportDtlsPacket(packet); } // Send from server to client embedded in STUN void SendServerToClientEmbedded(const std::vector& packet, StunMessageType type) { if (!packet.empty()) { - server_.CapturePacket(packet); - server_.Flush(); + server_->CapturePacket(packet); + server_->Flush(); } else { - server_.ClearCachedPacketForTesting(); + server_->ClearCachedPacketForTesting(); } std::unique_ptr attr_data; std::optional> view_data; - if (auto data = server_.GetDataToPiggyback(type)) { + if (auto data = server_->GetDataToPiggyback(type)) { attr_data = WrapInStun(STUN_ATTR_META_DTLS_IN_STUN, *data); view_data = attr_data->array_view(); } std::unique_ptr attr_ack; std::optional> view_acks; - if (auto ack = server_.GetAckToPiggyback(type)) { + if (auto ack = server_->GetAckToPiggyback(type)) { attr_ack = WrapInStun(STUN_ATTR_META_DTLS_IN_STUN_ACK, *ack); view_acks = FromAckAttribute(attr_ack->array_view()); } - client_.ReportDataPiggybacked(view_data, view_acks); + client_->ReportDataPiggybacked(view_data, view_acks); MaybeSetHandshakeComplete(packet); } // Send from server to client as plain DTLS. void SendServerToClientDtls(const std::vector packet) { if (!packet.empty()) { - server_.CapturePacket(packet); - server_.Flush(); + server_->CapturePacket(packet); + server_->Flush(); } else { - server_.ClearCachedPacketForTesting(); + server_->ClearCachedPacketForTesting(); } - client_.ReportDtlsPacket(packet); + client_->ReportDtlsPacket(packet); MaybeSetHandshakeComplete(packet); } - void DisableSupport(DtlsStunPiggybackController& client_or_server) { + void DisableSupport(DtlsStunPiggybackControllerInterface& client_or_server) { ASSERT_EQ(client_or_server.state(), State::TENTATIVE); client_or_server.ReportDataPiggybacked(std::nullopt, std::nullopt); ASSERT_EQ(client_or_server.state(), State::OFF); } - DtlsStunPiggybackController client_; - DtlsStunPiggybackController server_; + std::unique_ptr client_; + std::unique_ptr server_; MOCK_METHOD(void, ClientPacketSink, (std::span)); MOCK_METHOD(void, ServerPacketSink, (std::span)); @@ -207,125 +235,99 @@ class DtlsStunPiggybackControllerTest : public ::testing::Test { // Note: this assumes DTLS 1.2 if (packet == dtls_flight4) { // After sending flight 4, the server handshake is complete. - server_.SetDtlsHandshakeComplete(/*is_client=*/false, - /*is_dtls13=*/false); + server_->SetDtlsHandshakeComplete(/*is_client=*/false, + /*is_dtls13=*/false); // When receiving flight 4, client handshake is complete. - client_.SetDtlsHandshakeComplete(/*is_client=*/true, /*is_dtls13=*/false); + client_->SetDtlsHandshakeComplete(/*is_client=*/true, + /*is_dtls13=*/false); } } }; -TEST_F(DtlsStunPiggybackControllerTest, BasicHandshake) { +// Behaviour shared by both variants. +class DtlsStunPiggybackControllerTest + : public DtlsStunPiggybackControllerTestBase, + public ::testing::WithParamInterface { + protected: + DtlsStunPiggybackControllerTest() + : DtlsStunPiggybackControllerTestBase(GetParam()) {} +}; + +INSTANTIATE_TEST_SUITE_P(All, + DtlsStunPiggybackControllerTest, + ::testing::Values(Variant::kGoogSped, Variant::kSped), + [](const ::testing::TestParamInfo& info) { + return info.param == Variant::kSped ? "Sped" + : "GoogSped"; + }); + +// Behaviour of the default variant only. +class DtlsStunPiggybackControllerGoogSpedTest + : public DtlsStunPiggybackControllerTestBase { + protected: + DtlsStunPiggybackControllerGoogSpedTest() + : DtlsStunPiggybackControllerTestBase(Variant::kGoogSped) {} +}; + +// Behaviour of the variant behind WebRTC-DtlsStunPiggybackControllerSped only. +class DtlsStunPiggybackControllerSpedTest + : public DtlsStunPiggybackControllerTestBase { + protected: + DtlsStunPiggybackControllerSpedTest() + : DtlsStunPiggybackControllerTestBase(Variant::kSped) {} +}; + +TEST_P(DtlsStunPiggybackControllerTest, BasicHandshake) { // Flight 1+2 SendClientToServerEmbedded(dtls_flight1, STUN_BINDING_REQUEST); - EXPECT_EQ(server_.state(), State::CONFIRMED); + EXPECT_EQ(server_->state(), State::CONFIRMED); SendServerToClientEmbedded(dtls_flight2, STUN_BINDING_RESPONSE); - EXPECT_EQ(client_.state(), State::CONFIRMED); + EXPECT_EQ(client_->state(), State::CONFIRMED); // Flight 3+4 SendClientToServerEmbedded(dtls_flight3, STUN_BINDING_REQUEST); SendServerToClientEmbedded(dtls_flight4, STUN_BINDING_RESPONSE); - EXPECT_EQ(server_.state(), State::PENDING); - EXPECT_EQ(client_.state(), State::PENDING); + EXPECT_EQ(server_->state(), State::PENDING); + EXPECT_EQ(client_->state(), State::PENDING); // Post-handshake ACK EXPECT_CALL(*this, ClientCompleteCallback(true)); SendServerToClientEmbedded(empty, STUN_BINDING_REQUEST); EXPECT_CALL(*this, ServerCompleteCallback(true)); SendClientToServerEmbedded(empty, STUN_BINDING_RESPONSE); - EXPECT_EQ(server_.state(), State::COMPLETE); - EXPECT_EQ(client_.state(), State::COMPLETE); -} - -TEST_F(DtlsStunPiggybackControllerTest, - BasicHandshakeCompleteWithDecryptedPacket) { - // Flight 1+2 - SendClientToServerEmbedded(dtls_flight1, STUN_BINDING_REQUEST); - EXPECT_EQ(server_.state(), State::CONFIRMED); - SendServerToClientEmbedded(dtls_flight2, STUN_BINDING_RESPONSE); - EXPECT_EQ(client_.state(), State::CONFIRMED); - - // Flight 3+4 - SendClientToServerEmbedded(dtls_flight3, STUN_BINDING_REQUEST); - SendServerToClientEmbedded(dtls_flight4, STUN_BINDING_RESPONSE); - EXPECT_EQ(server_.state(), State::PENDING); - EXPECT_EQ(client_.state(), State::PENDING); - - // Post-handshake ACK - EXPECT_CALL(*this, ClientCompleteCallback); - client_.ApplicationPacketReceived( - packet_.CopyAndSet(ReceivedIpPacket::kDtlsDecrypted)); - EXPECT_EQ(client_.state(), State::COMPLETE); - - EXPECT_CALL(*this, ServerCompleteCallback); - server_.ApplicationPacketReceived( - packet_.CopyAndSet(ReceivedIpPacket::kSrtpEncrypted)); - EXPECT_EQ(server_.state(), State::COMPLETE); -} - -TEST_F(DtlsStunPiggybackControllerTest, - BasicHandshakeEarlySrtpDoesNotComplete) { - // Flight 1+2 - SendClientToServerEmbedded(dtls_flight1, STUN_BINDING_REQUEST); - EXPECT_EQ(server_.state(), State::CONFIRMED); - SendServerToClientEmbedded(dtls_flight2, STUN_BINDING_RESPONSE); - EXPECT_EQ(client_.state(), State::CONFIRMED); - - // Flight 3 - SendClientToServerEmbedded(dtls_flight3, STUN_BINDING_REQUEST); - EXPECT_EQ(server_.state(), State::CONFIRMED); - - // An srtp packet arriving before reaching PENDING state. - server_.ApplicationPacketReceived( - packet_.CopyAndSet(ReceivedIpPacket::kSrtpEncrypted)); - EXPECT_EQ(server_.state(), State::CONFIRMED); - - // Flight 4 - SendServerToClientEmbedded(dtls_flight4, STUN_BINDING_RESPONSE); - EXPECT_EQ(server_.state(), State::PENDING); - EXPECT_EQ(client_.state(), State::PENDING); - - // Post-handshake ACK - EXPECT_CALL(*this, ClientCompleteCallback); - client_.ApplicationPacketReceived( - packet_.CopyAndSet(ReceivedIpPacket::kDtlsDecrypted)); - EXPECT_EQ(client_.state(), State::COMPLETE); - - EXPECT_CALL(*this, ServerCompleteCallback); - server_.ApplicationPacketReceived( - packet_.CopyAndSet(ReceivedIpPacket::kSrtpEncrypted)); - EXPECT_EQ(server_.state(), State::COMPLETE); + EXPECT_EQ(server_->state(), State::COMPLETE); + EXPECT_EQ(client_->state(), State::COMPLETE); } -TEST_F(DtlsStunPiggybackControllerTest, FirstClientPacketLost) { +TEST_P(DtlsStunPiggybackControllerTest, FirstClientPacketLost) { // Client to server got lost (or arrives late) // Flight 1 SendServerToClientEmbedded(empty, STUN_BINDING_REQUEST); SendClientToServerEmbedded(dtls_flight1, STUN_BINDING_RESPONSE); - EXPECT_EQ(server_.state(), State::CONFIRMED); - EXPECT_EQ(client_.state(), State::CONFIRMED); + EXPECT_EQ(server_->state(), State::CONFIRMED); + EXPECT_EQ(client_->state(), State::CONFIRMED); // Flight 2+3 SendServerToClientEmbedded(dtls_flight2, STUN_BINDING_REQUEST); SendClientToServerEmbedded(dtls_flight3, STUN_BINDING_RESPONSE); - EXPECT_EQ(server_.state(), State::CONFIRMED); - EXPECT_EQ(client_.state(), State::CONFIRMED); + EXPECT_EQ(server_->state(), State::CONFIRMED); + EXPECT_EQ(client_->state(), State::CONFIRMED); // Flight 4 SendServerToClientEmbedded(dtls_flight4, STUN_BINDING_REQUEST); EXPECT_CALL(*this, ServerCompleteCallback(true)); SendClientToServerEmbedded(empty, STUN_BINDING_RESPONSE); - EXPECT_EQ(server_.state(), State::COMPLETE); - EXPECT_EQ(client_.state(), State::PENDING); + EXPECT_EQ(server_->state(), State::COMPLETE); + EXPECT_EQ(client_->state(), State::PENDING); // Post-handshake ACK EXPECT_CALL(*this, ClientCompleteCallback(true)); SendServerToClientEmbedded(empty, STUN_BINDING_REQUEST); - EXPECT_EQ(client_.state(), State::COMPLETE); + EXPECT_EQ(client_->state(), State::COMPLETE); } -TEST_F(DtlsStunPiggybackControllerTest, NotSupportedByServer) { - DisableSupport(server_); +TEST_P(DtlsStunPiggybackControllerTest, NotSupportedByServer) { + DisableSupport(*server_); // Flight 1 SendClientToServerEmbedded(dtls_flight1, STUN_BINDING_REQUEST); @@ -333,111 +335,126 @@ TEST_F(DtlsStunPiggybackControllerTest, NotSupportedByServer) { // callback in this case which currently causes a sleuth of test failures. // EXPECT_CALL(*this, ClientCompleteCallback()); SendServerToClientEmbedded(empty, STUN_BINDING_RESPONSE); - EXPECT_EQ(client_.state(), State::OFF); + EXPECT_EQ(client_->state(), State::OFF); } -TEST_F(DtlsStunPiggybackControllerTest, NotSupportedByServerClientReceives) { - DisableSupport(server_); +TEST_P(DtlsStunPiggybackControllerTest, NotSupportedByServerClientReceives) { + DisableSupport(*server_); // Client to server got lost (or arrives late) SendServerToClientEmbedded(empty, STUN_BINDING_REQUEST); - EXPECT_EQ(client_.state(), State::OFF); + EXPECT_EQ(client_->state(), State::OFF); } -TEST_F(DtlsStunPiggybackControllerTest, NotSupportedByClient) { - DisableSupport(client_); +TEST_P(DtlsStunPiggybackControllerTest, NotSupportedByClient) { + DisableSupport(*client_); SendServerToClientEmbedded(empty, STUN_BINDING_REQUEST); SendClientToServerEmbedded(empty, STUN_BINDING_RESPONSE); - EXPECT_EQ(server_.state(), State::OFF); + EXPECT_EQ(server_->state(), State::OFF); } -TEST_F(DtlsStunPiggybackControllerTest, SomeRequestsDoNotGoThrough) { +TEST_P(DtlsStunPiggybackControllerTest, SomeRequestsDoNotGoThrough) { // Client to server got lost (or arrives late) // Flight 1 SendServerToClientEmbedded(empty, STUN_BINDING_REQUEST); SendClientToServerEmbedded(dtls_flight1, STUN_BINDING_RESPONSE); - EXPECT_EQ(server_.state(), State::CONFIRMED); - EXPECT_EQ(client_.state(), State::CONFIRMED); + EXPECT_EQ(server_->state(), State::CONFIRMED); + EXPECT_EQ(client_->state(), State::CONFIRMED); // Flight 1+2, server sent request got lost. SendClientToServerEmbedded(dtls_flight1, STUN_BINDING_REQUEST); SendServerToClientEmbedded(dtls_flight2, STUN_BINDING_RESPONSE); - EXPECT_EQ(server_.state(), State::CONFIRMED); - EXPECT_EQ(client_.state(), State::CONFIRMED); + EXPECT_EQ(server_->state(), State::CONFIRMED); + EXPECT_EQ(client_->state(), State::CONFIRMED); // Flight 3+4 SendClientToServerEmbedded(dtls_flight3, STUN_BINDING_REQUEST); SendServerToClientEmbedded(dtls_flight4, STUN_BINDING_RESPONSE); - EXPECT_EQ(server_.state(), State::PENDING); - EXPECT_EQ(client_.state(), State::PENDING); + EXPECT_EQ(server_->state(), State::PENDING); + EXPECT_EQ(client_->state(), State::PENDING); // Post-handshake ACK EXPECT_CALL(*this, ServerCompleteCallback(true)); SendClientToServerEmbedded(empty, STUN_BINDING_REQUEST); EXPECT_CALL(*this, ClientCompleteCallback(true)); SendServerToClientEmbedded(empty, STUN_BINDING_RESPONSE); - EXPECT_EQ(server_.state(), State::COMPLETE); - EXPECT_EQ(client_.state(), State::COMPLETE); + EXPECT_EQ(server_->state(), State::COMPLETE); + EXPECT_EQ(client_->state(), State::COMPLETE); } -TEST_F(DtlsStunPiggybackControllerTest, LossOnPostHandshakeAck) { +TEST_P(DtlsStunPiggybackControllerTest, LossOnPostHandshakeAck) { // Flight 1+2 SendClientToServerEmbedded(dtls_flight1, STUN_BINDING_REQUEST); - EXPECT_EQ(server_.state(), State::CONFIRMED); + EXPECT_EQ(server_->state(), State::CONFIRMED); SendServerToClientEmbedded(dtls_flight2, STUN_BINDING_RESPONSE); - EXPECT_EQ(client_.state(), State::CONFIRMED); + EXPECT_EQ(client_->state(), State::CONFIRMED); // Flight 3+4 SendClientToServerEmbedded(dtls_flight3, STUN_BINDING_REQUEST); SendServerToClientEmbedded(dtls_flight4, STUN_BINDING_RESPONSE); - EXPECT_EQ(server_.state(), State::PENDING); - EXPECT_EQ(client_.state(), State::PENDING); + EXPECT_EQ(server_->state(), State::PENDING); + EXPECT_EQ(client_->state(), State::PENDING); // Post-handshake ACK. Client to server gets lost EXPECT_CALL(*this, ClientCompleteCallback(true)); SendServerToClientEmbedded(empty, STUN_BINDING_REQUEST); EXPECT_CALL(*this, ServerCompleteCallback(true)); SendClientToServerEmbedded(empty, STUN_BINDING_RESPONSE); - EXPECT_EQ(server_.state(), State::COMPLETE); - EXPECT_EQ(client_.state(), State::COMPLETE); + EXPECT_EQ(server_->state(), State::COMPLETE); + EXPECT_EQ(client_->state(), State::COMPLETE); } -TEST_F(DtlsStunPiggybackControllerTest, +TEST_P(DtlsStunPiggybackControllerTest, UnsupportedStateAfterFallbackHandshakeRemainsOff) { - DisableSupport(client_); - DisableSupport(server_); + DisableSupport(*client_); + DisableSupport(*server_); // Set DTLS complete after normal handshake. - client_.SetDtlsHandshakeComplete(/*is_client=*/true, /*is_dtls13=*/false); - EXPECT_EQ(client_.state(), State::OFF); - server_.SetDtlsHandshakeComplete(/*is_client=*/false, /*is_dtls13=*/false); - EXPECT_EQ(server_.state(), State::OFF); + client_->SetDtlsHandshakeComplete(/*is_client=*/true, /*is_dtls13=*/false); + EXPECT_EQ(client_->state(), State::OFF); + server_->SetDtlsHandshakeComplete(/*is_client=*/false, /*is_dtls13=*/false); + EXPECT_EQ(server_->state(), State::OFF); +} + +TEST_P(DtlsStunPiggybackControllerTest, DtlsFailed) { + SendClientToServerEmbedded(dtls_flight1, STUN_BINDING_REQUEST); + server_->CapturePacket(dtls_flight2); + server_->Flush(); + ASSERT_EQ(server_->state(), State::CONFIRMED); + ASSERT_THAT(server_->GetPending(), SizeIs(1)); + + EXPECT_CALL(*this, ServerCompleteCallback(false)); + server_->SetDtlsFailed(); + EXPECT_EQ(server_->state(), State::OFF); + EXPECT_THAT(server_->GetPending(), IsEmpty()); + EXPECT_EQ(server_->GetDataToPiggyback(STUN_BINDING_RESPONSE), std::nullopt); + EXPECT_EQ(server_->GetAckToPiggyback(STUN_BINDING_RESPONSE), std::nullopt); } -TEST_F(DtlsStunPiggybackControllerTest, BasicHandshakeAckData) { - EXPECT_EQ(server_.GetAckToPiggyback(STUN_BINDING_RESPONSE), +TEST_P(DtlsStunPiggybackControllerTest, BasicHandshakeAckData) { + EXPECT_EQ(server_->GetAckToPiggyback(STUN_BINDING_RESPONSE), std::vector({})); - EXPECT_EQ(client_.GetAckToPiggyback(STUN_BINDING_RESPONSE), + EXPECT_EQ(client_->GetAckToPiggyback(STUN_BINDING_RESPONSE), std::vector({})); // Flight 1+2 SendClientToServerEmbedded(dtls_flight1, STUN_BINDING_REQUEST); SendServerToClientEmbedded(dtls_flight2, STUN_BINDING_RESPONSE); - EXPECT_THAT(*server_.GetAckToPiggyback(STUN_BINDING_REQUEST), + EXPECT_THAT(*server_->GetAckToPiggyback(STUN_BINDING_REQUEST), ElementsAreArray({ComputeDtlsPacketHash(dtls_flight1)})); - EXPECT_THAT(*client_.GetAckToPiggyback(STUN_BINDING_RESPONSE), + EXPECT_THAT(*client_->GetAckToPiggyback(STUN_BINDING_RESPONSE), ElementsAreArray({ComputeDtlsPacketHash(dtls_flight2)})); // Flight 3+4 SendClientToServerEmbedded(dtls_flight3, STUN_BINDING_REQUEST); SendServerToClientEmbedded(dtls_flight4, STUN_BINDING_RESPONSE); - EXPECT_THAT(*server_.GetAckToPiggyback(STUN_BINDING_RESPONSE), + EXPECT_THAT(*server_->GetAckToPiggyback(STUN_BINDING_RESPONSE), ElementsAreArray({ ComputeDtlsPacketHash(dtls_flight1), ComputeDtlsPacketHash(dtls_flight3), })); - EXPECT_THAT(*client_.GetAckToPiggyback(STUN_BINDING_REQUEST), + EXPECT_THAT(*client_->GetAckToPiggyback(STUN_BINDING_REQUEST), ElementsAreArray({ ComputeDtlsPacketHash(dtls_flight2), ComputeDtlsPacketHash(dtls_flight4), @@ -448,35 +465,35 @@ TEST_F(DtlsStunPiggybackControllerTest, BasicHandshakeAckData) { SendServerToClientEmbedded(empty, STUN_BINDING_REQUEST); EXPECT_CALL(*this, ServerCompleteCallback); SendClientToServerEmbedded(empty, STUN_BINDING_RESPONSE); - EXPECT_EQ(server_.state(), State::COMPLETE); - EXPECT_EQ(client_.state(), State::COMPLETE); - EXPECT_EQ(server_.GetAckToPiggyback(STUN_BINDING_RESPONSE), std::nullopt); - EXPECT_EQ(client_.GetAckToPiggyback(STUN_BINDING_REQUEST), std::nullopt); + EXPECT_EQ(server_->state(), State::COMPLETE); + EXPECT_EQ(client_->state(), State::COMPLETE); + EXPECT_EQ(server_->GetAckToPiggyback(STUN_BINDING_RESPONSE), std::nullopt); + EXPECT_EQ(client_->GetAckToPiggyback(STUN_BINDING_REQUEST), std::nullopt); } -TEST_F(DtlsStunPiggybackControllerTest, UnwrappedHandshakeAckData) { - EXPECT_EQ(server_.GetAckToPiggyback(STUN_BINDING_RESPONSE), +TEST_P(DtlsStunPiggybackControllerTest, UnwrappedHandshakeAckData) { + EXPECT_EQ(server_->GetAckToPiggyback(STUN_BINDING_RESPONSE), std::vector({})); - EXPECT_EQ(client_.GetAckToPiggyback(STUN_BINDING_RESPONSE), + EXPECT_EQ(client_->GetAckToPiggyback(STUN_BINDING_RESPONSE), std::vector({})); // Flight 1+2 (embedded) SendClientToServerEmbedded(dtls_flight1, STUN_BINDING_REQUEST); SendServerToClientEmbedded(dtls_flight2, STUN_BINDING_RESPONSE); - EXPECT_THAT(*server_.GetAckToPiggyback(STUN_BINDING_REQUEST), + EXPECT_THAT(*server_->GetAckToPiggyback(STUN_BINDING_REQUEST), ElementsAreArray({ComputeDtlsPacketHash(dtls_flight1)})); - EXPECT_THAT(*client_.GetAckToPiggyback(STUN_BINDING_RESPONSE), + EXPECT_THAT(*client_->GetAckToPiggyback(STUN_BINDING_RESPONSE), ElementsAreArray({ComputeDtlsPacketHash(dtls_flight2)})); // Flight 3+4 (not embedded) SendClientToServerDtls(dtls_flight3); SendServerToClientDtls(dtls_flight4); - EXPECT_THAT(*server_.GetAckToPiggyback(STUN_BINDING_REQUEST), + EXPECT_THAT(*server_->GetAckToPiggyback(STUN_BINDING_REQUEST), ElementsAreArray({ ComputeDtlsPacketHash(dtls_flight1), ComputeDtlsPacketHash(dtls_flight3), })); - EXPECT_THAT(*client_.GetAckToPiggyback(STUN_BINDING_REQUEST), + EXPECT_THAT(*client_->GetAckToPiggyback(STUN_BINDING_REQUEST), ElementsAreArray({ ComputeDtlsPacketHash(dtls_flight2), ComputeDtlsPacketHash(dtls_flight4), @@ -487,19 +504,19 @@ TEST_F(DtlsStunPiggybackControllerTest, UnwrappedHandshakeAckData) { SendServerToClientEmbedded(empty, STUN_BINDING_REQUEST); EXPECT_CALL(*this, ServerCompleteCallback); SendClientToServerEmbedded(empty, STUN_BINDING_RESPONSE); - EXPECT_EQ(server_.state(), State::COMPLETE); - EXPECT_EQ(client_.state(), State::COMPLETE); - EXPECT_EQ(server_.GetAckToPiggyback(STUN_BINDING_RESPONSE), std::nullopt); - EXPECT_EQ(client_.GetAckToPiggyback(STUN_BINDING_REQUEST), std::nullopt); + EXPECT_EQ(server_->state(), State::COMPLETE); + EXPECT_EQ(client_->state(), State::COMPLETE); + EXPECT_EQ(server_->GetAckToPiggyback(STUN_BINDING_RESPONSE), std::nullopt); + EXPECT_EQ(client_->GetAckToPiggyback(STUN_BINDING_REQUEST), std::nullopt); } -TEST_F(DtlsStunPiggybackControllerTest, AckDataNoDuplicates) { +TEST_P(DtlsStunPiggybackControllerTest, AckDataNoDuplicates) { // Flight 1+2 SendClientToServerEmbedded(dtls_flight1, STUN_BINDING_REQUEST); - EXPECT_THAT(*server_.GetAckToPiggyback(STUN_BINDING_REQUEST), + EXPECT_THAT(*server_->GetAckToPiggyback(STUN_BINDING_REQUEST), ElementsAreArray({ComputeDtlsPacketHash(dtls_flight1)})); SendClientToServerEmbedded(dtls_flight3, STUN_BINDING_REQUEST); - EXPECT_THAT(*server_.GetAckToPiggyback(STUN_BINDING_REQUEST), + EXPECT_THAT(*server_->GetAckToPiggyback(STUN_BINDING_REQUEST), ElementsAreArray({ ComputeDtlsPacketHash(dtls_flight1), ComputeDtlsPacketHash(dtls_flight3), @@ -507,76 +524,77 @@ TEST_F(DtlsStunPiggybackControllerTest, AckDataNoDuplicates) { // Receive Flight 1 again, no change expected. SendClientToServerEmbedded(dtls_flight1, STUN_BINDING_REQUEST); - EXPECT_THAT(*server_.GetAckToPiggyback(STUN_BINDING_REQUEST), + EXPECT_THAT(*server_->GetAckToPiggyback(STUN_BINDING_REQUEST), ElementsAreArray({ ComputeDtlsPacketHash(dtls_flight1), ComputeDtlsPacketHash(dtls_flight3), })); } -TEST_F(DtlsStunPiggybackControllerTest, AckDataNoDuplicatesFromDualReporting) { +TEST_P(DtlsStunPiggybackControllerTest, AckDataNoDuplicatesFromDualReporting) { std::unique_ptr attr_data = WrapInStun(STUN_ATTR_META_DTLS_IN_STUN, dtls_flight1); std::unique_ptr attr_ack; - if (auto ack = client_.GetAckToPiggyback(STUN_BINDING_REQUEST)) { + if (auto ack = client_->GetAckToPiggyback(STUN_BINDING_REQUEST)) { attr_ack = WrapInStun(STUN_ATTR_META_DTLS_IN_STUN_ACK, *ack); } ASSERT_THAT(attr_ack, NotNull()); - server_.ReportDataPiggybacked(attr_data->array_view(), - FromAckAttribute(attr_ack->array_view())); - server_.ReportDtlsPacket(dtls_flight1); - EXPECT_THAT(*server_.GetAckToPiggyback(STUN_BINDING_REQUEST), + server_->ReportDataPiggybacked(attr_data->array_view(), + FromAckAttribute(attr_ack->array_view())); + server_->ReportDtlsPacket(dtls_flight1); + EXPECT_THAT(*server_->GetAckToPiggyback(STUN_BINDING_REQUEST), ElementsAreArray({ComputeDtlsPacketHash(dtls_flight1)})); } -TEST_F(DtlsStunPiggybackControllerTest, IgnoresNonDtlsData) { +TEST_P(DtlsStunPiggybackControllerTest, IgnoresNonDtlsData) { std::vector ascii = {0x64, 0x72, 0x6f, 0x70, 0x6d, 0x65}; EXPECT_CALL(*this, ServerPacketSink).Times(0); - server_.ReportDataPiggybacked( + server_->ReportDataPiggybacked( WrapInStun(STUN_ATTR_META_DTLS_IN_STUN, ascii)->array_view(), std::nullopt); - EXPECT_EQ(0, server_.GetCountOfReceivedData()); + EXPECT_EQ(0, server_->GetCountOfReceivedData()); } -TEST_F(DtlsStunPiggybackControllerTest, DontSendAckedPackets) { - server_.CapturePacket(dtls_flight1); - server_.Flush(); - EXPECT_TRUE(server_.GetDataToPiggyback(STUN_BINDING_REQUEST).has_value()); - server_.ReportDataPiggybacked( +TEST_P(DtlsStunPiggybackControllerTest, DontSendAckedPackets) { + server_->CapturePacket(dtls_flight1); + server_->Flush(); + EXPECT_TRUE(server_->GetDataToPiggyback(STUN_BINDING_REQUEST).has_value()); + server_->ReportDataPiggybacked( std::nullopt, std::vector({ComputeDtlsPacketHash(dtls_flight1)})); - // No unacked packet exists. - EXPECT_FALSE(server_.GetDataToPiggyback(STUN_BINDING_REQUEST).has_value()); + // No unacked packet exists, i.e. empty response. + auto response = server_->GetDataToPiggyback(STUN_BINDING_REQUEST); + EXPECT_TRUE(response && response->empty()); } -TEST_F(DtlsStunPiggybackControllerTest, LimitAckSize) { +TEST_P(DtlsStunPiggybackControllerTest, LimitAckSize) { std::vector dtls_flight5 = FakeDtlsPacket(0x5487); - server_.ReportDataPiggybacked( + server_->ReportDataPiggybacked( WrapInStun(STUN_ATTR_META_DTLS_IN_STUN, dtls_flight1)->array_view(), std::nullopt); - EXPECT_EQ(server_.GetAckToPiggyback(STUN_BINDING_REQUEST)->size(), 1u); - server_.ReportDataPiggybacked( + EXPECT_EQ(server_->GetAckToPiggyback(STUN_BINDING_REQUEST)->size(), 1u); + server_->ReportDataPiggybacked( WrapInStun(STUN_ATTR_META_DTLS_IN_STUN, dtls_flight2)->array_view(), std::nullopt); - EXPECT_EQ(server_.GetAckToPiggyback(STUN_BINDING_REQUEST)->size(), 2u); - server_.ReportDataPiggybacked( + EXPECT_EQ(server_->GetAckToPiggyback(STUN_BINDING_REQUEST)->size(), 2u); + server_->ReportDataPiggybacked( WrapInStun(STUN_ATTR_META_DTLS_IN_STUN, dtls_flight3)->array_view(), std::nullopt); - EXPECT_EQ(server_.GetAckToPiggyback(STUN_BINDING_REQUEST)->size(), 3u); - server_.ReportDataPiggybacked( + EXPECT_EQ(server_->GetAckToPiggyback(STUN_BINDING_REQUEST)->size(), 3u); + server_->ReportDataPiggybacked( WrapInStun(STUN_ATTR_META_DTLS_IN_STUN, dtls_flight4)->array_view(), std::nullopt); - EXPECT_EQ(server_.GetAckToPiggyback(STUN_BINDING_REQUEST)->size(), 4u); + EXPECT_EQ(server_->GetAckToPiggyback(STUN_BINDING_REQUEST)->size(), 4u); // Limit size of ack so that it does not grow unbounded. - server_.ReportDataPiggybacked( + server_->ReportDataPiggybacked( WrapInStun(STUN_ATTR_META_DTLS_IN_STUN, dtls_flight5)->array_view(), std::nullopt); - EXPECT_EQ(server_.GetAckToPiggyback(STUN_BINDING_REQUEST)->size(), + EXPECT_EQ(server_->GetAckToPiggyback(STUN_BINDING_REQUEST)->size(), DtlsStunPiggybackController::kMaxAckSize); - EXPECT_THAT(*server_.GetAckToPiggyback(STUN_BINDING_REQUEST), + EXPECT_THAT(*server_->GetAckToPiggyback(STUN_BINDING_REQUEST), ElementsAreArray({ ComputeDtlsPacketHash(dtls_flight2), ComputeDtlsPacketHash(dtls_flight3), @@ -585,65 +603,265 @@ TEST_F(DtlsStunPiggybackControllerTest, LimitAckSize) { })); } -TEST_F(DtlsStunPiggybackControllerTest, EmptyDataDoesNotClearAck) { +TEST_P(DtlsStunPiggybackControllerTest, EmptyDataDoesNotClearAck) { std::vector dtls_flight5 = FakeDtlsPacket(0x5487); - server_.ReportDataPiggybacked( + server_->ReportDataPiggybacked( WrapInStun(STUN_ATTR_META_DTLS_IN_STUN, dtls_flight1)->array_view(), std::nullopt); - EXPECT_EQ(server_.GetAckToPiggyback(STUN_BINDING_REQUEST)->size(), 1u); + EXPECT_EQ(server_->GetAckToPiggyback(STUN_BINDING_REQUEST)->size(), 1u); // The fact that we don't get any data does not mean that // we can clear the ack list. // a) packets can be arbitrary reordered. // b) the peer might be needing 2 packets (ie. pqc) to produce // a return packet and only one of them has arrived. - server_.ReportDataPiggybacked( + server_->ReportDataPiggybacked( std::nullopt, std::vector({ComputeDtlsPacketHash(dtls_flight1)})); - EXPECT_EQ(server_.GetAckToPiggyback(STUN_BINDING_REQUEST)->size(), 1u); + EXPECT_EQ(server_->GetAckToPiggyback(STUN_BINDING_REQUEST)->size(), 1u); } -TEST_F(DtlsStunPiggybackControllerTest, MultiPacketRoundRobin) { +TEST_P(DtlsStunPiggybackControllerTest, NoEmptyDataInPending) { + std::vector packet = FakeDtlsPacket(0x5487); + + server_->ReportDataPiggybacked( + WrapInStun(STUN_ATTR_META_DTLS_IN_STUN, packet)->array_view(), + std::nullopt); + // If this is one of the two packets of a PQC client hello the server + // does not have a response yet. + auto response = server_->GetDataToPiggyback(STUN_BINDING_REQUEST); + EXPECT_TRUE(response && response->empty()); + EXPECT_EQ(server_->GetAckToPiggyback(STUN_BINDING_REQUEST)->size(), 1u); +} + +TEST_P(DtlsStunPiggybackControllerTest, MultiPacketRoundRobin) { // Let's pretend that a flight is 3 packets... - server_.CapturePacket(dtls_flight1); - server_.CapturePacket(dtls_flight2); - server_.CapturePacket(dtls_flight3); - server_.Flush(); - EXPECT_EQ(server_.GetDataToPiggyback(STUN_BINDING_REQUEST), + server_->CapturePacket(dtls_flight1); + server_->CapturePacket(dtls_flight2); + server_->CapturePacket(dtls_flight3); + server_->Flush(); + EXPECT_EQ(server_->GetDataToPiggyback(STUN_BINDING_REQUEST), std::string(dtls_flight1.begin(), dtls_flight1.end())); - EXPECT_EQ(server_.GetDataToPiggyback(STUN_BINDING_REQUEST), + EXPECT_EQ(server_->GetDataToPiggyback(STUN_BINDING_REQUEST), std::string(dtls_flight2.begin(), dtls_flight2.end())); - EXPECT_EQ(server_.GetDataToPiggyback(STUN_BINDING_REQUEST), + EXPECT_EQ(server_->GetDataToPiggyback(STUN_BINDING_REQUEST), std::string(dtls_flight3.begin(), dtls_flight3.end())); - server_.ReportDataPiggybacked( + server_->ReportDataPiggybacked( std::nullopt, std::vector({ComputeDtlsPacketHash(dtls_flight1)})); - EXPECT_EQ(server_.GetDataToPiggyback(STUN_BINDING_REQUEST), + EXPECT_EQ(server_->GetDataToPiggyback(STUN_BINDING_REQUEST), std::string(dtls_flight2.begin(), dtls_flight2.end())); - EXPECT_EQ(server_.GetDataToPiggyback(STUN_BINDING_REQUEST), + EXPECT_EQ(server_->GetDataToPiggyback(STUN_BINDING_REQUEST), std::string(dtls_flight3.begin(), dtls_flight3.end())); - server_.ReportDataPiggybacked( + server_->ReportDataPiggybacked( std::nullopt, std::vector({ComputeDtlsPacketHash(dtls_flight3)})); - EXPECT_EQ(server_.GetDataToPiggyback(STUN_BINDING_REQUEST), + EXPECT_EQ(server_->GetDataToPiggyback(STUN_BINDING_REQUEST), std::string(dtls_flight2.begin(), dtls_flight2.end())); - EXPECT_EQ(server_.GetDataToPiggyback(STUN_BINDING_REQUEST), + EXPECT_EQ(server_->GetDataToPiggyback(STUN_BINDING_REQUEST), std::string(dtls_flight2.begin(), dtls_flight2.end())); } -TEST_F(DtlsStunPiggybackControllerTest, DuplicateAck) { - server_.CapturePacket(dtls_flight1); - server_.Flush(); - server_.ReportDataPiggybacked( +TEST_P(DtlsStunPiggybackControllerTest, DuplicateAck) { + server_->CapturePacket(dtls_flight1); + server_->Flush(); + server_->ReportDataPiggybacked( std::nullopt, std::vector({ComputeDtlsPacketHash(dtls_flight1), ComputeDtlsPacketHash(dtls_flight1)})); } +// In DTLS 1.3 the last flight is sent by the client, so the roles in the +// termination are swapped compared to DTLS 1.2. The server is done once it +// has sent its own flight and must keep retransmitting it until acked. +TEST_P(DtlsStunPiggybackControllerTest, + Dtls13KeepsPendingPacketsOnHandshakeComplete) { + SendClientToServerEmbedded(dtls_flight1, STUN_BINDING_REQUEST); + server_->CapturePacket(dtls_flight2); + server_->Flush(); + ASSERT_THAT(server_->GetPending(), SizeIs(1)); + + server_->SetDtlsHandshakeComplete(/*is_client=*/false, /*is_dtls13=*/true); + EXPECT_EQ(server_->state(), State::PENDING); + EXPECT_THAT(server_->GetPending(), SizeIs(1)); + + // Same for the client, whose flight 3 is the last one of the handshake. + SendServerToClientEmbedded(dtls_flight2, STUN_BINDING_RESPONSE); + client_->CapturePacket(dtls_flight3); + client_->Flush(); + client_->SetDtlsHandshakeComplete(/*is_client=*/true, /*is_dtls13=*/true); + EXPECT_EQ(client_->state(), State::PENDING); + EXPECT_THAT(client_->GetPending(), SizeIs(1)); +} + +TEST_F(DtlsStunPiggybackControllerSpedTest, Dtls13Handshake) { + // Flight 1+2. The 1.3 server can send application data once it has sent its + // own flight, i.e. before it has seen the client Finished. + SendClientToServerEmbedded(dtls_flight1, STUN_BINDING_REQUEST); + SendServerToClientEmbedded(dtls_flight2, STUN_BINDING_RESPONSE); + server_->SetDtlsHandshakeComplete(/*is_client=*/false, /*is_dtls13=*/true); + EXPECT_EQ(server_->state(), State::PENDING); + EXPECT_EQ(client_->state(), State::CONFIRMED); + + // Flight 3, the last one, acks flight 2 and empties the server stash. + SendClientToServerEmbedded(dtls_flight3, STUN_BINDING_REQUEST); + client_->SetDtlsHandshakeComplete(/*is_client=*/true, /*is_dtls13=*/true); + EXPECT_EQ(server_->state(), State::PENDING); + EXPECT_EQ(client_->state(), State::PENDING); + + // Closing handshake, ack-only in this direction. + EXPECT_CALL(*this, ClientCompleteCallback(true)); + SendServerToClientEmbedded(empty, STUN_BINDING_RESPONSE); + EXPECT_EQ(client_->state(), State::COMPLETE); + + EXPECT_CALL(*this, ServerCompleteCallback(true)); + SendClientToServerEmbedded(empty, STUN_BINDING_REQUEST); + EXPECT_EQ(server_->state(), State::COMPLETE); +} + +TEST_F(DtlsStunPiggybackControllerGoogSpedTest, Dtls13Handshake) { + SendClientToServerEmbedded(dtls_flight1, STUN_BINDING_REQUEST); + SendServerToClientEmbedded(dtls_flight2, STUN_BINDING_RESPONSE); + server_->SetDtlsHandshakeComplete(/*is_client=*/false, /*is_dtls13=*/true); + EXPECT_EQ(server_->state(), State::PENDING); + + // The ack for flight 2 arrives with flight 3, so the server empties its + // stash and completes inside the same call that processes the last flight. + EXPECT_CALL(*this, ServerCompleteCallback(true)); + SendClientToServerEmbedded(dtls_flight3, STUN_BINDING_REQUEST); + client_->SetDtlsHandshakeComplete(/*is_client=*/true, /*is_dtls13=*/true); + EXPECT_EQ(server_->state(), State::COMPLETE); + EXPECT_EQ(client_->state(), State::PENDING); + + // A COMPLETE peer sends neither attribute, so the ack for flight 3 never + // goes on the wire and the client keeps retransmitting it. + SendServerToClientEmbedded(empty, STUN_BINDING_RESPONSE); + EXPECT_EQ(client_->state(), State::PENDING); + EXPECT_THAT(client_->GetPending(), SizeIs(1)); + + // Only an application packet unblocks the client. + EXPECT_CALL(*this, ClientCompleteCallback(true)); + client_->ApplicationPacketReceived( + packet_.CopyAndSet(ReceivedIpPacket::kDtlsDecrypted)); + EXPECT_EQ(client_->state(), State::COMPLETE); +} + +TEST_F(DtlsStunPiggybackControllerGoogSpedTest, + BasicHandshakeCompleteWithDecryptedPacket) { + // Flight 1+2 + SendClientToServerEmbedded(dtls_flight1, STUN_BINDING_REQUEST); + EXPECT_EQ(server_->state(), State::CONFIRMED); + SendServerToClientEmbedded(dtls_flight2, STUN_BINDING_RESPONSE); + EXPECT_EQ(client_->state(), State::CONFIRMED); + + // Flight 3+4 + SendClientToServerEmbedded(dtls_flight3, STUN_BINDING_REQUEST); + SendServerToClientEmbedded(dtls_flight4, STUN_BINDING_RESPONSE); + EXPECT_EQ(server_->state(), State::PENDING); + EXPECT_EQ(client_->state(), State::PENDING); + + // Post-handshake ACK + EXPECT_CALL(*this, ClientCompleteCallback); + client_->ApplicationPacketReceived( + packet_.CopyAndSet(ReceivedIpPacket::kDtlsDecrypted)); + EXPECT_EQ(client_->state(), State::COMPLETE); + + EXPECT_CALL(*this, ServerCompleteCallback); + server_->ApplicationPacketReceived( + packet_.CopyAndSet(ReceivedIpPacket::kSrtpEncrypted)); + EXPECT_EQ(server_->state(), State::COMPLETE); +} + +TEST_F(DtlsStunPiggybackControllerGoogSpedTest, + BasicHandshakeEarlySrtpDoesNotComplete) { + // Flight 1+2 + SendClientToServerEmbedded(dtls_flight1, STUN_BINDING_REQUEST); + EXPECT_EQ(server_->state(), State::CONFIRMED); + SendServerToClientEmbedded(dtls_flight2, STUN_BINDING_RESPONSE); + EXPECT_EQ(client_->state(), State::CONFIRMED); + + // Flight 3 + SendClientToServerEmbedded(dtls_flight3, STUN_BINDING_REQUEST); + EXPECT_EQ(server_->state(), State::CONFIRMED); + + // An srtp packet arriving before reaching PENDING state. + server_->ApplicationPacketReceived( + packet_.CopyAndSet(ReceivedIpPacket::kSrtpEncrypted)); + EXPECT_EQ(server_->state(), State::CONFIRMED); + + // Flight 4 + SendServerToClientEmbedded(dtls_flight4, STUN_BINDING_RESPONSE); + EXPECT_EQ(server_->state(), State::PENDING); + EXPECT_EQ(client_->state(), State::PENDING); + + // Post-handshake ACK + EXPECT_CALL(*this, ClientCompleteCallback); + client_->ApplicationPacketReceived( + packet_.CopyAndSet(ReceivedIpPacket::kDtlsDecrypted)); + EXPECT_EQ(client_->state(), State::COMPLETE); + + EXPECT_CALL(*this, ServerCompleteCallback); + server_->ApplicationPacketReceived( + packet_.CopyAndSet(ReceivedIpPacket::kSrtpEncrypted)); + EXPECT_EQ(server_->state(), State::COMPLETE); +} + +// Application packets are not a completion signal here, the closing handshake +// is. +TEST_F(DtlsStunPiggybackControllerSpedTest, + ApplicationPacketDoesNotCompleteHandshake) { + SendClientToServerEmbedded(dtls_flight1, STUN_BINDING_REQUEST); + SendServerToClientEmbedded(dtls_flight2, STUN_BINDING_RESPONSE); + SendClientToServerEmbedded(dtls_flight3, STUN_BINDING_REQUEST); + SendServerToClientEmbedded(dtls_flight4, STUN_BINDING_RESPONSE); + ASSERT_EQ(server_->state(), State::PENDING); + ASSERT_EQ(client_->state(), State::PENDING); + + EXPECT_CALL(*this, ClientCompleteCallback).Times(0); + EXPECT_CALL(*this, ServerCompleteCallback).Times(0); + client_->ApplicationPacketReceived( + packet_.CopyAndSet(ReceivedIpPacket::kDtlsDecrypted)); + server_->ApplicationPacketReceived( + packet_.CopyAndSet(ReceivedIpPacket::kSrtpEncrypted)); + EXPECT_EQ(client_->state(), State::PENDING); + EXPECT_EQ(server_->state(), State::PENDING); +} + +// The DTLS 1.2 client is complete when it receives flight 4, which the server +// only sends after flight 3 arrived. Its stash can therefore be dropped. +TEST_F(DtlsStunPiggybackControllerSpedTest, + Dtls12ClientClearsPendingPacketsOnHandshakeComplete) { + SendClientToServerEmbedded(dtls_flight1, STUN_BINDING_REQUEST); + SendServerToClientEmbedded(dtls_flight2, STUN_BINDING_RESPONSE); + client_->CapturePacket(dtls_flight3); + client_->Flush(); + ASSERT_THAT(client_->GetPending(), SizeIs(1)); + + client_->SetDtlsHandshakeComplete(/*is_client=*/true, /*is_dtls13=*/false); + EXPECT_EQ(client_->state(), State::PENDING); + EXPECT_THAT(client_->GetPending(), IsEmpty()); + EXPECT_EQ(client_->GetDataToPiggyback(STUN_BINDING_REQUEST), std::nullopt); +} + +// The default variant keeps retransmitting flight 3 instead, since an empty +// stash is what terminates the session there. +TEST_F(DtlsStunPiggybackControllerGoogSpedTest, + Dtls12ClientKeepsPendingPacketsOnHandshakeComplete) { + SendClientToServerEmbedded(dtls_flight1, STUN_BINDING_REQUEST); + SendServerToClientEmbedded(dtls_flight2, STUN_BINDING_RESPONSE); + client_->CapturePacket(dtls_flight3); + client_->Flush(); + ASSERT_THAT(client_->GetPending(), SizeIs(1)); + + client_->SetDtlsHandshakeComplete(/*is_client=*/true, /*is_dtls13=*/false); + EXPECT_EQ(client_->state(), State::PENDING); + EXPECT_THAT(client_->GetPending(), SizeIs(1)); +} + } // namespace webrtc diff --git a/p2p/dtls/dtls_transport.cc b/p2p/dtls/dtls_transport.cc index 57ccb59ee3..6aa5eeb25e 100644 --- a/p2p/dtls/dtls_transport.cc +++ b/p2p/dtls/dtls_transport.cc @@ -40,6 +40,8 @@ #include "p2p/base/packet_transport_internal.h" #include "p2p/dtls/dtls_stun_piggyback_callbacks.h" #include "p2p/dtls/dtls_stun_piggyback_controller.h" +#include "p2p/dtls/dtls_stun_piggyback_controller_interface.h" +#include "p2p/dtls/dtls_stun_piggyback_controller_sped.h" #include "p2p/dtls/dtls_transport_internal.h" #include "p2p/dtls/dtls_utils.h" #include "rtc_base/async_packet_socket.h" @@ -144,7 +146,7 @@ StreamInterfaceChannel::StreamInterfaceChannel( packets_(kMaxPendingPackets, kMaxDtlsPacketLen) {} void StreamInterfaceChannel::SetDtlsStunPiggybackController( - DtlsStunPiggybackController* dtls_stun_piggyback_controller) { + DtlsStunPiggybackControllerInterface* dtls_stun_piggyback_controller) { dtls_stun_piggyback_controller_ = dtls_stun_piggyback_controller; } @@ -257,17 +259,29 @@ DtlsTransportInternalImpl::DtlsTransportInternalImpl( srtp_ciphers_(crypto_options.GetSupportedDtlsSrtpCryptoSuites()), ephemeral_key_exchange_cipher_groups_( crypto_options.ephemeral_key_exchange_cipher_groups.GetEnabled()), - ssl_max_version_(max_version), - dtls_stun_piggyback_controller_( - [this](std::span piggybacked_dtls_packet) { - if (piggybacked_dtls_callback_ == nullptr) { - return; - } - piggybacked_dtls_callback_( - this, - ReceivedIpPacket(piggybacked_dtls_packet, SocketAddress())); - }, - [this](bool success) { CompleteDtlsInStun(success); }) { + ssl_max_version_(max_version) { + auto dtls_data_callback = + [this](std::span piggybacked_dtls_packet) { + if (piggybacked_dtls_callback_ == nullptr) { + return; + } + piggybacked_dtls_callback_( + this, ReceivedIpPacket(piggybacked_dtls_packet, SocketAddress())); + }; + auto piggyback_complete_callback = [this](bool success) { + CompleteDtlsInStun(success); + }; + if (env_.field_trials().IsEnabled("WebRTC-DtlsStunPiggybackControllerSped")) { + dtls_stun_piggyback_controller_ = + std::make_unique( + std::move(dtls_data_callback), + std::move(piggyback_complete_callback)); + } else { + dtls_stun_piggyback_controller_ = + std::make_unique( + std::move(dtls_data_callback), + std::move(piggyback_complete_callback)); + } RTC_DCHECK(ice_transport_); ConnectToIceTransport(); if (SSLStreamAdapter::IsBoringSsl()) { @@ -303,7 +317,7 @@ void DtlsTransportInternalImpl::CompleteDtlsInStun(bool success) { } ice_transport()->ResetDtlsStunPiggybackCallbacks(); - DeregisterReceivedPacketCallback(&dtls_stun_piggyback_controller_); + DeregisterReceivedPacketCallback(dtls_stun_piggyback_controller_.get()); } DtlsTransportState DtlsTransportInternalImpl::dtls_state() const { @@ -523,7 +537,7 @@ bool DtlsTransportInternalImpl::AppendSrtpKeyingMaterial( bool DtlsTransportInternalImpl::SetupDtls() { RTC_DCHECK(dtls_role_); - if (SSLStreamAdapter::IsBoringSsl()) { + if (SSLStreamAdapter::IsBoringSsl() && !dtls_in_stun_disabled_) { dtls_in_stun_ = ice_transport()->config().dtls_handshake_in_stun; } @@ -533,13 +547,13 @@ bool DtlsTransportInternalImpl::SetupDtls() { if (dtls_in_stun_ && !dtls_in_stun_complete_) { downward_ptr->SetDtlsStunPiggybackController( - &dtls_stun_piggyback_controller_); + dtls_stun_piggyback_controller_.get()); RegisterReceivedPacketCallback( - &dtls_stun_piggyback_controller_, + dtls_stun_piggyback_controller_.get(), [this](webrtc::PacketTransportInternal* transport, const ReceivedIpPacket& packet) { - dtls_stun_piggyback_controller_.ApplicationPacketReceived(packet); + dtls_stun_piggyback_controller_->ApplicationPacketReceived(packet); }); } if (ssl_stream_factory_) { @@ -708,6 +722,37 @@ int DtlsTransportInternalImpl::SendPacket( } } +void DtlsTransportInternalImpl::DisableDtlsInStun() { + RTC_DCHECK_RUN_ON(&thread_checker_); + // The remote description may arrive after the handshake started. + if (dtls_state() != DtlsTransportState::kNew) { + return; + } + dtls_in_stun_disabled_ = true; + dtls_in_stun_ = false; + peer_supports_dtls_in_stun_ = false; + if (ice_transport_) { + ice_transport_->internal()->ResetDtlsStunPiggybackCallbacks(); + } + if (downward_) { + downward_->SetDtlsStunPiggybackController(nullptr); + } +} + +void DtlsTransportInternalImpl::MaybeStartDtlsInStun() { + RTC_DCHECK_RUN_ON(&thread_checker_); + if (peer_supports_dtls_in_stun_) { + return; + } + peer_supports_dtls_in_stun_ = true; + dtls_in_stun_disabled_ = false; + // The remote description may arrive after ICE became writable and DTLS + // already started. + if (dtls_state() == DtlsTransportState::kNew) { + MaybeStartDtls(); + } +} + IceTransportInternal* DtlsTransportInternalImpl::ice_transport() { return ice_transport_->internal(); } @@ -775,9 +820,9 @@ void DtlsTransportInternalImpl::ConnectToIceTransport() { std::optional data; std::optional> ack; if (dtls_in_stun_) { - data = dtls_stun_piggyback_controller_.GetDataToPiggyback( + data = dtls_stun_piggyback_controller_->GetDataToPiggyback( stun_message_type); - ack = dtls_stun_piggyback_controller_.GetAckToPiggyback( + ack = dtls_stun_piggyback_controller_->GetAckToPiggyback( stun_message_type); } return std::make_pair(data, ack); @@ -787,7 +832,7 @@ void DtlsTransportInternalImpl::ConnectToIceTransport() { if (!dtls_in_stun_) { return; } - dtls_stun_piggyback_controller_.ReportDataPiggybacked(data, acks); + dtls_stun_piggyback_controller_->ReportDataPiggybacked(data, acks); })); SetPiggybackDtlsDataCallback([this](PacketTransportInternal* transport, const ReceivedIpPacket& packet) { @@ -893,6 +938,10 @@ void DtlsTransportInternalImpl::OnReadPacket(PacketTransportInternal* transport, RTC_DCHECK_RUN_ON(&thread_checker_); RTC_DCHECK(transport == ice_transport()); + if (piggybacked) { + peer_supports_dtls_in_stun_ = true; + } + if (!dtls_active_) { // Not doing DTLS. NotifyPacketReceived(packet); @@ -999,7 +1048,7 @@ void DtlsTransportInternalImpl::OnDtlsEvent(int sig, int err) { int ssl_version_bytes; bool ret = dtls_->GetSslVersionBytes(&ssl_version_bytes); RTC_DCHECK(ret); - dtls_stun_piggyback_controller_.SetDtlsHandshakeComplete( + dtls_stun_piggyback_controller_->SetDtlsHandshakeComplete( dtls_role_ == SSL_CLIENT, ssl_version_bytes == kDtls13VersionBytes); set_dtls_state(DtlsTransportState::kConnected); set_writable(true); @@ -1063,7 +1112,8 @@ void DtlsTransportInternalImpl::OnNetworkRouteChanged( void DtlsTransportInternalImpl::MaybeStartDtls() { // When adding the DTLS handshake in STUN we want to call StartSSL even // before the ICE transport is ready. - if (dtls_ && (ice_transport()->writable() || dtls_in_stun_)) { + if (dtls_ && (ice_transport()->writable() || + (dtls_in_stun_ && peer_supports_dtls_in_stun_))) { ConfigureHandshakeTimeout(); RTC_LOG(LS_INFO) @@ -1200,7 +1250,7 @@ void DtlsTransportInternalImpl::set_dtls_state(DtlsTransportState state) { } } if (dtls_state_ == DtlsTransportState::kFailed) { - dtls_stun_piggyback_controller_.SetDtlsFailed(); + dtls_stun_piggyback_controller_->SetDtlsFailed(); } SendDtlsState(this, state); } @@ -1238,8 +1288,8 @@ void DtlsTransportInternalImpl::UpdateHandshakeTimeout() { const auto rtt_ms = ice_transport()->GetRttEstimate(); int delay_ms = ComputeRetransmissionTimeout( rtt_ms.value_or(kDefaultHandshakeEstimateRttMs)); - if (dtls_stun_piggyback_controller_.state() == - DtlsStunPiggybackController::State::OFF && + if (dtls_stun_piggyback_controller_->state() == + DtlsStunPiggybackControllerInterface::State::OFF && dtls_role_ == SSL_CLIENT) { // We sent one STUN BINDING request with an embedded DTLS packet and // discovered that peer does not support DtlsInStun. The DTLS packet will be @@ -1264,33 +1314,30 @@ void DtlsTransportInternalImpl::SetPiggybackDtlsDataCallback( bool DtlsTransportInternalImpl::IsDtlsPiggybackSupportedByPeer() { RTC_DCHECK_RUN_ON(&thread_checker_); - return dtls_in_stun_ && (dtls_stun_piggyback_controller_.state() != - DtlsStunPiggybackController::State::OFF); + return dtls_in_stun_ && (dtls_stun_piggyback_controller_->state() != + DtlsStunPiggybackControllerInterface::State::OFF); } bool DtlsTransportInternalImpl::WasDtlsCompletedByPiggybacking() { RTC_DCHECK_RUN_ON(&thread_checker_); - return dtls_in_stun_ && (dtls_stun_piggyback_controller_.state() == - DtlsStunPiggybackController::State::COMPLETE || - dtls_stun_piggyback_controller_.state() == - DtlsStunPiggybackController::State::PENDING); + return dtls_in_stun_ && + (dtls_stun_piggyback_controller_->state() == + DtlsStunPiggybackControllerInterface::State::COMPLETE || + dtls_stun_piggyback_controller_->state() == + DtlsStunPiggybackControllerInterface::State::PENDING); } void DtlsTransportInternalImpl::FlushPendingDtlsPacket() { RTC_DCHECK_RUN_ON(&thread_checker_); - if (dtls_stun_piggyback_controller_.state() == - DtlsStunPiggybackController::State::COMPLETE) { + if (dtls_stun_piggyback_controller_->state() == + DtlsStunPiggybackControllerInterface::State::COMPLETE) { // We're done. return; } if (ice_transport()->writable() && dtls_in_stun_) { - auto data_to_send = dtls_stun_piggyback_controller_.GetPending(); - if (data_to_send.empty()) { - // No data to send, we're done. - return; - } + auto data_to_send = dtls_stun_piggyback_controller_->GetPending(); for (const auto& packet : data_to_send) { AsyncSocketPacketOptions packet_options; ice_transport()->SendPacket(reinterpret_cast(packet.data()), @@ -1311,7 +1358,7 @@ int DtlsTransportInternalImpl::GetStunDataCount() const { if (!dtls_in_stun_) { return 0; } - return dtls_stun_piggyback_controller_.GetCountOfReceivedData(); + return dtls_stun_piggyback_controller_->GetCountOfReceivedData(); } } // namespace webrtc diff --git a/p2p/dtls/dtls_transport.h b/p2p/dtls/dtls_transport.h index 7c8b86df41..3450d6222a 100644 --- a/p2p/dtls/dtls_transport.h +++ b/p2p/dtls/dtls_transport.h @@ -33,7 +33,7 @@ #include "api/units/timestamp.h" #include "p2p/base/ice_transport_internal.h" #include "p2p/base/packet_transport_internal.h" -#include "p2p/dtls/dtls_stun_piggyback_controller.h" +#include "p2p/dtls/dtls_stun_piggyback_controller_interface.h" #include "p2p/dtls/dtls_transport_internal.h" #include "p2p/dtls/dtls_utils.h" #include "rtc_base/async_packet_socket.h" @@ -69,7 +69,7 @@ class StreamInterfaceChannel : public StreamInterface { explicit StreamInterfaceChannel(IceTransportInternal* ice_transport); void SetDtlsStunPiggybackController( - DtlsStunPiggybackController* dtls_stun_piggyback_controller); + DtlsStunPiggybackControllerInterface* dtls_stun_piggyback_controller); StreamInterfaceChannel(const StreamInterfaceChannel&) = delete; StreamInterfaceChannel& operator=(const StreamInterfaceChannel&) = delete; @@ -98,7 +98,7 @@ class StreamInterfaceChannel : public StreamInterface { private: IceTransportInternal* const ice_transport_; // owned by DtlsTransport - DtlsStunPiggybackController* dtls_stun_piggyback_controller_ = + DtlsStunPiggybackControllerInterface* dtls_stun_piggyback_controller_ = nullptr; // owned by DtlsTransport StreamState state_ RTC_GUARDED_BY(callback_sequence_); BufferQueue packets_ RTC_GUARDED_BY(callback_sequence_); @@ -240,6 +240,11 @@ class DtlsTransportInternalImpl : public DtlsTransportInternal { bool AppendSrtpKeyingMaterial( ZeroOnFreeBuffer& keying_material) override; + // Disable DTLS-in-STUN. + void DisableDtlsInStun() override; + + void MaybeStartDtlsInStun() override; + IceTransportInternal* ice_transport() override; // For informational purposes. Tells if the DTLS handshake has finished. @@ -356,12 +361,17 @@ class DtlsTransportInternalImpl : public DtlsTransportInternal { // (so that we return PIGGYBACK_ACK to client if we get STUN_BINDING_REQUEST // directly). Maybe disabled in SetupDtls has been called. bool dtls_in_stun_ = false; + bool peer_supports_dtls_in_stun_ = false; + // Set when the remote description did not signal support; makes the + // decision survive SetupDtls() re-reading the ICE config. + bool dtls_in_stun_disabled_ = false; // Has DtlsInStun Complete been run? // This variable is used to prevent reinitializing after dtls-restart. bool dtls_in_stun_complete_ = false; // A controller for piggybacking DTLS in STUN. - DtlsStunPiggybackController dtls_stun_piggyback_controller_; + std::unique_ptr + dtls_stun_piggyback_controller_; absl::AnyInvocable piggybacked_dtls_callback_; diff --git a/p2p/dtls/dtls_transport_internal.h b/p2p/dtls/dtls_transport_internal.h index df287ee013..360daac5e6 100644 --- a/p2p/dtls/dtls_transport_internal.h +++ b/p2p/dtls/dtls_transport_internal.h @@ -100,6 +100,12 @@ class DtlsTransportInternal : public PacketTransportInternal { return false; } + // Disable DTLS-in-STUN. + virtual void DisableDtlsInStun() = 0; + + // Enable early DTLS start because the peer supports DTLS-in-STUN. + virtual void MaybeStartDtlsInStun() = 0; + // Set DTLS remote fingerprint and role. Must be after local identity set. virtual RTCError SetRemoteParameters(absl::string_view digest_alg, const uint8_t* digest, diff --git a/p2p/dtls/dtls_transport_unittest.cc b/p2p/dtls/dtls_transport_unittest.cc index 181550a288..ce0cd318bf 100644 --- a/p2p/dtls/dtls_transport_unittest.cc +++ b/p2p/dtls/dtls_transport_unittest.cc @@ -1093,15 +1093,18 @@ class DtlsTransportInternalImplVersionTest client2_.dtls_transport()->SetDtlsRole( config2.ssl_role.value_or(SSL_SERVER)); + // No SDP exchange in this fixture; stand in for JsepTransport's call. if (config1.dtls_in_stun) { auto config = client1_.fake_ice_transport()->config(); config.dtls_handshake_in_stun = true; client1_.fake_ice_transport()->SetIceConfig(config); + client1_.dtls_transport()->MaybeStartDtlsInStun(); } if (config2.dtls_in_stun) { auto config = client2_.fake_ice_transport()->config(); config.dtls_handshake_in_stun = true; client2_.fake_ice_transport()->SetIceConfig(config); + client2_.dtls_transport()->MaybeStartDtlsInStun(); } SetRemoteFingerprintFromCert(client1_.dtls_transport(), @@ -1598,6 +1601,59 @@ TEST_F(DtlsTransportInternalImplTest, TestImplicitRoleDetection) { IsRtcOk()); } +TEST_F(DtlsTransportInternalImplTest, + NoEarlyDtlsInStunStartWithoutPeerSupport) { + if (!SSLStreamAdapter::IsBoringSsl()) { + GTEST_SKIP() << "DTLS-in-STUN requires BoringSSL."; + } + PrepareDtls(KT_DEFAULT); + + client1_.SetupTransports(env_, ICEROLE_CONTROLLING); + + auto ice_config = client1_.fake_ice_transport()->config(); + ice_config.dtls_handshake_in_stun = true; + client1_.fake_ice_transport()->SetIceConfig(ice_config); + ASSERT_FALSE(client1_.fake_ice_transport()->writable()); + + client1_.dtls_transport()->SetDtlsRole(SSL_SERVER); + SetRemoteFingerprintFromCert(client1_.dtls_transport(), + client2_.certificate()); + + // Without DTLS-in-STUN, DTLS starts after ICE becomes writable. + EXPECT_EQ(client1_.dtls_transport()->dtls_state(), DtlsTransportState::kNew); + + client1_.fake_ice_transport()->SetWritable(true); + EXPECT_EQ(client1_.dtls_transport()->dtls_state(), + DtlsTransportState::kConnecting); +} + +// MaybeStartDtlsInStun() confirms peer support; DTLS starts before ICE is +// writable. +TEST_F(DtlsTransportInternalImplTest, + EarlyDtlsInStunStartAfterPeerSupportConfirmed) { + if (!SSLStreamAdapter::IsBoringSsl()) { + GTEST_SKIP() << "DTLS-in-STUN requires BoringSSL."; + } + PrepareDtls(KT_DEFAULT); + + client1_.SetupTransports(env_, ICEROLE_CONTROLLING); + + auto ice_config = client1_.fake_ice_transport()->config(); + ice_config.dtls_handshake_in_stun = true; + client1_.fake_ice_transport()->SetIceConfig(ice_config); + ASSERT_FALSE(client1_.fake_ice_transport()->writable()); + + client1_.dtls_transport()->MaybeStartDtlsInStun(); + + client1_.dtls_transport()->SetDtlsRole(SSL_SERVER); + SetRemoteFingerprintFromCert(client1_.dtls_transport(), + client2_.certificate()); + + // With DTLS-in-STUN, DTLS starts immediately. + EXPECT_EQ(client1_.dtls_transport()->dtls_state(), + DtlsTransportState::kConnecting); +} + // Test that packets are retransmitted according to the expected schedule. // Each time a timeout occurs, the retransmission timer should be doubled up to // 60 seconds. The timer defaults to 1 second, but for WebRTC we should be diff --git a/p2p/dtls/fake_dtls_transport.h b/p2p/dtls/fake_dtls_transport.h index bec2ea28cb..a12022287f 100644 --- a/p2p/dtls/fake_dtls_transport.h +++ b/p2p/dtls/fake_dtls_transport.h @@ -106,6 +106,9 @@ class FakeDtlsTransport : public DtlsTransportInternal { ice_transport_->DeregisterReceivedPacketCallback(this); } + void DisableDtlsInStun() override {} + void MaybeStartDtlsInStun() override {} + // Get inner fake ICE transport. FakeIceTransportInternal* fake_ice_transport() { return ice_transport_; } diff --git a/p2p/test/test_port.cc b/p2p/test/test_port.cc index 83cc7ecdc0..d34eb812fa 100644 --- a/p2p/test/test_port.cc +++ b/p2p/test/test_port.cc @@ -103,7 +103,7 @@ int TestPort::SendTo(std::span data, bool payload) { if (!payload) { auto msg = std::make_unique(); - auto buf = std::make_unique>(data); + auto buf = std::make_unique(data); ByteBufferReader read_buf(*buf); if (!msg->Read(&read_buf)) { return -1; diff --git a/p2p/test/test_port.h b/p2p/test/test_port.h index 13340279d2..169231888e 100644 --- a/p2p/test/test_port.h +++ b/p2p/test/test_port.h @@ -76,7 +76,7 @@ class TestPort : public Port { private: void OnSentPacket(AsyncPacketSocket* socket, const SentPacketInfo& sent_packet) override; - std::unique_ptr> last_stun_buf_; + std::unique_ptr last_stun_buf_; std::unique_ptr last_stun_msg_; int type_preference_ = 0; }; diff --git a/pc/BUILD.gn b/pc/BUILD.gn index c4847e0675..3caff2eb3b 100644 --- a/pc/BUILD.gn +++ b/pc/BUILD.gn @@ -184,6 +184,7 @@ rtc_library("ice_transport") { deps = [ "../api:ice_transport_interface", "../api:sequence_checker", + "../p2p:ice_transport_internal", "../rtc_base:checks", "../rtc_base:macromagic", "../rtc_base:threading", @@ -3650,6 +3651,7 @@ if (rtc_include_tests && !build_with_chromium) { "../p2p:p2p_constants", "../p2p:transport_description", "../rtc_base:rtc_event", + "../rtc_base:ssl_adapter", "../rtc_base:stringutils", "../rtc_base:threading", "../system_wrappers:metrics", @@ -3695,6 +3697,7 @@ if (rtc_include_tests && !build_with_chromium) { "../media:codec", "../media:media_constants", "../media:stream_params", + "../rtc_base:ssl_adapter", "../rtc_base:threading", "../system_wrappers:metrics", "../test:create_test_field_trials", diff --git a/pc/jsep_transport.cc b/pc/jsep_transport.cc index 7605a706f2..845b713c91 100644 --- a/pc/jsep_transport.cc +++ b/pc/jsep_transport.cc @@ -234,6 +234,19 @@ RTCError JsepTransport::SetRemoteJsepTransportDescription( remote_description_.reset(new JsepTransportDescription(jsep_description)); RTC_DCHECK(rtp_dtls_transport()); + bool has_dtls_in_stun = + jsep_description.transport_desc.HasOption(ICE_OPTION_SPED); + if (!has_dtls_in_stun) { + rtp_dtls_transport()->DisableDtlsInStun(); + if (rtcp_dtls_transport() != nullptr) { + rtcp_dtls_transport()->DisableDtlsInStun(); + } + } else { + rtp_dtls_transport()->MaybeStartDtlsInStun(); + if (rtcp_dtls_transport() != nullptr) { + rtcp_dtls_transport()->MaybeStartDtlsInStun(); + } + } SetRemoteIceParameters(ice_parameters, rtp_dtls_transport()->ice_transport()); if (rtcp_dtls_transport()) { diff --git a/pc/peer_connection.cc b/pc/peer_connection.cc index 0755ed6fff..5ea57b419b 100644 --- a/pc/peer_connection.cc +++ b/pc/peer_connection.cc @@ -3124,7 +3124,7 @@ PeerConnection::InitializeUnDemuxablePacketHandler() { }; } -bool PeerConnection::CanAttemptDtlsStunPiggybacking() { +bool PeerConnection::CanAttemptDtlsStunPiggybacking() const { return dtls_enabled_ && SSLStreamAdapter::IsBoringSsl() && env_.field_trials().IsEnabled("WebRTC-IceHandshakeDtls"); } diff --git a/pc/peer_connection.h b/pc/peer_connection.h index 67748f5526..1707750d61 100644 --- a/pc/peer_connection.h +++ b/pc/peer_connection.h @@ -467,6 +467,7 @@ class PeerConnection : public PeerConnectionInternal, RTC_DCHECK_RUN_ON(signaling_thread()); sdp_handler_->DisableSdpMungingChecksForTesting(); } + bool CanAttemptDtlsStunPiggybacking() const override; protected: // Available for webrtc::scoped_refptr creation @@ -629,8 +630,6 @@ class PeerConnection : public PeerConnectionInternal, absl::AnyInvocable InitializeUnDemuxablePacketHandler(); - bool CanAttemptDtlsStunPiggybacking(); - // Runs a task on the signaling thread. If the current thread is the signaling // thread, the task will run immediately. Otherwise it will be posted to the // signaling thread and run asynchronously behind the diff --git a/pc/peer_connection_integrationtest.cc b/pc/peer_connection_integrationtest.cc index c64467e687..e2b1961f45 100644 --- a/pc/peer_connection_integrationtest.cc +++ b/pc/peer_connection_integrationtest.cc @@ -15,6 +15,7 @@ // do NOT add it here, but instead add it to the file // slow_peer_connection_integrationtest.cc +#include #include #include #include @@ -5141,6 +5142,50 @@ TEST_P(PeerConnectionIntegrationTest, DtlsPqcFieldTrial) { EXPECT_EQ(caller()->dtls_transport_information().ssl_group_id(), expected); } +TEST_P(PeerConnectionIntegrationTest, + SpedWireTriggerStartsDtlsWithoutRemoteDescription) { + if (!SSLStreamAdapter::IsBoringSsl()) { + GTEST_SKIP() << "DTLS-in-STUN requires BoringSSL."; + } + SetFieldTrials("WebRTC-IceHandshakeDtls/Enabled/"); + ASSERT_TRUE(CreatePeerConnectionWrappers()); + ConnectFakeSignaling(); + // Suppress SetRemoteDescription and ICE candidates from callee. + caller()->SetReceivedSdpMunger( + [](std::unique_ptr& desc) { desc.reset(); }); + callee()->set_signal_ice_candidates(false); + + caller()->CreateDataChannel(); + caller()->CreateAndSetAndSignalOffer(); + + ASSERT_THAT( + WaitUntil( + [&] { + const auto& history = caller()->peer_connection_state_history(); + return std::find(history.begin(), history.end(), + PeerConnectionInterface::PeerConnectionState:: + kConnecting) != history.end(); + }, + IsTrue()), + IsRtcOk()); +} + +TEST_P(PeerConnectionIntegrationTest, NoEarlyDtlsStartWhenSpedNotInAnswer) { + if (!SSLStreamAdapter::IsBoringSsl()) { + GTEST_SKIP() << "DTLS-in-STUN requires BoringSSL."; + } + SetFieldTrials(kCallerName, "WebRTC-IceHandshakeDtls/Enabled/"); + ASSERT_TRUE(CreatePeerConnectionWrappers()); + ConnectFakeSignalingForSdpOnly(); + caller()->CreateDataChannel(); + caller()->CreateAndSetAndSignalOffer(); + ASSERT_THAT(WaitUntil([&] { return SignalingStateStable(); }, IsTrue()), + IsRtcOk()); + EXPECT_THAT(caller()->peer_connection_state_history(), + ::testing::Not(::testing::Contains( + PeerConnectionInterface::PeerConnectionState::kConnecting))); +} + #endif // WEBRTC_HAVE_SCTP TEST_P(PeerConnectionIntegrationTest, PerPeerConnectionHeaderExtensions) { diff --git a/pc/peer_connection_internal.h b/pc/peer_connection_internal.h index cdffa213e3..c943485046 100644 --- a/pc/peer_connection_internal.h +++ b/pc/peer_connection_internal.h @@ -154,6 +154,9 @@ class PeerConnectionSdpMethods { // Keeps track of assigned payload types and comes up with reasonable // suggestions when new PTs need to be assigned. virtual PayloadTypePicker& payload_type_picker() = 0; + + // Determine whether DTLS-in-STUN is configured. + virtual bool CanAttemptDtlsStunPiggybacking() const = 0; }; // Functions defined in this class are called by other objects, diff --git a/pc/sdp_munging_detector.cc b/pc/sdp_munging_detector.cc index 83d8edbbcb..6fabf6c4d2 100644 --- a/pc/sdp_munging_detector.cc +++ b/pc/sdp_munging_detector.cc @@ -99,6 +99,19 @@ SdpMungingType DetermineTransportModification( if (created_trickle && !set_trickle) { return SdpMungingType::kIceOptionsTrickle; } + // Munging new features is not allowed. + bool created_sped = + absl::c_find( + last_created_transport_infos[i].description.transport_options, + ICE_OPTION_SPED) != + last_created_transport_infos[i].description.transport_options.end(); + bool set_sped = + absl::c_find(transport_infos_to_set[i].description.transport_options, + ICE_OPTION_SPED) != + transport_infos_to_set[i].description.transport_options.end(); + if (created_sped != set_sped) { + return SdpMungingType::kIceOptionsSped; + } return SdpMungingType::kIceOptions; } } @@ -770,6 +783,7 @@ bool IsSdpMungingAllowed(SdpMungingType sdp_munging_type, case SdpMungingType::kSframe: return false; case SdpMungingType::kDataChannelSctpInit: + case SdpMungingType::kIceOptionsSped: return false; case SdpMungingType::kCryptex: return false; diff --git a/pc/sdp_munging_detector_unittest.cc b/pc/sdp_munging_detector_unittest.cc index 58fc3ff8e9..ad025684cc 100644 --- a/pc/sdp_munging_detector_unittest.cc +++ b/pc/sdp_munging_detector_unittest.cc @@ -63,6 +63,7 @@ #include "pc/test/integration_test_helpers.h" #include "pc/test/mock_peer_connection_observers.h" #include "rtc_base/event.h" +#include "rtc_base/ssl_stream_adapter.h" #include "rtc_base/strings/string_format.h" #include "rtc_base/thread.h" #include "system_wrappers/include/metrics.h" @@ -679,6 +680,56 @@ TEST_F(SdpMungingTest, IceOptionsTrickle) { ElementsAre(Pair(SdpMungingType::kIceOptionsTrickle, 1))); } +TEST_F(SdpMungingTest, IceOptionsAddSpedDisallowed) { + auto pc = CreatePeerConnection("WebRTC-IceHandshakeDtls/Disabled/"); + pc->AddAudioTrack("audio_track", {}); + + std::unique_ptr offer = pc->CreateOffer(); + auto& transport_infos = offer->description()->transport_infos(); + ASSERT_EQ(transport_infos.size(), 1u); + ASSERT_THAT(transport_infos[0].description.transport_options, + ElementsAre("trickle")); + transport_infos[0].description.transport_options.push_back("sped"); + RTCError error; + EXPECT_FALSE(pc->SetLocalDescription(std::move(offer), &error)); + EXPECT_THAT( + metrics::Samples("WebRTC.PeerConnection.SdpMunging.Offer.Initial"), + ElementsAre(Pair(SdpMungingType::kIceOptionsSped, 1))); + EXPECT_THAT( + metrics::Samples("WebRTC.PeerConnection.SdpMunging.SdpOutcome.Rejected"), + ElementsAre(Pair(SdpMungingType::kIceOptionsSped, 1))); + EXPECT_THAT( + metrics::Samples("WebRTC.PeerConnection.SdpMunging.Outcome"), + ElementsAre(Pair(static_cast(SdpMungingOutcome::kRejected), 1))); +} + +TEST_F(SdpMungingTest, IceOptionsRemoveSpedDisallowed) { + if (!SSLStreamAdapter::IsBoringSsl()) { + GTEST_SKIP() << "DTLS-in-STUN requires BoringSSL."; + } + auto pc = CreatePeerConnection("WebRTC-IceHandshakeDtls/Enabled/"); + pc->AddAudioTrack("audio_track", {}); + + std::unique_ptr offer = pc->CreateOffer(); + auto& transport_infos = offer->description()->transport_infos(); + ASSERT_EQ(transport_infos.size(), 1u); + ASSERT_THAT(transport_infos[0].description.transport_options, + ElementsAre("trickle", "sped", "goog-sped-v1")); + auto& options = transport_infos[0].description.transport_options; + options.erase(std::find(options.begin(), options.end(), "sped")); + RTCError error; + EXPECT_FALSE(pc->SetLocalDescription(std::move(offer), &error)); + EXPECT_THAT( + metrics::Samples("WebRTC.PeerConnection.SdpMunging.Offer.Initial"), + ElementsAre(Pair(SdpMungingType::kIceOptionsSped, 1))); + EXPECT_THAT( + metrics::Samples("WebRTC.PeerConnection.SdpMunging.SdpOutcome.Rejected"), + ElementsAre(Pair(SdpMungingType::kIceOptionsSped, 1))); + EXPECT_THAT( + metrics::Samples("WebRTC.PeerConnection.SdpMunging.Outcome"), + ElementsAre(Pair(static_cast(SdpMungingOutcome::kRejected), 1))); +} + TEST_F(SdpMungingTest, DtlsRole) { auto pc = CreatePeerConnection(); pc->AddAudioTrack("audio_track", {}); diff --git a/pc/sdp_offer_answer.cc b/pc/sdp_offer_answer.cc index fcebc56db9..06e85784f1 100644 --- a/pc/sdp_offer_answer.cc +++ b/pc/sdp_offer_answer.cc @@ -4731,12 +4731,14 @@ void SdpOfferAnswerHandler::GetOptionsForOffer( RTC_ALLOW_PLAN_B_DEPRECATION_END(); } - // Apply ICE restart flag and renomination flag. + // Apply ICE restart flag and renomination and dtls-in-stun ICE options. bool ice_restart = offer_answer_options.ice_restart || HasNewIceCredentials(); for (auto& options : session_options->media_description_options) { options.transport_options.ice_restart = ice_restart; options.transport_options.enable_ice_renomination = pc_->configuration()->enable_ice_renomination; + options.transport_options.dtls_handshake_in_stun = + pc_->CanAttemptDtlsStunPiggybacking(); } session_options->rtcp_cname = rtcp_cname_; @@ -5026,10 +5028,12 @@ void SdpOfferAnswerHandler::GetOptionsForAnswer( RTC_ALLOW_PLAN_B_DEPRECATION_END(); } - // Apply ICE renomination flag. + // Apply renomination and dtls-in-stun ICE options. for (auto& options : session_options->media_description_options) { options.transport_options.enable_ice_renomination = pc_->configuration()->enable_ice_renomination; + options.transport_options.dtls_handshake_in_stun = + pc_->CanAttemptDtlsStunPiggybacking(); } session_options->rtcp_cname = rtcp_cname_; diff --git a/pc/sdp_offer_answer_unittest.cc b/pc/sdp_offer_answer_unittest.cc index ff06d90659..40ce2aff16 100644 --- a/pc/sdp_offer_answer_unittest.cc +++ b/pc/sdp_offer_answer_unittest.cc @@ -54,6 +54,7 @@ #include "pc/test/fake_audio_capture_module.h" #include "pc/test/integration_test_helpers.h" #include "pc/test/mock_peer_connection_observers.h" +#include "rtc_base/ssl_stream_adapter.h" #include "rtc_base/thread.h" #include "system_wrappers/include/metrics.h" #include "test/create_test_field_trials.h" @@ -2111,6 +2112,22 @@ TEST_F(SdpOfferAnswerTest, EXPECT_TRUE(callee_transceiver->receptive()); } +TEST_F(SdpOfferAnswerTest, IceOptionsDtlsInStun) { + if (!SSLStreamAdapter::IsBoringSsl()) { + GTEST_SKIP() << "DTLS-in-STUN requires BoringSSL."; + } + auto pc1 = CreatePeerConnection("WebRTC-IceHandshakeDtls/Enabled/"); + pc1->AddAudioTrack("audio_track", {}); + + auto offer = pc1->CreateOfferAndSetAsLocal(); + ASSERT_NE(offer, nullptr); + + const auto& transport_infos = offer->description()->transport_infos(); + ASSERT_THAT(transport_infos, SizeIs(1)); + const auto& transport_description = transport_infos[0].description; + EXPECT_TRUE(transport_description.HasOption("sped")); +} + #ifdef WEBRTC_HAVE_SCTP TEST_F(SdpOfferAnswerTest, SctpInitDisabled) { auto pc1 = CreatePeerConnection("WebRTC-Sctp-Snap/Disabled/"); diff --git a/pc/test/fake_peer_connection_base.h b/pc/test/fake_peer_connection_base.h index d395e83f9c..9a127b3cfb 100644 --- a/pc/test/fake_peer_connection_base.h +++ b/pc/test/fake_peer_connection_base.h @@ -425,6 +425,7 @@ class FakePeerConnectionBase : public PeerConnectionInternal { } CandidateStatsList GetPooledCandidateStats() const override { return {}; } + bool CanAttemptDtlsStunPiggybacking() const override { return false; } protected: Environment env_; diff --git a/pc/test/mock_peer_connection_internal.h b/pc/test/mock_peer_connection_internal.h index 37bd49b211..827dff6497 100644 --- a/pc/test/mock_peer_connection_internal.h +++ b/pc/test/mock_peer_connection_internal.h @@ -385,6 +385,7 @@ class MockPeerConnectionInternal : public PeerConnectionInternal { (int channel_id, DataChannelInterface::DataState), (override)); MOCK_METHOD(PayloadTypePicker&, payload_type_picker, (), (override)); + MOCK_METHOD(bool, CanAttemptDtlsStunPiggybacking, (), (const, override)); }; } // namespace webrtc