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:
@@ -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" {
|
||||
|
||||
@@ -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 {
|
||||
|
||||
Reference in New Issue
Block a user