1
0
mirror of https://github.com/MetaCubeX/mihomo.git synced 2026-10-10 04:03:11 +08:00

fix: make shadowtls v3 read deadlines recoverable

This commit is contained in:
wwqgtxx
2026-07-15 00:31:14 +08:00
parent 3f22309b51
commit 7a3c37cb33
2 changed files with 175 additions and 3 deletions
+141
View File
@@ -6,9 +6,11 @@ import (
"crypto/hmac"
"crypto/sha1"
stdTLS "crypto/tls"
"encoding/binary"
"errors"
"io"
"net"
"os"
"testing"
"time"
@@ -134,6 +136,145 @@ func TestV3UnauthenticatedConnectionFallsBack(t *testing.T) {
}
}
func TestV3InterruptedReadDoesNotSendAlert(t *testing.T) {
clientSide, serverSide := net.Pipe()
defer clientSide.Close()
defer serverSide.Close()
serverRaw := &readStartedConn{Conn: serverSide, started: make(chan struct{}, 1)}
serverRandom := bytes.Repeat([]byte{1}, tlsRandomSize)
clientAdd := hmac.New(sha1.New, []byte(testPassword))
hmacReset(clientAdd, serverRandom, 'C')
clientVerify := hmac.New(sha1.New, []byte(testPassword))
hmacReset(clientVerify, serverRandom, 'S')
serverAdd := hmac.New(sha1.New, []byte(testPassword))
hmacReset(serverAdd, serverRandom, 'S')
serverVerify := hmac.New(sha1.New, []byte(testPassword))
hmacReset(serverVerify, serverRandom, 'C')
client := newVerifiedConn(clientSide, clientAdd, clientVerify, nil)
server := newVerifiedConn(serverRaw, serverAdd, serverVerify, nil)
type readResult struct {
data []byte
err error
}
clientRead := make(chan readResult, 1)
response := []byte("response after interrupted idle read")
go func() {
buffer := make([]byte, len(response))
_, err := io.ReadFull(client, buffer)
clientRead <- readResult{data: buffer, err: err}
}()
serverRead := make(chan error, 1)
go func() {
var buffer [1]byte
_, err := server.Read(buffer[:])
serverRead <- 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)
}
select {
case result := <-clientRead:
t.Fatalf("client received data before response: data=%x err=%v", result.data, result.err)
default:
}
request := []byte("request after interrupted partial read")
frame := makeV3ClientFrame(client, request)
buffer := make([]byte, len(request))
serverRead = make(chan error, 1)
go func() {
_, err := server.Read(buffer)
serverRead <- err
}()
<-serverRaw.started
if _, err := clientSide.Write(frame[:tlsHeaderSize]); err != nil {
t.Fatalf("write partial request: %v", err)
}
<-serverRaw.started
if err := server.SetReadDeadline(time.Now()); err != nil {
t.Fatalf("interrupt partial server read: %v", err)
}
if err := <-serverRead; !errors.Is(err, os.ErrDeadlineExceeded) {
t.Fatalf("partial server read error = %v, want deadline exceeded", err)
}
if err := server.SetReadDeadline(time.Time{}); err != nil {
t.Fatalf("clear partial server read deadline: %v", err)
}
select {
case result := <-clientRead:
t.Fatalf("client received data after interrupted partial read: data=%x err=%v", result.data, result.err)
default:
}
clientWrite := make(chan error, 1)
go func() {
_, err := clientSide.Write(frame[tlsHeaderSize:])
clientWrite <- err
}()
if _, err := io.ReadFull(server, buffer); err != nil {
t.Fatalf("server read request: %v", err)
}
if !bytes.Equal(buffer, request) {
t.Fatalf("server request = %q, want %q", buffer, request)
}
if err := <-clientWrite; err != nil {
t.Fatalf("client write request: %v", err)
}
serverWrite := make(chan error, 1)
go func() {
_, err := server.Write(response)
serverWrite <- err
}()
result := <-clientRead
if result.err != nil {
t.Fatalf("client read response: %v", result.err)
}
if !bytes.Equal(result.data, response) {
t.Fatalf("client response = %q, want %q", result.data, response)
}
if err := <-serverWrite; err != nil {
t.Fatalf("server write response: %v", err)
}
}
type readStartedConn struct {
net.Conn
started chan struct{}
}
func (c *readStartedConn) Read(p []byte) (int, error) {
select {
case c.started <- struct{}{}:
default:
}
return c.Conn.Read(p)
}
func makeV3ClientFrame(client *verifiedConn, payload []byte) []byte {
frame := make([]byte, tlsHMACHeaderSize+len(payload))
frame[0] = applicationData
frame[1] = 3
frame[2] = 3
binary.BigEndian.PutUint16(frame[3:tlsHeaderSize], uint16(hmacSize+len(payload)))
_, _ = client.hmacAdd.Write(payload)
hmacHash := client.hmacAdd.Sum(nil)[:hmacSize]
_, _ = client.hmacAdd.Write(hmacHash)
copy(frame[tlsHeaderSize:tlsHMACHeaderSize], hmacHash)
copy(frame[tlsHMACHeaderSize:], payload)
return frame
}
func TestHandshakeSelectionByServerName(t *testing.T) {
frame := captureClientHello(t, "mapped.example")
if serverName, err := extractServerName(frame); err != nil || serverName != "mapped.example" {
+34 -3
View File
@@ -105,6 +105,8 @@ type verifiedConn struct {
hmacVerify hash.Hash
hmacIgnore hash.Hash
pending []byte
readBuffer []byte
readOffset int
}
func newVerifiedConn(conn net.Conn, hmacAdd, hmacVerify, hmacIgnore hash.Hash) *verifiedConn {
@@ -121,8 +123,12 @@ func (c *verifiedConn) Read(p []byte) (int, error) {
return c.readPending(p), nil
}
for {
frame, err := readFrame(c.Conn)
frame, err := c.readRecord()
if err != nil {
var netErr net.Error
if errors.As(err, &netErr) && netErr.Timeout() {
return 0, err
}
sendAlert(c.Conn)
return 0, err
}
@@ -149,6 +155,33 @@ func (c *verifiedConn) Read(p []byte) (int, error) {
}
}
func (c *verifiedConn) readRecord() ([]byte, error) {
// Keep an incomplete record so a read deadline only interrupts the current Read.
if c.readBuffer == nil {
c.readBuffer = make([]byte, tlsHeaderSize)
}
if c.readOffset < tlsHeaderSize {
n, err := io.ReadFull(c.Conn, c.readBuffer[c.readOffset:tlsHeaderSize])
c.readOffset += n
if err != nil {
return nil, err
}
length := int(binary.BigEndian.Uint16(c.readBuffer[3:]))
c.readBuffer = append(c.readBuffer, make([]byte, length)...)
}
if c.readOffset < len(c.readBuffer) {
n, err := io.ReadFull(c.Conn, c.readBuffer[c.readOffset:])
c.readOffset += n
if err != nil {
return nil, err
}
}
frame := c.readBuffer
c.readBuffer = nil
c.readOffset = 0
return frame, nil
}
func (c *verifiedConn) readPending(p []byte) int {
n := copy(p, c.pending)
c.pending = c.pending[n:]
@@ -185,8 +218,6 @@ func (c *verifiedConn) writeRecord(p []byte) error {
return writeBuffers(c.Conn, header[:], p)
}
func (c *verifiedConn) NeedAdditionalReadDeadline() bool { return true }
func (c *verifiedConn) Upstream() any { return c.Conn }
func verifyApplicationData(frame []byte, h hash.Hash, update bool) bool {