src/crypto/tls/handshake_messages.go | 4 +++- src/crypto/tls/handshake_messages_test.go | 99 ++++++++++++++++++++++++++++++++++++++++++++++++++--- src/crypto/tls/tls_test.go | 1 + diff --git a/src/crypto/tls/handshake_messages.go b/src/crypto/tls/handshake_messages.go index aa0b7db75dd493eed6b21460999bbd24289dfb00..511e073df57d724da7c9f31e15f3d4ff610a323b 100644 --- a/src/crypto/tls/handshake_messages.go +++ b/src/crypto/tls/handshake_messages.go @@ -5,6 +5,7 @@ package tls import ( + "bytes" "errors" "fmt" "slices" @@ -317,7 +318,8 @@ } }) }) } - if len(m.pskIdentities) > 0 { // pre_shared_key must be the last extension + // pre_shared_key must be the last extension + if len(m.pskIdentities) > 0 && (echInner || len(m.encryptedClientHello) == 0 || bytes.Equal(m.encryptedClientHello, []byte{byte(innerECHExt)})) { // RFC 8446, Section 4.2.11 exts.AddUint16(extensionPreSharedKey) exts.AddUint16LengthPrefixed(func(exts *cryptobyte.Builder) { diff --git a/src/crypto/tls/handshake_messages_test.go b/src/crypto/tls/handshake_messages_test.go index fa81a72b0de0181e43a8c1edd1eef453561251ee..80a7e76f816ffd7e067bb715d07a4e766b40ed45 100644 --- a/src/crypto/tls/handshake_messages_test.go +++ b/src/crypto/tls/handshake_messages_test.go @@ -215,7 +215,16 @@ m.pskModes = []uint8{pskModeDHE} case 2: m.pskModes = []uint8{pskModeDHE, pskModePlain} } - for i := 0; i < rand.Intn(5); i++ { + // clientHelloMsg.marshal uses echInner == false. If m.encryptedClientHello > 0 and + // does not equal []byte{innerECHExt}, then the psk extension will be omitted, + // so either only emit empty encryptedClientHello and psk, encryptedClientHello with + // []byte{innerECHExt} and psk, or random encryptedClientHello and no psk. + if rand.Intn(10) > 5 { + m.encryptedClientHello = randomBytes(rand.Intn(50)+1, rand) + } else { + if rand.Intn(10) > 5 { + m.encryptedClientHello = []byte{byte(innerECHExt)} + } var psk pskIdentity psk.obfuscatedTicketAge = uint32(rand.Intn(500000)) psk.label = randomBytes(rand.Intn(500)+1, rand) @@ -227,9 +236,6 @@ m.quicTransportParameters = randomBytes(rand.Intn(500), rand) } if rand.Intn(10) > 5 { m.earlyData = true - } - if rand.Intn(10) > 5 { - m.encryptedClientHello = randomBytes(rand.Intn(50)+1, rand) } return reflect.ValueOf(m) @@ -587,3 +593,88 @@ if serverHelloCopy.unmarshal(serverHelloBytes) { t.Fatal("Unmarshaled ServerHello with duplicate extensions") } } + +func TestECHRemoveOuterPSK(t *testing.T) { + r := rand.New(rand.NewSource(0)) + + for _, tc := range []struct { + name string + echInner bool + echExt []byte + expectRemoved bool + }{ + { + name: "echInner true", + echInner: true, + expectRemoved: false, + }, + { + name: "echInner true, no ech ext", + echInner: true, + expectRemoved: false, + }, + { + name: "echInner true, ech ext present", + echInner: true, + echExt: []byte{254}, + expectRemoved: false, + }, + { + name: "echInner true, ech ext present, inner ech sentinel", + echInner: true, + echExt: []byte{byte(innerECHExt)}, + expectRemoved: false, + }, + { + name: "echInner false, no ech ext", + echInner: false, + expectRemoved: false, + }, + { + name: "echInner false, ech ext present", + echInner: false, + echExt: []byte{254}, + expectRemoved: true, + }, + { + name: "echInner false, ech ext present, inner ech sentinel", + echInner: false, + echExt: []byte{byte(innerECHExt)}, + expectRemoved: false, + }, + } { + t.Run(tc.name, func(t *testing.T) { + ch := (&clientHelloMsg{}).Generate(r, 0).Interface().(*clientHelloMsg) + + ch.pskBinders = [][]byte{[]byte("test")} + ch.pskIdentities = []pskIdentity{{label: []byte("test")}} + ch.encryptedClientHello = tc.echExt + + b, err := ch.marshalMsg(tc.echInner) + if err != nil { + t.Fatal(err) + } + var rch clientHelloMsg + if !rch.unmarshal(b) { + t.Fatal("Failed to unmarshal ClientHello") + } + + if tc.expectRemoved { + if rch.pskIdentities != nil { + t.Error("expected PSK identities to be removed") + } + if rch.pskBinders != nil { + t.Error("expected PSK binders to be removed") + } + } else { + if rch.pskIdentities == nil { + t.Error("expected PSK identities to be present") + } + if rch.pskBinders == nil { + t.Error("expected PSK binders to be present") + } + } + }) + } + +} diff --git a/src/crypto/tls/tls_test.go b/src/crypto/tls/tls_test.go index 86322390862137772bfcce2e2ed472a7473a675c..655f58b747dbcabde08cd770f68427716fde0e7c 100644 --- a/src/crypto/tls/tls_test.go +++ b/src/crypto/tls/tls_test.go @@ -2379,6 +2379,7 @@ clientConfig.RootCAs = x509.NewCertPool() clientConfig.RootCAs.AddCert(secretCert) clientConfig.RootCAs.AddCert(publicCert) clientConfig.EncryptedClientHelloConfigList = echConfigList + clientConfig.ClientSessionCache = NewLRUClientSessionCache(2) serverConfig.InsecureSkipVerify = false serverConfig.Rand = rand.Reader serverConfig.Time = nil