diff --git a/quiche/quic/core/deterministic_connection_id_generator.cc b/quiche/quic/core/deterministic_connection_id_generator.cc index bd0ebf169..e1ba94e98 100644 --- a/quiche/quic/core/deterministic_connection_id_generator.cc +++ b/quiche/quic/core/deterministic_connection_id_generator.cc @@ -13,8 +13,13 @@ namespace quic { DeterministicConnectionIdGenerator::DeterministicConnectionIdGenerator( - uint8_t expected_connection_id_length) - : expected_connection_id_length_(expected_connection_id_length) { + uint8_t expected_connection_id_length) { + set_expected_connection_id_length(expected_connection_id_length); +} + +void DeterministicConnectionIdGenerator::set_expected_connection_id_length( + uint8_t expected_connection_id_length) { + expected_connection_id_length_ = expected_connection_id_length; if (expected_connection_id_length_ > kQuicMaxConnectionIdWithLengthPrefixLength) { QUIC_BUG(quic_bug_465151159_01) diff --git a/quiche/quic/core/deterministic_connection_id_generator.h b/quiche/quic/core/deterministic_connection_id_generator.h index fc8f8e467..30319824f 100644 --- a/quiche/quic/core/deterministic_connection_id_generator.h +++ b/quiche/quic/core/deterministic_connection_id_generator.h @@ -17,7 +17,11 @@ namespace quic { class QUICHE_EXPORT DeterministicConnectionIdGenerator : public ConnectionIdGeneratorInterface { public: - DeterministicConnectionIdGenerator(uint8_t expected_connection_id_length); + explicit DeterministicConnectionIdGenerator( + uint8_t expected_connection_id_length); + + // Sets the length of the connection IDs generated from now on. + void set_expected_connection_id_length(uint8_t expected_connection_id_length); // Hashes |original| to create a new connection ID. std::optional GenerateNextConnectionId( @@ -32,7 +36,7 @@ class QUICHE_EXPORT DeterministicConnectionIdGenerator } private: - const uint8_t expected_connection_id_length_; + uint8_t expected_connection_id_length_; }; } // namespace quic diff --git a/quiche/quic/core/http/end_to_end_test.cc b/quiche/quic/core/http/end_to_end_test.cc index 6cf05b7eb..ac9ac4d89 100644 --- a/quiche/quic/core/http/end_to_end_test.cc +++ b/quiche/quic/core/http/end_to_end_test.cc @@ -1141,6 +1141,9 @@ class EndToEndTest : public QuicTestWithParam { void TestMultiPacketChaosProtection(int num_packets, bool drop_first_packet, bool kyber = false); + void ConnectionMigrationNonZeroConnectionIDClientIPChange( + uint8_t client_connection_id_length); + quiche::test::ScopedEnvironmentForThreads environment_; bool initialized_; // If true, the Initialize() function will create |client_| and starts to @@ -3613,12 +3616,24 @@ TEST_P(EndToEndTest, IetfConnectionMigrationClientIPChangedMultipleTimes) { TEST_P(EndToEndTest, ConnectionMigrationWithNonZeroConnectionIDClientIPChangedMultipleTimes) { + ConnectionMigrationNonZeroConnectionIDClientIPChange( + kQuicDefaultConnectionIdLength); +} + +TEST_P(EndToEndTest, + ConnectionMigrationWithMaxConnectionIDClientIPChangedMultipleTimes) { + ConnectionMigrationNonZeroConnectionIDClientIPChange( + kQuicMaxConnectionIdWithLengthPrefixLength); +} + +void EndToEndTest::ConnectionMigrationNonZeroConnectionIDClientIPChange( + uint8_t client_connection_id_length) { if (!version_.IsIetfQuic() || GetQuicFlag(quic_enforce_strict_amplification_factor)) { ASSERT_TRUE(Initialize()); return; } - override_client_connection_id_length_ = kQuicDefaultConnectionIdLength; + override_client_connection_id_length_ = client_connection_id_length; ASSERT_TRUE(Initialize()); SendSynchronousFooRequestAndCheckResponse(); diff --git a/quiche/quic/test_tools/quic_test_client.cc b/quiche/quic/test_tools/quic_test_client.cc index 01dedee97..f8ccaa687 100644 --- a/quiche/quic/test_tools/quic_test_client.cc +++ b/quiche/quic/test_tools/quic_test_client.cc @@ -223,9 +223,7 @@ MockableQuicClient::MockableQuicClient( std::make_unique(event_loop, this), std::make_unique(std::move(proof_verifier)), - std::move(session_cache)), - override_client_connection_id_(EmptyQuicConnectionId()), - client_connection_id_overridden_(false) {} + std::move(session_cache)) {} MockableQuicClient::~MockableQuicClient() { if (connected()) { @@ -245,17 +243,6 @@ MockableQuicClient::mockable_network_helper() const { default_network_helper()); } -QuicConnectionId MockableQuicClient::GetClientConnectionId() { - if (client_connection_id_overridden_) { - return override_client_connection_id_; - } - if (override_client_connection_id_length_ >= 0) { - return QuicUtils::CreateRandomConnectionId( - override_client_connection_id_length_); - } - return QuicDefaultClient::GetClientConnectionId(); -} - std::unique_ptr MockableQuicClient::CreateQuicMigrationHelper() { auto migration_helper = std::make_unique(*this); @@ -263,11 +250,6 @@ MockableQuicClient::CreateQuicMigrationHelper() { return migration_helper; } -void MockableQuicClient::UseClientConnectionIdLength( - int client_connection_id_length) { - override_client_connection_id_length_ = client_connection_id_length; -} - void MockableQuicClient::UseWriter(QuicPacketWriterWrapper* writer) { mockable_network_helper()->UseWriter(writer); } @@ -785,7 +767,7 @@ void QuicTestClient::UseConnectionIdLength( void QuicTestClient::UseClientConnectionIdLength( uint8_t client_connection_id_length) { QUICHE_DCHECK(!connected()); - client_->UseClientConnectionIdLength(client_connection_id_length); + client_->set_client_connection_id_length(client_connection_id_length); } bool QuicTestClient::MigrateSocket(const QuicIpAddress& new_host) { diff --git a/quiche/quic/test_tools/quic_test_client.h b/quiche/quic/test_tools/quic_test_client.h index d9cd3feb2..a3db70336 100644 --- a/quiche/quic/test_tools/quic_test_client.h +++ b/quiche/quic/test_tools/quic_test_client.h @@ -182,9 +182,7 @@ class MockableQuicClient : public QuicDefaultClient { enable_web_transport()); } - QuicConnectionId GetClientConnectionId() override; std::unique_ptr CreateQuicMigrationHelper() override; - void UseClientConnectionIdLength(int client_connection_id_length); void UseWriter(QuicPacketWriterWrapper* writer); void set_peer_address(const QuicSocketAddress& address); @@ -208,11 +206,6 @@ class MockableQuicClient : public QuicDefaultClient { } private: - // Client connection ID to use, if client_connection_id_overridden_. - // TODO(wub): Move client_connection_id_(length_) overrides to QuicClientBase. - QuicConnectionId override_client_connection_id_; - bool client_connection_id_overridden_; - int override_client_connection_id_length_ = -1; CachedNetworkParameters cached_network_paramaters_; // Owned by the base class. QuicTestMigrationHelper* migration_helper_ = nullptr; @@ -375,9 +368,6 @@ class QuicTestClient : public QuicSpdyStream::Visitor { // Configures client_ to use a specific server connection ID length instead // of the default of kQuicDefaultConnectionIdLength. void UseConnectionIdLength(uint8_t server_connection_id_length); - // Configures client_ to use a specific client connection ID instead of an - // empty one. - void UseClientConnectionId(QuicConnectionId client_connection_id); // Configures client_ to use a specific client connection ID length instead // of the default of zero. void UseClientConnectionIdLength(uint8_t client_connection_id_length); diff --git a/quiche/quic/tools/quic_client_base.cc b/quiche/quic/tools/quic_client_base.cc index 7b41d70fe..20bcaf51a 100644 --- a/quiche/quic/tools/quic_client_base.cc +++ b/quiche/quic/tools/quic_client_base.cc @@ -484,6 +484,14 @@ QuicConnectionId QuicClientBase::GetClientConnectionId() { return QuicUtils::CreateRandomConnectionId(client_connection_id_length_); } +void QuicClientBase::set_client_connection_id_length( + uint8_t client_connection_id_length) { + QUICHE_DCHECK(!connected()); + client_connection_id_length_ = client_connection_id_length; + connection_id_generator_.set_expected_connection_id_length( + client_connection_id_length); +} + bool QuicClientBase::CanReconnectWithDifferentVersion( ParsedQuicVersion* version) const { if (session_ == nullptr || session_->connection() == nullptr || diff --git a/quiche/quic/tools/quic_client_base.h b/quiche/quic/tools/quic_client_base.h index a38d2516c..c4344d0bc 100644 --- a/quiche/quic/tools/quic_client_base.h +++ b/quiche/quic/tools/quic_client_base.h @@ -317,9 +317,10 @@ class QuicClientBase : public QuicSession::Visitor { server_connection_id_length_ = server_connection_id_length; } - void set_client_connection_id_length(uint8_t client_connection_id_length) { - client_connection_id_length_ = client_connection_id_length; - } + // Sets the length of the client connection IDs this client generates, + // including the ones self-issued after the handshake. Must be called before + // the connection is created. + void set_client_connection_id_length(uint8_t client_connection_id_length); bool HasPendingPathValidation();