diff --git a/transport/shadowtls/server_test.go b/transport/shadowtls/server_test.go index b6581da2..6e031b2c 100644 --- a/transport/shadowtls/server_test.go +++ b/transport/shadowtls/server_test.go @@ -136,6 +136,58 @@ func TestV3UnauthenticatedConnectionFallsBack(t *testing.T) { } } +func TestV2InterruptedHeaderRead(t *testing.T) { + clientSide, serverSide := net.Pipe() + defer clientSide.Close() + defer serverSide.Close() + + serverRaw := &readStartedConn{Conn: serverSide, started: make(chan struct{}, 1)} + server := newConn(serverRaw) + payload := []byte("payload after interrupted v2 header") + frame := make([]byte, tlsHeaderSize+len(payload)) + frame[0] = applicationData + frame[1] = 3 + frame[2] = 3 + binary.BigEndian.PutUint16(frame[3:tlsHeaderSize], uint16(len(payload))) + copy(frame[tlsHeaderSize:], payload) + + buffer := make([]byte, len(payload)) + serverRead := make(chan error, 1) + go func() { + _, err := server.Read(buffer) + serverRead <- err + }() + <-serverRaw.started + if _, err := clientSide.Write(frame[:2]); err != nil { + t.Fatalf("write partial header: %v", err) + } + <-serverRaw.started + if err := server.SetReadDeadline(time.Now()); err != nil { + t.Fatalf("interrupt server read: %v", err) + } + if err := <-serverRead; !errors.Is(err, os.ErrDeadlineExceeded) { + t.Fatalf("server read error = %v, want deadline exceeded", err) + } + if err := server.SetReadDeadline(time.Time{}); err != nil { + t.Fatalf("clear server read deadline: %v", err) + } + + clientWrite := make(chan error, 1) + go func() { + _, err := clientSide.Write(frame[2:]) + clientWrite <- err + }() + if _, err := io.ReadFull(server, buffer); err != nil { + t.Fatalf("server read completed frame: %v", err) + } + if !bytes.Equal(buffer, payload) { + t.Fatalf("server payload = %q, want %q", buffer, payload) + } + if err := <-clientWrite; err != nil { + t.Fatalf("client write remaining frame: %v", err) + } +} + func TestV3InterruptedReadDoesNotSendAlert(t *testing.T) { clientSide, serverSide := net.Pipe() defer clientSide.Close() diff --git a/transport/shadowtls/v2.go b/transport/shadowtls/v2.go index ad922e07..7b8fd3cc 100644 --- a/transport/shadowtls/v2.go +++ b/transport/shadowtls/v2.go @@ -80,8 +80,10 @@ func (c *hashWriteConn) Fallback() { type shadowConn struct { net.Conn - readRemaining int - writeMu sync.Mutex + readRemaining int + readHeader [tlsHeaderSize]byte + readHeaderOffset int + writeMu sync.Mutex } func newConn(conn net.Conn) *shadowConn { @@ -97,19 +99,22 @@ func (c *shadowConn) Read(p []byte) (int, error) { c.readRemaining -= n return n, err } - var header [tlsHeaderSize]byte - if _, err := io.ReadFull(c.Conn, header[:]); err != nil { + // Keep an incomplete header so a read deadline only interrupts the current Read. + n, err := io.ReadFull(c.Conn, c.readHeader[c.readHeaderOffset:]) + c.readHeaderOffset += n + if err != nil { return 0, err } - if header[0] != applicationData { - return 0, fmt.Errorf("shadow-tls: unexpected TLS record type: %d", header[0]) + c.readHeaderOffset = 0 + if c.readHeader[0] != applicationData { + return 0, fmt.Errorf("shadow-tls: unexpected TLS record type: %d", c.readHeader[0]) } - length := int(binary.BigEndian.Uint16(header[3:])) + length := int(binary.BigEndian.Uint16(c.readHeader[3:])) readLength := len(p) if readLength > length { readLength = length } - n, err := c.Conn.Read(p[:readLength]) + n, err = c.Conn.Read(p[:readLength]) c.readRemaining = length - n return n, err } @@ -146,8 +151,6 @@ func (c *shadowConn) writeRecord(p []byte) error { return writeBuffers(c.Conn, header[:], p) } -func (c *shadowConn) NeedAdditionalReadDeadline() bool { return true } - func (c *shadowConn) Upstream() any { return c.Conn } type clientConn struct {