Source file src/crypto/tls/conn_test.go

     1  // Copyright 2010 The Go Authors. All rights reserved.
     2  // Use of this source code is governed by a BSD-style
     3  // license that can be found in the LICENSE file.
     4  
     5  package tls
     6  
     7  import (
     8  	"bytes"
     9  	"io"
    10  	"net"
    11  	"testing"
    12  )
    13  
    14  func TestRoundUp(t *testing.T) {
    15  	if roundUp(0, 16) != 0 ||
    16  		roundUp(1, 16) != 16 ||
    17  		roundUp(15, 16) != 16 ||
    18  		roundUp(16, 16) != 16 ||
    19  		roundUp(17, 16) != 32 {
    20  		t.Error("roundUp broken")
    21  	}
    22  }
    23  
    24  // will be initialized with {0, 255, 255, ..., 255}
    25  var padding255Bad = [256]byte{}
    26  
    27  // will be initialized with {255, 255, 255, ..., 255}
    28  var padding255Good = [256]byte{255}
    29  
    30  var paddingTests = []struct {
    31  	in          []byte
    32  	good        bool
    33  	expectedLen int
    34  }{
    35  	{[]byte{1, 2, 3, 4, 0}, true, 4},
    36  	{[]byte{1, 2, 3, 4, 0, 1}, false, 0},
    37  	{[]byte{1, 2, 3, 4, 99, 99}, false, 0},
    38  	{[]byte{1, 2, 3, 4, 1, 1}, true, 4},
    39  	{[]byte{1, 2, 3, 2, 2, 2}, true, 3},
    40  	{[]byte{1, 2, 3, 3, 3, 3}, true, 2},
    41  	{[]byte{1, 2, 3, 4, 3, 3}, false, 0},
    42  	{[]byte{1, 4, 4, 4, 4, 4}, true, 1},
    43  	{[]byte{5, 5, 5, 5, 5, 5}, true, 0},
    44  	{[]byte{6, 6, 6, 6, 6, 6}, false, 0},
    45  	{padding255Bad[:], false, 0},
    46  	{padding255Good[:], true, 0},
    47  }
    48  
    49  func TestRemovePadding(t *testing.T) {
    50  	for i := 1; i < len(padding255Bad); i++ {
    51  		padding255Bad[i] = 255
    52  		padding255Good[i] = 255
    53  	}
    54  	for i, test := range paddingTests {
    55  		paddingLen, good := extractPadding(test.in)
    56  		expectedGood := byte(255)
    57  		if !test.good {
    58  			expectedGood = 0
    59  		}
    60  		if good != expectedGood {
    61  			t.Errorf("#%d: wrong validity, want:%d got:%d", i, expectedGood, good)
    62  		}
    63  		if good == 255 && len(test.in)-paddingLen != test.expectedLen {
    64  			t.Errorf("#%d: got %d, want %d", i, len(test.in)-paddingLen, test.expectedLen)
    65  		}
    66  	}
    67  }
    68  
    69  func TestCertificateSelection(t *testing.T) {
    70  	var certExampleCom = `308201713082011ba003020102021005a75ddf21014d5f417083b7a010ba2e300d06092a864886f70d01010b050030123110300e060355040a130741636d6520436f301e170d3136303831373231343135335a170d3137303831373231343135335a30123110300e060355040a130741636d6520436f305c300d06092a864886f70d0101010500034b003048024100b37f0fdd67e715bf532046ac34acbd8fdc4dabe2b598588f3f58b1f12e6219a16cbfe54d2b4b665396013589262360b6721efa27d546854f17cc9aeec6751db10203010001a34d304b300e0603551d0f0101ff0404030205a030130603551d25040c300a06082b06010505070301300c0603551d130101ff0402300030160603551d11040f300d820b6578616d706c652e636f6d300d06092a864886f70d01010b050003410059fc487866d3d855503c8e064ca32aac5e9babcece89ec597f8b2b24c17867f4a5d3b4ece06e795bfc5448ccbd2ffca1b3433171ebf3557a4737b020565350a0`
    71  
    72  	var certWildcardExampleCom = `308201743082011ea003020102021100a7aa6297c9416a4633af8bec2958c607300d06092a864886f70d01010b050030123110300e060355040a130741636d6520436f301e170d3136303831373231343231395a170d3137303831373231343231395a30123110300e060355040a130741636d6520436f305c300d06092a864886f70d0101010500034b003048024100b105afc859a711ee864114e7d2d46c2dcbe392d3506249f6c2285b0eb342cc4bf2d803677c61c0abde443f084745c1a6d62080e5664ef2cc8f50ad8a0ab8870b0203010001a34f304d300e0603551d0f0101ff0404030205a030130603551d25040c300a06082b06010505070301300c0603551d130101ff0402300030180603551d110411300f820d2a2e6578616d706c652e636f6d300d06092a864886f70d01010b0500034100af26088584d266e3f6566360cf862c7fecc441484b098b107439543144a2b93f20781988281e108c6d7656934e56950e1e5f2bcf38796b814ccb729445856c34`
    73  
    74  	var certFooExampleCom = `308201753082011fa00302010202101bbdb6070b0aeffc49008cde74deef29300d06092a864886f70d01010b050030123110300e060355040a130741636d6520436f301e170d3136303831373231343234345a170d3137303831373231343234345a30123110300e060355040a130741636d6520436f305c300d06092a864886f70d0101010500034b003048024100f00ac69d8ca2829f26216c7b50f1d4bbabad58d447706476cd89a2f3e1859943748aa42c15eedc93ac7c49e40d3b05ed645cb6b81c4efba60d961f44211a54eb0203010001a351304f300e0603551d0f0101ff0404030205a030130603551d25040c300a06082b06010505070301300c0603551d130101ff04023000301a0603551d1104133011820f666f6f2e6578616d706c652e636f6d300d06092a864886f70d01010b0500034100a0957fca6d1e0f1ef4b247348c7a8ca092c29c9c0ecc1898ea6b8065d23af6d922a410dd2335a0ea15edd1394cef9f62c9e876a21e35250a0b4fe1ddceba0f36`
    75  
    76  	config := Config{
    77  		Certificates: []Certificate{
    78  			{
    79  				Certificate: [][]byte{fromHex(certExampleCom)},
    80  			},
    81  			{
    82  				Certificate: [][]byte{fromHex(certWildcardExampleCom)},
    83  			},
    84  			{
    85  				Certificate: [][]byte{fromHex(certFooExampleCom)},
    86  			},
    87  		},
    88  	}
    89  
    90  	config.BuildNameToCertificate()
    91  
    92  	pointerToIndex := func(c *Certificate) int {
    93  		for i := range config.Certificates {
    94  			if c == &config.Certificates[i] {
    95  				return i
    96  			}
    97  		}
    98  		return -1
    99  	}
   100  
   101  	certificateForName := func(name string) *Certificate {
   102  		clientHello := &ClientHelloInfo{
   103  			ServerName: name,
   104  		}
   105  		if cert, err := config.getCertificate(clientHello); err != nil {
   106  			t.Errorf("unable to get certificate for name '%s': %s", name, err)
   107  			return nil
   108  		} else {
   109  			return cert
   110  		}
   111  	}
   112  
   113  	if n := pointerToIndex(certificateForName("example.com")); n != 0 {
   114  		t.Errorf("example.com returned certificate %d, not 0", n)
   115  	}
   116  	if n := pointerToIndex(certificateForName("bar.example.com")); n != 1 {
   117  		t.Errorf("bar.example.com returned certificate %d, not 1", n)
   118  	}
   119  	if n := pointerToIndex(certificateForName("foo.example.com")); n != 2 {
   120  		t.Errorf("foo.example.com returned certificate %d, not 2", n)
   121  	}
   122  	if n := pointerToIndex(certificateForName("foo.bar.example.com")); n != 0 {
   123  		t.Errorf("foo.bar.example.com returned certificate %d, not 0", n)
   124  	}
   125  }
   126  
   127  // TestBrokenCertificateSkipped checks that a Certificate in Config.Certificates
   128  // whose leaf doesn't parse as X.509 doesn't prevent the next, valid certificate
   129  // from being selected. It exercises both the legacy BuildNameToCertificate path
   130  // and the SupportsCertificate-based selection.
   131  func TestBrokenCertificateSkipped(t *testing.T) {
   132  	brokenCert := Certificate{Certificate: [][]byte{[]byte("not a valid X.509 certificate")}}
   133  	for _, test := range []struct {
   134  		name       string
   135  		buildIndex bool
   136  	}{
   137  		{name: "BuildNameToCertificate", buildIndex: true},
   138  		{name: "SupportsCertificate", buildIndex: false},
   139  	} {
   140  		t.Run(test.name, func(t *testing.T) {
   141  			serverConfig := testConfigServer.Clone()
   142  			serverConfig.Certificates = []Certificate{brokenCert, testECDSAP256Cert}
   143  			if test.buildIndex {
   144  				serverConfig.BuildNameToCertificate()
   145  			}
   146  			clientConfig := testConfigClient.Clone()
   147  			_, cs, err := testHandshake(t, clientConfig, serverConfig)
   148  			if err != nil {
   149  				t.Fatalf("handshake failed: %v", err)
   150  			}
   151  			if !cs.PeerCertificates[0].Equal(testECDSAP256Cert.Leaf) {
   152  				t.Fatalf("handshake succeeded but wrong certificate was used")
   153  			}
   154  		})
   155  	}
   156  }
   157  
   158  // Run with multiple crypto configs to test the logic for computing TLS record overheads.
   159  func runDynamicRecordSizingTest(t *testing.T, serverConfig *Config) {
   160  	clientConn, serverConn := localPipe(t)
   161  
   162  	serverConfig = serverConfig.Clone()
   163  	serverConfig.DynamicRecordSizingDisabled = false
   164  	tlsConn := Server(serverConn, serverConfig)
   165  
   166  	clientConfig := testConfigClient.Clone()
   167  	clientConfig.MinVersion = serverConfig.MinVersion
   168  	clientConfig.MaxVersion = serverConfig.MaxVersion
   169  	clientConfig.CipherSuites = serverConfig.CipherSuites
   170  
   171  	handshakeDone := make(chan struct{})
   172  	recordSizesChan := make(chan []int, 1)
   173  	defer func() { <-recordSizesChan }() // wait for the goroutine to exit
   174  	go func() {
   175  		// This goroutine performs a TLS handshake over clientConn and
   176  		// then reads TLS records until EOF. It writes a slice that
   177  		// contains all the record sizes to recordSizesChan.
   178  		defer close(recordSizesChan)
   179  		defer clientConn.Close()
   180  
   181  		tlsConn := Client(clientConn, clientConfig)
   182  		if err := tlsConn.Handshake(); err != nil {
   183  			t.Errorf("Error from client handshake: %v", err)
   184  			return
   185  		}
   186  		close(handshakeDone)
   187  
   188  		var recordHeader [recordHeaderLen]byte
   189  		var record []byte
   190  		var recordSizes []int
   191  
   192  		for {
   193  			n, err := io.ReadFull(clientConn, recordHeader[:])
   194  			if err == io.EOF {
   195  				break
   196  			}
   197  			if err != nil || n != len(recordHeader) {
   198  				t.Errorf("io.ReadFull = %d, %v", n, err)
   199  				return
   200  			}
   201  
   202  			length := int(recordHeader[3])<<8 | int(recordHeader[4])
   203  			if len(record) < length {
   204  				record = make([]byte, length)
   205  			}
   206  
   207  			n, err = io.ReadFull(clientConn, record[:length])
   208  			if err != nil || n != length {
   209  				t.Errorf("io.ReadFull = %d, %v", n, err)
   210  				return
   211  			}
   212  
   213  			recordSizes = append(recordSizes, recordHeaderLen+length)
   214  		}
   215  
   216  		recordSizesChan <- recordSizes
   217  	}()
   218  
   219  	if err := tlsConn.Handshake(); err != nil {
   220  		t.Fatalf("Error from server handshake: %s", err)
   221  	}
   222  	<-handshakeDone
   223  
   224  	// The server writes these plaintexts in order.
   225  	plaintext := bytes.Join([][]byte{
   226  		bytes.Repeat([]byte("x"), recordSizeBoostThreshold),
   227  		bytes.Repeat([]byte("y"), maxPlaintext*2),
   228  		bytes.Repeat([]byte("z"), maxPlaintext),
   229  	}, nil)
   230  
   231  	if _, err := tlsConn.Write(plaintext); err != nil {
   232  		t.Fatalf("Error from server write: %s", err)
   233  	}
   234  	if err := tlsConn.Close(); err != nil {
   235  		t.Fatalf("Error from server close: %s", err)
   236  	}
   237  
   238  	recordSizes := <-recordSizesChan
   239  	if recordSizes == nil {
   240  		t.Fatalf("Client encountered an error")
   241  	}
   242  
   243  	// Drop the size of the second to last record, which is likely to be
   244  	// truncated, and the last record, which is a close_notify alert.
   245  	recordSizes = recordSizes[:len(recordSizes)-2]
   246  
   247  	// recordSizes should contain a series of records smaller than
   248  	// tcpMSSEstimate followed by some larger than maxPlaintext.
   249  	seenLargeRecord := false
   250  	for i, size := range recordSizes {
   251  		if !seenLargeRecord {
   252  			if size > (i+1)*tcpMSSEstimate {
   253  				t.Fatalf("Record #%d has size %d, which is too large too soon", i, size)
   254  			}
   255  			if size >= maxPlaintext {
   256  				seenLargeRecord = true
   257  			}
   258  		} else if size <= maxPlaintext {
   259  			t.Fatalf("Record #%d has size %d but should be full sized", i, size)
   260  		}
   261  	}
   262  
   263  	if !seenLargeRecord {
   264  		t.Fatalf("No large records observed")
   265  	}
   266  }
   267  
   268  func TestDynamicRecordSizingWithStreamCipher(t *testing.T) {
   269  	skipFIPS(t) // No RC4 in FIPS mode.
   270  
   271  	config := testConfigServer.Clone()
   272  	config.MaxVersion = VersionTLS12
   273  	config.CipherSuites = []uint16{TLS_RSA_WITH_RC4_128_SHA}
   274  	runDynamicRecordSizingTest(t, config)
   275  }
   276  
   277  func TestDynamicRecordSizingWithCBC(t *testing.T) {
   278  	skipFIPS(t) // No CBC cipher suites in defaultCipherSuitesFIPS.
   279  
   280  	config := testConfigServer.Clone()
   281  	config.MaxVersion = VersionTLS12
   282  	config.CipherSuites = []uint16{TLS_RSA_WITH_AES_256_CBC_SHA}
   283  	runDynamicRecordSizingTest(t, config)
   284  }
   285  
   286  func TestDynamicRecordSizingWithAEAD(t *testing.T) {
   287  	config := testConfigServer.Clone()
   288  	config.MaxVersion = VersionTLS12
   289  	config.CipherSuites = []uint16{TLS_ECDHE_RSA_WITH_AES_128_GCM_SHA256}
   290  	runDynamicRecordSizingTest(t, config)
   291  }
   292  
   293  func TestDynamicRecordSizingWithTLSv13(t *testing.T) {
   294  	config := testConfigServer.Clone()
   295  	runDynamicRecordSizingTest(t, config)
   296  }
   297  
   298  // hairpinConn is a net.Conn that makes a “hairpin” call when closed, back into
   299  // the tls.Conn which is calling it.
   300  type hairpinConn struct {
   301  	net.Conn
   302  	tlsConn *Conn
   303  }
   304  
   305  func (conn *hairpinConn) Close() error {
   306  	conn.tlsConn.ConnectionState()
   307  	return nil
   308  }
   309  
   310  func TestHairpinInClose(t *testing.T) {
   311  	// This tests that the underlying net.Conn can call back into the
   312  	// tls.Conn when being closed without deadlocking.
   313  	client, server := localPipe(t)
   314  	defer server.Close()
   315  	defer client.Close()
   316  
   317  	conn := &hairpinConn{client, nil}
   318  	tlsConn := Server(conn, &Config{
   319  		GetCertificate: func(*ClientHelloInfo) (*Certificate, error) {
   320  			panic("unreachable")
   321  		},
   322  	})
   323  	conn.tlsConn = tlsConn
   324  
   325  	// This call should not deadlock.
   326  	tlsConn.Close()
   327  }
   328  
   329  func TestRecordBadVersionTLS13(t *testing.T) {
   330  	client, server := localPipe(t)
   331  	defer server.Close()
   332  	defer client.Close()
   333  
   334  	clientConfig := testConfigClient.Clone()
   335  	clientConfig.MinVersion, clientConfig.MaxVersion = VersionTLS13, VersionTLS13
   336  	serverConfig := testConfigServer.Clone()
   337  	serverConfig.MinVersion, serverConfig.MaxVersion = VersionTLS13, VersionTLS13
   338  
   339  	go func() {
   340  		tlsConn := Client(client, clientConfig)
   341  		if err := tlsConn.Handshake(); err != nil {
   342  			t.Errorf("Error from client handshake: %v", err)
   343  			return
   344  		}
   345  		tlsConn.vers = 0x1111
   346  		tlsConn.Write([]byte{1})
   347  	}()
   348  
   349  	tlsConn := Server(server, serverConfig)
   350  	if err := tlsConn.Handshake(); err != nil {
   351  		t.Errorf("Error from client handshake: %v", err)
   352  		return
   353  	}
   354  
   355  	expectedErr := "tls: received record with version 1111 when expecting version 303"
   356  
   357  	_, err := tlsConn.Read(make([]byte, 10))
   358  	if err.Error() != expectedErr {
   359  		t.Fatalf("unexpected error: got %q, want %q", err, expectedErr)
   360  	}
   361  }
   362  
   363  // TestKeyUpdateSpamPostHandshakeTLS13 tests that TLS 1.3 KeyUpdate messages do
   364  // not count as a state-advancing message after a handshake has been completed.
   365  func TestKeyUpdateSpamPostHandshakeTLS13(t *testing.T) {
   366  	client, server := localPipe(t)
   367  	defer server.Close()
   368  	defer client.Close()
   369  
   370  	go func() {
   371  		c := Client(client, testConfigClient.Clone())
   372  		if err := c.Handshake(); err != nil {
   373  			t.Error(err)
   374  			return
   375  		}
   376  		ku, err := (&keyUpdateMsg{}).marshal()
   377  		if err != nil {
   378  			t.Error(err)
   379  			return
   380  		}
   381  		cs := cipherSuiteTLS13ByID(c.cipherSuite)
   382  		for i := 0; i <= maxUselessRecords; i++ {
   383  			c.writeRecordLocked(recordTypeHandshake, ku)
   384  			c.setWriteTrafficSecret(cs, QUICEncryptionLevelInitial, cs.nextTrafficSecret(c.out.trafficSecret))
   385  		}
   386  	}()
   387  	s := Server(server, testConfigServer.Clone())
   388  	if err := s.Handshake(); err != nil {
   389  		t.Fatal(err)
   390  	}
   391  	expectedErr := "tls: too many non-advancing records"
   392  	if _, err := s.Read(make([]byte, 1)); err == nil || err.Error() != expectedErr {
   393  		t.Fatalf("unexpected error: got %v, want %q", err, expectedErr)
   394  	}
   395  }
   396  

View as plain text