Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
42 changes: 25 additions & 17 deletions core/internal/integration_tests/hook_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -86,14 +86,19 @@ func TestClientServerHookUDP(t *testing.T) {
auth := mocks.NewMockAuthenticator(t)
auth.EXPECT().Authenticate(mock.Anything, mock.Anything, mock.Anything).Return(true, "nobody")
hook := mocks.NewMockRequestHook(t)
hook.EXPECT().Check(true, fakeEchoAddr).Return(true).Once()
hook.EXPECT().UDP(mock.Anything, mock.Anything).RunAndReturn(func(bytes []byte, s *string) error {
hook.EXPECT().Check(true, fakeEchoAddr).Return(true).Twice()
hook.EXPECT().UDP(mock.Anything, mock.Anything).RunAndReturn(func(packets [][]byte, s *string) (bool, error) {
assert.Equal(t, fakeEchoAddr, *s)
assert.Equal(t, []byte("hello world"), bytes)
if len(packets) == 1 {
// Hold the first packet back
assert.Equal(t, [][]byte{[]byte("hello")}, packets)
return false, nil
}
assert.Equal(t, [][]byte{[]byte("hello"), []byte(" world")}, packets)
// Change the address
*s = realEchoAddr
return nil
}).Once()
return true, nil
}).Twice()
s, err := server.NewServer(&server.Config{
TLSConfig: serverTLSConfig(),
Conn: udpConn,
Expand Down Expand Up @@ -124,22 +129,25 @@ func TestClientServerHookUDP(t *testing.T) {
assert.NoError(t, err)
defer conn.Close()

// Send and receive data
sData := []byte("hello world")
err = conn.Send(sData, fakeEchoAddr)
assert.NoError(t, err)
rData, rAddr, err := conn.Receive()
assert.NoError(t, err)
assert.Equal(t, sData, rData)
// Hook address change is transparent,
// the client should still see the fake echo address it sent packets to
assert.Equal(t, fakeEchoAddr, rAddr)
// Send and receive data, both held packets should reach the real echo server
for _, sData := range [][]byte{[]byte("hello"), []byte(" world")} {
err = conn.Send(sData, fakeEchoAddr)
assert.NoError(t, err)
}
for _, sData := range [][]byte{[]byte("hello"), []byte(" world")} {
rData, rAddr, err := conn.Receive()
assert.NoError(t, err)
assert.Equal(t, sData, rData)
// Hook address change is transparent,
// the client should still see the fake echo address it sent packets to
assert.Equal(t, fakeEchoAddr, rAddr)
}

// Subsequent packets should also be sent to the real echo server
sData = []byte("never stop fighting")
sData := []byte("never stop fighting")
err = conn.Send(sData, fakeEchoAddr)
assert.NoError(t, err)
rData, rAddr, err = conn.Receive()
rData, rAddr, err := conn.Receive()
assert.NoError(t, err)
assert.Equal(t, sData, rData)
assert.Equal(t, fakeEchoAddr, rAddr)
Expand Down
42 changes: 26 additions & 16 deletions core/internal/integration_tests/mocks/mock_RequestHook.go

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

8 changes: 5 additions & 3 deletions core/server/config.go
Original file line number Diff line number Diff line change
Expand Up @@ -146,12 +146,14 @@ type CongestionConfig struct {
// The returned byte slice, if not empty, will be sent to the remote before proxying - this is
// mainly for "putting back" the content read from the client for sniffing, etc.
// Return a non-nil error to abort the connection.
// Note that due to the current architectural limitations, it can only inspect the first packet
// of a UDP connection. It also cannot put back any data as the first packet is always sent as-is.
// For UDP, the first packets of a session are held back until the hook is done with them:
// UDP is called with all the packets held so far each time a new one arrives, and returns
// false to wait for the next one. The hook must give up eventually, as held packets are only
// sent (as-is, to the possibly modified reqAddr) once it returns true.
type RequestHook interface {
Check(isUDP bool, reqAddr string) bool
TCP(stream HyStream, reqAddr *string) ([]byte, error)
UDP(data []byte, reqAddr *string) error
UDP(packets [][]byte, reqAddr *string) (done bool, err error)
}

// Outbound provides the implementation of how the server should connect to remote servers.
Expand Down
42 changes: 26 additions & 16 deletions core/server/mock_udpIO.go

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

6 changes: 3 additions & 3 deletions core/server/server.go
Original file line number Diff line number Diff line change
Expand Up @@ -406,11 +406,11 @@ func (io *udpIOImpl) SendMessage(buf []byte, msg *protocol.UDPMessage) error {
return io.Conn.SendDatagram(buf[:msgN])
}

func (io *udpIOImpl) Hook(data []byte, reqAddr *string) error {
func (io *udpIOImpl) Hook(packets [][]byte, reqAddr *string) (bool, error) {
if io.RequestHook != nil && io.RequestHook.Check(true, *reqAddr) {
return io.RequestHook.UDP(data, reqAddr)
return io.RequestHook.UDP(packets, reqAddr)
} else {
return nil
return true, nil
}
}

Expand Down
77 changes: 51 additions & 26 deletions core/server/udp.go
Original file line number Diff line number Diff line change
Expand Up @@ -21,7 +21,7 @@ const (
type udpIO interface {
ReceiveMessage() (*protocol.UDPMessage, error)
SendMessage([]byte, *protocol.UDPMessage) error
Hook(data []byte, reqAddr *string) error
Hook(packets [][]byte, reqAddr *string) (bool, error)
UDP(reqAddr string) (UDPConn, error)
CheckUDP(reqAddr string) error
}
Expand All @@ -39,9 +39,10 @@ type udpSessionEntry struct {
Last *utils.AtomicTime
IO udpIO

DialFunc func(addr string, firstMsgData []byte) (conn UDPConn, actualAddr string, err error)
DialFunc func(addr string) (UDPConn, error)
ExitFunc func(err error)

held []*protocol.UDPMessage // Messages held back until the hook is done with them
conn UDPConn
connLock sync.Mutex
closed bool
Expand All @@ -51,7 +52,7 @@ type udpSessionEntry struct {

func newUDPSessionEntry(
id uint32, io udpIO,
dialFunc func(string, []byte) (UDPConn, string, error),
dialFunc func(string) (UDPConn, error),
exitFunc func(error),
) (e *udpSessionEntry) {
e = &udpSessionEntry{
Expand Down Expand Up @@ -91,25 +92,55 @@ func (e *udpSessionEntry) CloseWithErr(err error) {
// Feed feeds a UDP message to the session.
// If the message itself is a complete message, or it completes a fragmented message,
// the message is written to the session's UDP connection, and the number of bytes
// written is returned.
// Otherwise, 0 and nil are returned.
// written is returned. Otherwise, 0 and nil are returned.
// Until the connection is established, messages are held back instead, and are
// written together once the hook is done with them.
func (e *udpSessionEntry) Feed(msg *protocol.UDPMessage) (int, error) {
e.Last.Set(time.Now())
dfMsg := e.D.Feed(msg)
if dfMsg == nil {
return 0, nil
}
if e.conn != nil {
return e.write(dfMsg)
}

if e.conn == nil {
err := e.initConn(dfMsg)
e.held = append(e.held, dfMsg)
packets := make([][]byte, len(e.held))
for i, m := range e.held {
packets[i] = m.Data
}
firstAddr := e.held[0].Addr
addr := firstAddr
done, err := e.IO.Hook(packets, &addr)
if err != nil {
e.CloseWithErr(err)
return 0, err
}
if !done {
return 0, nil
}
if err := e.initConn(firstAddr, addr); err != nil {
return 0, err
}
if e.OverrideAddr == "" {
e.aclCache = map[string]error{firstAddr: nil}
}
held := e.held
e.held = nil
total := 0
for _, m := range held {
n, err := e.write(m)
if err != nil {
return 0, err
}
if e.OverrideAddr == "" {
e.aclCache = map[string]error{dfMsg.Addr: nil}
return total, err
}
total += n
}
return total, nil
}

// write writes a message to the session's UDP connection.
func (e *udpSessionEntry) write(dfMsg *protocol.UDPMessage) (int, error) {
addr := dfMsg.Addr
if e.OverrideAddr != "" {
addr = e.OverrideAddr
Expand Down Expand Up @@ -140,9 +171,10 @@ func (e *udpSessionEntry) checkAddr(addr string) error {
return decision
}

// initConn initializes the UDP connection of the session.
// initConn initializes the UDP connection of the session to addr,
// which is what the hook turned the address of the first message into.
// If no error is returned, the e.conn is set to the new connection.
func (e *udpSessionEntry) initConn(firstMsg *protocol.UDPMessage) error {
func (e *udpSessionEntry) initConn(firstAddr, addr string) error {
// We need this lock to ensure not to create conn after session exit
e.connLock.Lock()

Expand All @@ -151,7 +183,7 @@ func (e *udpSessionEntry) initConn(firstMsg *protocol.UDPMessage) error {
return errors.New("session is closed")
}

conn, actualAddr, err := e.DialFunc(firstMsg.Addr, firstMsg.Data)
conn, err := e.DialFunc(addr)
if err != nil {
// Fail fast if DialFunc failed
// (usually indicates the connection has been rejected by the ACL)
Expand All @@ -163,10 +195,10 @@ func (e *udpSessionEntry) initConn(firstMsg *protocol.UDPMessage) error {

e.conn = conn

if firstMsg.Addr != actualAddr {
if firstAddr != addr {
// Hook changed the address, enable address override
e.OverrideAddr = actualAddr
e.OriginalAddr = firstMsg.Addr
e.OverrideAddr = addr
e.OriginalAddr = firstAddr
}
go e.receiveLoop()

Expand Down Expand Up @@ -313,18 +345,11 @@ func (m *udpSessionManager) feed(msg *protocol.UDPMessage) {

// Create a new session if not exists
if entry == nil {
dialFunc := func(addr string, firstMsgData []byte) (conn UDPConn, actualAddr string, err error) {
// Call the hook
err = m.io.Hook(firstMsgData, &addr)
if err != nil {
return conn, actualAddr, err
}
actualAddr = addr
dialFunc := func(addr string) (UDPConn, error) {
// Log the event
m.eventLogger.New(msg.SessionID, addr)
// Dial target
conn, err = m.io.UDP(addr)
return conn, actualAddr, err
return m.io.UDP(addr)
}
exitFunc := func(err error) {
// Log the event
Expand Down
Loading
Loading