Source file src/crypto/tls/handshake_test.go

     1  // Copyright 2013 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  	"bufio"
     9  	"bytes"
    10  	"context"
    11  	"crypto/internal/boring"
    12  	"encoding/hex"
    13  	"errors"
    14  	"flag"
    15  	"fmt"
    16  	"internal/testenv"
    17  	"io"
    18  	"net"
    19  	"os"
    20  	"os/exec"
    21  	"runtime"
    22  	"strconv"
    23  	"strings"
    24  	"sync"
    25  	"testing"
    26  	"testing/cryptotest"
    27  	"time"
    28  )
    29  
    30  // TLS reference tests run a connection against a reference implementation
    31  // (OpenSSL) of TLS and record the bytes of the resulting connection. The Go
    32  // code, during a test, is configured with deterministic randomness and so the
    33  // reference test can be reproduced exactly in the future.
    34  //
    35  // In order to save everyone who wishes to run the tests from needing the
    36  // reference implementation installed, the reference connections are saved in
    37  // files in the testdata directory. Thus running the tests involves nothing
    38  // external, but creating and updating them requires the reference
    39  // implementation.
    40  //
    41  // Tests can be updated by running them with the -update flag. This will cause
    42  // the test files for failing tests to be regenerated. Since the reference
    43  // implementation will always generate fresh random numbers, large parts of the
    44  // reference connection will always change.
    45  
    46  var (
    47  	update       = flag.Bool("update", false, "update golden files on failure")
    48  	keyFile      = flag.String("keylog", "", "destination file for KeyLogWriter")
    49  	bogoMode     = flag.Bool("bogo-mode", false, "Enabled bogo shim mode, ignore everything else")
    50  	bogoFilter   = flag.String("bogo-filter", "", "BoGo test filter")
    51  	bogoLocalDir = flag.String("bogo-local-dir", "",
    52  		"If not-present, checkout BoGo into this dir, or otherwise use it as a pre-existing checkout")
    53  	bogoReport = flag.String("bogo-html-report", "", "File path to render an HTML report with BoGo results")
    54  )
    55  
    56  func runTestAndUpdateIfNeeded(t *testing.T, name string, run func(t *testing.T, update bool)) {
    57  	skipFIPS(t) // FIPS 140-3 mode changes the advertised parameters.
    58  
    59  	// Go+BoringCrypto's boring.RandReader ignores the testing override set by
    60  	// cryptotest.SetGlobalRandom, so e.g. ECDH key generation would be
    61  	// non-deterministic. Setting cryptocustomrand=1 makes rand.CustomReader
    62  	// pass the caller's reader (the testing source) through instead.
    63  	if boring.Enabled {
    64  		testenv.SetGODEBUG(t, "cryptocustomrand=1")
    65  	}
    66  
    67  	success := t.Run(name, func(t *testing.T) {
    68  		cryptotest.SetGlobalRandom(t, 0)
    69  		run(t, false)
    70  	})
    71  
    72  	if !success && *update {
    73  		t.Run(name+"#update", func(t *testing.T) {
    74  			cryptotest.SetGlobalRandom(t, 0)
    75  			run(t, true)
    76  		})
    77  	}
    78  }
    79  
    80  // checkOpenSSLVersion ensures that the version of OpenSSL looks reasonable
    81  // before updating the test data.
    82  func checkOpenSSLVersion() error {
    83  	if !*update {
    84  		return nil
    85  	}
    86  
    87  	openssl := exec.Command("openssl", "version")
    88  	output, err := openssl.CombinedOutput()
    89  	if err != nil {
    90  		return err
    91  	}
    92  
    93  	version := string(output)
    94  	if strings.HasPrefix(version, "OpenSSL 1.1.1") {
    95  		return nil
    96  	}
    97  
    98  	println("***********************************************")
    99  	println("")
   100  	println("You need to build OpenSSL 1.1.1 from source in order")
   101  	println("to update the test data.")
   102  	println("")
   103  	println("Configure it with:")
   104  	println("./Configure enable-weak-ssl-ciphers no-shared")
   105  	println("and then add the apps/ directory at the front of your PATH.")
   106  	println("***********************************************")
   107  
   108  	return errors.New("version of OpenSSL does not appear to be suitable for updating test data")
   109  }
   110  
   111  // recordingConn is a net.Conn that records the traffic that passes through it.
   112  // WriteTo can be used to produce output that can be later be loaded with
   113  // ParseTestData.
   114  type recordingConn struct {
   115  	net.Conn
   116  	sync.Mutex
   117  	flows   [][]byte
   118  	reading bool
   119  }
   120  
   121  func (r *recordingConn) Read(b []byte) (n int, err error) {
   122  	if n, err = r.Conn.Read(b); n == 0 {
   123  		return
   124  	}
   125  	b = b[:n]
   126  
   127  	r.Lock()
   128  	defer r.Unlock()
   129  
   130  	if l := len(r.flows); l == 0 || !r.reading {
   131  		buf := make([]byte, len(b))
   132  		copy(buf, b)
   133  		r.flows = append(r.flows, buf)
   134  	} else {
   135  		r.flows[l-1] = append(r.flows[l-1], b[:n]...)
   136  	}
   137  	r.reading = true
   138  	return
   139  }
   140  
   141  func (r *recordingConn) Write(b []byte) (n int, err error) {
   142  	if n, err = r.Conn.Write(b); n == 0 {
   143  		return
   144  	}
   145  	b = b[:n]
   146  
   147  	r.Lock()
   148  	defer r.Unlock()
   149  
   150  	if l := len(r.flows); l == 0 || r.reading {
   151  		buf := make([]byte, len(b))
   152  		copy(buf, b)
   153  		r.flows = append(r.flows, buf)
   154  	} else {
   155  		r.flows[l-1] = append(r.flows[l-1], b[:n]...)
   156  	}
   157  	r.reading = false
   158  	return
   159  }
   160  
   161  // WriteTo writes Go source code to w that contains the recorded traffic.
   162  func (r *recordingConn) WriteTo(w io.Writer) (int64, error) {
   163  	// TLS always starts with a client to server flow.
   164  	clientToServer := true
   165  	var written int64
   166  	for i, flow := range r.flows {
   167  		source, dest := "client", "server"
   168  		if !clientToServer {
   169  			source, dest = dest, source
   170  		}
   171  		n, err := fmt.Fprintf(w, ">>> Flow %d (%s to %s)\n", i+1, source, dest)
   172  		written += int64(n)
   173  		if err != nil {
   174  			return written, err
   175  		}
   176  		dumper := hex.Dumper(w)
   177  		n, err = dumper.Write(flow)
   178  		written += int64(n)
   179  		if err != nil {
   180  			return written, err
   181  		}
   182  		err = dumper.Close()
   183  		if err != nil {
   184  			return written, err
   185  		}
   186  		clientToServer = !clientToServer
   187  	}
   188  	return written, nil
   189  }
   190  
   191  func parseTestData(r io.Reader) (flows [][]byte, err error) {
   192  	var currentFlow []byte
   193  
   194  	scanner := bufio.NewScanner(r)
   195  	for scanner.Scan() {
   196  		line := scanner.Text()
   197  		// If the line starts with ">>> " then it marks the beginning
   198  		// of a new flow.
   199  		if strings.HasPrefix(line, ">>> ") {
   200  			if len(currentFlow) > 0 || len(flows) > 0 {
   201  				flows = append(flows, currentFlow)
   202  				currentFlow = nil
   203  			}
   204  			continue
   205  		}
   206  
   207  		// Otherwise the line is a line of hex dump that looks like:
   208  		// 00000170  fc f5 06 bf (...)  |.....X{&?......!|
   209  		// (Some bytes have been omitted from the middle section.)
   210  		_, after, ok := strings.Cut(line, " ")
   211  		if !ok {
   212  			return nil, errors.New("invalid test data")
   213  		}
   214  		line = after
   215  
   216  		before, _, ok := strings.Cut(line, "|")
   217  		if !ok {
   218  			return nil, errors.New("invalid test data")
   219  		}
   220  		line = before
   221  
   222  		hexBytes := strings.Fields(line)
   223  		for _, hexByte := range hexBytes {
   224  			val, err := strconv.ParseUint(hexByte, 16, 8)
   225  			if err != nil {
   226  				return nil, errors.New("invalid hex byte in test data: " + err.Error())
   227  			}
   228  			currentFlow = append(currentFlow, byte(val))
   229  		}
   230  	}
   231  
   232  	if len(currentFlow) > 0 {
   233  		flows = append(flows, currentFlow)
   234  	}
   235  
   236  	return flows, nil
   237  }
   238  
   239  // replayingConn is a net.Conn that replays flows recorded by recordingConn.
   240  type replayingConn struct {
   241  	t testing.TB
   242  	sync.Mutex
   243  	flows   [][]byte
   244  	reading bool
   245  }
   246  
   247  var _ net.Conn = (*replayingConn)(nil)
   248  
   249  func (r *replayingConn) Read(b []byte) (n int, err error) {
   250  	r.Lock()
   251  	defer r.Unlock()
   252  
   253  	if !r.reading {
   254  		r.t.Errorf("expected write, got read")
   255  		return 0, fmt.Errorf("recording expected write, got read")
   256  	}
   257  
   258  	n = copy(b, r.flows[0])
   259  	r.flows[0] = r.flows[0][n:]
   260  	if len(r.flows[0]) == 0 {
   261  		r.flows = r.flows[1:]
   262  		if len(r.flows) == 0 {
   263  			return n, io.EOF
   264  		} else {
   265  			r.reading = false
   266  		}
   267  	}
   268  	return n, nil
   269  }
   270  
   271  func (r *replayingConn) Write(b []byte) (n int, err error) {
   272  	r.Lock()
   273  	defer r.Unlock()
   274  
   275  	if r.reading {
   276  		r.t.Errorf("expected read, got write")
   277  		return 0, fmt.Errorf("recording expected read, got write")
   278  	}
   279  
   280  	if !bytes.HasPrefix(r.flows[0], b) {
   281  		r.t.Errorf("write mismatch: expected %x, got %x", r.flows[0], b)
   282  		return 0, fmt.Errorf("write mismatch")
   283  	}
   284  	r.flows[0] = r.flows[0][len(b):]
   285  	if len(r.flows[0]) == 0 {
   286  		r.flows = r.flows[1:]
   287  		r.reading = true
   288  	}
   289  	return len(b), nil
   290  }
   291  
   292  func (r *replayingConn) Close() error {
   293  	r.Lock()
   294  	defer r.Unlock()
   295  
   296  	if len(r.flows) > 0 {
   297  		r.t.Errorf("closed with unfinished flows")
   298  		return fmt.Errorf("unexpected close")
   299  	}
   300  	return nil
   301  }
   302  
   303  func (r *replayingConn) LocalAddr() net.Addr                { return nil }
   304  func (r *replayingConn) RemoteAddr() net.Addr               { return nil }
   305  func (r *replayingConn) SetDeadline(t time.Time) error      { return nil }
   306  func (r *replayingConn) SetReadDeadline(t time.Time) error  { return nil }
   307  func (r *replayingConn) SetWriteDeadline(t time.Time) error { return nil }
   308  
   309  // tempFile creates a temp file containing contents and returns its path.
   310  func tempFile(contents string) string {
   311  	file, err := os.CreateTemp("", "go-tls-test")
   312  	if err != nil {
   313  		panic("failed to create temp file: " + err.Error())
   314  	}
   315  	path := file.Name()
   316  	file.WriteString(contents)
   317  	file.Close()
   318  	return path
   319  }
   320  
   321  // localListener is set up by TestMain and used by localPipe to create Conn
   322  // pairs like net.Pipe, but connected by an actual buffered TCP connection.
   323  var localListener struct {
   324  	mu   sync.Mutex
   325  	addr net.Addr
   326  	ch   chan net.Conn
   327  }
   328  
   329  const localFlakes = 0 // change to 1 or 2 to exercise localServer/localPipe handling of mismatches
   330  
   331  func localServer(l net.Listener) {
   332  	for n := 0; ; n++ {
   333  		c, err := l.Accept()
   334  		if err != nil {
   335  			return
   336  		}
   337  		if localFlakes == 1 && n%2 == 0 {
   338  			c.Close()
   339  			continue
   340  		}
   341  		localListener.ch <- c
   342  	}
   343  }
   344  
   345  var isConnRefused = func(err error) bool { return false }
   346  
   347  func localPipe(t testing.TB) (net.Conn, net.Conn) {
   348  	localListener.mu.Lock()
   349  	defer localListener.mu.Unlock()
   350  
   351  	addr := localListener.addr
   352  
   353  	var err error
   354  Dialing:
   355  	// We expect a rare mismatch, but probably not 5 in a row.
   356  	for i := 0; i < 5; i++ {
   357  		tooSlow := time.NewTimer(1 * time.Second)
   358  		defer tooSlow.Stop()
   359  		var c1 net.Conn
   360  		c1, err = net.Dial(addr.Network(), addr.String())
   361  		if err != nil {
   362  			if runtime.GOOS == "dragonfly" && (isConnRefused(err) || os.IsTimeout(err)) {
   363  				// golang.org/issue/29583: Dragonfly sometimes returns a spurious
   364  				// ECONNREFUSED or ETIMEDOUT.
   365  				<-tooSlow.C
   366  				continue
   367  			}
   368  			t.Fatalf("localPipe: %v", err)
   369  		}
   370  		if localFlakes == 2 && i == 0 {
   371  			c1.Close()
   372  			continue
   373  		}
   374  		for {
   375  			select {
   376  			case <-tooSlow.C:
   377  				t.Logf("localPipe: timeout waiting for %v", c1.LocalAddr())
   378  				c1.Close()
   379  				continue Dialing
   380  
   381  			case c2 := <-localListener.ch:
   382  				if c2.RemoteAddr().String() == c1.LocalAddr().String() {
   383  					t.Cleanup(func() { c1.Close() })
   384  					t.Cleanup(func() { c2.Close() })
   385  					return c1, c2
   386  				}
   387  				t.Logf("localPipe: unexpected connection: %v != %v", c2.RemoteAddr(), c1.LocalAddr())
   388  				c2.Close()
   389  			}
   390  		}
   391  	}
   392  
   393  	t.Fatalf("localPipe: failed to connect: %v", err)
   394  	panic("unreachable")
   395  }
   396  
   397  func TestMain(m *testing.M) {
   398  	flag.Usage = func() {
   399  		fmt.Fprintf(flag.CommandLine.Output(), "Usage of %s:\n", os.Args)
   400  		flag.PrintDefaults()
   401  		if *bogoMode {
   402  			os.Exit(89)
   403  		}
   404  	}
   405  
   406  	flag.Parse()
   407  
   408  	if *bogoMode {
   409  		bogoShim()
   410  		os.Exit(0)
   411  	}
   412  
   413  	os.Exit(runMain(m))
   414  }
   415  
   416  func runMain(m *testing.M) int {
   417  	// Cipher suites preferences change based on the architecture. Force them to
   418  	// the version without AES acceleration for test consistency.
   419  	hasAESGCMHardwareSupport = false
   420  
   421  	// Set up localPipe.
   422  	l, err := net.Listen("tcp", "127.0.0.1:0")
   423  	if err != nil {
   424  		l, err = net.Listen("tcp6", "[::1]:0")
   425  	}
   426  	if err != nil {
   427  		fmt.Fprintf(os.Stderr, "Failed to open local listener: %v", err)
   428  		os.Exit(1)
   429  	}
   430  	localListener.ch = make(chan net.Conn)
   431  	localListener.addr = l.Addr()
   432  	defer l.Close()
   433  	go localServer(l)
   434  
   435  	if err := checkOpenSSLVersion(); err != nil {
   436  		fmt.Fprintf(os.Stderr, "Error: %v", err)
   437  		os.Exit(1)
   438  	}
   439  
   440  	rootCAPath := tempFile(testRootCertPEM)
   441  	defer os.Remove(rootCAPath)
   442  	defaultClientCommand = []string{"openssl", "s_client", "-no_ticket",
   443  		"-verify", "1", "-verify_return_error", "-CAfile", rootCAPath,
   444  		"-servername", "test.golang.example", "-attime", fmt.Sprint(testTime().Unix())}
   445  
   446  	clientRootCAPath := tempFile(testClientRootCertPEM)
   447  	defer os.Remove(clientRootCAPath)
   448  	serverCommand = []string{"openssl", "s_server", "-no_ticket", "-num_tickets", "0",
   449  		"-naccept", "1", "-verify_return_error", "-verifyCAfile", clientRootCAPath,
   450  		"-attime", fmt.Sprint(testTime().Unix())}
   451  
   452  	if *keyFile != "" {
   453  		f, err := os.OpenFile(*keyFile, os.O_APPEND|os.O_CREATE|os.O_WRONLY, 0644)
   454  		if err != nil {
   455  			panic("failed to open -keylog file: " + err.Error())
   456  		}
   457  		testConfigClient.KeyLogWriter = f
   458  		testConfigServer.KeyLogWriter = f
   459  		testConfigFIPS140.KeyLogWriter = f
   460  		defer f.Close()
   461  	}
   462  
   463  	return m.Run()
   464  }
   465  
   466  func testHandshake(t *testing.T, clientConfig, serverConfig *Config) (serverState, clientState ConnectionState, err error) {
   467  	const sentinel = "SENTINEL\n"
   468  	c, s := localPipe(t)
   469  	errChan := make(chan error, 1)
   470  	go func() {
   471  		cli := Client(c, clientConfig)
   472  		err := cli.Handshake()
   473  		if err != nil {
   474  			errChan <- fmt.Errorf("client: %v", err)
   475  			c.Close()
   476  			return
   477  		}
   478  		clientState = cli.ConnectionState()
   479  		buf, err := io.ReadAll(cli)
   480  		if err != nil {
   481  			if serverConfig.ClientAuth != NoClientCert && clientState.Version == VersionTLS13 {
   482  				// In TLS 1.3, client certificates are sent after the server's
   483  				// handshake has completed, and the client only learns about it
   484  				// reading the alert after the handshake.
   485  				errChan <- fmt.Errorf("client (from Read): %v", err)
   486  				c.Close()
   487  				return
   488  			} else {
   489  				t.Errorf("failed to call cli.Read: %v", err)
   490  			}
   491  		}
   492  		defer func() { errChan <- nil }()
   493  		if got := string(buf); got != sentinel {
   494  			t.Errorf("read %q from TLS connection, but expected %q", got, sentinel)
   495  		}
   496  		// We discard the error because after ReadAll returns the server must
   497  		// have already closed the connection. Sending data (the closeNotify
   498  		// alert) can cause a reset, that will make Close return an error.
   499  		cli.Close()
   500  	}()
   501  	server := Server(s, serverConfig)
   502  	err = server.Handshake()
   503  	if err == nil {
   504  		serverState = server.ConnectionState()
   505  		if _, err := io.WriteString(server, sentinel); err != nil {
   506  			t.Errorf("failed to call server.Write: %v", err)
   507  		}
   508  		if err := server.Close(); err != nil {
   509  			t.Errorf("failed to call server.Close: %v", err)
   510  		}
   511  	} else {
   512  		err = fmt.Errorf("server: %v", err)
   513  		s.Close()
   514  	}
   515  	err = errors.Join(err, <-errChan)
   516  	return
   517  }
   518  
   519  func fromHex(s string) []byte {
   520  	b, _ := hex.DecodeString(s)
   521  	return b
   522  }
   523  
   524  func TestServerHelloTrailingMessage(t *testing.T) {
   525  	// In TLS 1.3 the change cipher spec message is optional. If a CCS message
   526  	// is not sent, after reading the ServerHello, the read traffic secret is
   527  	// set, and all following messages must be encrypted. If the server sends
   528  	// additional unencrypted messages in a record with the ServerHello, the
   529  	// client must either fail or ignore the additional messages.
   530  
   531  	c, s := localPipe(t)
   532  	go func() {
   533  		ctx := context.Background()
   534  		srv := Server(s, testConfigServer.Clone())
   535  		clientHello, _, err := srv.readClientHello(ctx)
   536  		if err != nil {
   537  			testFatal(t, err)
   538  		}
   539  
   540  		hs := serverHandshakeStateTLS13{
   541  			c:           srv,
   542  			ctx:         ctx,
   543  			clientHello: clientHello,
   544  		}
   545  		if err := hs.processClientHello(); err != nil {
   546  			testFatal(t, err)
   547  		}
   548  		if err := transcriptMsg(hs.clientHello, hs.transcript); err != nil {
   549  			testFatal(t, err)
   550  		}
   551  
   552  		record, err := concatHandshakeMessages(hs.hello, &encryptedExtensionsMsg{alpnProtocol: "h2"})
   553  		if err != nil {
   554  			testFatal(t, err)
   555  		}
   556  
   557  		if _, err := s.Write(record); err != nil {
   558  			testFatal(t, err)
   559  		}
   560  		srv.Close()
   561  	}()
   562  
   563  	cli := Client(c, testConfigClient.Clone())
   564  	expectedErr := "tls: handshake buffer not empty before setting read traffic secret"
   565  	if err := cli.Handshake(); err == nil {
   566  		t.Fatal("expected error from incomplete handshake, got nil")
   567  	} else if err.Error() != expectedErr {
   568  		t.Fatalf("expected error %q, got %q", expectedErr, err.Error())
   569  	}
   570  }
   571  
   572  func TestClientHelloTrailingMessage(t *testing.T) {
   573  	// Same as TestServerHelloTrailingMessage but for the client side.
   574  
   575  	c, s := localPipe(t)
   576  	go func() {
   577  		cli := Client(c, testConfigClient.Clone())
   578  
   579  		hello, _, _, err := cli.makeClientHello()
   580  		if err != nil {
   581  			testFatal(t, err)
   582  		}
   583  
   584  		record, err := concatHandshakeMessages(hello, &certificateMsgTLS13{})
   585  		if err != nil {
   586  			testFatal(t, err)
   587  		}
   588  
   589  		if _, err := c.Write(record); err != nil {
   590  			testFatal(t, err)
   591  		}
   592  		cli.Close()
   593  	}()
   594  
   595  	srv := Server(s, testConfigServer.Clone())
   596  	expectedErr := "tls: handshake buffer not empty before setting read traffic secret"
   597  	if err := srv.Handshake(); err == nil {
   598  		t.Fatal("expected error from incomplete handshake, got nil")
   599  	} else if err.Error() != expectedErr {
   600  		t.Fatalf("expected error %q, got %q", expectedErr, err.Error())
   601  	}
   602  }
   603  
   604  func TestDoubleClientHelloHRR(t *testing.T) {
   605  	// If a client sends two ClientHello messages in a single record, and the
   606  	// server sends a HRR after reading the first ClientHello, the server must
   607  	// either fail or ignore the trailing ClientHello.
   608  
   609  	c, s := localPipe(t)
   610  
   611  	go func() {
   612  		cli := Client(c, testConfigClient.Clone())
   613  
   614  		hello, _, _, err := cli.makeClientHello()
   615  		if err != nil {
   616  			testFatal(t, err)
   617  		}
   618  		hello.keyShares = nil
   619  
   620  		record, err := concatHandshakeMessages(hello, hello)
   621  		if err != nil {
   622  			testFatal(t, err)
   623  		}
   624  
   625  		if _, err := c.Write(record); err != nil {
   626  			testFatal(t, err)
   627  		}
   628  		cli.Close()
   629  	}()
   630  
   631  	srv := Server(s, testConfigServer.Clone())
   632  	expectedErr := "tls: handshake buffer not empty before HelloRetryRequest"
   633  	if err := srv.Handshake(); err == nil {
   634  		t.Fatal("expected error from incomplete handshake, got nil")
   635  	} else if err.Error() != expectedErr {
   636  		t.Fatalf("expected error %q, got %q", expectedErr, err.Error())
   637  	}
   638  }
   639  
   640  // concatHandshakeMessages marshals and concatenates the given handshake
   641  // messages into a single record.
   642  func concatHandshakeMessages(msgs ...handshakeMessage) ([]byte, error) {
   643  	var marshalled []byte
   644  	for _, msg := range msgs {
   645  		data, err := msg.marshal()
   646  		if err != nil {
   647  			return nil, err
   648  		}
   649  		marshalled = append(marshalled, data...)
   650  	}
   651  	m := len(marshalled)
   652  	outBuf := make([]byte, recordHeaderLen)
   653  	outBuf[0] = byte(recordTypeHandshake)
   654  	vers := VersionTLS12
   655  	outBuf[1] = byte(vers >> 8)
   656  	outBuf[2] = byte(vers)
   657  	outBuf[3] = byte(m >> 8)
   658  	outBuf[4] = byte(m)
   659  	outBuf = append(outBuf, marshalled...)
   660  	return outBuf, nil
   661  }
   662  
   663  func TestMultipleKeyUpdate(t *testing.T) {
   664  	for _, requestUpdate := range []bool{true, false} {
   665  		t.Run(fmt.Sprintf("requestUpdate=%t", requestUpdate), func(t *testing.T) {
   666  
   667  			c, s := localPipe(t)
   668  			clientConfig := testConfigClient.Clone()
   669  			clientConfig.MinVersion = VersionTLS13
   670  			clientConfig.MaxVersion = VersionTLS13
   671  			serverConfig := testConfigServer.Clone()
   672  			serverConfig.MinVersion = VersionTLS13
   673  			serverConfig.MaxVersion = VersionTLS13
   674  			client := Client(c, clientConfig)
   675  			server := Server(s, serverConfig)
   676  
   677  			clientHandshakeDone := make(chan struct{})
   678  			go func() {
   679  				if err := client.Handshake(); err != nil {
   680  				}
   681  				close(clientHandshakeDone)
   682  				io.Copy(io.Discard, server)
   683  			}()
   684  
   685  			if err := server.Handshake(); err != nil {
   686  				t.Fatalf("server handshake failed: %v\n", err)
   687  			}
   688  			<-clientHandshakeDone
   689  
   690  			c.SetReadDeadline(time.Now().Add(1 * time.Second))
   691  			s.SetReadDeadline(time.Now().Add(1 * time.Second))
   692  
   693  			kuMsg, err := (&keyUpdateMsg{updateRequested: requestUpdate}).marshal()
   694  			if err != nil {
   695  				t.Fatalf("failed to marshal key update message: %v", err)
   696  			}
   697  
   698  			client.out.Lock()
   699  			if _, err := client.writeRecordLocked(recordTypeHandshake, append(kuMsg, kuMsg...)); err != nil {
   700  				t.Fatalf("failed to write key update messages: %v", err)
   701  			}
   702  			client.out.Unlock()
   703  
   704  			_, err = io.Copy(io.Discard, client)
   705  			if err == nil {
   706  				t.Fatal("expected multiple key update messages to cause an error, got nil")
   707  			} else if !strings.HasSuffix(err.Error(), "tls: unexpected message") {
   708  				t.Fatalf("unexpected error: %v", err)
   709  			}
   710  		})
   711  	}
   712  }
   713  

View as plain text