1
0
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:
wwqgtxx
2026-07-15 00:31:49 +08:00
parent 7a3c37cb33
commit 2b0d8411fa
2 changed files with 65 additions and 10 deletions
+52
View File
@@ -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
View File
@@ -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 {