diff --git a/core/internal/integration_tests/mocks/mock_Conn.go b/core/internal/integration_tests/mocks/mock_Conn.go index d75f34b..d068033 100644 --- a/core/internal/integration_tests/mocks/mock_Conn.go +++ b/core/internal/integration_tests/mocks/mock_Conn.go @@ -4,10 +4,9 @@ package mocks import ( net "net" + time "time" mock "github.com/stretchr/testify/mock" - - time "time" ) // MockConn is an autogenerated mock type for the Conn type diff --git a/core/internal/integration_tests/mocks/mock_Outbound.go b/core/internal/integration_tests/mocks/mock_Outbound.go index 6167bc5..6fda640 100644 --- a/core/internal/integration_tests/mocks/mock_Outbound.go +++ b/core/internal/integration_tests/mocks/mock_Outbound.go @@ -5,9 +5,8 @@ package mocks import ( net "net" - mock "github.com/stretchr/testify/mock" - server "github.com/apernet/hysteria/core/v2/server" + mock "github.com/stretchr/testify/mock" ) // MockOutbound is an autogenerated mock type for the Outbound type @@ -23,6 +22,52 @@ func (_m *MockOutbound) EXPECT() *MockOutbound_Expecter { return &MockOutbound_Expecter{mock: &_m.Mock} } +// CheckUDP provides a mock function with given fields: reqAddr +func (_m *MockOutbound) CheckUDP(reqAddr string) error { + ret := _m.Called(reqAddr) + + if len(ret) == 0 { + panic("no return value specified for CheckUDP") + } + + var r0 error + if rf, ok := ret.Get(0).(func(string) error); ok { + r0 = rf(reqAddr) + } else { + r0 = ret.Error(0) + } + + return r0 +} + +// MockOutbound_CheckUDP_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'CheckUDP' +type MockOutbound_CheckUDP_Call struct { + *mock.Call +} + +// CheckUDP is a helper method to define mock.On call +// - reqAddr string +func (_e *MockOutbound_Expecter) CheckUDP(reqAddr interface{}) *MockOutbound_CheckUDP_Call { + return &MockOutbound_CheckUDP_Call{Call: _e.mock.On("CheckUDP", reqAddr)} +} + +func (_c *MockOutbound_CheckUDP_Call) Run(run func(reqAddr string)) *MockOutbound_CheckUDP_Call { + _c.Call.Run(func(args mock.Arguments) { + run(args[0].(string)) + }) + return _c +} + +func (_c *MockOutbound_CheckUDP_Call) Return(_a0 error) *MockOutbound_CheckUDP_Call { + _c.Call.Return(_a0) + return _c +} + +func (_c *MockOutbound_CheckUDP_Call) RunAndReturn(run func(string) error) *MockOutbound_CheckUDP_Call { + _c.Call.Return(run) + return _c +} + // TCP provides a mock function with given fields: reqAddr func (_m *MockOutbound) TCP(reqAddr string) (net.Conn, error) { ret := _m.Called(reqAddr) diff --git a/core/internal/integration_tests/udp_acl_test.go b/core/internal/integration_tests/udp_acl_test.go new file mode 100644 index 0000000..b8bda01 --- /dev/null +++ b/core/internal/integration_tests/udp_acl_test.go @@ -0,0 +1,177 @@ +package integration_tests + +import ( + "errors" + "net" + "sync/atomic" + "testing" + "time" + + "github.com/apernet/hysteria/core/v2/client" + "github.com/apernet/hysteria/core/v2/internal/integration_tests/mocks" + "github.com/apernet/hysteria/core/v2/server" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/mock" +) + +type gatedOutbound struct { + blocked string + checkCalls atomic.Int32 + dialedAddrs atomic.Int32 +} + +func (o *gatedOutbound) TCP(reqAddr string) (net.Conn, error) { + return net.Dial("tcp", reqAddr) +} + +func (o *gatedOutbound) UDP(reqAddr string) (server.UDPConn, error) { + if reqAddr == o.blocked { + return nil, errors.New("rejected") + } + o.dialedAddrs.Add(1) + c, err := net.ListenUDP("udp", nil) + if err != nil { + return nil, err + } + return &gatedUDPConn{UDPConn: c}, nil +} + +func (o *gatedOutbound) CheckUDP(reqAddr string) error { + o.checkCalls.Add(1) + if reqAddr == o.blocked { + return errors.New("rejected") + } + return nil +} + +type gatedUDPConn struct { + *net.UDPConn +} + +func (c *gatedUDPConn) ReadFrom(b []byte) (int, string, error) { + n, addr, err := c.UDPConn.ReadFrom(b) + if addr != nil { + return n, addr.String(), err + } + return n, "", err +} + +func (c *gatedUDPConn) WriteTo(b []byte, addr string) (int, error) { + uAddr, err := net.ResolveUDPAddr("udp", addr) + if err != nil { + return 0, err + } + return c.UDPConn.WriteTo(b, uAddr) +} + +func TestClientServerUDPACLBypass(t *testing.T) { + const allowed, blocked = "127.0.0.1:22444", "127.0.0.1:22445" + ob := &gatedOutbound{blocked: blocked} + + udpConn, udpAddr, err := serverConn() + assert.NoError(t, err) + auth := mocks.NewMockAuthenticator(t) + auth.EXPECT().Authenticate(mock.Anything, mock.Anything, mock.Anything).Return(true, "nobody") + s, err := server.NewServer(&server.Config{ + TLSConfig: serverTLSConfig(), + Conn: udpConn, + Authenticator: auth, + Outbound: ob, + }) + assert.NoError(t, err) + defer s.Close() + go s.Serve() + + allowedConn, err := net.ListenPacket("udp", allowed) + assert.NoError(t, err) + defer allowedConn.Close() + go (&udpEchoServer{Conn: allowedConn}).Serve() + + blockedConn, err := net.ListenPacket("udp", blocked) + assert.NoError(t, err) + defer blockedConn.Close() + go (&udpEchoServer{Conn: blockedConn}).Serve() + + c, _, err := client.NewClient(&client.Config{ + ServerAddr: udpAddr, + TLSConfig: client.TLSConfig{InsecureSkipVerify: true}, + }) + assert.NoError(t, err) + defer c.Close() + + conn, err := c.UDP() + assert.NoError(t, err) + defer conn.Close() + + assert.NoError(t, conn.Send([]byte("hello"), allowed)) + rData, rAddr, err := conn.Receive() + assert.NoError(t, err) + assert.Equal(t, []byte("hello"), rData) + assert.Equal(t, allowed, rAddr) + + assert.NoError(t, conn.Send([]byte("ssrf"), blocked)) + + done := make(chan struct{}) + var leakedAddr string + go func() { + _, addr, err := conn.Receive() + if err == nil { + leakedAddr = addr + } + close(done) + }() + select { + case <-done: + assert.NotEqual(t, blocked, leakedAddr, "ACL bypass: blocked destination relayed") + case <-time.After(500 * time.Millisecond): + } + + assert.GreaterOrEqual(t, ob.checkCalls.Load(), int32(1), "CheckUDP not invoked for subsequent packet") + assert.Equal(t, int32(1), ob.dialedAddrs.Load(), "outbound dial must happen only on first allowed destination") +} + +func TestClientServerUDPACLMultiDestAllowed(t *testing.T) { + const dest1, dest2 = "127.0.0.1:22448", "127.0.0.1:22449" + ob := &gatedOutbound{blocked: ""} + + udpConn, udpAddr, err := serverConn() + assert.NoError(t, err) + auth := mocks.NewMockAuthenticator(t) + auth.EXPECT().Authenticate(mock.Anything, mock.Anything, mock.Anything).Return(true, "nobody") + s, err := server.NewServer(&server.Config{ + TLSConfig: serverTLSConfig(), + Conn: udpConn, + Authenticator: auth, + Outbound: ob, + }) + assert.NoError(t, err) + defer s.Close() + go s.Serve() + + for _, addr := range []string{dest1, dest2} { + ec, err := net.ListenPacket("udp", addr) + assert.NoError(t, err) + defer ec.Close() + go (&udpEchoServer{Conn: ec}).Serve() + } + + c, _, err := client.NewClient(&client.Config{ + ServerAddr: udpAddr, + TLSConfig: client.TLSConfig{InsecureSkipVerify: true}, + }) + assert.NoError(t, err) + defer c.Close() + + conn, err := c.UDP() + assert.NoError(t, err) + defer conn.Close() + + for _, addr := range []string{dest1, dest2} { + assert.NoError(t, conn.Send([]byte("hi"), addr)) + rData, rAddr, err := conn.Receive() + assert.NoError(t, err) + assert.Equal(t, []byte("hi"), rData) + assert.Equal(t, addr, rAddr) + } +} diff --git a/core/server/config.go b/core/server/config.go index 1d5c568..32c3ffd 100644 --- a/core/server/config.go +++ b/core/server/config.go @@ -153,9 +153,11 @@ type RequestHook interface { // Although UDP includes a reqAddr, the implementation does not necessarily have to use it // to make a "connected" UDP connection that does not accept packets from other addresses. // In fact, the default implementation simply uses net.ListenUDP for a "full-cone" behavior. +// CheckUDP is used to check if a UDP packet to reqAddr is permitted (useful for e.g. ACL). type Outbound interface { TCP(reqAddr string) (net.Conn, error) UDP(reqAddr string) (UDPConn, error) + CheckUDP(reqAddr string) error } // UDPConn is like net.PacketConn, but uses string for addresses. @@ -183,6 +185,10 @@ func (o *defaultOutbound) UDP(reqAddr string) (UDPConn, error) { return &defaultUDPConn{conn}, nil } +func (o *defaultOutbound) CheckUDP(reqAddr string) error { + return nil +} + type defaultUDPConn struct { *net.UDPConn } diff --git a/core/server/mock_udpIO.go b/core/server/mock_udpIO.go index a472203..bb512c0 100644 --- a/core/server/mock_udpIO.go +++ b/core/server/mock_udpIO.go @@ -20,6 +20,52 @@ func (_m *mockUDPIO) EXPECT() *mockUDPIO_Expecter { return &mockUDPIO_Expecter{mock: &_m.Mock} } +// CheckUDP provides a mock function with given fields: reqAddr +func (_m *mockUDPIO) CheckUDP(reqAddr string) error { + ret := _m.Called(reqAddr) + + if len(ret) == 0 { + panic("no return value specified for CheckUDP") + } + + var r0 error + if rf, ok := ret.Get(0).(func(string) error); ok { + r0 = rf(reqAddr) + } else { + r0 = ret.Error(0) + } + + return r0 +} + +// mockUDPIO_CheckUDP_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'CheckUDP' +type mockUDPIO_CheckUDP_Call struct { + *mock.Call +} + +// CheckUDP is a helper method to define mock.On call +// - reqAddr string +func (_e *mockUDPIO_Expecter) CheckUDP(reqAddr interface{}) *mockUDPIO_CheckUDP_Call { + return &mockUDPIO_CheckUDP_Call{Call: _e.mock.On("CheckUDP", reqAddr)} +} + +func (_c *mockUDPIO_CheckUDP_Call) Run(run func(reqAddr string)) *mockUDPIO_CheckUDP_Call { + _c.Call.Run(func(args mock.Arguments) { + run(args[0].(string)) + }) + return _c +} + +func (_c *mockUDPIO_CheckUDP_Call) Return(_a0 error) *mockUDPIO_CheckUDP_Call { + _c.Call.Return(_a0) + return _c +} + +func (_c *mockUDPIO_CheckUDP_Call) RunAndReturn(run func(string) error) *mockUDPIO_CheckUDP_Call { + _c.Call.Return(run) + 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) diff --git a/core/server/server.go b/core/server/server.go index e6c26c8..05a01e3 100644 --- a/core/server/server.go +++ b/core/server/server.go @@ -398,6 +398,10 @@ func (io *udpIOImpl) UDP(reqAddr string) (UDPConn, error) { return io.Outbound.UDP(reqAddr) } +func (io *udpIOImpl) CheckUDP(reqAddr string) error { + return io.Outbound.CheckUDP(reqAddr) +} + type udpEventLoggerImpl struct { Conn *quic.Conn AuthID string diff --git a/core/server/udp.go b/core/server/udp.go index 360a8cf..bfe6ae1 100644 --- a/core/server/udp.go +++ b/core/server/udp.go @@ -15,6 +15,7 @@ import ( const ( idleCleanupInterval = 1 * time.Second + maxSessionACLCache = 256 ) type udpIO interface { @@ -22,6 +23,7 @@ type udpIO interface { SendMessage([]byte, *protocol.UDPMessage) error Hook(data []byte, reqAddr *string) error UDP(reqAddr string) (UDPConn, error) + CheckUDP(reqAddr string) error } type udpEventLogger interface { @@ -43,6 +45,8 @@ type udpSessionEntry struct { conn UDPConn connLock sync.Mutex closed bool + + aclCache map[string]error } func newUDPSessionEntry( @@ -101,16 +105,41 @@ func (e *udpSessionEntry) Feed(msg *protocol.UDPMessage) (int, error) { if err != nil { return 0, err } + if e.OverrideAddr == "" { + e.aclCache = map[string]error{dfMsg.Addr: nil} + } } addr := dfMsg.Addr if e.OverrideAddr != "" { addr = e.OverrideAddr + } else if err := e.checkAddr(addr); err != nil { + return 0, err } return e.conn.WriteTo(dfMsg.Data, addr) } +// checkAddr checks outbound policy for the given address. +// The decision is cached in e.aclCache for future use. +func (e *udpSessionEntry) checkAddr(addr string) error { + if decision, ok := e.aclCache[addr]; ok { + return decision + } + decision := e.IO.CheckUDP(addr) + if len(e.aclCache) >= maxSessionACLCache { + for k := range e.aclCache { + delete(e.aclCache, k) + break + } + } + if e.aclCache == nil { + e.aclCache = make(map[string]error, 4) + } + e.aclCache[addr] = decision + return decision +} + // initConn initializes the UDP connection of the session. // If no error is returned, the e.conn is set to the new connection. func (e *udpSessionEntry) initConn(firstMsg *protocol.UDPMessage) error { diff --git a/extras/outbounds/acl.go b/extras/outbounds/acl.go index ecdeaaa..abb7523 100644 --- a/extras/outbounds/acl.go +++ b/extras/outbounds/acl.go @@ -110,6 +110,11 @@ func (a *aclEngine) UDP(reqAddr *AddrEx) (UDPConn, error) { return ob.UDP(reqAddr) } +func (a *aclEngine) CheckUDP(reqAddr *AddrEx) error { + ob := a.handle(reqAddr, acl.ProtocolUDP) + return ob.CheckUDP(reqAddr) +} + type aclRejectOutbound struct{} func (a *aclRejectOutbound) TCP(reqAddr *AddrEx) (net.Conn, error) { @@ -119,3 +124,7 @@ func (a *aclRejectOutbound) TCP(reqAddr *AddrEx) (net.Conn, error) { func (a *aclRejectOutbound) UDP(reqAddr *AddrEx) (UDPConn, error) { return nil, errRejected } + +func (a *aclRejectOutbound) CheckUDP(reqAddr *AddrEx) error { + return errRejected +} diff --git a/extras/outbounds/dns_https.go b/extras/outbounds/dns_https.go index c6817c1..212dbd0 100644 --- a/extras/outbounds/dns_https.go +++ b/extras/outbounds/dns_https.go @@ -87,3 +87,8 @@ func (r *dohResolver) UDP(reqAddr *AddrEx) (UDPConn, error) { r.resolve(reqAddr) return r.Next.UDP(reqAddr) } + +func (r *dohResolver) CheckUDP(reqAddr *AddrEx) error { + r.resolve(reqAddr) + return r.Next.CheckUDP(reqAddr) +} diff --git a/extras/outbounds/dns_standard.go b/extras/outbounds/dns_standard.go index a9df238..9dec606 100644 --- a/extras/outbounds/dns_standard.go +++ b/extras/outbounds/dns_standard.go @@ -219,3 +219,8 @@ func (r *standardResolver) UDP(reqAddr *AddrEx) (UDPConn, error) { r.resolve(reqAddr) return r.Next.UDP(reqAddr) } + +func (r *standardResolver) CheckUDP(reqAddr *AddrEx) error { + r.resolve(reqAddr) + return r.Next.CheckUDP(reqAddr) +} diff --git a/extras/outbounds/dns_system.go b/extras/outbounds/dns_system.go index 8f9a429..09abac9 100644 --- a/extras/outbounds/dns_system.go +++ b/extras/outbounds/dns_system.go @@ -39,3 +39,8 @@ func (r *systemResolver) UDP(reqAddr *AddrEx) (UDPConn, error) { r.resolve(reqAddr) return r.Next.UDP(reqAddr) } + +func (r *systemResolver) CheckUDP(reqAddr *AddrEx) error { + r.resolve(reqAddr) + return r.Next.CheckUDP(reqAddr) +} diff --git a/extras/outbounds/interface.go b/extras/outbounds/interface.go index bbd4dc2..a70902c 100644 --- a/extras/outbounds/interface.go +++ b/extras/outbounds/interface.go @@ -24,6 +24,7 @@ import ( type PluggableOutbound interface { TCP(reqAddr *AddrEx) (net.Conn, error) UDP(reqAddr *AddrEx) (UDPConn, error) + CheckUDP(reqAddr *AddrEx) error } type UDPConn interface { @@ -109,6 +110,21 @@ func (a *PluggableOutboundAdapter) UDP(reqAddr string) (server.UDPConn, error) { return &udpConnAdapter{conn}, nil } +func (a *PluggableOutboundAdapter) CheckUDP(reqAddr string) error { + host, port, err := net.SplitHostPort(reqAddr) + if err != nil { + return err + } + portUint, err := parsePortUint16(port) + if err != nil { + return err + } + return a.PluggableOutbound.CheckUDP(&AddrEx{ + Host: host, + Port: portUint, + }) +} + type udpConnAdapter struct { UDPConn } diff --git a/extras/outbounds/mock_PluggableOutbound.go b/extras/outbounds/mock_PluggableOutbound.go index e754fa6..f775cde 100644 --- a/extras/outbounds/mock_PluggableOutbound.go +++ b/extras/outbounds/mock_PluggableOutbound.go @@ -21,6 +21,52 @@ func (_m *mockPluggableOutbound) EXPECT() *mockPluggableOutbound_Expecter { return &mockPluggableOutbound_Expecter{mock: &_m.Mock} } +// CheckUDP provides a mock function with given fields: reqAddr +func (_m *mockPluggableOutbound) CheckUDP(reqAddr *AddrEx) error { + ret := _m.Called(reqAddr) + + if len(ret) == 0 { + panic("no return value specified for CheckUDP") + } + + var r0 error + if rf, ok := ret.Get(0).(func(*AddrEx) error); ok { + r0 = rf(reqAddr) + } else { + r0 = ret.Error(0) + } + + return r0 +} + +// mockPluggableOutbound_CheckUDP_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'CheckUDP' +type mockPluggableOutbound_CheckUDP_Call struct { + *mock.Call +} + +// CheckUDP is a helper method to define mock.On call +// - reqAddr *AddrEx +func (_e *mockPluggableOutbound_Expecter) CheckUDP(reqAddr interface{}) *mockPluggableOutbound_CheckUDP_Call { + return &mockPluggableOutbound_CheckUDP_Call{Call: _e.mock.On("CheckUDP", reqAddr)} +} + +func (_c *mockPluggableOutbound_CheckUDP_Call) Run(run func(reqAddr *AddrEx)) *mockPluggableOutbound_CheckUDP_Call { + _c.Call.Run(func(args mock.Arguments) { + run(args[0].(*AddrEx)) + }) + return _c +} + +func (_c *mockPluggableOutbound_CheckUDP_Call) Return(_a0 error) *mockPluggableOutbound_CheckUDP_Call { + _c.Call.Return(_a0) + return _c +} + +func (_c *mockPluggableOutbound_CheckUDP_Call) RunAndReturn(run func(*AddrEx) error) *mockPluggableOutbound_CheckUDP_Call { + _c.Call.Return(run) + return _c +} + // TCP provides a mock function with given fields: reqAddr func (_m *mockPluggableOutbound) TCP(reqAddr *AddrEx) (net.Conn, error) { ret := _m.Called(reqAddr) diff --git a/extras/outbounds/ob_direct.go b/extras/outbounds/ob_direct.go index de7ddd2..d6a27be 100644 --- a/extras/outbounds/ob_direct.go +++ b/extras/outbounds/ob_direct.go @@ -402,6 +402,10 @@ func (u *directOutboundUDPConn) Close() error { return u.UDPConn.Close() } +func (d *directOutbound) CheckUDP(reqAddr *AddrEx) error { + return nil +} + func (d *directOutbound) UDP(reqAddr *AddrEx) (UDPConn, error) { if d.BindIP4 == nil && d.BindIP6 == nil { // No bind address specified, use default dual stack implementation diff --git a/extras/outbounds/ob_http.go b/extras/outbounds/ob_http.go index 48d5aac..890cd7d 100644 --- a/extras/outbounds/ob_http.go +++ b/extras/outbounds/ob_http.go @@ -170,6 +170,10 @@ func (o *httpOutbound) UDP(reqAddr *AddrEx) (UDPConn, error) { return nil, errHTTPUDPNotSupported } +func (o *httpOutbound) CheckUDP(reqAddr *AddrEx) error { + return errHTTPUDPNotSupported +} + // cachedConn is a net.Conn wrapper that first Read()s from a buffer, // and then from the underlying net.Conn when the buffer is drained. type cachedConn struct { diff --git a/extras/outbounds/ob_socks5.go b/extras/outbounds/ob_socks5.go index 72dc567..a0e483c 100644 --- a/extras/outbounds/ob_socks5.go +++ b/extras/outbounds/ob_socks5.go @@ -171,6 +171,10 @@ func (s *socks5Outbound) TCP(reqAddr *AddrEx) (net.Conn, error) { return conn, nil } +func (s *socks5Outbound) CheckUDP(reqAddr *AddrEx) error { + return nil +} + func (s *socks5Outbound) UDP(reqAddr *AddrEx) (UDPConn, error) { conn, err := s.dialAndNegotiate() if err != nil { diff --git a/extras/outbounds/speedtest.go b/extras/outbounds/speedtest.go index 162f4dc..c17add6 100644 --- a/extras/outbounds/speedtest.go +++ b/extras/outbounds/speedtest.go @@ -34,3 +34,7 @@ func (s *speedtestHandler) TCP(reqAddr *AddrEx) (net.Conn, error) { func (s *speedtestHandler) UDP(reqAddr *AddrEx) (UDPConn, error) { return s.Next.UDP(reqAddr) } + +func (s *speedtestHandler) CheckUDP(reqAddr *AddrEx) error { + return s.Next.CheckUDP(reqAddr) +}