diff --git a/core/internal/integration_tests/hook_test.go b/core/internal/integration_tests/hook_test.go index db959958a..6393143ba 100644 --- a/core/internal/integration_tests/hook_test.go +++ b/core/internal/integration_tests/hook_test.go @@ -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, @@ -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) diff --git a/core/internal/integration_tests/mocks/mock_RequestHook.go b/core/internal/integration_tests/mocks/mock_RequestHook.go index 49e8c6c21..79c828f0d 100644 --- a/core/internal/integration_tests/mocks/mock_RequestHook.go +++ b/core/internal/integration_tests/mocks/mock_RequestHook.go @@ -126,22 +126,32 @@ func (_c *MockRequestHook_TCP_Call) RunAndReturn(run func(server.HyStream, *stri return _c } -// UDP provides a mock function with given fields: data, reqAddr -func (_m *MockRequestHook) UDP(data []byte, reqAddr *string) error { - ret := _m.Called(data, reqAddr) +// UDP provides a mock function with given fields: packets, reqAddr +func (_m *MockRequestHook) UDP(packets [][]byte, reqAddr *string) (bool, error) { + ret := _m.Called(packets, reqAddr) if len(ret) == 0 { panic("no return value specified for UDP") } - var r0 error - if rf, ok := ret.Get(0).(func([]byte, *string) error); ok { - r0 = rf(data, reqAddr) + var r0 bool + var r1 error + if rf, ok := ret.Get(0).(func([][]byte, *string) (bool, error)); ok { + return rf(packets, reqAddr) + } + if rf, ok := ret.Get(0).(func([][]byte, *string) bool); ok { + r0 = rf(packets, reqAddr) } else { - r0 = ret.Error(0) + r0 = ret.Get(0).(bool) } - return r0 + if rf, ok := ret.Get(1).(func([][]byte, *string) error); ok { + r1 = rf(packets, reqAddr) + } else { + r1 = ret.Error(1) + } + + return r0, r1 } // MockRequestHook_UDP_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'UDP' @@ -150,25 +160,25 @@ type MockRequestHook_UDP_Call struct { } // UDP is a helper method to define mock.On call -// - data []byte +// - packets [][]byte // - reqAddr *string -func (_e *MockRequestHook_Expecter) UDP(data interface{}, reqAddr interface{}) *MockRequestHook_UDP_Call { - return &MockRequestHook_UDP_Call{Call: _e.mock.On("UDP", data, reqAddr)} +func (_e *MockRequestHook_Expecter) UDP(packets interface{}, reqAddr interface{}) *MockRequestHook_UDP_Call { + return &MockRequestHook_UDP_Call{Call: _e.mock.On("UDP", packets, reqAddr)} } -func (_c *MockRequestHook_UDP_Call) Run(run func(data []byte, reqAddr *string)) *MockRequestHook_UDP_Call { +func (_c *MockRequestHook_UDP_Call) Run(run func(packets [][]byte, reqAddr *string)) *MockRequestHook_UDP_Call { _c.Call.Run(func(args mock.Arguments) { - run(args[0].([]byte), args[1].(*string)) + run(args[0].([][]byte), args[1].(*string)) }) return _c } -func (_c *MockRequestHook_UDP_Call) Return(_a0 error) *MockRequestHook_UDP_Call { - _c.Call.Return(_a0) +func (_c *MockRequestHook_UDP_Call) Return(done bool, err error) *MockRequestHook_UDP_Call { + _c.Call.Return(done, err) return _c } -func (_c *MockRequestHook_UDP_Call) RunAndReturn(run func([]byte, *string) error) *MockRequestHook_UDP_Call { +func (_c *MockRequestHook_UDP_Call) RunAndReturn(run func([][]byte, *string) (bool, error)) *MockRequestHook_UDP_Call { _c.Call.Return(run) return _c } diff --git a/core/server/config.go b/core/server/config.go index 365e9361a..1cf8bd391 100644 --- a/core/server/config.go +++ b/core/server/config.go @@ -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. diff --git a/core/server/mock_udpIO.go b/core/server/mock_udpIO.go index bb512c089..62e9ac04b 100644 --- a/core/server/mock_udpIO.go +++ b/core/server/mock_udpIO.go @@ -66,22 +66,32 @@ func (_c *mockUDPIO_CheckUDP_Call) RunAndReturn(run func(string) error) *mockUDP return _c } -// Hook provides a mock function with given fields: data, reqAddr -func (_m *mockUDPIO) Hook(data []byte, reqAddr *string) error { - ret := _m.Called(data, reqAddr) +// Hook provides a mock function with given fields: packets, reqAddr +func (_m *mockUDPIO) Hook(packets [][]byte, reqAddr *string) (bool, error) { + ret := _m.Called(packets, reqAddr) if len(ret) == 0 { panic("no return value specified for Hook") } - var r0 error - if rf, ok := ret.Get(0).(func([]byte, *string) error); ok { - r0 = rf(data, reqAddr) + var r0 bool + var r1 error + if rf, ok := ret.Get(0).(func([][]byte, *string) (bool, error)); ok { + return rf(packets, reqAddr) + } + if rf, ok := ret.Get(0).(func([][]byte, *string) bool); ok { + r0 = rf(packets, reqAddr) } else { - r0 = ret.Error(0) + r0 = ret.Get(0).(bool) } - return r0 + if rf, ok := ret.Get(1).(func([][]byte, *string) error); ok { + r1 = rf(packets, reqAddr) + } else { + r1 = ret.Error(1) + } + + return r0, r1 } // mockUDPIO_Hook_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Hook' @@ -90,25 +100,25 @@ type mockUDPIO_Hook_Call struct { } // Hook is a helper method to define mock.On call -// - data []byte +// - packets [][]byte // - reqAddr *string -func (_e *mockUDPIO_Expecter) Hook(data interface{}, reqAddr interface{}) *mockUDPIO_Hook_Call { - return &mockUDPIO_Hook_Call{Call: _e.mock.On("Hook", data, reqAddr)} +func (_e *mockUDPIO_Expecter) Hook(packets interface{}, reqAddr interface{}) *mockUDPIO_Hook_Call { + return &mockUDPIO_Hook_Call{Call: _e.mock.On("Hook", packets, reqAddr)} } -func (_c *mockUDPIO_Hook_Call) Run(run func(data []byte, reqAddr *string)) *mockUDPIO_Hook_Call { +func (_c *mockUDPIO_Hook_Call) Run(run func(packets [][]byte, reqAddr *string)) *mockUDPIO_Hook_Call { _c.Call.Run(func(args mock.Arguments) { - run(args[0].([]byte), args[1].(*string)) + run(args[0].([][]byte), args[1].(*string)) }) return _c } -func (_c *mockUDPIO_Hook_Call) Return(_a0 error) *mockUDPIO_Hook_Call { - _c.Call.Return(_a0) +func (_c *mockUDPIO_Hook_Call) Return(_a0 bool, _a1 error) *mockUDPIO_Hook_Call { + _c.Call.Return(_a0, _a1) return _c } -func (_c *mockUDPIO_Hook_Call) RunAndReturn(run func([]byte, *string) error) *mockUDPIO_Hook_Call { +func (_c *mockUDPIO_Hook_Call) RunAndReturn(run func([][]byte, *string) (bool, error)) *mockUDPIO_Hook_Call { _c.Call.Return(run) return _c } diff --git a/core/server/server.go b/core/server/server.go index d2b1e24a5..c3fb6d72a 100644 --- a/core/server/server.go +++ b/core/server/server.go @@ -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 } } diff --git a/core/server/udp.go b/core/server/udp.go index bfe6ae118..a0bbedfb3 100644 --- a/core/server/udp.go +++ b/core/server/udp.go @@ -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 } @@ -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 @@ -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{ @@ -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 @@ -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() @@ -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) @@ -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() @@ -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 diff --git a/core/server/udp_test.go b/core/server/udp_test.go index 8aa899f30..3772255ab 100644 --- a/core/server/udp_test.go +++ b/core/server/udp_test.go @@ -49,7 +49,7 @@ func TestUDPSessionManager(t *testing.T) { eventLogger.EXPECT().New(msg1.SessionID, msg1.Addr).Return().Once() udpConn1 := newMockUDPConn(t) udpConn1Ch := make(chan []byte, 1) - io.EXPECT().Hook(msg1.Data, &msg1.Addr).Return(nil).Once() + io.EXPECT().Hook([][]byte{msg1.Data}, &msg1.Addr).Return(true, nil).Once() io.EXPECT().UDP(msg1.Addr).Return(udpConn1, nil).Once() udpConn1.EXPECT().WriteTo(msg1.Data, msg1.Addr).Return(5, nil).Once() udpConn1.EXPECT().ReadFrom(mock.Anything).RunAndReturn(func(b []byte) (int, string, error) { @@ -88,7 +88,7 @@ func TestUDPSessionManager(t *testing.T) { udpConn2 := newMockUDPConn(t) udpConn2Ch := make(chan []byte, 1) // On fragmentation, make sure hook gets the whole message - io.EXPECT().Hook(msg2data, &msg2_1.Addr).Return(nil).Once() + io.EXPECT().Hook([][]byte{msg2data}, &msg2_1.Addr).Return(true, nil).Once() io.EXPECT().UDP(msg2_1.Addr).Return(udpConn2, nil).Once() udpConn2.EXPECT().WriteTo(msg2data, msg2_1.Addr).Return(11, nil).Once() udpConn2.EXPECT().ReadFrom(mock.Anything).RunAndReturn(func(b []byte) (int, string, error) { @@ -153,7 +153,7 @@ func TestUDPSessionManager(t *testing.T) { } eventLogger.EXPECT().New(msg4.SessionID, msg4.Addr).Return().Once() udpConn4 := newMockUDPConn(t) - io.EXPECT().Hook(msg4.Data, &msg4.Addr).Return(nil).Once() + io.EXPECT().Hook([][]byte{msg4.Data}, &msg4.Addr).Return(true, nil).Once() io.EXPECT().UDP(msg4.Addr).Return(udpConn4, nil).Once() udpConn4.EXPECT().WriteTo(msg4.Data, msg4.Addr).Return(12, nil).Once() udpConn4.EXPECT().ReadFrom(mock.Anything).Return(0, "", errUDPClosed).Once() @@ -175,7 +175,7 @@ func TestUDPSessionManager(t *testing.T) { Data: []byte("babe i miss you"), } eventLogger.EXPECT().New(msg5.SessionID, msg5.Addr).Return().Once() - io.EXPECT().Hook(msg5.Data, &msg5.Addr).Return(nil).Once() + io.EXPECT().Hook([][]byte{msg5.Data}, &msg5.Addr).Return(true, nil).Once() io.EXPECT().UDP(msg5.Addr).Return(nil, errUDPIO).Once() eventLogger.EXPECT().Close(msg5.SessionID, errUDPIO).Once() msgCh <- msg5 @@ -189,3 +189,73 @@ func TestUDPSessionManager(t *testing.T) { assert.Zero(t, sm.Count(), "session count should be 0") goleak.VerifyNone(t) } + +func TestUDPSessionManagerHookHold(t *testing.T) { + io := newMockUDPIO(t) + eventLogger := newMockUDPEventLogger(t) + sm := newUDPSessionManager(io, eventLogger, 2*time.Second) + + msgCh := make(chan *protocol.UDPMessage, 4) + io.EXPECT().ReceiveMessage().RunAndReturn(func() (*protocol.UDPMessage, error) { + m := <-msgCh + if m == nil { + return nil, errors.New("closed") + } + return m, nil + }) + + go sm.Run() + + msg1 := &protocol.UDPMessage{SessionID: 42, FragCount: 1, Addr: "1.2.3.4:443", Data: []byte("first")} + msg2 := &protocol.UDPMessage{SessionID: 42, FragCount: 1, Addr: "1.2.3.4:443", Data: []byte("second")} + // The hook holds the first message back, and rewrites the address once it sees the second + io.EXPECT().Hook([][]byte{msg1.Data}, &msg1.Addr).Return(false, nil).Once() + io.EXPECT().Hook([][]byte{msg1.Data, msg2.Data}, &msg1.Addr).RunAndReturn(func(_ [][]byte, addr *string) (bool, error) { + *addr = "example.com:443" + return true, nil + }).Once() + eventLogger.EXPECT().New(msg1.SessionID, "example.com:443").Return().Once() + udpConn := newMockUDPConn(t) + udpConnCh := make(chan []byte, 1) + io.EXPECT().UDP("example.com:443").Return(udpConn, nil).Once() + mock.InOrder( + udpConn.EXPECT().WriteTo(msg1.Data, "example.com:443").Return(5, nil).Call, + udpConn.EXPECT().WriteTo(msg2.Data, "example.com:443").Return(6, nil).Call, + ) + udpConn.EXPECT().ReadFrom(mock.Anything).RunAndReturn(func(b []byte) (int, string, error) { + bs := <-udpConnCh + if bs == nil { + return 0, "", errors.New("closed") + } + return copy(b, bs), "93.184.215.14:443", nil + }) + // Replies come from the original address + replied := make(chan struct{}) + io.EXPECT().SendMessage(mock.Anything, &protocol.UDPMessage{ + SessionID: msg1.SessionID, + FragCount: 1, + Addr: msg1.Addr, + Data: []byte("reply"), + }).RunAndReturn(func([]byte, *protocol.UDPMessage) error { + close(replied) + return nil + }).Once() + msgCh <- msg1 + msgCh <- msg2 + udpConnCh <- []byte("reply") + <-replied + + udpConn.EXPECT().Close().RunAndReturn(func() error { + close(udpConnCh) + return nil + }).Once() + eventLogger.EXPECT().Close(msg1.SessionID, nil).Once() + + // Wait for timeout + assert.Eventually(t, func() bool { return sm.Count() == 0 }, 5*time.Second, 100*time.Millisecond) + mock.AssertExpectationsForObjects(t, io, eventLogger, udpConn) + + close(msgCh) + time.Sleep(1 * time.Second) + goleak.VerifyNone(t) +} diff --git a/extras/go.mod b/extras/go.mod index 4d0381cdf..69837f2f9 100644 --- a/extras/go.mod +++ b/extras/go.mod @@ -10,7 +10,6 @@ require ( github.com/libp2p/go-nat v1.0.1-0.20250821073202-01afc089f138 github.com/miekg/dns v1.1.72 github.com/pion/stun/v3 v3.1.6 - github.com/refraction-networking/utls v1.8.2 github.com/stretchr/testify v1.12.1 github.com/txthinking/socks5 v0.0.0-20230325130024-4230056ae301 golang.org/x/crypto v0.54.0 @@ -32,6 +31,7 @@ require ( github.com/pion/logging v0.2.4 // indirect github.com/pion/transport/v4 v4.0.2 // indirect github.com/quic-go/qpack v0.6.0 // indirect + github.com/refraction-networking/utls v1.8.2 // indirect github.com/stretchr/objx v0.5.3 // indirect github.com/txthinking/runnergroup v0.0.0-20210608031112-152c7c4432bf // indirect github.com/wlynxg/anet v0.0.5 // indirect diff --git a/extras/sniff/.mockery.yaml b/extras/sniff/.mockery.yaml deleted file mode 100644 index c866d1da1..000000000 --- a/extras/sniff/.mockery.yaml +++ /dev/null @@ -1,12 +0,0 @@ -with-expecter: true -dir: . -outpkg: sniff -packages: - github.com/apernet/quic-go: - interfaces: - Stream: - config: - mockname: mockStream - replace-type: # internal package alias dirty fix - - github.com/apernet/quic-go/internal/protocol=github.com/apernet/quic-go - - github.com/apernet/quic-go/internal/qerr=github.com/apernet/quic-go diff --git a/extras/sniff/internal/quic/LICENSE b/extras/sniff/internal/quic/LICENSE deleted file mode 100644 index 43970c410..000000000 --- a/extras/sniff/internal/quic/LICENSE +++ /dev/null @@ -1,31 +0,0 @@ -Author:: Cuong Manh Le -Copyright:: Copyright (c) 2023, Cuong Manh Le -All rights reserved. - -Redistribution and use in source and binary forms, with or without -modification, are permitted provided that the following conditions are -met: - - * Redistributions of source code must retain the above copyright - notice, this list of conditions and the following disclaimer. - - * Redistributions in binary form must reproduce the above - copyright notice, this list of conditions and the following - disclaimer in the documentation and/or other materials provided - with the distribution. - - * Neither the name of the @organization@ nor the names of its - contributors may be used to endorse or promote products derived - from this software without specific prior written permission. - -THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS -"AS IS" AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT -LIMITED TO, THE IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR -A PARTICULAR PURPOSE ARE DISCLAIMED. IN NO EVENT SHALL LE MANH CUONG -BE LIABLE FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR -CONSEQUENTIAL DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF -SUBSTITUTE GOODS OR SERVICES; LOSS OF USE, DATA, OR PROFITS; OR -BUSINESS INTERRUPTION) HOWEVER CAUSED AND ON ANY THEORY OF LIABILITY, -WHETHER IN CONTRACT, STRICT LIABILITY, OR TORT (INCLUDING NEGLIGENCE -OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE OF THIS SOFTWARE, EVEN -IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. \ No newline at end of file diff --git a/extras/sniff/internal/quic/README.md b/extras/sniff/internal/quic/README.md deleted file mode 100644 index 8f3a5e2a6..000000000 --- a/extras/sniff/internal/quic/README.md +++ /dev/null @@ -1 +0,0 @@ -The code here is from https://github.com/cuonglm/quicsni with various modifications. \ No newline at end of file diff --git a/extras/sniff/internal/quic/header.go b/extras/sniff/internal/quic/header.go deleted file mode 100644 index c1a5e7c4c..000000000 --- a/extras/sniff/internal/quic/header.go +++ /dev/null @@ -1,105 +0,0 @@ -package quic - -import ( - "bytes" - "encoding/binary" - "errors" - "io" - - "github.com/apernet/quic-go/quicvarint" -) - -// The Header represents a QUIC header. -type Header struct { - Type uint8 - Version uint32 - SrcConnectionID []byte - DestConnectionID []byte - Length int64 - Token []byte -} - -// ParseInitialHeader parses the initial packet of a QUIC connection, -// return the initial header and number of bytes read so far. -func ParseInitialHeader(data []byte) (*Header, int64, error) { - br := bytes.NewReader(data) - hdr, err := parseLongHeader(br) - if err != nil { - return nil, 0, err - } - n := int64(len(data) - br.Len()) - return hdr, n, nil -} - -func parseLongHeader(b *bytes.Reader) (*Header, error) { - typeByte, err := b.ReadByte() - if err != nil { - return nil, err - } - h := &Header{} - ver, err := beUint32(b) - if err != nil { - return nil, err - } - h.Version = ver - if h.Version != 0 && typeByte&0x40 == 0 { - return nil, errors.New("not a QUIC packet") - } - destConnIDLen, err := b.ReadByte() - if err != nil { - return nil, err - } - h.DestConnectionID = make([]byte, int(destConnIDLen)) - if err := readConnectionID(b, h.DestConnectionID); err != nil { - return nil, err - } - srcConnIDLen, err := b.ReadByte() - if err != nil { - return nil, err - } - h.SrcConnectionID = make([]byte, int(srcConnIDLen)) - if err := readConnectionID(b, h.SrcConnectionID); err != nil { - return nil, err - } - - initialPacketType := byte(0b00) - if h.Version == V2 { - initialPacketType = 0b01 - } - if (typeByte >> 4 & 0b11) == initialPacketType { - tokenLen, err := quicvarint.Read(b) - if err != nil { - return nil, err - } - if tokenLen > uint64(b.Len()) { - return nil, io.EOF - } - h.Token = make([]byte, tokenLen) - if _, err := io.ReadFull(b, h.Token); err != nil { - return nil, err - } - } - - pl, err := quicvarint.Read(b) - if err != nil { - return nil, err - } - h.Length = int64(pl) - return h, err -} - -func readConnectionID(r io.Reader, cid []byte) error { - _, err := io.ReadFull(r, cid) - if err == io.ErrUnexpectedEOF { - return io.EOF - } - return nil -} - -func beUint32(r io.Reader) (uint32, error) { - b := make([]byte, 4) - if _, err := io.ReadFull(r, b); err != nil { - return 0, err - } - return binary.BigEndian.Uint32(b), nil -} diff --git a/extras/sniff/internal/quic/packet_protector.go b/extras/sniff/internal/quic/packet_protector.go deleted file mode 100644 index 42de84113..000000000 --- a/extras/sniff/internal/quic/packet_protector.go +++ /dev/null @@ -1,193 +0,0 @@ -package quic - -import ( - "crypto" - "crypto/aes" - "crypto/cipher" - "crypto/sha256" - "crypto/tls" - "encoding/binary" - "errors" - "fmt" - "hash" - - "golang.org/x/crypto/chacha20" - "golang.org/x/crypto/chacha20poly1305" - "golang.org/x/crypto/cryptobyte" - "golang.org/x/crypto/hkdf" -) - -// NewProtectionKey creates a new ProtectionKey. -func NewProtectionKey(suite uint16, secret []byte, v uint32) (*ProtectionKey, error) { - return newProtectionKey(suite, secret, v) -} - -// NewInitialProtectionKey is like NewProtectionKey, but the returned protection key -// is used for encrypt/decrypt Initial Packet only. -// -// See: https://datatracker.ietf.org/doc/html/draft-ietf-quic-tls-32#name-initial-secrets -func NewInitialProtectionKey(secret []byte, v uint32) (*ProtectionKey, error) { - return NewProtectionKey(tls.TLS_AES_128_GCM_SHA256, secret, v) -} - -// NewPacketProtector creates a new PacketProtector. -func NewPacketProtector(key *ProtectionKey) *PacketProtector { - return &PacketProtector{key: key} -} - -// PacketProtector is used for protecting a QUIC packet. -// -// See: https://www.rfc-editor.org/rfc/rfc9001.html#name-packet-protection -type PacketProtector struct { - key *ProtectionKey -} - -// UnProtect decrypts a QUIC packet. -func (pp *PacketProtector) UnProtect(packet []byte, pnOffset, pnMax int64) ([]byte, error) { - if isLongHeader(packet[0]) && int64(len(packet)) < pnOffset+4+16 { - return nil, errors.New("packet with long header is too small") - } - - // https://www.rfc-editor.org/rfc/rfc9001.html#name-header-protection-sample - sampleOffset := pnOffset + 4 - sample := packet[sampleOffset : sampleOffset+16] - - // https://www.rfc-editor.org/rfc/rfc9001.html#name-header-protection-applicati - mask := pp.key.headerProtection(sample) - if isLongHeader(packet[0]) { - // Long header: 4 bits masked - packet[0] ^= mask[0] & 0x0f - } else { - // Short header: 5 bits masked - packet[0] ^= mask[0] & 0x1f - } - - pnLen := packet[0]&0x3 + 1 - pn := int64(0) - for i := uint8(0); i < pnLen; i++ { - packet[pnOffset:][i] ^= mask[1+i] - pn = (pn << 8) | int64(packet[pnOffset:][i]) - } - pn = decodePacketNumber(pnMax, pn, pnLen) - hdr := packet[:pnOffset+int64(pnLen)] - payload := packet[pnOffset:][pnLen:] - dec, err := pp.key.aead.Open(payload[:0], pp.key.nonce(pn), payload, hdr) - if err != nil { - return nil, fmt.Errorf("decryption failed: %w", err) - } - return dec, nil -} - -// ProtectionKey is the key used to protect a QUIC packet. -type ProtectionKey struct { - aead cipher.AEAD - headerProtection func(sample []byte) (mask []byte) - iv []byte -} - -// https://datatracker.ietf.org/doc/html/draft-ietf-quic-tls-32#name-aead-usage -// -// "The 62 bits of the reconstructed QUIC packet number in network byte order are -// left-padded with zeros to the size of the IV. The exclusive OR of the padded -// packet number and the IV forms the AEAD nonce." -func (pk *ProtectionKey) nonce(pn int64) []byte { - nonce := make([]byte, len(pk.iv)) - binary.BigEndian.PutUint64(nonce[len(nonce)-8:], uint64(pn)) - for i := range pk.iv { - nonce[i] ^= pk.iv[i] - } - return nonce -} - -func newProtectionKey(suite uint16, secret []byte, v uint32) (*ProtectionKey, error) { - switch suite { - case tls.TLS_AES_128_GCM_SHA256: - key := hkdfExpandLabel(crypto.SHA256.New, secret, keyLabel(v), nil, 16) - c, err := aes.NewCipher(key) - if err != nil { - panic(err) - } - aead, err := cipher.NewGCM(c) - if err != nil { - panic(err) - } - iv := hkdfExpandLabel(crypto.SHA256.New, secret, ivLabel(v), nil, aead.NonceSize()) - hpKey := hkdfExpandLabel(crypto.SHA256.New, secret, headerProtectionLabel(v), nil, 16) - hp, err := aes.NewCipher(hpKey) - if err != nil { - panic(err) - } - k := &ProtectionKey{} - k.aead = aead - // https://datatracker.ietf.org/doc/html/draft-ietf-quic-tls-32#name-aes-based-header-protection - k.headerProtection = func(sample []byte) []byte { - mask := make([]byte, hp.BlockSize()) - hp.Encrypt(mask, sample) - return mask - } - k.iv = iv - return k, nil - case tls.TLS_CHACHA20_POLY1305_SHA256: - key := hkdfExpandLabel(crypto.SHA256.New, secret, keyLabel(v), nil, chacha20poly1305.KeySize) - aead, err := chacha20poly1305.New(key) - if err != nil { - return nil, err - } - iv := hkdfExpandLabel(crypto.SHA256.New, secret, ivLabel(v), nil, aead.NonceSize()) - hpKey := hkdfExpandLabel(sha256.New, secret, headerProtectionLabel(v), nil, chacha20.KeySize) - k := &ProtectionKey{} - k.aead = aead - // https://datatracker.ietf.org/doc/html/draft-ietf-quic-tls-32#name-chacha20-based-header-prote - k.headerProtection = func(sample []byte) []byte { - nonce := sample[4:16] - c, err := chacha20.NewUnauthenticatedCipher(hpKey, nonce) - if err != nil { - panic(err) - } - c.SetCounter(binary.LittleEndian.Uint32(sample[:4])) - mask := make([]byte, 5) - c.XORKeyStream(mask, mask) - return mask - } - k.iv = iv - return k, nil - } - return nil, errors.New("not supported cipher suite") -} - -// decodePacketNumber decode the packet number after header protection removed. -// -// See: https://datatracker.ietf.org/doc/html/draft-ietf-quic-transport-32#section-appendix.a -func decodePacketNumber(largest, truncated int64, nbits uint8) int64 { - expected := largest + 1 - win := int64(1 << (nbits * 8)) - hwin := win / 2 - mask := win - 1 - candidate := (expected &^ mask) | truncated - switch { - case candidate <= expected-hwin && candidate < (1<<62)-win: - return candidate + win - case candidate > expected+hwin && candidate >= win: - return candidate - win - } - return candidate -} - -// Copied from crypto/tls/key_schedule.go. -func hkdfExpandLabel(hash func() hash.Hash, secret []byte, label string, context []byte, length int) []byte { - var hkdfLabel cryptobyte.Builder - hkdfLabel.AddUint16(uint16(length)) - hkdfLabel.AddUint8LengthPrefixed(func(b *cryptobyte.Builder) { - b.AddBytes([]byte("tls13 ")) - b.AddBytes([]byte(label)) - }) - hkdfLabel.AddUint8LengthPrefixed(func(b *cryptobyte.Builder) { - b.AddBytes(context) - }) - out := make([]byte, length) - n, err := hkdf.Expand(hash, secret, hkdfLabel.BytesOrPanic()).Read(out) - if err != nil || n != length { - panic("quic: HKDF-Expand-Label invocation failed unexpectedly") - } - return out -} diff --git a/extras/sniff/internal/quic/packet_protector_test.go b/extras/sniff/internal/quic/packet_protector_test.go deleted file mode 100644 index bc355d218..000000000 --- a/extras/sniff/internal/quic/packet_protector_test.go +++ /dev/null @@ -1,94 +0,0 @@ -package quic - -import ( - "bytes" - "crypto" - "crypto/tls" - "encoding/hex" - "strings" - "testing" - "unicode" - - "golang.org/x/crypto/hkdf" -) - -func TestInitialPacketProtector_UnProtect(t *testing.T) { - // https://datatracker.ietf.org/doc/html/draft-ietf-quic-tls-32#name-server-initial - protect := mustHexDecodeString(` - c7ff0000200008f067a5502a4262b500 4075fb12ff07823a5d24534d906ce4c7 - 6782a2167e3479c0f7f6395dc2c91676 302fe6d70bb7cbeb117b4ddb7d173498 - 44fd61dae200b8338e1b932976b61d91 e64a02e9e0ee72e3a6f63aba4ceeeec5 - be2f24f2d86027572943533846caa13e 6f163fb257473d0eda5047360fd4a47e - fd8142fafc0f76 - `) - unProtect := mustHexDecodeString(` - 02000000000600405a020000560303ee fce7f7b37ba1d1632e96677825ddf739 - 88cfc79825df566dc5430b9a045a1200 130100002e00330024001d00209d3c94 - 0d89690b84d08a60993c144eca684d10 81287c834d5311bcf32bb9da1a002b00 - 020304 - `) - - connID := mustHexDecodeString(`8394c8f03e515708`) - - packet := append([]byte{}, protect...) - hdr, offset, err := ParseInitialHeader(packet) - if err != nil { - t.Fatal(err) - } - - initialSecret := hkdf.Extract(crypto.SHA256.New, connID, getSalt(hdr.Version)) - serverSecret := hkdfExpandLabel(crypto.SHA256.New, initialSecret, "server in", []byte{}, crypto.SHA256.Size()) - key, err := NewInitialProtectionKey(serverSecret, hdr.Version) - if err != nil { - t.Fatal(err) - } - pp := NewPacketProtector(key) - got, err := pp.UnProtect(protect, offset, 1) - if err != nil { - t.Fatal(err) - } - if !bytes.Equal(got, unProtect) { - t.Error("UnProtect returns wrong result") - } -} - -func TestPacketProtectorShortHeader_UnProtect(t *testing.T) { - // https://datatracker.ietf.org/doc/html/draft-ietf-quic-tls-32#name-chacha20-poly1305-short-hea - protect := mustHexDecodeString(`4cfe4189655e5cd55c41f69080575d7999c25a5bfb`) - unProtect := mustHexDecodeString(`01`) - hdr := mustHexDecodeString(`4200bff4`) - - secret := mustHexDecodeString(`9ac312a7f877468ebe69422748ad00a1 5443f18203a07d6060f688f30f21632b`) - k, err := NewProtectionKey(tls.TLS_CHACHA20_POLY1305_SHA256, secret, V1) - if err != nil { - t.Fatal(err) - } - - pnLen := int(hdr[0]&0x03) + 1 - offset := len(hdr) - pnLen - pp := NewPacketProtector(k) - got, err := pp.UnProtect(protect, int64(offset), 654360564) - if err != nil { - t.Fatal(err) - } - if !bytes.Equal(got, unProtect) { - t.Error("UnProtect returns wrong result") - } -} - -func mustHexDecodeString(s string) []byte { - b, err := hex.DecodeString(normalizeHex(s)) - if err != nil { - panic(err) - } - return b -} - -func normalizeHex(s string) string { - return strings.Map(func(c rune) rune { - if unicode.IsSpace(c) { - return -1 - } - return c - }, s) -} diff --git a/extras/sniff/internal/quic/payload.go b/extras/sniff/internal/quic/payload.go deleted file mode 100644 index 14e80c3c6..000000000 --- a/extras/sniff/internal/quic/payload.go +++ /dev/null @@ -1,148 +0,0 @@ -package quic - -import ( - "bytes" - "crypto" - "errors" - "fmt" - "io" - "math" - "sort" - - "github.com/apernet/quic-go/quicvarint" - "golang.org/x/crypto/hkdf" -) - -const ( - maxCryptoFrameDataLen = 256 * 1024 // 256 KiB - maxCryptoPayloadLen = 256 * 1024 // 256 KiB -) - -func ReadCryptoPayload(packet []byte) ([]byte, error) { - hdr, offset, err := ParseInitialHeader(packet) - if err != nil { - return nil, err - } - // Some sanity checks - if hdr.Version != V1 && hdr.Version != V2 { - return nil, fmt.Errorf("unsupported version: %x", hdr.Version) - } - if offset == 0 || hdr.Length == 0 { - return nil, errors.New("invalid packet") - } - - initialSecret := hkdf.Extract(crypto.SHA256.New, hdr.DestConnectionID, getSalt(hdr.Version)) - clientSecret := hkdfExpandLabel(crypto.SHA256.New, initialSecret, "client in", []byte{}, crypto.SHA256.Size()) - key, err := NewInitialProtectionKey(clientSecret, hdr.Version) - if err != nil { - return nil, fmt.Errorf("NewInitialProtectionKey: %w", err) - } - pp := NewPacketProtector(key) - // https://datatracker.ietf.org/doc/html/draft-ietf-quic-tls-32#name-client-initial - // - // "The unprotected header includes the connection ID and a 4-byte packet number encoding for a packet number of 2" - if int64(len(packet)) < offset+hdr.Length { - return nil, fmt.Errorf("packet is too short: %d < %d", len(packet), offset+hdr.Length) - } - unProtectedPayload, err := pp.UnProtect(packet[:offset+hdr.Length], offset, 2) - if err != nil { - return nil, err - } - frs, err := extractCryptoFrames(bytes.NewReader(unProtectedPayload)) - if err != nil { - return nil, err - } - data := assembleCryptoFrames(frs) - if data == nil { - return nil, errors.New("unable to assemble crypto frames") - } - return data, nil -} - -const ( - paddingFrameType = 0x00 - pingFrameType = 0x01 - cryptoFrameType = 0x06 -) - -type cryptoFrame struct { - Offset int64 - Data []byte -} - -func extractCryptoFrames(r *bytes.Reader) ([]cryptoFrame, error) { - var frames []cryptoFrame - for r.Len() > 0 { - typ, err := quicvarint.Read(r) - if err != nil { - return nil, err - } - if typ == paddingFrameType || typ == pingFrameType { - continue - } - if typ != cryptoFrameType { - return nil, fmt.Errorf("encountered unexpected frame type: %d", typ) - } - var frame cryptoFrame - offset, err := quicvarint.Read(r) - if err != nil { - return nil, err - } - if offset > uint64(math.MaxInt64) { - return nil, errors.New("invalid crypto frame offset") - } - frame.Offset = int64(offset) - dataLen, err := quicvarint.Read(r) - if err != nil { - return nil, err - } - if dataLen > maxCryptoFrameDataLen { - return nil, errors.New("crypto frame data too large") - } - if dataLen > uint64(r.Len()) { - return nil, io.ErrUnexpectedEOF - } - frame.Data = make([]byte, dataLen) - if _, err := io.ReadFull(r, frame.Data); err != nil { - return nil, err - } - frames = append(frames, frame) - } - return frames, nil -} - -// assembleCryptoFrames assembles multiple crypto frames into a single slice (if possible). -// It returns an error if the frames cannot be assembled. This can happen if the frames are not contiguous. -func assembleCryptoFrames(frames []cryptoFrame) []byte { - if len(frames) == 0 { - return nil - } - if len(frames) == 1 { - return frames[0].Data - } - // sort the frames by offset - sort.Slice(frames, func(i, j int) bool { return frames[i].Offset < frames[j].Offset }) - // check if the frames are contiguous - for i := 1; i < len(frames); i++ { - if frames[i].Offset != frames[i-1].Offset+int64(len(frames[i-1].Data)) { - return nil - } - } - // concatenate the frames - last := frames[len(frames)-1] - if last.Offset < 0 { - return nil - } - if last.Offset > maxCryptoPayloadLen { - return nil - } - end := last.Offset + int64(len(last.Data)) - if end < 0 || end > maxCryptoPayloadLen { - return nil - } - data := make([]byte, end) - for _, frame := range frames { - copy(data[frame.Offset:], frame.Data) - } - return data -} diff --git a/extras/sniff/internal/quic/quic.go b/extras/sniff/internal/quic/quic.go deleted file mode 100644 index 1cfa10386..000000000 --- a/extras/sniff/internal/quic/quic.go +++ /dev/null @@ -1,59 +0,0 @@ -package quic - -const ( - V1 uint32 = 0x1 - V2 uint32 = 0x6b3343cf - - hkdfLabelKeyV1 = "quic key" - hkdfLabelKeyV2 = "quicv2 key" - hkdfLabelIVV1 = "quic iv" - hkdfLabelIVV2 = "quicv2 iv" - hkdfLabelHPV1 = "quic hp" - hkdfLabelHPV2 = "quicv2 hp" -) - -var ( - quicSaltOld = []byte{0xaf, 0xbf, 0xec, 0x28, 0x99, 0x93, 0xd2, 0x4c, 0x9e, 0x97, 0x86, 0xf1, 0x9c, 0x61, 0x11, 0xe0, 0x43, 0x90, 0xa8, 0x99} - // https://www.rfc-editor.org/rfc/rfc9001.html#name-initial-secrets - quicSaltV1 = []byte{0x38, 0x76, 0x2c, 0xf7, 0xf5, 0x59, 0x34, 0xb3, 0x4d, 0x17, 0x9a, 0xe6, 0xa4, 0xc8, 0x0c, 0xad, 0xcc, 0xbb, 0x7f, 0x0a} - // https://www.ietf.org/archive/id/draft-ietf-quic-v2-10.html#name-initial-salt-2 - quicSaltV2 = []byte{0x0d, 0xed, 0xe3, 0xde, 0xf7, 0x00, 0xa6, 0xdb, 0x81, 0x93, 0x81, 0xbe, 0x6e, 0x26, 0x9d, 0xcb, 0xf9, 0xbd, 0x2e, 0xd9} -) - -// isLongHeader reports whether b is the first byte of a long header packet. -func isLongHeader(b byte) bool { - return b&0x80 > 0 -} - -func getSalt(v uint32) []byte { - switch v { - case V1: - return quicSaltV1 - case V2: - return quicSaltV2 - } - return quicSaltOld -} - -func keyLabel(v uint32) string { - kl := hkdfLabelKeyV1 - if v == V2 { - kl = hkdfLabelKeyV2 - } - return kl -} - -func ivLabel(v uint32) string { - ivl := hkdfLabelIVV1 - if v == V2 { - ivl = hkdfLabelIVV2 - } - return ivl -} - -func headerProtectionLabel(v uint32) string { - if v == V2 { - return hkdfLabelHPV2 - } - return hkdfLabelHPV1 -} diff --git a/extras/sniff/mock_Stream.go b/extras/sniff/mock_Stream.go deleted file mode 100644 index 8b21e953b..000000000 --- a/extras/sniff/mock_Stream.go +++ /dev/null @@ -1,492 +0,0 @@ -// Code generated by mockery v2.43.0. DO NOT EDIT. - -package sniff - -import ( - context "context" - - qerr "github.com/apernet/quic-go" - mock "github.com/stretchr/testify/mock" - - time "time" -) - -// mockStream is an autogenerated mock type for the Stream type -type mockStream struct { - mock.Mock -} - -type mockStream_Expecter struct { - mock *mock.Mock -} - -func (_m *mockStream) EXPECT() *mockStream_Expecter { - return &mockStream_Expecter{mock: &_m.Mock} -} - -// CancelRead provides a mock function with given fields: _a0 -func (_m *mockStream) CancelRead(_a0 qerr.StreamErrorCode) { - _m.Called(_a0) -} - -// mockStream_CancelRead_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'CancelRead' -type mockStream_CancelRead_Call struct { - *mock.Call -} - -// CancelRead is a helper method to define mock.On call -// - _a0 qerr.StreamErrorCode -func (_e *mockStream_Expecter) CancelRead(_a0 interface{}) *mockStream_CancelRead_Call { - return &mockStream_CancelRead_Call{Call: _e.mock.On("CancelRead", _a0)} -} - -func (_c *mockStream_CancelRead_Call) Run(run func(_a0 qerr.StreamErrorCode)) *mockStream_CancelRead_Call { - _c.Call.Run(func(args mock.Arguments) { - run(args[0].(qerr.StreamErrorCode)) - }) - return _c -} - -func (_c *mockStream_CancelRead_Call) Return() *mockStream_CancelRead_Call { - _c.Call.Return() - return _c -} - -func (_c *mockStream_CancelRead_Call) RunAndReturn(run func(qerr.StreamErrorCode)) *mockStream_CancelRead_Call { - _c.Call.Return(run) - return _c -} - -// CancelWrite provides a mock function with given fields: _a0 -func (_m *mockStream) CancelWrite(_a0 qerr.StreamErrorCode) { - _m.Called(_a0) -} - -// mockStream_CancelWrite_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'CancelWrite' -type mockStream_CancelWrite_Call struct { - *mock.Call -} - -// CancelWrite is a helper method to define mock.On call -// - _a0 qerr.StreamErrorCode -func (_e *mockStream_Expecter) CancelWrite(_a0 interface{}) *mockStream_CancelWrite_Call { - return &mockStream_CancelWrite_Call{Call: _e.mock.On("CancelWrite", _a0)} -} - -func (_c *mockStream_CancelWrite_Call) Run(run func(_a0 qerr.StreamErrorCode)) *mockStream_CancelWrite_Call { - _c.Call.Run(func(args mock.Arguments) { - run(args[0].(qerr.StreamErrorCode)) - }) - return _c -} - -func (_c *mockStream_CancelWrite_Call) Return() *mockStream_CancelWrite_Call { - _c.Call.Return() - return _c -} - -func (_c *mockStream_CancelWrite_Call) RunAndReturn(run func(qerr.StreamErrorCode)) *mockStream_CancelWrite_Call { - _c.Call.Return(run) - return _c -} - -// Close provides a mock function with given fields: -func (_m *mockStream) Close() error { - ret := _m.Called() - - if len(ret) == 0 { - panic("no return value specified for Close") - } - - var r0 error - if rf, ok := ret.Get(0).(func() error); ok { - r0 = rf() - } else { - r0 = ret.Error(0) - } - - return r0 -} - -// mockStream_Close_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Close' -type mockStream_Close_Call struct { - *mock.Call -} - -// Close is a helper method to define mock.On call -func (_e *mockStream_Expecter) Close() *mockStream_Close_Call { - return &mockStream_Close_Call{Call: _e.mock.On("Close")} -} - -func (_c *mockStream_Close_Call) Run(run func()) *mockStream_Close_Call { - _c.Call.Run(func(args mock.Arguments) { - run() - }) - return _c -} - -func (_c *mockStream_Close_Call) Return(_a0 error) *mockStream_Close_Call { - _c.Call.Return(_a0) - return _c -} - -func (_c *mockStream_Close_Call) RunAndReturn(run func() error) *mockStream_Close_Call { - _c.Call.Return(run) - return _c -} - -// Context provides a mock function with given fields: -func (_m *mockStream) Context() context.Context { - ret := _m.Called() - - if len(ret) == 0 { - panic("no return value specified for Context") - } - - var r0 context.Context - if rf, ok := ret.Get(0).(func() context.Context); ok { - r0 = rf() - } else { - if ret.Get(0) != nil { - r0 = ret.Get(0).(context.Context) - } - } - - return r0 -} - -// mockStream_Context_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Context' -type mockStream_Context_Call struct { - *mock.Call -} - -// Context is a helper method to define mock.On call -func (_e *mockStream_Expecter) Context() *mockStream_Context_Call { - return &mockStream_Context_Call{Call: _e.mock.On("Context")} -} - -func (_c *mockStream_Context_Call) Run(run func()) *mockStream_Context_Call { - _c.Call.Run(func(args mock.Arguments) { - run() - }) - return _c -} - -func (_c *mockStream_Context_Call) Return(_a0 context.Context) *mockStream_Context_Call { - _c.Call.Return(_a0) - return _c -} - -func (_c *mockStream_Context_Call) RunAndReturn(run func() context.Context) *mockStream_Context_Call { - _c.Call.Return(run) - return _c -} - -// Read provides a mock function with given fields: p -func (_m *mockStream) Read(p []byte) (int, error) { - ret := _m.Called(p) - - if len(ret) == 0 { - panic("no return value specified for Read") - } - - var r0 int - var r1 error - if rf, ok := ret.Get(0).(func([]byte) (int, error)); ok { - return rf(p) - } - if rf, ok := ret.Get(0).(func([]byte) int); ok { - r0 = rf(p) - } else { - r0 = ret.Get(0).(int) - } - - if rf, ok := ret.Get(1).(func([]byte) error); ok { - r1 = rf(p) - } else { - r1 = ret.Error(1) - } - - return r0, r1 -} - -// mockStream_Read_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Read' -type mockStream_Read_Call struct { - *mock.Call -} - -// Read is a helper method to define mock.On call -// - p []byte -func (_e *mockStream_Expecter) Read(p interface{}) *mockStream_Read_Call { - return &mockStream_Read_Call{Call: _e.mock.On("Read", p)} -} - -func (_c *mockStream_Read_Call) Run(run func(p []byte)) *mockStream_Read_Call { - _c.Call.Run(func(args mock.Arguments) { - run(args[0].([]byte)) - }) - return _c -} - -func (_c *mockStream_Read_Call) Return(n int, err error) *mockStream_Read_Call { - _c.Call.Return(n, err) - return _c -} - -func (_c *mockStream_Read_Call) RunAndReturn(run func([]byte) (int, error)) *mockStream_Read_Call { - _c.Call.Return(run) - return _c -} - -// SetDeadline provides a mock function with given fields: t -func (_m *mockStream) SetDeadline(t time.Time) error { - ret := _m.Called(t) - - if len(ret) == 0 { - panic("no return value specified for SetDeadline") - } - - var r0 error - if rf, ok := ret.Get(0).(func(time.Time) error); ok { - r0 = rf(t) - } else { - r0 = ret.Error(0) - } - - return r0 -} - -// mockStream_SetDeadline_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'SetDeadline' -type mockStream_SetDeadline_Call struct { - *mock.Call -} - -// SetDeadline is a helper method to define mock.On call -// - t time.Time -func (_e *mockStream_Expecter) SetDeadline(t interface{}) *mockStream_SetDeadline_Call { - return &mockStream_SetDeadline_Call{Call: _e.mock.On("SetDeadline", t)} -} - -func (_c *mockStream_SetDeadline_Call) Run(run func(t time.Time)) *mockStream_SetDeadline_Call { - _c.Call.Run(func(args mock.Arguments) { - run(args[0].(time.Time)) - }) - return _c -} - -func (_c *mockStream_SetDeadline_Call) Return(_a0 error) *mockStream_SetDeadline_Call { - _c.Call.Return(_a0) - return _c -} - -func (_c *mockStream_SetDeadline_Call) RunAndReturn(run func(time.Time) error) *mockStream_SetDeadline_Call { - _c.Call.Return(run) - return _c -} - -// SetReadDeadline provides a mock function with given fields: t -func (_m *mockStream) SetReadDeadline(t time.Time) error { - ret := _m.Called(t) - - if len(ret) == 0 { - panic("no return value specified for SetReadDeadline") - } - - var r0 error - if rf, ok := ret.Get(0).(func(time.Time) error); ok { - r0 = rf(t) - } else { - r0 = ret.Error(0) - } - - return r0 -} - -// mockStream_SetReadDeadline_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'SetReadDeadline' -type mockStream_SetReadDeadline_Call struct { - *mock.Call -} - -// SetReadDeadline is a helper method to define mock.On call -// - t time.Time -func (_e *mockStream_Expecter) SetReadDeadline(t interface{}) *mockStream_SetReadDeadline_Call { - return &mockStream_SetReadDeadline_Call{Call: _e.mock.On("SetReadDeadline", t)} -} - -func (_c *mockStream_SetReadDeadline_Call) Run(run func(t time.Time)) *mockStream_SetReadDeadline_Call { - _c.Call.Run(func(args mock.Arguments) { - run(args[0].(time.Time)) - }) - return _c -} - -func (_c *mockStream_SetReadDeadline_Call) Return(_a0 error) *mockStream_SetReadDeadline_Call { - _c.Call.Return(_a0) - return _c -} - -func (_c *mockStream_SetReadDeadline_Call) RunAndReturn(run func(time.Time) error) *mockStream_SetReadDeadline_Call { - _c.Call.Return(run) - return _c -} - -// SetWriteDeadline provides a mock function with given fields: t -func (_m *mockStream) SetWriteDeadline(t time.Time) error { - ret := _m.Called(t) - - if len(ret) == 0 { - panic("no return value specified for SetWriteDeadline") - } - - var r0 error - if rf, ok := ret.Get(0).(func(time.Time) error); ok { - r0 = rf(t) - } else { - r0 = ret.Error(0) - } - - return r0 -} - -// mockStream_SetWriteDeadline_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'SetWriteDeadline' -type mockStream_SetWriteDeadline_Call struct { - *mock.Call -} - -// SetWriteDeadline is a helper method to define mock.On call -// - t time.Time -func (_e *mockStream_Expecter) SetWriteDeadline(t interface{}) *mockStream_SetWriteDeadline_Call { - return &mockStream_SetWriteDeadline_Call{Call: _e.mock.On("SetWriteDeadline", t)} -} - -func (_c *mockStream_SetWriteDeadline_Call) Run(run func(t time.Time)) *mockStream_SetWriteDeadline_Call { - _c.Call.Run(func(args mock.Arguments) { - run(args[0].(time.Time)) - }) - return _c -} - -func (_c *mockStream_SetWriteDeadline_Call) Return(_a0 error) *mockStream_SetWriteDeadline_Call { - _c.Call.Return(_a0) - return _c -} - -func (_c *mockStream_SetWriteDeadline_Call) RunAndReturn(run func(time.Time) error) *mockStream_SetWriteDeadline_Call { - _c.Call.Return(run) - return _c -} - -// StreamID provides a mock function with given fields: -func (_m *mockStream) StreamID() qerr.StreamID { - ret := _m.Called() - - if len(ret) == 0 { - panic("no return value specified for StreamID") - } - - var r0 qerr.StreamID - if rf, ok := ret.Get(0).(func() qerr.StreamID); ok { - r0 = rf() - } else { - r0 = ret.Get(0).(qerr.StreamID) - } - - return r0 -} - -// mockStream_StreamID_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'StreamID' -type mockStream_StreamID_Call struct { - *mock.Call -} - -// StreamID is a helper method to define mock.On call -func (_e *mockStream_Expecter) StreamID() *mockStream_StreamID_Call { - return &mockStream_StreamID_Call{Call: _e.mock.On("StreamID")} -} - -func (_c *mockStream_StreamID_Call) Run(run func()) *mockStream_StreamID_Call { - _c.Call.Run(func(args mock.Arguments) { - run() - }) - return _c -} - -func (_c *mockStream_StreamID_Call) Return(_a0 qerr.StreamID) *mockStream_StreamID_Call { - _c.Call.Return(_a0) - return _c -} - -func (_c *mockStream_StreamID_Call) RunAndReturn(run func() qerr.StreamID) *mockStream_StreamID_Call { - _c.Call.Return(run) - return _c -} - -// Write provides a mock function with given fields: p -func (_m *mockStream) Write(p []byte) (int, error) { - ret := _m.Called(p) - - if len(ret) == 0 { - panic("no return value specified for Write") - } - - var r0 int - var r1 error - if rf, ok := ret.Get(0).(func([]byte) (int, error)); ok { - return rf(p) - } - if rf, ok := ret.Get(0).(func([]byte) int); ok { - r0 = rf(p) - } else { - r0 = ret.Get(0).(int) - } - - if rf, ok := ret.Get(1).(func([]byte) error); ok { - r1 = rf(p) - } else { - r1 = ret.Error(1) - } - - return r0, r1 -} - -// mockStream_Write_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Write' -type mockStream_Write_Call struct { - *mock.Call -} - -// Write is a helper method to define mock.On call -// - p []byte -func (_e *mockStream_Expecter) Write(p interface{}) *mockStream_Write_Call { - return &mockStream_Write_Call{Call: _e.mock.On("Write", p)} -} - -func (_c *mockStream_Write_Call) Run(run func(p []byte)) *mockStream_Write_Call { - _c.Call.Run(func(args mock.Arguments) { - run(args[0].([]byte)) - }) - return _c -} - -func (_c *mockStream_Write_Call) Return(n int, err error) *mockStream_Write_Call { - _c.Call.Return(n, err) - return _c -} - -func (_c *mockStream_Write_Call) RunAndReturn(run func([]byte) (int, error)) *mockStream_Write_Call { - _c.Call.Return(run) - return _c -} - -// newMockStream creates a new instance of mockStream. It also registers a testing interface on the mock and a cleanup function to assert the mocks expectations. -// The first argument is typically a *testing.T value. -func newMockStream(t interface { - mock.TestingT - Cleanup(func()) -}) *mockStream { - mock := &mockStream{} - mock.Mock.Test(t) - - t.Cleanup(func() { mock.AssertExpectations(t) }) - - return mock -} diff --git a/extras/sniff/quic.go b/extras/sniff/quic.go new file mode 100644 index 000000000..789646826 --- /dev/null +++ b/extras/sniff/quic.go @@ -0,0 +1,194 @@ +package sniff + +import ( + "bytes" + "cmp" + "crypto/aes" + "crypto/cipher" + "crypto/hkdf" + "crypto/sha256" + "encoding/binary" + "slices" + + "golang.org/x/crypto/cryptobyte" + + "github.com/apernet/quic-go/quicvarint" +) + +// quicVersions are the QUIC versions whose Initial packets we can decrypt (RFC 9001, RFC 9369). +var quicVersions = map[uint32]struct { + salt []byte + initialType uint8 + labelPrefix string +}{ + 0x00000001: { + salt: []byte{0x38, 0x76, 0x2c, 0xf7, 0xf5, 0x59, 0x34, 0xb3, 0x4d, 0x17, 0x9a, 0xe6, 0xa4, 0xc8, 0x0c, 0xad, 0xcc, 0xbb, 0x7f, 0x0a}, + initialType: 0b00, + labelPrefix: "quic ", + }, + 0x6b3343cf: { + salt: []byte{0x0d, 0xed, 0xe3, 0xde, 0xf7, 0x00, 0xa6, 0xdb, 0x81, 0x93, 0x81, 0xbe, 0x6e, 0x26, 0x9d, 0xcb, 0xf9, 0xbd, 0x2e, 0xd9}, + initialType: 0b01, + labelPrefix: "quicv2 ", + }, +} + +// quicCrypto collects the CRYPTO frames in the client Initial packets of a QUIC connection. +// Clients like Chrome split the ClientHello across packets, and shuffle its fragments. +type quicCrypto struct { + dcid []byte + frames []quicCryptoFrame +} + +type quicCryptoFrame struct { + offset uint64 + data []byte +} + +// feed collects the CRYPTO frames of the Initial packets coalesced in a datagram, +// and reports whether it had any. +func (c *quicCrypto) feed(b []byte) bool { + found := false + for len(b) > 0 { + s := cryptobyte.String(b) + var first uint8 + var version uint32 + var dcid, scid, token cryptobyte.String + var length uint64 + if !s.ReadUint8(&first) || first&0x80 == 0 || // Long header + !s.ReadUint32(&version) || + !s.ReadUint8LengthPrefixed(&dcid) || + !s.ReadUint8LengthPrefixed(&scid) { + break + } + v, ok := quicVersions[version] + if !ok { + break + } + initial := (first>>4)&0b11 == v.initialType + if initial && !readVarintPrefixed(&s, &token) { + break + } + if !readVarint(&s, &length) || uint64(len(s)) < length { + break + } + pnOffset := len(b) - len(s) + packet := b[:pnOffset+int(length)] + b = b[len(packet):] + if !initial || c.dcid != nil && !bytes.Equal(dcid, c.dcid) { + continue + } + payload := decryptInitial(packet, pnOffset, dcid, v.salt, v.labelPrefix) + if payload == nil { + continue + } + c.dcid = dcid + c.frames = appendCryptoFrames(c.frames, payload) + found = true + } + return found +} + +// stream returns the contiguous start of the CRYPTO stream. +func (c *quicCrypto) stream() []byte { + slices.SortFunc(c.frames, func(a, b quicCryptoFrame) int { return cmp.Compare(a.offset, b.offset) }) + var stream []byte + for _, f := range c.frames { + n := uint64(len(stream)) + if f.offset > n { + break + } + if f.offset+uint64(len(f.data)) > n { + stream = append(stream, f.data[n-f.offset:]...) + } + } + return stream +} + +// decryptInitial removes the protection of a client Initial packet (RFC 9001, Section 5), +// and returns its payload, or nil if it can't. +func decryptInitial(packet []byte, pnOffset int, dcid, salt []byte, labelPrefix string) []byte { + if len(packet) < pnOffset+4+16 { + return nil + } + secret, _ := hkdf.Extract(sha256.New, dcid, salt) + secret = hkdfExpandLabel(secret, "client in", sha256.Size) + + // Header protection + hp, _ := aes.NewCipher(hkdfExpandLabel(secret, labelPrefix+"hp", 16)) + mask := make([]byte, 16) + hp.Encrypt(mask, packet[pnOffset+4:pnOffset+4+16]) + header := slices.Clone(packet[:pnOffset+4]) + header[0] ^= mask[0] & 0x0f + pnLen := int(header[0]&0x03) + 1 + header = header[:pnOffset+pnLen] + var pn uint64 + for i := range pnLen { + header[pnOffset+i] ^= mask[1+i] + pn = pn<<8 | uint64(header[pnOffset+i]) + } + + // Packet protection. The packet number is the truncated one, as nothing has been + // received on the connection that could make it larger. + block, _ := aes.NewCipher(hkdfExpandLabel(secret, labelPrefix+"key", 16)) + aead, _ := cipher.NewGCM(block) + nonce := hkdfExpandLabel(secret, labelPrefix+"iv", aead.NonceSize()) + binary.BigEndian.PutUint64(nonce[4:], binary.BigEndian.Uint64(nonce[4:])^pn) + payload, err := aead.Open(nil, nonce, packet[pnOffset+pnLen:], header) + if err != nil { + return nil + } + return payload +} + +// appendCryptoFrames appends the CRYPTO frames in a decrypted Initial packet payload. +// It stops at the first frame that isn't PADDING, PING or CRYPTO, as a client doesn't +// send any other before it hears from the server. +func appendCryptoFrames(frames []quicCryptoFrame, payload []byte) []quicCryptoFrame { + s := cryptobyte.String(payload) + for !s.Empty() { + var typ, offset uint64 + var data cryptobyte.String + if !readVarint(&s, &typ) { + break + } + switch typ { + case 0x00, 0x01: // PADDING, PING + case 0x06: // CRYPTO + if !readVarint(&s, &offset) || !readVarintPrefixed(&s, &data) { + return frames + } + frames = append(frames, quicCryptoFrame{offset, data}) + default: + return frames + } + } + return frames +} + +func readVarint(s *cryptobyte.String, v *uint64) bool { + n, l, err := quicvarint.Parse(*s) + if err != nil { + return false + } + *v = n + return s.Skip(l) +} + +func readVarintPrefixed(s, out *cryptobyte.String) bool { + var n uint64 + if !readVarint(s, &n) || uint64(len(*s)) < n { + return false + } + return s.ReadBytes((*[]byte)(out), int(n)) +} + +// hkdfExpandLabel implements HKDF-Expand-Label from RFC 8446, Section 7.1, with an empty context. +func hkdfExpandLabel(secret []byte, label string, length int) []byte { + info := []byte{byte(length >> 8), byte(length), byte(len("tls13 ") + len(label))} + info = append(info, "tls13 "...) + info = append(info, label...) + info = append(info, 0) + out, _ := hkdf.Expand(sha256.New, secret, string(info), length) + return out +} diff --git a/extras/sniff/sniff.go b/extras/sniff/sniff.go index ff5173df0..30202b9bf 100644 --- a/extras/sniff/sniff.go +++ b/extras/sniff/sniff.go @@ -9,16 +9,14 @@ import ( "strings" "time" - utls "github.com/refraction-networking/utls" - "github.com/apernet/hysteria/core/v2/server" - quicInternal "github.com/apernet/hysteria/extras/v2/sniff/internal/quic" "github.com/apernet/hysteria/extras/v2/utils" ) const ( - sniffDefaultTimeout = 4 * time.Second - sniffMaxHTTPHeaderBytes = 256 * 1024 + sniffDefaultTimeout = 4 * time.Second + sniffMaxTCPBytes = 64 * 1024 + sniffMaxUDPPackets = 8 ) var _ server.RequestHook = (*Sniffer)(nil) @@ -35,35 +33,6 @@ type Sniffer struct { UDPPorts utils.PortUnion } -func (h *Sniffer) isDomain(addr string) bool { - host, _, err := net.SplitHostPort(addr) - if err != nil { - return false - } - return net.ParseIP(host) == nil -} - -func (h *Sniffer) isHTTP(buf []byte) bool { - if len(buf) < 3 { - return false - } - // First 3 bytes should be English letters (whatever HTTP method) - for _, b := range buf[:3] { - if (b < 'A' || b > 'Z') && (b < 'a' || b > 'z') { - return false - } - } - return true -} - -func (h *Sniffer) isTLS(buf []byte) bool { - if len(buf) < 3 { - return false - } - return buf[0] >= 0x16 && buf[0] <= 0x17 && - buf[1] == 0x03 && buf[2] <= 0x09 -} - func (h *Sniffer) Check(isUDP bool, reqAddr string) bool { // @ means it's internal (e.g. speed test) if strings.HasPrefix(reqAddr, "@") { @@ -88,112 +57,108 @@ func (h *Sniffer) Check(isUDP bool, reqAddr string) bool { } } +// TCP reads from the stream until it has an HTTP request header or a TLS ClientHello, +// or can tell it's neither, and returns everything read. func (h *Sniffer) TCP(stream server.HyStream, reqAddr *string) ([]byte, error) { - var err error - if h.Timeout == 0 { - err = stream.SetReadDeadline(time.Now().Add(sniffDefaultTimeout)) - } else { - err = stream.SetReadDeadline(time.Now().Add(h.Timeout)) + timeout := h.Timeout + if timeout == 0 { + timeout = sniffDefaultTimeout } - if err != nil { + if err := stream.SetReadDeadline(time.Now().Add(timeout)); err != nil { return nil, err } // Make sure to reset the deadline after sniffing defer stream.SetReadDeadline(time.Time{}) - // Read 3 bytes to determine the protocol - pre := make([]byte, 3) - n, err := io.ReadFull(stream, pre) - if err != nil { - // Not enough within the timeout, just return what we have - return pre[:n], nil - } - if h.isHTTP(pre) { - // HTTP - tr := &teeReader{Stream: stream, Pre: pre} - req, _ := http.ReadRequest(bufio.NewReader(io.LimitReader(tr, sniffMaxHTTPHeaderBytes))) - if req != nil && req.Host != "" { - // req.Host can be host:port, in which case we need to extract the host part - host, _, err := net.SplitHostPort(req.Host) - if err != nil { - // No port, just use the whole string - host = req.Host - } - _, port, err := net.SplitHostPort(*reqAddr) - if err != nil { - return nil, err - } - *reqAddr = net.JoinHostPort(host, port) + + rec := &recorder{r: io.LimitReader(stream, sniffMaxTCPBytes)} + rewrite(reqAddr, sniffStream(bufio.NewReader(rec))) + return rec.buf, nil +} + +// UDP looks for a QUIC ClientHello, which can span multiple packets. +func (h *Sniffer) UDP(packets [][]byte, reqAddr *string) (bool, error) { + var c quicCrypto + for i, p := range packets { + if !c.feed(p) && i == 0 { + // Not QUIC + return true, nil } - return tr.Buffer(), nil - } else if h.isTLS(pre) { - // TLS - // Need to read 2 more bytes (content length) - pre = append(pre, make([]byte, 2)...) - n, err = io.ReadFull(stream, pre[3:]) + } + hello, more := clientHello(c.stream()) + if more && len(packets) < sniffMaxUDPPackets { + return false, nil + } + rewrite(reqAddr, serverName(hello)) + return true, nil +} + +// sniffStream returns the domain in an HTTP request or a TLS ClientHello read from r. +func sniffStream(r *bufio.Reader) string { + b, err := r.Peek(1) + switch { + case err != nil: + return "" + case b[0] == 0x16: // TLS handshake record + return serverName(readClientHello(r)) + case isHTTP(r): + req, err := http.ReadRequest(r) if err != nil { - // Not enough within the timeout, just return what we have - return pre[:3+n], nil + return "" } - contentLength := int(pre[3])<<8 | int(pre[4]) - pre = append(pre, make([]byte, contentLength)...) - n, err = io.ReadFull(stream, pre[5:]) + return req.Host + } + return "" +} + +// isHTTP reports whether r starts with an HTTP method, like GET or M-SEARCH, and a space. +func isHTTP(r *bufio.Reader) bool { + for i := 1; ; i++ { + b, err := r.Peek(i) if err != nil { - // Not enough within the timeout, just return what we have - return pre[:5+n], nil + return false } - clientHello := utls.UnmarshalClientHello(pre[5:]) - if clientHello != nil && clientHello.ServerName != "" { - _, port, err := net.SplitHostPort(*reqAddr) - if err != nil { - return nil, err - } - *reqAddr = net.JoinHostPort(clientHello.ServerName, port) + switch c := b[i-1]; { + case c == ' ': + return i > 1 + case (c < 'A' || c > 'Z') && c != '-' && c != '_': + return false } - return pre, nil - } else { - // Unrecognized protocol, just return what we have - return pre, nil } } -func (h *Sniffer) UDP(data []byte, reqAddr *string) error { - pl, err := quicInternal.ReadCryptoPayload(data) - if err != nil || len(pl) < 4 || pl[0] != 0x01 { - // Unrecognized protocol, incomplete payload or not a client hello - return nil +// rewrite replaces the host of reqAddr with a sniffed domain, keeping the port. +func rewrite(reqAddr *string, host string) { + if h, _, err := net.SplitHostPort(host); err == nil { + // HTTP Host can have a port + host = h } - clientHello := utls.UnmarshalClientHello(pl) - if clientHello != nil && clientHello.ServerName != "" { - _, port, err := net.SplitHostPort(*reqAddr) - if err != nil { - return err - } - *reqAddr = net.JoinHostPort(clientHello.ServerName, port) + _, port, err := net.SplitHostPort(*reqAddr) + if err != nil || !isDomain(host) { + return } - return nil + *reqAddr = net.JoinHostPort(host, port) } -type teeReader struct { - Stream server.HyStream - Pre []byte +func isDomain(s string) bool { + if s == "" || len(s) > 253 || net.ParseIP(s) != nil { + return false + } + for _, c := range []byte(s) { + if !('a' <= c && c <= 'z' || 'A' <= c && c <= 'Z' || '0' <= c && c <= '9' || c == '-' || c == '.' || c == '_') { + return false + } + } + return true +} +// recorder records everything read through it. +type recorder struct { + r io.Reader buf []byte } -func (c *teeReader) Read(b []byte) (n int, err error) { - if len(c.Pre) > 0 { - n = copy(b, c.Pre) - c.Pre = c.Pre[n:] - c.buf = append(c.buf, b[:n]...) - return n, nil - } - n, err = c.Stream.Read(b) - if n > 0 { - c.buf = append(c.buf, b[:n]...) - } +func (r *recorder) Read(p []byte) (int, error) { + n, err := r.r.Read(p) + r.buf = append(r.buf, p[:n]...) return n, err } - -func (c *teeReader) Buffer() []byte { - return append(c.Pre, c.buf...) -} diff --git a/extras/sniff/sniff_test.go b/extras/sniff/sniff_test.go index 445660bb0..5461ea65d 100644 --- a/extras/sniff/sniff_test.go +++ b/extras/sniff/sniff_test.go @@ -1,15 +1,21 @@ package sniff import ( - "encoding/base64" + "context" + "crypto/tls" "io" + "net" + "os" + "slices" + "strings" "testing" "time" - "github.com/apernet/hysteria/extras/v2/utils" - + "github.com/apernet/quic-go" "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/mock" + "github.com/stretchr/testify/require" + + "github.com/apernet/hysteria/extras/v2/utils" ) func TestSnifferCheck(t *testing.T) { @@ -36,112 +42,361 @@ func TestSnifferCheck(t *testing.T) { } func TestSnifferTCP(t *testing.T) { - sniffer := &Sniffer{ - Timeout: 1 * time.Second, - RewriteDomain: false, + goHello := goTLSClientHello(t, "go.sniff.test") + chromeHello := readTestdata(t, "tls-chrome153.bin") + + tests := []struct { + name string + data []byte + chunk int // Write the data in chunks of this size, 0 = all at once + reqAddr string + wantAddr string + }{ + { + name: "HTTP", + data: []byte("POST /hello HTTP/1.1\r\nHost: example.com\r\nContent-Length: 27\r\n\r\nparam1=value1¶m2=value2"), + reqAddr: "111.111.111.111:80", + wantAddr: "example.com:80", + }, + { + name: "HTTP host with port", + data: []byte("GET / HTTP/1.1\r\nHost: example.com:8080\r\nAccept: */*\r\n\r\n"), + reqAddr: "222.222.222.222:10086", + wantAddr: "example.com:10086", + }, + { + name: "HTTP absolute URI", + data: []byte("GET http://absolute.example.com/x HTTP/1.1\r\n\r\n"), + reqAddr: "1.2.3.4:80", + wantAddr: "absolute.example.com:80", + }, + { + name: "HTTP byte by byte", + data: []byte("GET / HTTP/1.1\r\nHost: slow.example.com\r\n\r\n"), + chunk: 1, + reqAddr: "1.2.3.4:80", + wantAddr: "slow.example.com:80", + }, + { + name: "HTTP Chrome 153", + data: readTestdata(t, "http-chrome153.txt"), + reqAddr: "1.2.3.4:80", + wantAddr: "chrome.sniff.test:80", + }, + { + name: "HTTP IPv6 host", + data: []byte("GET / HTTP/1.1\r\nHost: [2001:db8::1]:8080\r\n\r\n"), + reqAddr: "1.2.3.4:80", + wantAddr: "1.2.3.4:80", + }, + { + name: "TLS Go", + data: goHello, + reqAddr: "1.2.3.4:443", + wantAddr: "go.sniff.test:443", + }, + { + name: "TLS Chrome 153", + data: chromeHello, + reqAddr: "1.2.3.4:443", + wantAddr: "chrome.sniff.test:443", + }, + { + name: "TLS Firefox 153 ESR", + data: readTestdata(t, "tls-firefox153esr.bin"), + reqAddr: "1.2.3.4:443", + wantAddr: "firefox.sniff.test:443", + }, + { + name: "TLS curl 8.18 (OpenSSL 3.5)", + data: readTestdata(t, "tls-curl8.18-openssl3.5.bin"), + reqAddr: "1.2.3.4:443", + wantAddr: "curl.sniff.test:443", + }, + { + name: "TLS byte by byte", + data: chromeHello, + chunk: 1, + reqAddr: "1.2.3.4:443", + wantAddr: "chrome.sniff.test:443", + }, + { + name: "TLS ClientHello fragmented across records", + data: fragmentTLSRecords(chromeHello, 100), + chunk: 333, + reqAddr: "1.2.3.4:443", + wantAddr: "chrome.sniff.test:443", + }, + { + name: "TLS ClientHello followed by other records", + data: append(slices.Clone(goHello), 0x14, 0x03, 0x03, 0x00, 0x01, 0x01), + reqAddr: "1.2.3.4:443", + wantAddr: "go.sniff.test:443", + }, + { + name: "TLS ClientHello interrupted by another record", + data: append(fragmentTLSRecords(goHello, 100)[:105], 0x14, 0x03, 0x03, 0x00, 0x01, 0x01), + reqAddr: "1.2.3.4:443", + wantAddr: "1.2.3.4:443", + }, + { + name: "Unrecognized text", + data: []byte("Wait It's All Ohio? Always Has Been."), + reqAddr: "123.123.123.123:123", + wantAddr: "123.123.123.123:123", + }, + { + name: "Unrecognized SSH", + data: []byte("SSH-2.0-OpenSSH_9.6\r\n"), + reqAddr: "123.123.123.123:22", + wantAddr: "123.123.123.123:22", + }, + { + name: "Unrecognized binary", + data: []byte("\x01\x02\x03\x04\x05\x06\x07\x08\x09\x0a"), + reqAddr: "45.45.45.45:45", + wantAddr: "45.45.45.45:45", + }, } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + // The stream stays open, so the sniffer must stop as soon as it has seen enough, + // or it times out and doesn't rewrite anything. + stream, client := newPipeStream() + defer client.Close() + go writeChunks(client, tt.data, tt.chunk) - buf := &[]byte{} - - // Test HTTP - *buf = []byte("POST /hello HTTP/1.1\r\n" + - "Host: example.com\r\n" + - "User-Agent: mamamiya\r\n" + - "Content-Length: 27\r\n" + - "Connection: keep-alive\r\n\r\n" + - "param1=value1¶m2=value2") - index := 0 - stream := &mockStream{} - stream.EXPECT().SetReadDeadline(mock.Anything).Return(nil) - stream.EXPECT().Read(mock.Anything).RunAndReturn(func(bs []byte) (int, error) { - if index < len(*buf) { - n := copy(bs, (*buf)[index:]) - index += n - return n, nil - } else { - return 0, io.EOF - } - }) + reqAddr := tt.reqAddr + putback, err := (&Sniffer{Timeout: 3 * time.Second}).TCP(stream, &reqAddr) + require.NoError(t, err) + assert.Equal(t, tt.wantAddr, reqAddr) - // Rewrite IP to domain - reqAddr := "111.111.111.111:80" - putback, err := sniffer.TCP(stream, &reqAddr) - assert.NoError(t, err) - assert.Equal(t, *buf, putback) - assert.Equal(t, "example.com:80", reqAddr) - - // Test HTTP with Host as host:port - *buf = []byte("GET / HTTP/1.1\r\n" + - "Host: example.com:8080\r\n" + - "User-Agent: test-agent\r\n" + - "Accept: */*\r\n\r\n") - index = 0 - reqAddr = "222.222.222.222:10086" - putback, err = sniffer.TCP(stream, &reqAddr) - assert.NoError(t, err) - assert.Equal(t, *buf, putback) - assert.Equal(t, "example.com:10086", reqAddr) + // Nothing is lost: what's put back and what's left make up the whole data + rest := make([]byte, len(tt.data)-len(putback)) + _, err = io.ReadFull(stream, rest) + require.NoError(t, err) + assert.Equal(t, tt.data, append(putback, rest...)) + }) + } +} - // Test TLS - *buf, err = base64.StdEncoding.DecodeString("FgMBARcBAAETAwPJL2jlt1OAo+Rslkjv/aqKiTthKMaCKg2Gvd+uALDbDCDdY+UIk8ouadEB9fC3j52Y1i7SJZqGIgBRIS6kKieYrAAoEwITAcAswCvAMMAvwCTAI8AowCfACsAJwBTAEwCdAJwAPQA8ADUALwEAAKIAAAAOAAwAAAlpcGluZm8uaW8ABQAFAQAAAAAAKwAJCAMEAwMDAgMBAA0AGgAYCAQIBQgGBAEFAQIBBAMFAwIDAgIGAQYDACMAAAAKAAgABgAdABcAGAAQAAsACQhodHRwLzEuMQAzACYAJAAdACBguQbqNJNyamYxYcrBFpBP7pWv5TgZsP9gwGtMYNKVBQAxAAAAFwAA/wEAAQAALQACAQE=") - assert.NoError(t, err) - index = 0 - reqAddr = "222.222.222.222:443" - putback, err = sniffer.TCP(stream, &reqAddr) - assert.NoError(t, err) - assert.Equal(t, *buf, putback) - assert.Equal(t, "ipinfo.io:443", reqAddr) - - // Test unrecognized 1 - *buf = []byte("Wait It's All Ohio? Always Has Been.") - index = 0 - reqAddr = "123.123.123.123:123" - putback, err = sniffer.TCP(stream, &reqAddr) - assert.NoError(t, err) - assert.Equal(t, *buf, putback) - assert.Equal(t, "123.123.123.123:123", reqAddr) - - // Test unrecognized 2 - *buf = []byte("\x01\x02\x03\x04\x05\x06\x07\x08\x09\x0a") - index = 0 - reqAddr = "45.45.45.45:45" - putback, err = sniffer.TCP(stream, &reqAddr) +func TestSnifferTCPTimeout(t *testing.T) { + stream, client := newPipeStream() + defer client.Close() + go client.Write([]byte("GET / HTTP/1.1\r\nHost: example.com\r\n")) + + reqAddr := "66.66.66.66:80" + start := time.Now() + putback, err := (&Sniffer{Timeout: 500 * time.Millisecond}).TCP(stream, &reqAddr) assert.NoError(t, err) - assert.Equal(t, []byte("\x01\x02\x03"), putback) - assert.Equal(t, "45.45.45.45:45", reqAddr) - - // Test timeout - blockStream := &mockStream{} - blockStream.EXPECT().SetReadDeadline(mock.Anything).Return(nil) - blockStream.EXPECT().Read(mock.Anything).RunAndReturn(func(bs []byte) (int, error) { - time.Sleep(2 * time.Second) - return 0, io.EOF - }) - reqAddr = "66.66.66.66:66" - putback, err = sniffer.TCP(blockStream, &reqAddr) + assert.Less(t, time.Since(start), 2*time.Second) + assert.Equal(t, []byte("GET / HTTP/1.1\r\nHost: example.com\r\n"), putback) + assert.Equal(t, "66.66.66.66:80", reqAddr) +} + +func TestSnifferTCPMaxBytes(t *testing.T) { + stream, client := newPipeStream() + defer client.Close() + data := []byte("GET / HTTP/1.1\r\nX-Big: " + strings.Repeat("a", 2*sniffMaxTCPBytes) + "\r\nHost: example.com\r\n\r\n") + go client.Write(data) + + reqAddr := "66.66.66.66:80" + putback, err := (&Sniffer{Timeout: 3 * time.Second}).TCP(stream, &reqAddr) assert.NoError(t, err) - assert.Equal(t, []byte{}, putback) - assert.Equal(t, "66.66.66.66:66", reqAddr) + assert.Equal(t, data[:sniffMaxTCPBytes], putback) + assert.Equal(t, "66.66.66.66:80", reqAddr) } func TestSnifferUDP(t *testing.T) { - sniffer := &Sniffer{ - Timeout: 1 * time.Second, - RewriteDomain: false, + chrome := readTestdataPackets(t, "quic-chrome153", 3) + firefox := readTestdataPackets(t, "quic-firefox153esr", 2) + curl := readTestdataPackets(t, "quic-curl8.14-openssl3.5", 2) + + tests := []struct { + name string + packets [][]byte + wantAddr string + wantN int // Packets needed before the sniffer is done + }{ + // Chrome shuffles the ClientHello fragments across two packets, and retransmits them split differently + {"Chrome 153", chrome, "chrome.sniff.test:443", 2}, + {"Chrome 153 reordered", [][]byte{chrome[1], chrome[0]}, "chrome.sniff.test:443", 2}, + {"Chrome 153 retransmitted", [][]byte{chrome[1], chrome[2]}, "chrome.sniff.test:443", 2}, + {"Firefox 153 ESR", firefox, "firefox.sniff.test:443", 2}, + {"Firefox 153 ESR reordered", [][]byte{firefox[1], firefox[0]}, "firefox.sniff.test:443", 2}, + {"curl 8.14 (OpenSSL 3.5)", curl, "curl.sniff.test:443", 2}, + // quiche retransmits the first packet before sending the second + {"quiche", readTestdataPackets(t, "quic-quiche", 3), "quiche.sniff.test:443", 3}, + {"ngtcp2 1.11", readTestdataPackets(t, "quic-ngtcp2-1.11", 1), "ngtcp2.sniff.test:443", 1}, + {"aioquic 1.2", readTestdataPackets(t, "quic-aioquic1.2", 1), "aioquic.sniff.test:443", 1}, + {"Not QUIC", [][]byte{[]byte("oh my sweet summer child")}, "1.2.3.4:443", 1}, + {"Unsupported version", [][]byte{append([]byte{0xc0, 0xff, 0x00, 0x00, 0x1d}, chrome[0][5:]...)}, "1.2.3.4:443", 1}, + {"Other connections ignored", [][]byte{chrome[0], firefox[1], curl[1], chrome[1]}, "chrome.sniff.test:443", 4}, + {"Gives up", slices.Repeat([][]byte{chrome[0]}, sniffMaxUDPPackets), "1.2.3.4:443", sniffMaxUDPPackets}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + assertSniffUDP(t, tt.packets, tt.wantAddr, tt.wantN) + }) } +} - // Test QUIC - reqAddr := "2.3.4.5:443" - pkt, err := base64.StdEncoding.DecodeString("ygAAAAEIwugWgPS7ulYAAES8hY891uwgGE9GG4CPOLd+nsDe28raso24lCSFmlFwYQG1uF39ikbL13/R9ZTghYmTl+jEbr6F9TxxRiOgpTmKRmh6aKZiIiVfy5pVRckovaI8lq0WRoW9xoFNTyYtQP8TVJ3bLCK+zUqpquEQSyWf7CE43ywayyMpE9UlIoPXFWCoopXLM1SvzdQ+17P51N9KR7m4emti4DWWTBLMQOvrwd2HEEkbiZdRO1wf6ZXJlIat5dN0R/6uod60OFPO+u+awvq67MoMReC7+5I/xWI+xx6o4JpnZNn6YPG8Gqi8hS6doNcAAdtD8h5eMLuHCCgkpX3QVjjfWtcOhtw9xKjU43HhUPwzUTv+JDLgwuTQCTmlfYlb3B+pk4b2I9si0tJ0SBuYaZ2VQPtZbj2hpGXw3gn11pbN8xsbKkQL50+Scd4dGJxWQlGaJHeaU5WOCkxLXc635z8m5XO/CBHVYPGp4pfwfwNUgbe5WF+3MaUIlDB8dMfsnrO0BmZPo379jVx0SFLTAiS8wAdHib1WNEY8qKYnTWuiyxYg1GZEhJt0nXmI+8f0eJq42DgHBWC+Rf5rRBr/Sf25o3mFAmTUaul0Woo9/CIrpT73B63N91xd9A77i4ru995YG8l9Hen+eLtpDU9Q9376nwMDYBzeYG9U/Rn0Urbm6q4hmAgV/xlNJ2rAyDS+yLnwqD6I0PRy8bZJEttcidb/SkOyrpgMiAzWeT+SO+c/k+Y8H0UTRa05faZUrhuUaym9wAcaIVRA6nFI+fejfjVp+7afFv+kWn3vCqQEij+CRHuxkltrixZMD2rfYj6NUW7TTYBtPRtuV/V0ZIDjRR26vr4K+0D84+l3c0mA/l6nmpP5kkco3nmpdjtQN6sGXL7+5o0nnsftX5d6/n5mLyEpP+AEDl1zk3iqkS62RsITwql6DMMoGbSDdUpMclCIeM0vlo3CkxGMO7QA9ruVeNddkL3EWMivl+uxO43sXEEqYQHVl4N75y63t05GOf7/gm9Kb/BJ8MpG9ViEkVYaskQCzi3D8bVpzo8FfTj8te8B6c3ikc/cm7r8k0ZcZpr+YiLGDYq+0ilHxpqJfmq8dPkSvxdzLcUSvy7+LMQ/TTobRSF7L4JhtDKck0+00vl9H35Tkh9N+MsVtpKdWyoqZ4XaK2Nx1M6AieczXpdFc0y7lYPoUfF4IeW8WzeVUclol5ElYjkyFz/lDOGAe1bF2g5AYaGWCPiGleVZknNdD5ihB8W8Mfkt1pEwq2S97AHrppqkf/VoIfZzeqH8wUFw8fDDrZIpnoa0rW7HfwIQaqJhPCyB9Z6TVbV4x9UWmaHfVAcinCK/7o10dtaj3rvEqcUC/iPceGq3Tqv/p9GGNJ+Ci2JBjXqNxYr893Llk75VdPD9pM6y1SM0P80oXNy32VMtafkFFST8GpvvqWcxUJ93kzaY8RmU1g3XFOImSU2utU6+FUQ2Pn5uLwcfT2cTYfTpPGh+WXjSbZ6trqdEMEsLHybuPo2UN4WpVLXVQma3kSaHQggcLlEip8GhEUAy/xCb2eKqhI4HkDpDjwDnDVKufWlnRaOHf58cc8Woi+WT8JTOkHC+nBEG6fKRPHDG08U5yayIQIjI") - assert.NoError(t, err) - err = sniffer.UDP(pkt, &reqAddr) - assert.NoError(t, err) - assert.Equal(t, "www.notion.so:443", reqAddr) +// TestSnifferUDPQUICGo sniffs the first flight of live quic-go clients. +func TestSnifferUDPQUICGo(t *testing.T) { + tests := []struct { + name string + conf *quic.Config + }{ + {"v1", &quic.Config{}}, + {"v2", &quic.Config{Versions: []quic.Version{quic.Version2}}}, + {"Chrome parrot", &quic.Config{ChromeParrot: true}}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + packets := quicGoFirstFlight(t, tt.conf, "quic-go.sniff.test", 3) + assertSniffUDP(t, packets, "quic-go.sniff.test:443", 2) + }) + } +} - // Test unrecognized - pkt = []byte("oh my sweet summer child") - reqAddr = "90.90.90.90:90" - err = sniffer.UDP(pkt, &reqAddr) - assert.NoError(t, err) - assert.Equal(t, "90.90.90.90:90", reqAddr) +func TestRewrite(t *testing.T) { + tests := []struct { + host string + want string + }{ + {"example.com", "example.com:443"}, + {"example.com:8443", "example.com:443"}, + {"under_score.example.com.", "under_score.example.com.:443"}, + {"", "1.2.3.4:443"}, + {"5.6.7.8", "1.2.3.4:443"}, + {"[2001:db8::1]:80", "1.2.3.4:443"}, + {"evil.com/path", "1.2.3.4:443"}, + {"evil.com\x00", "1.2.3.4:443"}, + {strings.Repeat("a", 254), "1.2.3.4:443"}, + } + for _, tt := range tests { + reqAddr := "1.2.3.4:443" + rewrite(&reqAddr, tt.host) + assert.Equal(t, tt.want, reqAddr, "host %q", tt.host) + } +} + +// assertSniffUDP feeds packets to the sniffer one by one, and checks that it's done +// after exactly n packets, with the address rewritten to want. +func assertSniffUDP(t *testing.T, packets [][]byte, want string, n int) { + t.Helper() + sniffer := &Sniffer{} + for i := 1; i <= n; i++ { + held := make([][]byte, i) + for j := range held { + held[j] = slices.Clone(packets[j]) + } + reqAddr := "1.2.3.4:443" + done, err := sniffer.UDP(held, &reqAddr) + require.NoError(t, err) + if i < n { + require.False(t, done, "done after %d packets", i) + require.Equal(t, "1.2.3.4:443", reqAddr) + continue + } + require.True(t, done, "not done after %d packets", i) + require.Equal(t, want, reqAddr) + for j := range held { + require.Equal(t, packets[j], held[j], "packet %d was modified", j) + } + } +} + +// quicGoFirstFlight returns the first n datagrams a quic-go client sends. +func quicGoFirstFlight(t *testing.T, conf *quic.Config, serverName string, n int) [][]byte { + t.Helper() + ln, err := net.ListenPacket("udp", "127.0.0.1:0") + require.NoError(t, err) + defer ln.Close() + + ctx, cancel := context.WithCancel(context.Background()) + dialDone := make(chan struct{}) + defer func() { + cancel() + <-dialDone + }() + go func() { + defer close(dialDone) + tlsConf := &tls.Config{ServerName: serverName, NextProtos: []string{"h3"}} + _, _ = quic.DialAddr(ctx, ln.LocalAddr().String(), tlsConf, conf) + }() + + var packets [][]byte + buf := make([]byte, 2048) + for len(packets) < n { + require.NoError(t, ln.SetReadDeadline(time.Now().Add(5*time.Second))) + l, _, err := ln.ReadFrom(buf) + require.NoError(t, err) + packets = append(packets, slices.Clone(buf[:l])) + } + return packets +} + +// goTLSClientHello returns the TLS records carrying a crypto/tls ClientHello. +func goTLSClientHello(t *testing.T, serverName string) []byte { + t.Helper() + client, server := net.Pipe() + defer server.Close() + go func() { + _ = tls.Client(client, &tls.Config{ServerName: serverName}).Handshake() + client.Close() + }() + buf := make([]byte, 16384) + n, err := server.Read(buf) + require.NoError(t, err) + return buf[:n] +} + +// fragmentTLSRecords splits the handshake messages in a TLS record into records of size n. +func fragmentTLSRecords(record []byte, n int) []byte { + var out []byte + for p := range slices.Chunk(record[5:], n) { + out = append(out, 0x16, 0x03, 0x01, byte(len(p)>>8), byte(len(p))) + out = append(out, p...) + } + return out +} + +func readTestdata(t *testing.T, name string) []byte { + t.Helper() + b, err := os.ReadFile("testdata/" + name) + require.NoError(t, err) + return b +} + +func readTestdataPackets(t *testing.T, prefix string, n int) [][]byte { + t.Helper() + packets := make([][]byte, n) + for i := range packets { + packets[i] = readTestdata(t, prefix+"-"+string(rune('0'+i))+".bin") + } + return packets +} + +// pipeStream is a server.HyStream backed by a net.Pipe. +type pipeStream struct { + net.Conn +} + +func (pipeStream) StreamID() quic.StreamID { return 0 } + +func newPipeStream() (pipeStream, net.Conn) { + s, c := net.Pipe() + return pipeStream{s}, c +} + +func writeChunks(w io.Writer, data []byte, chunk int) { + if chunk == 0 { + chunk = len(data) + } + for p := range slices.Chunk(data, chunk) { + if _, err := w.Write(p); err != nil { + return + } + } } diff --git a/extras/sniff/testdata/http-chrome153.txt b/extras/sniff/testdata/http-chrome153.txt new file mode 100644 index 000000000..e229bca86 --- /dev/null +++ b/extras/sniff/testdata/http-chrome153.txt @@ -0,0 +1,9 @@ +GET / HTTP/1.1 +Host: chrome.sniff.test +Connection: keep-alive +Upgrade-Insecure-Requests: 1 +User-Agent: Mozilla/5.0 (X11; Linux x86_64) AppleWebKit/537.36 (KHTML, like Gecko) HeadlessChrome/153.0.0.0 Safari/537.36 +Accept: text/html,application/xhtml+xml,application/xml;q=0.9,image/avif,image/webp,image/apng,*/*;q=0.8,application/signed-exchange;v=b3;q=0.7 +Accept-Encoding: gzip, deflate +Accept-Language: en-US,en;q=0.9 + diff --git a/extras/sniff/testdata/quic-aioquic1.2-0.bin b/extras/sniff/testdata/quic-aioquic1.2-0.bin new file mode 100644 index 000000000..f255ed9ff Binary files /dev/null and b/extras/sniff/testdata/quic-aioquic1.2-0.bin differ diff --git a/extras/sniff/testdata/quic-chrome153-0.bin b/extras/sniff/testdata/quic-chrome153-0.bin new file mode 100644 index 000000000..b5e2d514c Binary files /dev/null and b/extras/sniff/testdata/quic-chrome153-0.bin differ diff --git a/extras/sniff/testdata/quic-chrome153-1.bin b/extras/sniff/testdata/quic-chrome153-1.bin new file mode 100644 index 000000000..f7f8dfdb8 Binary files /dev/null and b/extras/sniff/testdata/quic-chrome153-1.bin differ diff --git a/extras/sniff/testdata/quic-chrome153-2.bin b/extras/sniff/testdata/quic-chrome153-2.bin new file mode 100644 index 000000000..c984edf5c Binary files /dev/null and b/extras/sniff/testdata/quic-chrome153-2.bin differ diff --git a/extras/sniff/testdata/quic-curl8.14-openssl3.5-0.bin b/extras/sniff/testdata/quic-curl8.14-openssl3.5-0.bin new file mode 100644 index 000000000..669abac1a Binary files /dev/null and b/extras/sniff/testdata/quic-curl8.14-openssl3.5-0.bin differ diff --git a/extras/sniff/testdata/quic-curl8.14-openssl3.5-1.bin b/extras/sniff/testdata/quic-curl8.14-openssl3.5-1.bin new file mode 100644 index 000000000..dfd01bc41 Binary files /dev/null and b/extras/sniff/testdata/quic-curl8.14-openssl3.5-1.bin differ diff --git a/extras/sniff/testdata/quic-firefox153esr-0.bin b/extras/sniff/testdata/quic-firefox153esr-0.bin new file mode 100644 index 000000000..a6698b18c Binary files /dev/null and b/extras/sniff/testdata/quic-firefox153esr-0.bin differ diff --git a/extras/sniff/testdata/quic-firefox153esr-1.bin b/extras/sniff/testdata/quic-firefox153esr-1.bin new file mode 100644 index 000000000..7f0dea9c7 Binary files /dev/null and b/extras/sniff/testdata/quic-firefox153esr-1.bin differ diff --git a/extras/sniff/testdata/quic-ngtcp2-1.11-0.bin b/extras/sniff/testdata/quic-ngtcp2-1.11-0.bin new file mode 100644 index 000000000..f2c28e509 Binary files /dev/null and b/extras/sniff/testdata/quic-ngtcp2-1.11-0.bin differ diff --git a/extras/sniff/testdata/quic-quiche-0.bin b/extras/sniff/testdata/quic-quiche-0.bin new file mode 100644 index 000000000..78e6d615f Binary files /dev/null and b/extras/sniff/testdata/quic-quiche-0.bin differ diff --git a/extras/sniff/testdata/quic-quiche-1.bin b/extras/sniff/testdata/quic-quiche-1.bin new file mode 100644 index 000000000..7207ea84f Binary files /dev/null and b/extras/sniff/testdata/quic-quiche-1.bin differ diff --git a/extras/sniff/testdata/quic-quiche-2.bin b/extras/sniff/testdata/quic-quiche-2.bin new file mode 100644 index 000000000..a19a85b6e Binary files /dev/null and b/extras/sniff/testdata/quic-quiche-2.bin differ diff --git a/extras/sniff/testdata/tls-chrome153.bin b/extras/sniff/testdata/tls-chrome153.bin new file mode 100644 index 000000000..b50117469 Binary files /dev/null and b/extras/sniff/testdata/tls-chrome153.bin differ diff --git a/extras/sniff/testdata/tls-curl8.18-openssl3.5.bin b/extras/sniff/testdata/tls-curl8.18-openssl3.5.bin new file mode 100644 index 000000000..9c2ef61e7 Binary files /dev/null and b/extras/sniff/testdata/tls-curl8.18-openssl3.5.bin differ diff --git a/extras/sniff/testdata/tls-firefox153esr.bin b/extras/sniff/testdata/tls-firefox153esr.bin new file mode 100644 index 000000000..1e66ff962 Binary files /dev/null and b/extras/sniff/testdata/tls-firefox153esr.bin differ diff --git a/extras/sniff/tls.go b/extras/sniff/tls.go new file mode 100644 index 000000000..c1406b68d --- /dev/null +++ b/extras/sniff/tls.go @@ -0,0 +1,81 @@ +package sniff + +import ( + "io" + + "golang.org/x/crypto/cryptobyte" +) + +// readClientHello reads TLS records from r until it has a ClientHello, which can be +// fragmented across records, and returns its body. +func readClientHello(r io.Reader) []byte { + var stream []byte + header := make([]byte, 5) + for { + if _, err := io.ReadFull(r, header); err != nil || header[0] != 0x16 || header[1] != 0x03 { + return nil + } + fragment := make([]byte, int(header[3])<<8|int(header[4])) + if _, err := io.ReadFull(r, fragment); err != nil { + return nil + } + stream = append(stream, fragment...) + if hello, more := clientHello(stream); !more { + return hello + } + } +} + +// clientHello returns the body of the ClientHello message at the start of a TLS handshake stream. +// more reports that the stream is too short to tell. +func clientHello(stream []byte) (hello []byte, more bool) { + if len(stream) < 4 { + return nil, true + } + if stream[0] != 1 { // client_hello + return nil, false + } + n := int(stream[1])<<16 | int(stream[2])<<8 | int(stream[3]) + if len(stream) < 4+n { + return nil, true + } + return stream[4 : 4+n], false +} + +// serverName returns the host_name in the server_name extension of a ClientHello body. +func serverName(hello []byte) string { + s := cryptobyte.String(hello) + var skipped, exts cryptobyte.String + if !s.Skip(2+32) || // legacy_version, random + !s.ReadUint8LengthPrefixed(&skipped) || // legacy_session_id + !s.ReadUint16LengthPrefixed(&skipped) || // cipher_suites + !s.ReadUint8LengthPrefixed(&skipped) || // legacy_compression_methods + !s.ReadUint16LengthPrefixed(&exts) { + return "" + } + for !exts.Empty() { + var typ uint16 + var ext, names cryptobyte.String + if !exts.ReadUint16(&typ) || !exts.ReadUint16LengthPrefixed(&ext) { + return "" + } + if typ != 0 { // server_name + continue + } + if !ext.ReadUint16LengthPrefixed(&names) { + return "" + } + for !names.Empty() { + var nameType uint8 + var name cryptobyte.String + if !names.ReadUint8(&nameType) || !names.ReadUint16LengthPrefixed(&name) { + return "" + } + if nameType == 0 { // host_name + return string(name) + } + } + return "" + } + return "" +}