Rate limit QUIC reset sent to the clients.

Protected by FLAGS_quic_restart_flag_quic_use_recent_reset_addresses.

PiperOrigin-RevId: 394466549
diff --git a/quic/core/quic_connection_test.cc b/quic/core/quic_connection_test.cc
index 28a87ea..b7fba67 100644
--- a/quic/core/quic_connection_test.cc
+++ b/quic/core/quic_connection_test.cc
@@ -179,33 +179,6 @@
   SimpleBufferAllocator buffer_allocator_;
 };
 
-class TestAlarmFactory : public QuicAlarmFactory {
- public:
-  class TestAlarm : public QuicAlarm {
-   public:
-    explicit TestAlarm(QuicArenaScopedPtr<QuicAlarm::Delegate> delegate)
-        : QuicAlarm(std::move(delegate)) {}
-
-    void SetImpl() override {}
-    void CancelImpl() override {}
-    using QuicAlarm::Fire;
-  };
-
-  TestAlarmFactory() {}
-  TestAlarmFactory(const TestAlarmFactory&) = delete;
-  TestAlarmFactory& operator=(const TestAlarmFactory&) = delete;
-
-  QuicAlarm* CreateAlarm(QuicAlarm::Delegate* delegate) override {
-    return new TestAlarm(QuicArenaScopedPtr<QuicAlarm::Delegate>(delegate));
-  }
-
-  QuicArenaScopedPtr<QuicAlarm> CreateAlarm(
-      QuicArenaScopedPtr<QuicAlarm::Delegate> delegate,
-      QuicConnectionArena* arena) override {
-    return arena->New<TestAlarm>(std::move(delegate));
-  }
-};
-
 class TestConnection : public QuicConnection {
  public:
   TestConnection(QuicConnectionId connection_id,
diff --git a/quic/core/quic_dispatcher.cc b/quic/core/quic_dispatcher.cc
index 89904c5..00bd616 100644
--- a/quic/core/quic_dispatcher.cc
+++ b/quic/core/quic_dispatcher.cc
@@ -54,6 +54,24 @@
   QuicDispatcher* dispatcher_;
 };
 
+// An alarm that informs the QuicDispatcher to clear
+// recent_stateless_reset_addresses_.
+class ClearStatelessResetAddressesAlarm
+    : public QuicAlarm::DelegateWithoutContext {
+ public:
+  explicit ClearStatelessResetAddressesAlarm(QuicDispatcher* dispatcher)
+      : dispatcher_(dispatcher) {}
+  ClearStatelessResetAddressesAlarm(const DeleteSessionsAlarm&) = delete;
+  ClearStatelessResetAddressesAlarm& operator=(const DeleteSessionsAlarm&) =
+      delete;
+
+  void OnAlarm() override { dispatcher_->ClearStatelessResetAddresses(); }
+
+ private:
+  // Not owned.
+  QuicDispatcher* dispatcher_;
+};
+
 // Collects packets serialized by a QuicPacketCreator in order
 // to be handed off to the time wait list manager.
 class PacketCollector : public QuicPacketCreator::DelegateInterface,
@@ -311,8 +329,7 @@
 }  // namespace
 
 QuicDispatcher::QuicDispatcher(
-    const QuicConfig* config,
-    const QuicCryptoServerConfig* crypto_config,
+    const QuicConfig* config, const QuicCryptoServerConfig* crypto_config,
     QuicVersionManager* version_manager,
     std::unique_ptr<QuicConnectionHelperInterface> helper,
     std::unique_ptr<QuicCryptoServerStreamBase::Helper> session_helper,
@@ -335,6 +352,8 @@
       allow_short_initial_server_connection_ids_(false),
       expected_server_connection_id_length_(
           expected_server_connection_id_length),
+      clear_stateless_reset_addresses_alarm_(alarm_factory_->CreateAlarm(
+          new ClearStatelessResetAddressesAlarm(this))),
       should_update_expected_server_connection_id_length_(false) {
   QUIC_BUG_IF(quic_bug_12724_1, GetSupportedVersions().empty())
       << "Trying to create dispatcher without any supported versions";
@@ -348,6 +367,9 @@
       delete_sessions_alarm_ != nullptr) {
     delete_sessions_alarm_->PermanentCancel();
   }
+  if (clear_stateless_reset_addresses_alarm_ != nullptr) {
+    clear_stateless_reset_addresses_alarm_->PermanentCancel();
+  }
   reference_counted_session_map_.clear();
   closed_ref_counted_session_list_.clear();
   if (support_multiple_cid_per_connection_) {
@@ -941,6 +963,10 @@
   closed_ref_counted_session_list_.clear();
 }
 
+void QuicDispatcher::ClearStatelessResetAddresses() {
+  recent_stateless_reset_addresses_.clear();
+}
+
 void QuicDispatcher::OnCanWrite() {
   // The socket is now writable.
   writer_->SetWritable();
@@ -1366,6 +1392,13 @@
 void QuicDispatcher::MaybeResetPacketsWithNoVersion(
     const ReceivedPacketInfo& packet_info) {
   QUICHE_DCHECK(!packet_info.version_flag);
+  // Do not send a stateless reset if a reset has been sent to this address
+  // recently.
+  if (recent_stateless_reset_addresses_.contains(packet_info.peer_address)) {
+    QUIC_CODE_COUNT(quic_donot_send_reset_repeatedly);
+    QUICHE_DCHECK(use_recent_reset_addresses_);
+    return;
+  }
   if (packet_info.form != GOOGLE_QUIC_PACKET) {
     // Drop IETF packets smaller than the minimal stateless reset length.
     if (packet_info.packet.length() <=
@@ -1382,8 +1415,24 @@
       QUIC_CODE_COUNT(drop_too_small_packets);
       return;
     }
-    // TODO(fayang): Consider rate limiting reset packets if reset packet size >
-    // packet_length.
+  }
+  if (use_recent_reset_addresses_) {
+    QUIC_RESTART_FLAG_COUNT(quic_use_recent_reset_addresses);
+    // Do not send a stateless reset if there are too many stateless reset
+    // addresses.
+    if (recent_stateless_reset_addresses_.size() >=
+        GetQuicFlag(FLAGS_quic_max_recent_stateless_reset_addresses)) {
+      QUIC_CODE_COUNT(quic_too_many_recent_reset_addresses);
+      return;
+    }
+    if (recent_stateless_reset_addresses_.empty()) {
+      clear_stateless_reset_addresses_alarm_->Update(
+          helper()->GetClock()->ApproximateNow() +
+              QuicTime::Delta::FromMilliseconds(GetQuicFlag(
+                  FLAGS_quic_recent_stateless_reset_addresses_lifetime_ms)),
+          QuicTime::Delta::Zero());
+    }
+    recent_stateless_reset_addresses_.emplace(packet_info.peer_address);
   }
 
   time_wait_list_manager()->SendPublicReset(
diff --git a/quic/core/quic_dispatcher.h b/quic/core/quic_dispatcher.h
index a6ba1a6..3d36812 100644
--- a/quic/core/quic_dispatcher.h
+++ b/quic/core/quic_dispatcher.h
@@ -131,6 +131,9 @@
   // Deletes all sessions on the closed session list and clears the list.
   virtual void DeleteSessions();
 
+  // Clear recent_stateless_reset_addresses_.
+  void ClearStatelessResetAddresses();
+
   using ConnectionIdMap = absl::
       flat_hash_map<QuicConnectionId, QuicConnectionId, QuicConnectionIdHash>;
 
@@ -473,10 +476,20 @@
   // version does not allow variable length connection ID.
   uint8_t expected_server_connection_id_length_;
 
+  // Records client addresses that have been recently reset.
+  absl::flat_hash_set<QuicSocketAddress, QuicSocketAddressHash>
+      recent_stateless_reset_addresses_;
+
+  // An alarm which clear recent_stateless_reset_addresses_.
+  std::unique_ptr<QuicAlarm> clear_stateless_reset_addresses_alarm_;
+
   // If true, change expected_server_connection_id_length_ to be the received
   // destination connection ID length of all IETF long headers.
   bool should_update_expected_server_connection_id_length_;
 
+  const bool use_recent_reset_addresses_ =
+      GetQuicRestartFlag(quic_use_recent_reset_addresses);
+
   const bool support_multiple_cid_per_connection_ =
       GetQuicRestartFlag(quic_time_wait_list_support_multiple_cid_v2) &&
       GetQuicRestartFlag(
diff --git a/quic/core/quic_dispatcher_test.cc b/quic/core/quic_dispatcher_test.cc
index e927cda..a5bf226 100644
--- a/quic/core/quic_dispatcher_test.cc
+++ b/quic/core/quic_dispatcher_test.cc
@@ -121,15 +121,12 @@
  public:
   TestDispatcher(const QuicConfig* config,
                  const QuicCryptoServerConfig* crypto_config,
-                 QuicVersionManager* version_manager,
-                 QuicRandom* random)
-      : QuicDispatcher(config,
-                       crypto_config,
-                       version_manager,
+                 QuicVersionManager* version_manager, QuicRandom* random)
+      : QuicDispatcher(config, crypto_config, version_manager,
                        std::make_unique<MockQuicConnectionHelper>(),
                        std::unique_ptr<QuicCryptoServerStreamBase::Helper>(
                            new QuicSimpleCryptoServerStreamHelper()),
-                       std::make_unique<MockAlarmFactory>(),
+                       std::make_unique<TestAlarmFactory>(),
                        kQuicDefaultConnectionIdLength),
         random_(random) {}
 
@@ -500,6 +497,11 @@
       const QuicConnectionId& server_connection_id,
       const QuicConnectionId& client_connection_id);
 
+  TestAlarmFactory::TestAlarm* GetClearResetAddressesAlarm() {
+    return reinterpret_cast<TestAlarmFactory::TestAlarm*>(
+        QuicDispatcherPeer::GetClearResetAddressesAlarm(dispatcher_.get()));
+  }
+
   ParsedQuicVersion version_;
   MockQuicConnectionHelper mock_helper_;
   MockAlarmFactory mock_alarm_factory_;
@@ -965,6 +967,104 @@
   dispatcher_->ProcessPacket(server_address_, client_address, packet);
 }
 
+TEST_P(QuicDispatcherTestAllVersions, LimitResetsToSameClientAddress) {
+  CreateTimeWaitListManager();
+
+  QuicSocketAddress client_address(QuicIpAddress::Loopback4(), 1);
+  QuicSocketAddress client_address2(QuicIpAddress::Loopback4(), 2);
+  QuicSocketAddress client_address3(QuicIpAddress::Loopback6(), 1);
+  QuicConnectionId connection_id = TestConnectionId(1);
+
+  if (GetQuicRestartFlag(quic_use_recent_reset_addresses)) {
+    // Verify only one reset is sent to the address, although multiple packets
+    // are received.
+    EXPECT_CALL(*time_wait_list_manager_, SendPublicReset(_, _, _, _, _, _))
+        .Times(1);
+  } else {
+    EXPECT_CALL(*time_wait_list_manager_, SendPublicReset(_, _, _, _, _, _))
+        .Times(3);
+  }
+  ProcessPacket(client_address, connection_id, /*has_version_flag=*/false,
+                "data");
+  ProcessPacket(client_address, connection_id, /*has_version_flag=*/false,
+                "data2");
+  ProcessPacket(client_address, connection_id, /*has_version_flag=*/false,
+                "data3");
+
+  EXPECT_CALL(*time_wait_list_manager_, SendPublicReset(_, _, _, _, _, _))
+      .Times(2);
+  ProcessPacket(client_address2, connection_id, /*has_version_flag=*/false,
+                "data");
+  ProcessPacket(client_address3, connection_id, /*has_version_flag=*/false,
+                "data");
+}
+
+TEST_P(QuicDispatcherTestAllVersions,
+       StopSendingResetOnTooManyRecentAddresses) {
+  SetQuicFlag(FLAGS_quic_max_recent_stateless_reset_addresses, 2);
+  const size_t kTestLifeTimeMs = 10;
+  SetQuicFlag(FLAGS_quic_recent_stateless_reset_addresses_lifetime_ms,
+              kTestLifeTimeMs);
+  CreateTimeWaitListManager();
+
+  QuicSocketAddress client_address(QuicIpAddress::Loopback4(), 1);
+  QuicSocketAddress client_address2(QuicIpAddress::Loopback4(), 2);
+  QuicSocketAddress client_address3(QuicIpAddress::Loopback6(), 1);
+  QuicConnectionId connection_id = TestConnectionId(1);
+
+  EXPECT_CALL(*time_wait_list_manager_, SendPublicReset(_, _, _, _, _, _))
+      .Times(2);
+  EXPECT_FALSE(GetClearResetAddressesAlarm()->IsSet());
+  ProcessPacket(client_address, connection_id, /*has_version_flag=*/false,
+                "data");
+  const QuicTime expected_deadline =
+      mock_helper_.GetClock()->Now() +
+      QuicTime::Delta::FromMilliseconds(kTestLifeTimeMs);
+  if (GetQuicRestartFlag(quic_use_recent_reset_addresses)) {
+    ASSERT_TRUE(GetClearResetAddressesAlarm()->IsSet());
+    EXPECT_EQ(expected_deadline, GetClearResetAddressesAlarm()->deadline());
+  } else {
+    EXPECT_FALSE(GetClearResetAddressesAlarm()->IsSet());
+  }
+  // Received no version packet 2 after 5ms.
+  mock_helper_.AdvanceTime(QuicTime::Delta::FromMilliseconds(5));
+  ProcessPacket(client_address2, connection_id, /*has_version_flag=*/false,
+                "data");
+  if (GetQuicRestartFlag(quic_use_recent_reset_addresses)) {
+    ASSERT_TRUE(GetClearResetAddressesAlarm()->IsSet());
+    // Verify deadline does not change.
+    EXPECT_EQ(expected_deadline, GetClearResetAddressesAlarm()->deadline());
+  } else {
+    EXPECT_FALSE(GetClearResetAddressesAlarm()->IsSet());
+  }
+  if (GetQuicRestartFlag(quic_use_recent_reset_addresses)) {
+    // Verify reset gets throttled since there are too many recent addresses.
+    EXPECT_CALL(*time_wait_list_manager_, SendPublicReset(_, _, _, _, _, _))
+        .Times(0);
+  } else {
+    EXPECT_CALL(*time_wait_list_manager_, SendPublicReset(_, _, _, _, _, _))
+        .Times(1);
+  }
+  ProcessPacket(client_address3, connection_id, /*has_version_flag=*/false,
+                "data");
+
+  mock_helper_.AdvanceTime(QuicTime::Delta::FromMilliseconds(5));
+  if (GetQuicRestartFlag(quic_use_recent_reset_addresses)) {
+    GetClearResetAddressesAlarm()->Fire();
+    EXPECT_CALL(*time_wait_list_manager_, SendPublicReset(_, _, _, _, _, _))
+        .Times(2);
+  } else {
+    EXPECT_CALL(*time_wait_list_manager_, SendPublicReset(_, _, _, _, _, _))
+        .Times(3);
+  }
+  ProcessPacket(client_address, connection_id, /*has_version_flag=*/false,
+                "data");
+  ProcessPacket(client_address2, connection_id, /*has_version_flag=*/false,
+                "data");
+  ProcessPacket(client_address3, connection_id, /*has_version_flag=*/false,
+                "data");
+}
+
 // Makes sure nine-byte connection IDs are replaced by 8-byte ones.
 TEST_P(QuicDispatcherTestAllVersions, LongConnectionIdLengthReplaced) {
   if (!version_.AllowsVariableLengthConnectionIds()) {
diff --git a/quic/core/quic_flags_list.h b/quic/core/quic_flags_list.h
index 8054ebf..7e85017 100644
--- a/quic/core/quic_flags_list.h
+++ b/quic/core/quic_flags_list.h
@@ -87,6 +87,8 @@
 QUIC_FLAG(FLAGS_quic_restart_flag_quic_dispatcher_support_multiple_cid_per_connection_v2, true)
 // If true, receiving server push stream will trigger QUIC connection close.
 QUIC_FLAG(FLAGS_quic_reloadable_flag_quic_decline_server_push_stream, true)
+// If true, record addresses that server has sent reset to recently, and do not send reset if the address lives in the set.
+QUIC_FLAG(FLAGS_quic_restart_flag_quic_use_recent_reset_addresses, false)
 // If true, refactor how QUIC TLS server disables resumption. No behavior change.
 QUIC_FLAG(FLAGS_quic_reloadable_flag_quic_tls_disable_resumption_refactor, false)
 // If true, require handshake confirmation for QUIC connections, functionally disabling 0-rtt handshakes.
diff --git a/quic/core/quic_protocol_flags_list.h b/quic/core/quic_protocol_flags_list.h
index 4b2e818..9fe0af9 100644
--- a/quic/core/quic_protocol_flags_list.h
+++ b/quic/core/quic_protocol_flags_list.h
@@ -50,6 +50,20 @@
     uint64_t, quic_time_wait_list_max_pending_packets, 1024,
     "Upper limit of pending packets in time wait list when writer is blocked.")
 
+// Stop sending a reset if the recorded number of addresses that server has
+// recently sent stateless reset to exceeds this limit.
+QUIC_PROTOCOL_FLAG(uint64_t, quic_max_recent_stateless_reset_addresses, 1024,
+                   "Max number of recorded recent reset addresses.");
+
+// After this timeout, recent reset addresses will be cleared.
+// FLAGS_quic_max_recent_stateless_reset_addresses * (1000ms /
+// FLAGS_quic_recent_stateless_reset_addresses_lifetime_ms) is roughly the max
+// reset per second. For example, 1024 * (1000ms / 1000ms) = 1K reset per
+// second.
+QUIC_PROTOCOL_FLAG(
+    uint64_t, quic_recent_stateless_reset_addresses_lifetime_ms, 1000,
+    "Max time that a client address lives in recent reset addresses set.");
+
 QUIC_PROTOCOL_FLAG(double,
                    quic_bbr_cwnd_gain,
                    2.0f,
diff --git a/quic/platform/api/quic_socket_address.cc b/quic/platform/api/quic_socket_address.cc
index 4c60c1e..f0b667c 100644
--- a/quic/platform/api/quic_socket_address.cc
+++ b/quic/platform/api/quic_socket_address.cc
@@ -15,6 +15,23 @@
 
 namespace quic {
 
+namespace {
+
+uint32_t HashIP(const QuicIpAddress& ip) {
+  if (ip.IsIPv4()) {
+    return ip.GetIPv4().s_addr;
+  }
+  if (ip.IsIPv6()) {
+    auto v6addr = ip.GetIPv6();
+    const uint32_t* v6_as_ints =
+        reinterpret_cast<const uint32_t*>(&v6addr.s6_addr);
+    return v6_as_ints[0] ^ v6_as_ints[1] ^ v6_as_ints[2] ^ v6_as_ints[3];
+  }
+  return 0;
+}
+
+}  // namespace
+
 QuicSocketAddress::QuicSocketAddress(QuicIpAddress address, uint16_t port)
     : host_(address), port_(port) {}
 
@@ -131,4 +148,11 @@
   return result.storage;
 }
 
+uint32_t QuicSocketAddress::Hash() const {
+  uint32_t value = 0;
+  value ^= HashIP(host_);
+  value ^= port_ | (port_ << 16);
+  return value;
+}
+
 }  // namespace quic
diff --git a/quic/platform/api/quic_socket_address.h b/quic/platform/api/quic_socket_address.h
index 0831985..df6b9b7 100644
--- a/quic/platform/api/quic_socket_address.h
+++ b/quic/platform/api/quic_socket_address.h
@@ -37,6 +37,9 @@
   uint16_t port() const;
   sockaddr_storage generic_address() const;
 
+  // Hashes this address to an uint32_t.
+  uint32_t Hash() const;
+
  private:
   QuicIpAddress host_;
   uint16_t port_ = 0;
@@ -48,6 +51,13 @@
   return os;
 }
 
+class QUIC_EXPORT_PRIVATE QuicSocketAddressHash {
+ public:
+  size_t operator()(QuicSocketAddress const& address) const noexcept {
+    return address.Hash();
+  }
+};
+
 }  // namespace quic
 
 #endif  // QUICHE_QUIC_PLATFORM_API_QUIC_SOCKET_ADDRESS_H_
diff --git a/quic/test_tools/quic_dispatcher_peer.cc b/quic/test_tools/quic_dispatcher_peer.cc
index 8002170..341d5ac 100644
--- a/quic/test_tools/quic_dispatcher_peer.cc
+++ b/quic/test_tools/quic_dispatcher_peer.cc
@@ -133,5 +133,11 @@
              : it->second.get();
 }
 
+// static
+QuicAlarm* QuicDispatcherPeer::GetClearResetAddressesAlarm(
+    QuicDispatcher* dispatcher) {
+  return dispatcher->clear_stateless_reset_addresses_alarm_.get();
+}
+
 }  // namespace test
 }  // namespace quic
diff --git a/quic/test_tools/quic_dispatcher_peer.h b/quic/test_tools/quic_dispatcher_peer.h
index fc8ab20..ec83be4 100644
--- a/quic/test_tools/quic_dispatcher_peer.h
+++ b/quic/test_tools/quic_dispatcher_peer.h
@@ -76,6 +76,8 @@
   // Find the corresponding session if exsits.
   static const QuicSession* FindSession(const QuicDispatcher* dispatcher,
                                         QuicConnectionId id);
+
+  static QuicAlarm* GetClearResetAddressesAlarm(QuicDispatcher* dispatcher);
 };
 
 }  // namespace test
diff --git a/quic/test_tools/quic_test_utils.h b/quic/test_tools/quic_test_utils.h
index d513d0a..429b564 100644
--- a/quic/test_tools/quic_test_utils.h
+++ b/quic/test_tools/quic_test_utils.h
@@ -672,6 +672,33 @@
   }
 };
 
+class TestAlarmFactory : public QuicAlarmFactory {
+ public:
+  class TestAlarm : public QuicAlarm {
+   public:
+    explicit TestAlarm(QuicArenaScopedPtr<QuicAlarm::Delegate> delegate)
+        : QuicAlarm(std::move(delegate)) {}
+
+    void SetImpl() override {}
+    void CancelImpl() override {}
+    using QuicAlarm::Fire;
+  };
+
+  TestAlarmFactory() {}
+  TestAlarmFactory(const TestAlarmFactory&) = delete;
+  TestAlarmFactory& operator=(const TestAlarmFactory&) = delete;
+
+  QuicAlarm* CreateAlarm(QuicAlarm::Delegate* delegate) override {
+    return new TestAlarm(QuicArenaScopedPtr<QuicAlarm::Delegate>(delegate));
+  }
+
+  QuicArenaScopedPtr<QuicAlarm> CreateAlarm(
+      QuicArenaScopedPtr<QuicAlarm::Delegate> delegate,
+      QuicConnectionArena* arena) override {
+    return arena->New<TestAlarm>(std::move(delegate));
+  }
+};
+
 class MockQuicConnection : public QuicConnection {
  public:
   // Uses a ConnectionId of 42 and 127.0.0.1:123.