1
2
3
4
5
6
7 package tls
8
9 import (
10 "bytes"
11 "context"
12 "crypto/cipher"
13 "crypto/subtle"
14 "crypto/x509"
15 "errors"
16 "fmt"
17 "hash"
18 "io"
19 "net"
20 "sync"
21 "sync/atomic"
22 "time"
23 )
24
25
26
27 type Conn struct {
28
29 conn net.Conn
30 isClient bool
31 handshakeFn func(context.Context) error
32 quic *quicState
33
34
35
36
37 isHandshakeComplete atomic.Bool
38
39 handshakeMutex sync.Mutex
40 handshakeErr error
41 vers uint16
42 haveVers bool
43 config *Config
44
45
46
47 handshakes int
48 extMasterSecret bool
49 didResume bool
50 didHRR bool
51 cipherSuite uint16
52 curveID CurveID
53 peerSigAlg SignatureScheme
54 ocspResponse []byte
55 scts [][]byte
56 peerCertificates []*x509.Certificate
57 localCertificate [][]byte
58
59
60 verifiedChains [][]*x509.Certificate
61
62 serverName string
63
64
65
66 secureRenegotiation bool
67
68 ekm func(label string, context []byte, length int) ([]byte, error)
69
70
71 resumptionSecret []byte
72 echAccepted bool
73
74
75
76
77 ticketKeys []ticketKey
78
79
80
81
82
83 clientFinishedIsFirst bool
84
85
86 closeNotifyErr error
87
88
89 closeNotifySent bool
90
91
92
93
94
95 clientFinished [12]byte
96 serverFinished [12]byte
97
98
99 clientProtocol string
100
101
102 in, out halfConn
103 rawInput bytes.Buffer
104 input bytes.Reader
105 hand bytes.Buffer
106 buffering bool
107 sendBuf []byte
108
109
110
111 bytesSent int64
112 packetsSent int64
113
114
115
116
117 retryCount int
118
119
120
121 activeCall atomic.Int32
122
123 tmp [16]byte
124 }
125
126
127
128
129
130
131 func (c *Conn) LocalAddr() net.Addr {
132 return c.conn.LocalAddr()
133 }
134
135
136 func (c *Conn) RemoteAddr() net.Addr {
137 return c.conn.RemoteAddr()
138 }
139
140
141
142
143 func (c *Conn) SetDeadline(t time.Time) error {
144 return c.conn.SetDeadline(t)
145 }
146
147
148
149 func (c *Conn) SetReadDeadline(t time.Time) error {
150 return c.conn.SetReadDeadline(t)
151 }
152
153
154
155
156 func (c *Conn) SetWriteDeadline(t time.Time) error {
157 return c.conn.SetWriteDeadline(t)
158 }
159
160
161
162
163 func (c *Conn) NetConn() net.Conn {
164 return c.conn
165 }
166
167
168
169 type halfConn struct {
170 sync.Mutex
171
172 err error
173 version uint16
174 cipher any
175 mac hash.Hash
176 seq [8]byte
177
178 scratchBuf [13]byte
179
180 nextCipher any
181 nextMac hash.Hash
182
183 level QUICEncryptionLevel
184 trafficSecret []byte
185 }
186
187 type permanentError struct {
188 err net.Error
189 }
190
191 func (e *permanentError) Error() string { return e.err.Error() }
192 func (e *permanentError) Unwrap() error { return e.err }
193 func (e *permanentError) Timeout() bool { return e.err.Timeout() }
194 func (e *permanentError) Temporary() bool { return false }
195
196 func (hc *halfConn) setErrorLocked(err error) error {
197 if e, ok := err.(net.Error); ok {
198 hc.err = &permanentError{err: e}
199 } else {
200 hc.err = err
201 }
202 return hc.err
203 }
204
205
206
207 func (hc *halfConn) prepareCipherSpec(version uint16, cipher any, mac hash.Hash) {
208 hc.version = version
209 hc.nextCipher = cipher
210 hc.nextMac = mac
211 }
212
213
214
215 func (hc *halfConn) changeCipherSpec() error {
216 if hc.nextCipher == nil || hc.version == VersionTLS13 {
217 return alertInternalError
218 }
219 hc.cipher = hc.nextCipher
220 hc.mac = hc.nextMac
221 hc.nextCipher = nil
222 hc.nextMac = nil
223 clear(hc.seq[:])
224 return nil
225 }
226
227
228
229
230 func (hc *halfConn) setTrafficSecret(suite *cipherSuiteTLS13, level QUICEncryptionLevel, secret []byte) {
231 hc.trafficSecret = secret
232 hc.level = level
233 key, iv := suite.trafficKey(secret)
234 hc.cipher = suite.aead(key, iv)
235 clear(hc.seq[:])
236 }
237
238
239 func (hc *halfConn) incSeq() {
240 for i := 7; i >= 0; i-- {
241 hc.seq[i]++
242 if hc.seq[i] != 0 {
243 return
244 }
245 }
246
247
248
249
250 panic("TLS: sequence number wraparound")
251 }
252
253
254
255
256 func (hc *halfConn) explicitNonceLen() int {
257 if hc.cipher == nil {
258 return 0
259 }
260
261 switch c := hc.cipher.(type) {
262 case cipher.Stream:
263 return 0
264 case aead:
265 return c.explicitNonceLen()
266 case cbcMode:
267
268 if hc.version >= VersionTLS11 {
269 return c.BlockSize()
270 }
271 return 0
272 default:
273 panic("unknown cipher type")
274 }
275 }
276
277
278
279
280 func extractPadding(payload []byte) (toRemove int, good byte) {
281 if len(payload) < 1 {
282 return 0, 0
283 }
284
285 paddingLen := payload[len(payload)-1]
286 t := uint(len(payload)-1) - uint(paddingLen)
287
288 good = byte(int32(^t) >> 31)
289
290
291 toCheck := 256
292
293 if toCheck > len(payload) {
294 toCheck = len(payload)
295 }
296
297 for i := 0; i < toCheck; i++ {
298 t := uint(paddingLen) - uint(i)
299
300 mask := byte(int32(^t) >> 31)
301 b := payload[len(payload)-1-i]
302 good &^= mask&paddingLen ^ mask&b
303 }
304
305
306
307 good &= good << 4
308 good &= good << 2
309 good &= good << 1
310 good = uint8(int8(good) >> 7)
311
312
313
314
315
316
317
318
319
320
321 paddingLen &= good
322
323 toRemove = int(paddingLen) + 1
324 return
325 }
326
327 func roundUp(a, b int) int {
328 return a + (b-a%b)%b
329 }
330
331
332 type cbcMode interface {
333 cipher.BlockMode
334 SetIV([]byte)
335 }
336
337
338
339 func (hc *halfConn) decrypt(record []byte) ([]byte, recordType, error) {
340 var plaintext []byte
341 typ := recordType(record[0])
342 payload := record[recordHeaderLen:]
343
344
345
346 if hc.version == VersionTLS13 && typ == recordTypeChangeCipherSpec {
347 return payload, typ, nil
348 }
349
350 paddingGood := byte(255)
351 paddingLen := 0
352
353 explicitNonceLen := hc.explicitNonceLen()
354
355 if hc.cipher != nil {
356 switch c := hc.cipher.(type) {
357 case cipher.Stream:
358 c.XORKeyStream(payload, payload)
359 case aead:
360 if len(payload) < explicitNonceLen {
361 return nil, 0, alertBadRecordMAC
362 }
363 nonce := payload[:explicitNonceLen]
364 if len(nonce) == 0 {
365 nonce = hc.seq[:]
366 }
367 payload = payload[explicitNonceLen:]
368
369 var additionalData []byte
370 if hc.version == VersionTLS13 {
371 additionalData = record[:recordHeaderLen]
372 } else {
373 additionalData = append(hc.scratchBuf[:0], hc.seq[:]...)
374 additionalData = append(additionalData, record[:3]...)
375 n := len(payload) - c.Overhead()
376 additionalData = append(additionalData, byte(n>>8), byte(n))
377 }
378
379 var err error
380 plaintext, err = c.Open(payload[:0], nonce, payload, additionalData)
381 if err != nil {
382 return nil, 0, alertBadRecordMAC
383 }
384 case cbcMode:
385 blockSize := c.BlockSize()
386 minPayload := explicitNonceLen + roundUp(hc.mac.Size()+1, blockSize)
387 if len(payload)%blockSize != 0 || len(payload) < minPayload {
388 return nil, 0, alertBadRecordMAC
389 }
390
391 if explicitNonceLen > 0 {
392 c.SetIV(payload[:explicitNonceLen])
393 payload = payload[explicitNonceLen:]
394 }
395 c.CryptBlocks(payload, payload)
396
397
398
399
400
401
402
403 paddingLen, paddingGood = extractPadding(payload)
404 default:
405 panic("unknown cipher type")
406 }
407
408 if hc.version == VersionTLS13 {
409 if typ != recordTypeApplicationData {
410 return nil, 0, alertUnexpectedMessage
411 }
412 if len(plaintext) > maxPlaintext+1 {
413 return nil, 0, alertRecordOverflow
414 }
415
416 for i := len(plaintext) - 1; i >= 0; i-- {
417 if plaintext[i] != 0 {
418 typ = recordType(plaintext[i])
419 plaintext = plaintext[:i]
420 break
421 }
422 if i == 0 {
423 return nil, 0, alertUnexpectedMessage
424 }
425 }
426 }
427 } else {
428 plaintext = payload
429 }
430
431 if hc.mac != nil {
432 macSize := hc.mac.Size()
433 if len(payload) < macSize {
434 return nil, 0, alertBadRecordMAC
435 }
436
437 n := len(payload) - macSize - paddingLen
438 n = subtle.ConstantTimeSelect(int(uint32(n)>>31), 0, n)
439 record[3] = byte(n >> 8)
440 record[4] = byte(n)
441 remoteMAC := payload[n : n+macSize]
442 localMAC := tls10MAC(hc.mac, hc.scratchBuf[:0], hc.seq[:], record[:recordHeaderLen], payload[:n], payload[n+macSize:])
443
444
445
446
447
448
449
450
451 macAndPaddingGood := subtle.ConstantTimeCompare(localMAC, remoteMAC) & int(paddingGood)
452 if macAndPaddingGood != 1 {
453 return nil, 0, alertBadRecordMAC
454 }
455
456 plaintext = payload[:n]
457 }
458
459 hc.incSeq()
460 return plaintext, typ, nil
461 }
462
463
464
465
466 func sliceForAppend(in []byte, n int) (head, tail []byte) {
467 if total := len(in) + n; cap(in) >= total {
468 head = in[:total]
469 } else {
470 head = make([]byte, total)
471 copy(head, in)
472 }
473 tail = head[len(in):]
474 return
475 }
476
477
478
479 func (hc *halfConn) encrypt(record, payload []byte, rand io.Reader) ([]byte, error) {
480 if hc.cipher == nil {
481 return append(record, payload...), nil
482 }
483
484 var explicitNonce []byte
485 if explicitNonceLen := hc.explicitNonceLen(); explicitNonceLen > 0 {
486 record, explicitNonce = sliceForAppend(record, explicitNonceLen)
487 if _, isCBC := hc.cipher.(cbcMode); !isCBC && explicitNonceLen < 16 {
488
489
490
491
492
493
494
495
496
497 copy(explicitNonce, hc.seq[:])
498 } else {
499 if _, err := io.ReadFull(rand, explicitNonce); err != nil {
500 return nil, err
501 }
502 }
503 }
504
505 var dst []byte
506 switch c := hc.cipher.(type) {
507 case cipher.Stream:
508 mac := tls10MAC(hc.mac, hc.scratchBuf[:0], hc.seq[:], record[:recordHeaderLen], payload, nil)
509 record, dst = sliceForAppend(record, len(payload)+len(mac))
510 c.XORKeyStream(dst[:len(payload)], payload)
511 c.XORKeyStream(dst[len(payload):], mac)
512 case aead:
513 nonce := explicitNonce
514 if len(nonce) == 0 {
515 nonce = hc.seq[:]
516 }
517
518 if hc.version == VersionTLS13 {
519 record = append(record, payload...)
520
521
522 record = append(record, record[0])
523 record[0] = byte(recordTypeApplicationData)
524
525 n := len(payload) + 1 + c.Overhead()
526 record[3] = byte(n >> 8)
527 record[4] = byte(n)
528
529 record = c.Seal(record[:recordHeaderLen],
530 nonce, record[recordHeaderLen:], record[:recordHeaderLen])
531 } else {
532 additionalData := append(hc.scratchBuf[:0], hc.seq[:]...)
533 additionalData = append(additionalData, record[:recordHeaderLen]...)
534 record = c.Seal(record, nonce, payload, additionalData)
535 }
536 case cbcMode:
537 mac := tls10MAC(hc.mac, hc.scratchBuf[:0], hc.seq[:], record[:recordHeaderLen], payload, nil)
538 blockSize := c.BlockSize()
539 plaintextLen := len(payload) + len(mac)
540 paddingLen := blockSize - plaintextLen%blockSize
541 record, dst = sliceForAppend(record, plaintextLen+paddingLen)
542 copy(dst, payload)
543 copy(dst[len(payload):], mac)
544 for i := plaintextLen; i < len(dst); i++ {
545 dst[i] = byte(paddingLen - 1)
546 }
547 if len(explicitNonce) > 0 {
548 c.SetIV(explicitNonce)
549 }
550 c.CryptBlocks(dst, dst)
551 default:
552 panic("unknown cipher type")
553 }
554
555
556 n := len(record) - recordHeaderLen
557 record[3] = byte(n >> 8)
558 record[4] = byte(n)
559 hc.incSeq()
560
561 return record, nil
562 }
563
564
565 type RecordHeaderError struct {
566
567 Msg string
568
569
570 RecordHeader [5]byte
571
572
573
574
575 Conn net.Conn
576 }
577
578 func (e RecordHeaderError) Error() string { return "tls: " + e.Msg }
579
580 func (c *Conn) newRecordHeaderError(conn net.Conn, msg string) (err RecordHeaderError) {
581 err.Msg = msg
582 err.Conn = conn
583 copy(err.RecordHeader[:], c.rawInput.Bytes())
584 return err
585 }
586
587 func (c *Conn) readRecord() error {
588 return c.readRecordOrCCS(false)
589 }
590
591 func (c *Conn) readChangeCipherSpec() error {
592 return c.readRecordOrCCS(true)
593 }
594
595
596
597
598
599
600
601
602
603
604
605
606
607
608
609 func (c *Conn) readRecordOrCCS(expectChangeCipherSpec bool) error {
610 if c.in.err != nil {
611 return c.in.err
612 }
613 handshakeComplete := c.isHandshakeComplete.Load()
614
615
616 if c.input.Len() != 0 {
617 return c.in.setErrorLocked(errors.New("tls: internal error: attempted to read record with pending application data"))
618 }
619 c.input.Reset(nil)
620
621 if c.quic != nil {
622 return c.in.setErrorLocked(errors.New("tls: internal error: attempted to read record with QUIC transport"))
623 }
624
625
626 if err := c.readFromUntil(c.conn, recordHeaderLen); err != nil {
627
628
629
630 if err == io.ErrUnexpectedEOF && c.rawInput.Len() == 0 {
631 err = io.EOF
632 }
633 if e, ok := err.(net.Error); !ok || !e.Temporary() {
634 c.in.setErrorLocked(err)
635 }
636 return err
637 }
638 hdr := c.rawInput.Bytes()[:recordHeaderLen]
639 typ := recordType(hdr[0])
640
641
642
643
644
645 if !handshakeComplete && typ == 0x80 {
646 c.sendAlert(alertProtocolVersion)
647 return c.in.setErrorLocked(c.newRecordHeaderError(nil, "unsupported SSLv2 handshake received"))
648 }
649
650 vers := uint16(hdr[1])<<8 | uint16(hdr[2])
651 expectedVers := c.vers
652 if expectedVers == VersionTLS13 {
653
654
655 expectedVers = VersionTLS12
656 }
657 n := int(hdr[3])<<8 | int(hdr[4])
658 if c.haveVers && vers != expectedVers {
659 c.sendAlert(alertProtocolVersion)
660 msg := fmt.Sprintf("received record with version %x when expecting version %x", vers, expectedVers)
661 return c.in.setErrorLocked(c.newRecordHeaderError(nil, msg))
662 }
663 if !c.haveVers {
664
665
666
667
668 if (typ != recordTypeAlert && typ != recordTypeHandshake) || vers >= 0x1000 {
669 return c.in.setErrorLocked(c.newRecordHeaderError(c.conn, "first record does not look like a TLS handshake"))
670 }
671 }
672 if c.vers == VersionTLS13 && n > maxCiphertextTLS13 || n > maxCiphertext {
673 c.sendAlert(alertRecordOverflow)
674 msg := fmt.Sprintf("oversized record received with length %d", n)
675 return c.in.setErrorLocked(c.newRecordHeaderError(nil, msg))
676 }
677 if err := c.readFromUntil(c.conn, recordHeaderLen+n); err != nil {
678 if e, ok := err.(net.Error); !ok || !e.Temporary() {
679 c.in.setErrorLocked(err)
680 }
681 return err
682 }
683
684
685 record := c.rawInput.Next(recordHeaderLen + n)
686 data, typ, err := c.in.decrypt(record)
687 if err != nil {
688 return c.in.setErrorLocked(c.sendAlert(err.(alert)))
689 }
690 if len(data) > maxPlaintext {
691 return c.in.setErrorLocked(c.sendAlert(alertRecordOverflow))
692 }
693
694
695 if c.in.cipher == nil && typ == recordTypeApplicationData {
696 return c.in.setErrorLocked(c.sendAlert(alertUnexpectedMessage))
697 }
698
699 if (typ == recordTypeApplicationData || (typ == recordTypeHandshake && !handshakeComplete)) && len(data) > 0 {
700
701 c.retryCount = 0
702 }
703
704
705 if c.vers == VersionTLS13 && typ != recordTypeHandshake && c.hand.Len() > 0 {
706 return c.in.setErrorLocked(c.sendAlert(alertUnexpectedMessage))
707 }
708
709 switch typ {
710 default:
711 return c.in.setErrorLocked(c.sendAlert(alertUnexpectedMessage))
712
713 case recordTypeAlert:
714 if c.quic != nil {
715 return c.in.setErrorLocked(c.sendAlert(alertUnexpectedMessage))
716 }
717 if len(data) != 2 {
718 return c.in.setErrorLocked(c.sendAlert(alertUnexpectedMessage))
719 }
720 if alert(data[1]) == alertCloseNotify {
721 return c.in.setErrorLocked(io.EOF)
722 }
723 if c.vers == VersionTLS13 {
724
725
726
727
728
729 if alert(data[1]) == alertUserCanceled {
730
731 return c.retryReadRecord(expectChangeCipherSpec)
732 }
733 return c.in.setErrorLocked(&net.OpError{Op: "remote error", Err: alert(data[1])})
734 }
735 switch data[0] {
736 case alertLevelWarning:
737
738 return c.retryReadRecord(expectChangeCipherSpec)
739 case alertLevelError:
740 return c.in.setErrorLocked(&net.OpError{Op: "remote error", Err: alert(data[1])})
741 default:
742 return c.in.setErrorLocked(c.sendAlert(alertUnexpectedMessage))
743 }
744
745 case recordTypeChangeCipherSpec:
746 if len(data) != 1 || data[0] != 1 {
747 return c.in.setErrorLocked(c.sendAlert(alertDecodeError))
748 }
749
750 if c.hand.Len() > 0 {
751 return c.in.setErrorLocked(c.sendAlert(alertUnexpectedMessage))
752 }
753
754
755
756
757
758 if c.vers == VersionTLS13 {
759 return c.retryReadRecord(expectChangeCipherSpec)
760 }
761 if !expectChangeCipherSpec {
762 return c.in.setErrorLocked(c.sendAlert(alertUnexpectedMessage))
763 }
764 if err := c.in.changeCipherSpec(); err != nil {
765 return c.in.setErrorLocked(c.sendAlert(err.(alert)))
766 }
767
768 case recordTypeApplicationData:
769 if !handshakeComplete || expectChangeCipherSpec {
770 return c.in.setErrorLocked(c.sendAlert(alertUnexpectedMessage))
771 }
772
773
774 if len(data) == 0 {
775 return c.retryReadRecord(expectChangeCipherSpec)
776 }
777
778
779
780 c.input.Reset(data)
781
782 case recordTypeHandshake:
783 if len(data) == 0 || expectChangeCipherSpec {
784 return c.in.setErrorLocked(c.sendAlert(alertUnexpectedMessage))
785 }
786 c.hand.Write(data)
787 }
788
789 return nil
790 }
791
792
793
794 func (c *Conn) retryReadRecord(expectChangeCipherSpec bool) error {
795 c.retryCount++
796 if c.retryCount > maxUselessRecords {
797 c.sendAlert(alertUnexpectedMessage)
798 return c.in.setErrorLocked(errors.New("tls: too many ignored records"))
799 }
800 return c.readRecordOrCCS(expectChangeCipherSpec)
801 }
802
803
804
805 func (c *Conn) readFromUntil(r io.Reader, n int) error {
806 if c.rawInput.Len() >= n {
807 return nil
808 }
809 needs := n - c.rawInput.Len()
810
811
812
813
814
815
816
817 c.rawInput.Grow(needs + bytes.MinRead)
818 for {
819 buf := c.rawInput.AvailableBuffer()[:c.rawInput.Available()]
820 n, err := r.Read(buf)
821
822
823 c.rawInput.Write(buf[:n])
824 needs -= n
825 if needs <= 0 {
826 if err == io.EOF {
827 err = nil
828 }
829 return err
830 }
831 if err == io.EOF {
832 return io.ErrUnexpectedEOF
833 }
834 if err != nil {
835 return err
836 }
837 }
838 }
839
840
841 func (c *Conn) sendAlertLocked(err alert) error {
842 if c.quic != nil {
843 return c.out.setErrorLocked(&net.OpError{Op: "local error", Err: err})
844 }
845
846 switch err {
847 case alertNoRenegotiation, alertCloseNotify:
848 c.tmp[0] = alertLevelWarning
849 default:
850 c.tmp[0] = alertLevelError
851 }
852 c.tmp[1] = byte(err)
853
854 _, writeErr := c.writeRecordLocked(recordTypeAlert, c.tmp[0:2])
855 if err == alertCloseNotify {
856
857 return writeErr
858 }
859
860 return c.out.setErrorLocked(&net.OpError{Op: "local error", Err: err})
861 }
862
863
864 func (c *Conn) sendAlert(err alert) error {
865 c.out.Lock()
866 defer c.out.Unlock()
867 return c.sendAlertLocked(err)
868 }
869
870 const (
871
872
873
874
875
876 tcpMSSEstimate = 1208
877
878
879
880
881 recordSizeBoostThreshold = 128 * 1024
882 )
883
884
885
886
887
888
889
890
891
892
893
894
895
896
897
898
899
900 func (c *Conn) maxPayloadSizeForWrite(typ recordType) int {
901 if c.config.DynamicRecordSizingDisabled || typ != recordTypeApplicationData {
902 return maxPlaintext
903 }
904
905 if c.bytesSent >= recordSizeBoostThreshold {
906 return maxPlaintext
907 }
908
909
910 payloadBytes := tcpMSSEstimate - recordHeaderLen - c.out.explicitNonceLen()
911 if c.out.cipher != nil {
912 switch ciph := c.out.cipher.(type) {
913 case cipher.Stream:
914 payloadBytes -= c.out.mac.Size()
915 case cipher.AEAD:
916 payloadBytes -= ciph.Overhead()
917 case cbcMode:
918 blockSize := ciph.BlockSize()
919
920
921 payloadBytes = (payloadBytes & ^(blockSize - 1)) - 1
922
923
924 payloadBytes -= c.out.mac.Size()
925 default:
926 panic("unknown cipher type")
927 }
928 }
929 if c.vers == VersionTLS13 {
930 payloadBytes--
931 }
932
933
934 pkt := c.packetsSent
935 c.packetsSent++
936 if pkt > 1000 {
937 return maxPlaintext
938 }
939
940 n := payloadBytes * int(pkt+1)
941 if n > maxPlaintext {
942 n = maxPlaintext
943 }
944 return n
945 }
946
947 func (c *Conn) write(data []byte) (int, error) {
948 if c.buffering {
949 c.sendBuf = append(c.sendBuf, data...)
950 return len(data), nil
951 }
952
953 n, err := c.conn.Write(data)
954 c.bytesSent += int64(n)
955 return n, err
956 }
957
958 func (c *Conn) flush() (int, error) {
959 if len(c.sendBuf) == 0 {
960 return 0, nil
961 }
962
963 n, err := c.conn.Write(c.sendBuf)
964 c.bytesSent += int64(n)
965 c.sendBuf = nil
966 c.buffering = false
967 return n, err
968 }
969
970
971 var outBufPool = sync.Pool{
972 New: func() any {
973 return new([]byte)
974 },
975 }
976
977
978
979 func (c *Conn) writeRecordLocked(typ recordType, data []byte) (int, error) {
980 if c.quic != nil {
981 if typ != recordTypeHandshake {
982 return 0, errors.New("tls: internal error: sending non-handshake message to QUIC transport")
983 }
984 c.quicWriteCryptoData(c.out.level, data)
985 if !c.buffering {
986 if _, err := c.flush(); err != nil {
987 return 0, err
988 }
989 }
990 return len(data), nil
991 }
992
993 outBufPtr := outBufPool.Get().(*[]byte)
994 outBuf := *outBufPtr
995 defer func() {
996
997
998
999
1000
1001 *outBufPtr = outBuf
1002 outBufPool.Put(outBufPtr)
1003 }()
1004
1005 var n int
1006 for len(data) > 0 {
1007 m := len(data)
1008 if maxPayload := c.maxPayloadSizeForWrite(typ); m > maxPayload {
1009 m = maxPayload
1010 }
1011
1012 _, outBuf = sliceForAppend(outBuf[:0], recordHeaderLen)
1013 outBuf[0] = byte(typ)
1014 vers := c.vers
1015 if vers == 0 {
1016
1017
1018 vers = VersionTLS10
1019 } else if vers == VersionTLS13 {
1020
1021
1022 vers = VersionTLS12
1023 }
1024 outBuf[1] = byte(vers >> 8)
1025 outBuf[2] = byte(vers)
1026 outBuf[3] = byte(m >> 8)
1027 outBuf[4] = byte(m)
1028
1029 var err error
1030 outBuf, err = c.out.encrypt(outBuf, data[:m], c.config.rand())
1031 if err != nil {
1032 return n, err
1033 }
1034 if _, err := c.write(outBuf); err != nil {
1035 return n, err
1036 }
1037 n += m
1038 data = data[m:]
1039 }
1040
1041 if typ == recordTypeChangeCipherSpec && c.vers != VersionTLS13 {
1042 if err := c.out.changeCipherSpec(); err != nil {
1043 return n, c.sendAlertLocked(err.(alert))
1044 }
1045 }
1046
1047 return n, nil
1048 }
1049
1050
1051
1052
1053 func (c *Conn) writeHandshakeRecord(msg handshakeMessage, transcript transcriptHash) (int, error) {
1054 c.out.Lock()
1055 defer c.out.Unlock()
1056
1057 data, err := msg.marshal()
1058 if err != nil {
1059 return 0, err
1060 }
1061 if transcript != nil {
1062 transcript.Write(data)
1063 }
1064
1065 return c.writeRecordLocked(recordTypeHandshake, data)
1066 }
1067
1068
1069
1070 func (c *Conn) writeChangeCipherRecord() error {
1071 c.out.Lock()
1072 defer c.out.Unlock()
1073 _, err := c.writeRecordLocked(recordTypeChangeCipherSpec, []byte{1})
1074 return err
1075 }
1076
1077
1078 func (c *Conn) readHandshakeBytes(n int) error {
1079 if c.quic != nil {
1080 return c.quicReadHandshakeBytes(n)
1081 }
1082 for c.hand.Len() < n {
1083 if err := c.readRecord(); err != nil {
1084 return err
1085 }
1086 }
1087 return nil
1088 }
1089
1090
1091
1092
1093 func (c *Conn) readHandshake(transcript transcriptHash) (any, error) {
1094 if err := c.readHandshakeBytes(4); err != nil {
1095 return nil, err
1096 }
1097 data := c.hand.Bytes()
1098
1099 maxHandshakeSize := maxHandshake
1100
1101
1102
1103 if c.haveVers && data[0] == typeCertificate {
1104
1105
1106
1107 maxHandshakeSize = maxHandshakeCertificateMsg
1108 }
1109
1110 n := int(data[1])<<16 | int(data[2])<<8 | int(data[3])
1111 if n > maxHandshakeSize {
1112 c.sendAlertLocked(alertInternalError)
1113 return nil, c.in.setErrorLocked(fmt.Errorf("tls: handshake message of length %d bytes exceeds maximum of %d bytes", n, maxHandshakeSize))
1114 }
1115 if err := c.readHandshakeBytes(4 + n); err != nil {
1116 return nil, err
1117 }
1118 data = c.hand.Next(4 + n)
1119 return c.unmarshalHandshakeMessage(data, transcript)
1120 }
1121
1122 func (c *Conn) unmarshalHandshakeMessage(data []byte, transcript transcriptHash) (handshakeMessage, error) {
1123 var m handshakeMessage
1124 switch data[0] {
1125 case typeHelloRequest:
1126 m = new(helloRequestMsg)
1127 case typeClientHello:
1128 m = new(clientHelloMsg)
1129 case typeServerHello:
1130 m = new(serverHelloMsg)
1131 case typeNewSessionTicket:
1132 if c.vers == VersionTLS13 {
1133 m = new(newSessionTicketMsgTLS13)
1134 } else {
1135 m = new(newSessionTicketMsg)
1136 }
1137 case typeCertificate:
1138 if c.vers == VersionTLS13 {
1139 m = new(certificateMsgTLS13)
1140 } else {
1141 m = new(certificateMsg)
1142 }
1143 case typeCertificateRequest:
1144 if c.vers == VersionTLS13 {
1145 m = new(certificateRequestMsgTLS13)
1146 } else {
1147 m = &certificateRequestMsg{
1148 hasSignatureAlgorithm: c.vers >= VersionTLS12,
1149 }
1150 }
1151 case typeCertificateStatus:
1152 m = new(certificateStatusMsg)
1153 case typeServerKeyExchange:
1154 m = new(serverKeyExchangeMsg)
1155 case typeServerHelloDone:
1156 m = new(serverHelloDoneMsg)
1157 case typeClientKeyExchange:
1158 m = new(clientKeyExchangeMsg)
1159 case typeCertificateVerify:
1160 m = &certificateVerifyMsg{
1161 hasSignatureAlgorithm: c.vers >= VersionTLS12,
1162 }
1163 case typeFinished:
1164 m = new(finishedMsg)
1165 case typeEncryptedExtensions:
1166 m = new(encryptedExtensionsMsg)
1167 case typeEndOfEarlyData:
1168 m = new(endOfEarlyDataMsg)
1169 case typeKeyUpdate:
1170 m = new(keyUpdateMsg)
1171 default:
1172 return nil, c.in.setErrorLocked(c.sendAlert(alertUnexpectedMessage))
1173 }
1174
1175
1176
1177
1178 data = append([]byte(nil), data...)
1179
1180 if !m.unmarshal(data) {
1181 return nil, c.in.setErrorLocked(c.sendAlert(alertDecodeError))
1182 }
1183
1184 if transcript != nil {
1185 transcript.Write(data)
1186 }
1187
1188 return m, nil
1189 }
1190
1191 var (
1192 errShutdown = errors.New("tls: protocol is shutdown")
1193 )
1194
1195
1196
1197
1198
1199
1200
1201 func (c *Conn) Write(b []byte) (int, error) {
1202
1203 for {
1204 x := c.activeCall.Load()
1205 if x&1 != 0 {
1206 return 0, net.ErrClosed
1207 }
1208 if c.activeCall.CompareAndSwap(x, x+2) {
1209 break
1210 }
1211 }
1212 defer c.activeCall.Add(-2)
1213
1214 if err := c.Handshake(); err != nil {
1215 return 0, err
1216 }
1217
1218 c.out.Lock()
1219 defer c.out.Unlock()
1220
1221 if err := c.out.err; err != nil {
1222 return 0, err
1223 }
1224
1225 if !c.isHandshakeComplete.Load() {
1226 return 0, alertInternalError
1227 }
1228
1229 if c.closeNotifySent {
1230 return 0, errShutdown
1231 }
1232
1233
1234
1235
1236
1237
1238
1239
1240
1241
1242 var m int
1243 if len(b) > 1 && c.vers == VersionTLS10 {
1244 if _, ok := c.out.cipher.(cipher.BlockMode); ok {
1245 n, err := c.writeRecordLocked(recordTypeApplicationData, b[:1])
1246 if err != nil {
1247 return n, c.out.setErrorLocked(err)
1248 }
1249 m, b = 1, b[1:]
1250 }
1251 }
1252
1253 n, err := c.writeRecordLocked(recordTypeApplicationData, b)
1254 return n + m, c.out.setErrorLocked(err)
1255 }
1256
1257
1258 func (c *Conn) handleRenegotiation() error {
1259 if c.vers == VersionTLS13 {
1260 return errors.New("tls: internal error: unexpected renegotiation")
1261 }
1262
1263 msg, err := c.readHandshake(nil)
1264 if err != nil {
1265 return err
1266 }
1267
1268 helloReq, ok := msg.(*helloRequestMsg)
1269 if !ok {
1270 c.sendAlert(alertUnexpectedMessage)
1271 return unexpectedMessageError(helloReq, msg)
1272 }
1273
1274 if !c.isClient {
1275 return c.sendAlert(alertNoRenegotiation)
1276 }
1277
1278 switch c.config.Renegotiation {
1279 case RenegotiateNever:
1280 return c.sendAlert(alertNoRenegotiation)
1281 case RenegotiateOnceAsClient:
1282 if c.handshakes > 1 {
1283 return c.sendAlert(alertNoRenegotiation)
1284 }
1285 case RenegotiateFreelyAsClient:
1286
1287 default:
1288 c.sendAlert(alertInternalError)
1289 return errors.New("tls: unknown Renegotiation value")
1290 }
1291
1292 c.handshakeMutex.Lock()
1293 defer c.handshakeMutex.Unlock()
1294
1295 c.isHandshakeComplete.Store(false)
1296 if c.handshakeErr = c.clientHandshake(context.Background()); c.handshakeErr == nil {
1297 c.handshakes++
1298 }
1299 return c.handshakeErr
1300 }
1301
1302
1303
1304 func (c *Conn) handlePostHandshakeMessage() error {
1305 if c.vers != VersionTLS13 {
1306 return c.handleRenegotiation()
1307 }
1308
1309 msg, err := c.readHandshake(nil)
1310 if err != nil {
1311 return err
1312 }
1313 c.retryCount++
1314 if c.retryCount > maxUselessRecords {
1315 c.sendAlert(alertUnexpectedMessage)
1316 return c.in.setErrorLocked(errors.New("tls: too many non-advancing records"))
1317 }
1318
1319 switch msg := msg.(type) {
1320 case *newSessionTicketMsgTLS13:
1321 return c.handleNewSessionTicket(msg)
1322 case *keyUpdateMsg:
1323 return c.handleKeyUpdate(msg)
1324 }
1325
1326
1327
1328
1329 c.sendAlert(alertUnexpectedMessage)
1330 return fmt.Errorf("tls: received unexpected handshake message of type %T", msg)
1331 }
1332
1333 func (c *Conn) handleKeyUpdate(keyUpdate *keyUpdateMsg) error {
1334 if c.quic != nil {
1335 c.sendAlert(alertUnexpectedMessage)
1336 return c.in.setErrorLocked(errors.New("tls: received unexpected key update message"))
1337 }
1338
1339 cipherSuite := cipherSuiteTLS13ByID(c.cipherSuite)
1340 if cipherSuite == nil {
1341 return c.in.setErrorLocked(c.sendAlert(alertInternalError))
1342 }
1343
1344 if keyUpdate.updateRequested {
1345 c.out.Lock()
1346 defer c.out.Unlock()
1347
1348 msg := &keyUpdateMsg{}
1349 msgBytes, err := msg.marshal()
1350 if err != nil {
1351 return err
1352 }
1353 _, err = c.writeRecordLocked(recordTypeHandshake, msgBytes)
1354 if err != nil {
1355
1356 c.out.setErrorLocked(err)
1357 return nil
1358 }
1359
1360 newSecret := cipherSuite.nextTrafficSecret(c.out.trafficSecret)
1361 c.setWriteTrafficSecret(cipherSuite, QUICEncryptionLevelInitial, newSecret)
1362 }
1363
1364 newSecret := cipherSuite.nextTrafficSecret(c.in.trafficSecret)
1365 if err := c.setReadTrafficSecret(cipherSuite, QUICEncryptionLevelInitial, newSecret, keyUpdate.updateRequested); err != nil {
1366 return err
1367 }
1368
1369 return nil
1370 }
1371
1372
1373
1374
1375
1376
1377
1378 func (c *Conn) Read(b []byte) (int, error) {
1379 if err := c.Handshake(); err != nil {
1380 return 0, err
1381 }
1382 if len(b) == 0 {
1383
1384
1385 return 0, nil
1386 }
1387
1388 c.in.Lock()
1389 defer c.in.Unlock()
1390
1391 for c.input.Len() == 0 {
1392 if err := c.readRecord(); err != nil {
1393 return 0, err
1394 }
1395 for c.hand.Len() > 0 {
1396 if err := c.handlePostHandshakeMessage(); err != nil {
1397 return 0, err
1398 }
1399 }
1400 }
1401
1402 n, _ := c.input.Read(b)
1403
1404
1405
1406
1407
1408
1409
1410
1411 if n != 0 && c.input.Len() == 0 && c.rawInput.Len() > 0 &&
1412 recordType(c.rawInput.Bytes()[0]) == recordTypeAlert {
1413 if err := c.readRecord(); err != nil {
1414 return n, err
1415 }
1416 }
1417
1418 return n, nil
1419 }
1420
1421
1422 func (c *Conn) Close() error {
1423
1424 var x int32
1425 for {
1426 x = c.activeCall.Load()
1427 if x&1 != 0 {
1428 return net.ErrClosed
1429 }
1430 if c.activeCall.CompareAndSwap(x, x|1) {
1431 break
1432 }
1433 }
1434 if x != 0 {
1435
1436
1437
1438
1439
1440
1441 return c.conn.Close()
1442 }
1443
1444 var alertErr error
1445 if c.isHandshakeComplete.Load() {
1446 if err := c.closeNotify(); err != nil {
1447 alertErr = fmt.Errorf("tls: failed to send closeNotify alert (but connection was closed anyway): %w", err)
1448 }
1449 }
1450
1451 if err := c.conn.Close(); err != nil {
1452 return err
1453 }
1454 return alertErr
1455 }
1456
1457 var errEarlyCloseWrite = errors.New("tls: CloseWrite called before handshake complete")
1458
1459
1460
1461
1462 func (c *Conn) CloseWrite() error {
1463 if !c.isHandshakeComplete.Load() {
1464 return errEarlyCloseWrite
1465 }
1466
1467 return c.closeNotify()
1468 }
1469
1470 func (c *Conn) closeNotify() error {
1471 c.out.Lock()
1472 defer c.out.Unlock()
1473
1474 if !c.closeNotifySent {
1475
1476 c.SetWriteDeadline(time.Now().Add(time.Second * 5))
1477 c.closeNotifyErr = c.sendAlertLocked(alertCloseNotify)
1478 c.closeNotifySent = true
1479
1480 c.SetWriteDeadline(time.Now())
1481 }
1482 return c.closeNotifyErr
1483 }
1484
1485
1486
1487
1488
1489
1490
1491
1492
1493
1494
1495
1496
1497
1498 func (c *Conn) Handshake() error {
1499 return c.HandshakeContext(context.Background())
1500 }
1501
1502
1503
1504
1505
1506
1507
1508
1509
1510
1511
1512 func (c *Conn) HandshakeContext(ctx context.Context) error {
1513
1514
1515 return c.handshakeContext(ctx)
1516 }
1517
1518 func (c *Conn) handshakeContext(ctx context.Context) (ret error) {
1519
1520
1521
1522 if c.isHandshakeComplete.Load() {
1523 return nil
1524 }
1525
1526 handshakeCtx, cancel := context.WithCancel(ctx)
1527
1528
1529
1530 defer cancel()
1531
1532 if c.quic != nil {
1533 c.quic.ctx = handshakeCtx
1534 c.quic.cancel = cancel
1535 } else if ctx.Done() != nil {
1536
1537 stop := context.AfterFunc(ctx, func() {
1538 _ = c.conn.Close()
1539 })
1540 defer func() {
1541 if !stop() {
1542
1543 ret = ctx.Err()
1544 }
1545 }()
1546 }
1547
1548 c.handshakeMutex.Lock()
1549 defer c.handshakeMutex.Unlock()
1550
1551 if err := c.handshakeErr; err != nil {
1552 return err
1553 }
1554 if c.isHandshakeComplete.Load() {
1555 return nil
1556 }
1557
1558 c.in.Lock()
1559 defer c.in.Unlock()
1560
1561 c.handshakeErr = c.handshakeFn(handshakeCtx)
1562 if c.handshakeErr == nil {
1563 c.handshakes++
1564 } else {
1565
1566
1567 c.flush()
1568 }
1569
1570 if c.handshakeErr == nil && !c.isHandshakeComplete.Load() {
1571 c.handshakeErr = errors.New("tls: internal error: handshake should have had a result")
1572 }
1573 if c.handshakeErr != nil && c.isHandshakeComplete.Load() {
1574 panic("tls: internal error: handshake returned an error but is marked successful")
1575 }
1576
1577 if c.quic != nil {
1578 if c.handshakeErr == nil {
1579 c.quicHandshakeComplete()
1580
1581
1582
1583 if err := c.quicSetReadSecret(QUICEncryptionLevelApplication, c.cipherSuite, c.in.trafficSecret); err != nil {
1584 return err
1585 }
1586 } else {
1587 c.out.Lock()
1588 a, ok := errors.AsType[alert](c.out.err)
1589 if !ok {
1590 a = alertInternalError
1591 }
1592 c.out.Unlock()
1593
1594
1595
1596
1597 c.handshakeErr = fmt.Errorf("%w%.0w", c.handshakeErr, AlertError(a))
1598 }
1599 close(c.quic.blockedc)
1600 close(c.quic.signalc)
1601 }
1602
1603 return c.handshakeErr
1604 }
1605
1606
1607
1608
1609
1610
1611
1612
1613 func (c *Conn) ConnectionState() ConnectionState {
1614 c.handshakeMutex.Lock()
1615 defer c.handshakeMutex.Unlock()
1616 return c.connectionStateLocked()
1617 }
1618
1619 func (c *Conn) connectionStateLocked() ConnectionState {
1620 var state ConnectionState
1621 state.HandshakeComplete = c.isHandshakeComplete.Load()
1622 state.Version = c.vers
1623 state.NegotiatedProtocol = c.clientProtocol
1624 state.DidResume = c.didResume
1625 state.HelloRetryRequest = c.didHRR
1626 state.testingOnlyPeerSignatureAlgorithm = c.peerSigAlg
1627 state.CurveID = c.curveID
1628 state.NegotiatedProtocolIsMutual = true
1629 state.ServerName = c.serverName
1630 state.CipherSuite = c.cipherSuite
1631 state.PeerCertificates = c.peerCertificates
1632 state.LocalCertificate = c.localCertificate
1633 state.VerifiedChains = c.verifiedChains
1634 state.SignedCertificateTimestamps = c.scts
1635 state.OCSPResponse = c.ocspResponse
1636 if (!c.didResume || c.extMasterSecret) && c.vers != VersionTLS13 {
1637 if c.clientFinishedIsFirst {
1638 state.TLSUnique = c.clientFinished[:]
1639 } else {
1640 state.TLSUnique = c.serverFinished[:]
1641 }
1642 }
1643 if c.config.Renegotiation != RenegotiateNever {
1644 state.ekm = noEKMBecauseRenegotiation
1645 } else if c.vers != VersionTLS13 && !c.extMasterSecret {
1646 state.ekm = noEKMBecauseNoEMS
1647 } else {
1648 state.ekm = c.ekm
1649 }
1650 state.ECHAccepted = c.echAccepted
1651 return state
1652 }
1653
1654
1655
1656 func (c *Conn) OCSPResponse() []byte {
1657 c.handshakeMutex.Lock()
1658 defer c.handshakeMutex.Unlock()
1659
1660 return c.ocspResponse
1661 }
1662
1663
1664
1665
1666 func (c *Conn) VerifyHostname(host string) error {
1667 c.handshakeMutex.Lock()
1668 defer c.handshakeMutex.Unlock()
1669 if !c.isClient {
1670 return errors.New("tls: VerifyHostname called on TLS server connection")
1671 }
1672 if !c.isHandshakeComplete.Load() {
1673 return errors.New("tls: handshake has not yet been performed")
1674 }
1675 if len(c.verifiedChains) == 0 {
1676 return errors.New("tls: handshake did not verify certificate chain")
1677 }
1678 return c.peerCertificates[0].VerifyHostname(host)
1679 }
1680
1681
1682
1683
1684 func (c *Conn) setReadTrafficSecret(suite *cipherSuiteTLS13, level QUICEncryptionLevel, secret []byte, locked bool) error {
1685
1686
1687
1688 if c.hand.Len() != 0 {
1689 if locked {
1690 c.sendAlertLocked(alertUnexpectedMessage)
1691 } else {
1692 c.sendAlert(alertUnexpectedMessage)
1693 }
1694 return errors.New("tls: handshake buffer not empty before setting read traffic secret")
1695 }
1696 c.in.setTrafficSecret(suite, level, secret)
1697 return nil
1698 }
1699
1700
1701
1702
1703 func (c *Conn) setWriteTrafficSecret(suite *cipherSuiteTLS13, level QUICEncryptionLevel, secret []byte) {
1704 c.out.setTrafficSecret(suite, level, secret)
1705 }
1706
View as plain text