Source file
src/crypto/tls/handshake_test.go
1
2
3
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
31
32
33
34
35
36
37
38
39
40
41
42
43
44
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)
58
59
60
61
62
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
81
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
112
113
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
162 func (r *recordingConn) WriteTo(w io.Writer) (int64, error) {
163
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
198
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
208
209
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
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
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
322
323 var localListener struct {
324 mu sync.Mutex
325 addr net.Addr
326 ch chan net.Conn
327 }
328
329 const localFlakes = 0
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
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
364
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
418
419 hasAESGCMHardwareSupport = false
420
421
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
483
484
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
497
498
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
526
527
528
529
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
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
606
607
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
641
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