mirror of
https://github.com/MetaCubeX/mihomo.git
synced 2026-10-10 04:03:11 +08:00
fix: make shadowtls v2 read deadlines recoverable
This commit is contained in:
@@ -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()
|
||||
|
||||
+13
-10
@@ -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 {
|
||||
|
||||
Reference in New Issue
Block a user