Internal change

PiperOrigin-RevId: 378464164
diff --git a/quic/core/crypto/proof_source.h b/quic/core/crypto/proof_source.h
index e4047a9..80d337a 100644
--- a/quic/core/crypto/proof_source.h
+++ b/quic/core/crypto/proof_source.h
@@ -294,7 +294,8 @@
       const std::string& alpn,
       absl::optional<std::string> alps,
       const std::vector<uint8_t>& quic_transport_params,
-      const absl::optional<std::vector<uint8_t>>& early_data_context) = 0;
+      const absl::optional<std::vector<uint8_t>>& early_data_context,
+      const QuicSSLConfig& ssl_config) = 0;
 
   // Starts a compute signature operation. If the operation is not cancelled
   // when it completes, callback()->OnComputeSignatureDone will be invoked.
diff --git a/quic/core/crypto/tls_client_connection.cc b/quic/core/crypto/tls_client_connection.cc
index c6a45c4..dd53ee7 100644
--- a/quic/core/crypto/tls_client_connection.cc
+++ b/quic/core/crypto/tls_client_connection.cc
@@ -6,8 +6,12 @@
 
 namespace quic {
 
-TlsClientConnection::TlsClientConnection(SSL_CTX* ssl_ctx, Delegate* delegate)
-    : TlsConnection(ssl_ctx, delegate->ConnectionDelegate()),
+TlsClientConnection::TlsClientConnection(SSL_CTX* ssl_ctx,
+                                         Delegate* delegate,
+                                         QuicSSLConfig ssl_config)
+    : TlsConnection(ssl_ctx,
+                    delegate->ConnectionDelegate(),
+                    std::move(ssl_config)),
       delegate_(delegate) {}
 
 // static
@@ -24,6 +28,8 @@
       ssl_ctx.get(), SSL_SESS_CACHE_CLIENT | SSL_SESS_CACHE_NO_INTERNAL);
   SSL_CTX_sess_set_new_cb(ssl_ctx.get(), NewSessionCallback);
 
+  // TODO(wub): Always enable early data on the SSL_CTX, but allow it to be
+  // overridden on the SSL object, via QuicSSLConfig.
   SSL_CTX_set_early_data_enabled(ssl_ctx.get(), enable_early_data);
   return ssl_ctx;
 }
diff --git a/quic/core/crypto/tls_client_connection.h b/quic/core/crypto/tls_client_connection.h
index c947153..ce4b948 100644
--- a/quic/core/crypto/tls_client_connection.h
+++ b/quic/core/crypto/tls_client_connection.h
@@ -30,7 +30,9 @@
     friend class TlsClientConnection;
   };
 
-  TlsClientConnection(SSL_CTX* ssl_ctx, Delegate* delegate);
+  TlsClientConnection(SSL_CTX* ssl_ctx,
+                      Delegate* delegate,
+                      QuicSSLConfig ssl_config);
 
   // Creates and configures an SSL_CTX that is appropriate for clients to use.
   static bssl::UniquePtr<SSL_CTX> CreateSslCtx(bool enable_early_data);
diff --git a/quic/core/crypto/tls_connection.cc b/quic/core/crypto/tls_connection.cc
index 381de7e..3a1e652 100644
--- a/quic/core/crypto/tls_connection.cc
+++ b/quic/core/crypto/tls_connection.cc
@@ -88,12 +88,20 @@
 }
 
 TlsConnection::TlsConnection(SSL_CTX* ssl_ctx,
-                             TlsConnection::Delegate* delegate)
-    : delegate_(delegate), ssl_(SSL_new(ssl_ctx)) {
+                             TlsConnection::Delegate* delegate,
+                             QuicSSLConfig ssl_config)
+    : delegate_(delegate),
+      ssl_(SSL_new(ssl_ctx)),
+      ssl_config_(std::move(ssl_config)) {
   SSL_set_ex_data(
       ssl(), SslIndexSingleton::GetInstance()->ssl_ex_data_index_connection(),
       this);
+  if (ssl_config_.early_data_enabled.has_value()) {
+    const int early_data_enabled = *ssl_config_.early_data_enabled ? 1 : 0;
+    SSL_set_early_data_enabled(ssl(), early_data_enabled);
+  }
 }
+
 // static
 bssl::UniquePtr<SSL_CTX> TlsConnection::CreateSslCtx(int cert_verify_mode) {
   CRYPTO_library_init();
diff --git a/quic/core/crypto/tls_connection.h b/quic/core/crypto/tls_connection.h
index 28b5684..329bb33 100644
--- a/quic/core/crypto/tls_connection.h
+++ b/quic/core/crypto/tls_connection.h
@@ -89,10 +89,12 @@
 
   SSL* ssl() const { return ssl_.get(); }
 
+  const QuicSSLConfig& ssl_config() const { return ssl_config_; }
+
  protected:
-  // TlsConnection does not take ownership of any of its arguments; they must
+  // TlsConnection does not take ownership of |ssl_ctx| or |delegate|; they must
   // outlive the TlsConnection object.
-  TlsConnection(SSL_CTX* ssl_ctx, Delegate* delegate);
+  TlsConnection(SSL_CTX* ssl_ctx, Delegate* delegate, QuicSSLConfig ssl_config);
 
   // Creates an SSL_CTX and configures it with the options that are appropriate
   // for both client and server. The caller is responsible for ownership of the
@@ -141,6 +143,7 @@
 
   Delegate* delegate_;
   bssl::UniquePtr<SSL> ssl_;
+  const QuicSSLConfig ssl_config_;
 };
 
 }  // namespace quic
diff --git a/quic/core/crypto/tls_server_connection.cc b/quic/core/crypto/tls_server_connection.cc
index 6e9901b..2042c15 100644
--- a/quic/core/crypto/tls_server_connection.cc
+++ b/quic/core/crypto/tls_server_connection.cc
@@ -12,8 +12,12 @@
 
 namespace quic {
 
-TlsServerConnection::TlsServerConnection(SSL_CTX* ssl_ctx, Delegate* delegate)
-    : TlsConnection(ssl_ctx, delegate->ConnectionDelegate()),
+TlsServerConnection::TlsServerConnection(SSL_CTX* ssl_ctx,
+                                         Delegate* delegate,
+                                         QuicSSLConfig ssl_config)
+    : TlsConnection(ssl_ctx,
+                    delegate->ConnectionDelegate(),
+                    std::move(ssl_config)),
       delegate_(delegate) {}
 
 // static
diff --git a/quic/core/crypto/tls_server_connection.h b/quic/core/crypto/tls_server_connection.h
index 774bb44..6c775b8 100644
--- a/quic/core/crypto/tls_server_connection.h
+++ b/quic/core/crypto/tls_server_connection.h
@@ -120,7 +120,9 @@
     friend class TlsServerConnection;
   };
 
-  TlsServerConnection(SSL_CTX* ssl_ctx, Delegate* delegate);
+  TlsServerConnection(SSL_CTX* ssl_ctx,
+                      Delegate* delegate,
+                      QuicSSLConfig ssl_config);
 
   // Creates and configures an SSL_CTX that is appropriate for servers to use.
   static bssl::UniquePtr<SSL_CTX> CreateSslCtx(ProofSource* proof_source);
diff --git a/quic/core/quic_session.h b/quic/core/quic_session.h
index b6c1367..f00d840 100644
--- a/quic/core/quic_session.h
+++ b/quic/core/quic_session.h
@@ -17,6 +17,7 @@
 #include "absl/container/flat_hash_map.h"
 #include "absl/strings/string_view.h"
 #include "absl/types/optional.h"
+#include "quic/core/crypto/tls_connection.h"
 #include "quic/core/frames/quic_ack_frequency_frame.h"
 #include "quic/core/handshaker_delegate_interface.h"
 #include "quic/core/legacy_quic_stream_id_manager.h"
@@ -619,6 +620,8 @@
            use_write_or_buffer_data_at_level_;
   }
 
+  virtual QuicSSLConfig GetSSLConfig() const { return QuicSSLConfig(); }
+
  protected:
   using StreamMap =
       absl::flat_hash_map<QuicStreamId, std::unique_ptr<QuicStream>>;
diff --git a/quic/core/quic_types.h b/quic/core/quic_types.h
index 2e65e5f..dc1fc69 100644
--- a/quic/core/quic_types.h
+++ b/quic/core/quic_types.h
@@ -826,6 +826,13 @@
 
 QUIC_EXPORT_PRIVATE std::string KeyUpdateReasonString(KeyUpdateReason reason);
 
+// QuicSSLConfig contains configurations to be applied on a SSL object, which
+// overrides the configurations in SSL_CTX.
+struct QUIC_NO_EXPORT QuicSSLConfig {
+  // Whether TLS early data should be enabled. If not set, default to enabled.
+  absl::optional<bool> early_data_enabled;
+};
+
 }  // namespace quic
 
 #endif  // QUICHE_QUIC_CORE_QUIC_TYPES_H_
diff --git a/quic/core/tls_client_handshaker.cc b/quic/core/tls_client_handshaker.cc
index 1b9de03..ab8ba64 100644
--- a/quic/core/tls_client_handshaker.cc
+++ b/quic/core/tls_client_handshaker.cc
@@ -41,7 +41,7 @@
       crypto_negotiated_params_(new QuicCryptoNegotiatedParameters),
       has_application_state_(has_application_state),
       crypto_config_(crypto_config),
-      tls_connection_(crypto_config->ssl_ctx(), this) {
+      tls_connection_(crypto_config->ssl_ctx(), this, session->GetSSLConfig()) {
   std::string token =
       crypto_config->LookupOrCreate(server_id)->source_address_token();
   if (!token.empty()) {
diff --git a/quic/core/tls_server_handshaker.cc b/quic/core/tls_server_handshaker.cc
index d7a850e..d0257f9 100644
--- a/quic/core/tls_server_handshaker.cc
+++ b/quic/core/tls_server_handshaker.cc
@@ -62,7 +62,8 @@
     const std::string& /*alpn*/,
     absl::optional<std::string> /*alps*/,
     const std::vector<uint8_t>& /*quic_transport_params*/,
-    const absl::optional<std::vector<uint8_t>>& /*early_data_context*/) {
+    const absl::optional<std::vector<uint8_t>>& /*early_data_context*/,
+    const QuicSSLConfig& /*ssl_config*/) {
   if (!handshaker_ || !proof_source_) {
     QUIC_BUG(quic_bug_10341_1)
         << "SelectCertificate called on a detached handle";
@@ -164,7 +165,7 @@
       proof_source_(crypto_config->proof_source()),
       pre_shared_key_(crypto_config->pre_shared_key()),
       crypto_negotiated_params_(new QuicCryptoNegotiatedParameters),
-      tls_connection_(crypto_config->ssl_ctx(), this),
+      tls_connection_(crypto_config->ssl_ctx(), this, session->GetSSLConfig()),
       crypto_config_(crypto_config) {
   QUICHE_DCHECK_EQ(PROTOCOL_TLS1_3,
                    session->connection()->version().handshake_protocol);
@@ -843,7 +844,8 @@
           client_hello->client_hello_len),
       AlpnForVersion(session()->version()), std::move(alps),
       set_transport_params_result.quic_transport_params,
-      set_transport_params_result.early_data_context);
+      set_transport_params_result.early_data_context,
+      tls_connection_.ssl_config());
 
   QUICHE_DCHECK_EQ(status, select_cert_status().value());
 
diff --git a/quic/core/tls_server_handshaker.h b/quic/core/tls_server_handshaker.h
index 2715415..b5fff82 100644
--- a/quic/core/tls_server_handshaker.h
+++ b/quic/core/tls_server_handshaker.h
@@ -223,8 +223,8 @@
         const std::string& alpn,
         absl::optional<std::string> alps,
         const std::vector<uint8_t>& quic_transport_params,
-        const absl::optional<std::vector<uint8_t>>& early_data_context)
-        override;
+        const absl::optional<std::vector<uint8_t>>& early_data_context,
+        const QuicSSLConfig& ssl_config) override;
 
     // Delegates to proof_source_->ComputeTlsSignature.
     // Returns QUIC_SUCCESS, QUIC_FAILURE or QUIC_PENDING.
diff --git a/quic/core/tls_server_handshaker_test.cc b/quic/core/tls_server_handshaker_test.cc
index d67c879..6f8db12 100644
--- a/quic/core/tls_server_handshaker_test.cc
+++ b/quic/core/tls_server_handshaker_test.cc
@@ -561,6 +561,23 @@
   EXPECT_EQ(last_compute_signature_args().hostname, "test.example.com");
 }
 
+TEST_P(TlsServerHandshakerTest, SSLConfigForCertSelection) {
+  InitializeServerWithFakeProofSourceHandle();
+
+  // Disable early data.
+  server_session_->ssl_config()->early_data_enabled = false;
+
+  server_handshaker_->SetupProofSourceHandle(
+      /*select_cert_action=*/FakeProofSourceHandle::Action::DELEGATE_SYNC,
+      /*compute_signature_action=*/FakeProofSourceHandle::Action::
+          DELEGATE_SYNC);
+  InitializeFakeClient();
+  CompleteCryptoHandshake();
+  ExpectHandshakeSuccessful();
+
+  EXPECT_FALSE(last_select_cert_args().ssl_config.early_data_enabled);
+}
+
 TEST_P(TlsServerHandshakerTest, ConnectionClosedOnTlsError) {
   EXPECT_CALL(*server_connection_,
               CloseConnection(QUIC_HANDSHAKE_FAILED, _, _, _));
diff --git a/quic/test_tools/fake_proof_source_handle.cc b/quic/test_tools/fake_proof_source_handle.cc
index dc3a33a..a70ce40 100644
--- a/quic/test_tools/fake_proof_source_handle.cc
+++ b/quic/test_tools/fake_proof_source_handle.cc
@@ -78,11 +78,12 @@
     const std::string& alpn,
     absl::optional<std::string> alps,
     const std::vector<uint8_t>& quic_transport_params,
-    const absl::optional<std::vector<uint8_t>>& early_data_context) {
+    const absl::optional<std::vector<uint8_t>>& early_data_context,
+    const QuicSSLConfig& ssl_config) {
   QUICHE_CHECK(!closed_);
   all_select_cert_args_.push_back(SelectCertArgs(
       server_address, client_address, ssl_capabilities, hostname, client_hello,
-      alpn, alps, quic_transport_params, early_data_context));
+      alpn, alps, quic_transport_params, early_data_context, ssl_config));
 
   if (select_cert_action_ == Action::DELEGATE_ASYNC ||
       select_cert_action_ == Action::FAIL_ASYNC) {
diff --git a/quic/test_tools/fake_proof_source_handle.h b/quic/test_tools/fake_proof_source_handle.h
index ae06e42..3d038a4 100644
--- a/quic/test_tools/fake_proof_source_handle.h
+++ b/quic/test_tools/fake_proof_source_handle.h
@@ -46,7 +46,8 @@
       const std::string& alpn,
       absl::optional<std::string> alps,
       const std::vector<uint8_t>& quic_transport_params,
-      const absl::optional<std::vector<uint8_t>>& early_data_context) override;
+      const absl::optional<std::vector<uint8_t>>& early_data_context,
+      const QuicSSLConfig& ssl_config) override;
 
   QuicAsyncStatus ComputeSignature(const QuicSocketAddress& server_address,
                                    const QuicSocketAddress& client_address,
@@ -70,7 +71,8 @@
                    std::string alpn,
                    absl::optional<std::string> alps,
                    std::vector<uint8_t> quic_transport_params,
-                   absl::optional<std::vector<uint8_t>> early_data_context)
+                   absl::optional<std::vector<uint8_t>> early_data_context,
+                   QuicSSLConfig ssl_config)
         : server_address(server_address),
           client_address(client_address),
           ssl_capabilities(ssl_capabilities),
@@ -79,7 +81,8 @@
           alpn(alpn),
           alps(alps),
           quic_transport_params(quic_transport_params),
-          early_data_context(early_data_context) {}
+          early_data_context(early_data_context),
+          ssl_config(ssl_config) {}
 
     QuicSocketAddress server_address;
     QuicSocketAddress client_address;
@@ -90,6 +93,7 @@
     absl::optional<std::string> alps;
     std::vector<uint8_t> quic_transport_params;
     absl::optional<std::vector<uint8_t>> early_data_context;
+    QuicSSLConfig ssl_config;
   };
 
   struct ComputeSignatureArgs {
diff --git a/quic/test_tools/quic_test_utils.h b/quic/test_tools/quic_test_utils.h
index 9fbffe0..1200f01 100644
--- a/quic/test_tools/quic_test_utils.h
+++ b/quic/test_tools/quic_test_utils.h
@@ -1241,9 +1241,14 @@
 
   MockQuicCryptoServerStreamHelper* helper() { return &helper_; }
 
+  QuicSSLConfig GetSSLConfig() const override { return ssl_config_; }
+
+  QuicSSLConfig* ssl_config() { return &ssl_config_; }
+
  private:
   MockQuicSessionVisitor visitor_;
   MockQuicCryptoServerStreamHelper helper_;
+  QuicSSLConfig ssl_config_;
 };
 
 // A test implementation of QuicClientPushPromiseIndex::Delegate.