diff --git a/README.md b/README.md index 628efbe..7676eb9 100644 --- a/README.md +++ b/README.md @@ -36,6 +36,7 @@ The following environment variables can be configured in the `.env` file: | Variable | Description | Default | Required | |---------------------|-----------------------------------------------------------------------------|-------------------------|---------------------| | `DOMAIN` | Domain name for subdomain routing | `localhost` | No | +| `FRONTEND_URL` | URL for the frontend dashboard/landing page | `https://` | No | | `PORT` | SSH server port | `2200` | No | | `HTTP_PORT` | HTTP server port | `8080` | No | | `HTTPS_PORT` | HTTPS server port | `8443` | No | diff --git a/internal/bootstrap/bootstrap_test.go b/internal/bootstrap/bootstrap_test.go index f5a181e..80e9ebf 100644 --- a/internal/bootstrap/bootstrap_test.go +++ b/internal/bootstrap/bootstrap_test.go @@ -80,6 +80,7 @@ type MockConfig struct { } func (m *MockConfig) Domain() string { return m.Called().String(0) } +func (m *MockConfig) FrontendURL() string { return m.Called().String(0) } func (m *MockConfig) SSHPort() string { return m.Called().String(0) } func (m *MockConfig) HTTPPort() string { return m.Called().String(0) } func (m *MockConfig) HTTPSPort() string { return m.Called().String(0) } diff --git a/internal/config/config.go b/internal/config/config.go index f3876bf..d341269 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -6,6 +6,7 @@ import ( type Config interface { Domain() string + FrontendURL() string SSHPort() string HTTPPort() string @@ -50,6 +51,7 @@ func MustLoad() (Config, error) { } func (c *config) Domain() string { return c.domain } +func (c *config) FrontendURL() string { return c.frontendURL } func (c *config) SSHPort() string { return c.sshPort } func (c *config) HTTPPort() string { return c.httpPort } func (c *config) HTTPSPort() string { return c.httpsPort } diff --git a/internal/config/loader.go b/internal/config/loader.go index c6cb909..e8c7b1c 100644 --- a/internal/config/loader.go +++ b/internal/config/loader.go @@ -12,8 +12,9 @@ import ( ) type config struct { - domain string - sshPort string + domain string + frontendURL string + sshPort string httpPort string httpsPort string @@ -49,6 +50,7 @@ func parse() (*config, error) { } domain := getenv("DOMAIN", "localhost") + frontendURL := getenv("FRONTEND_URL", "https://"+domain) sshPort := getenv("PORT", "2200") httpPort := getenv("HTTP_PORT", "8080") @@ -89,6 +91,7 @@ func parse() (*config, error) { return &config{ domain: domain, + frontendURL: frontendURL, sshPort: sshPort, httpPort: httpPort, httpsPort: httpsPort, diff --git a/internal/grpc/client/client_test.go b/internal/grpc/client/client_test.go index 5009185..cb89f51 100644 --- a/internal/grpc/client/client_test.go +++ b/internal/grpc/client/client_test.go @@ -13,7 +13,6 @@ import ( "tunnel_pls/internal/session/slug" "tunnel_pls/internal/types" - "tunnel_pls/internal/port" "tunnel_pls/internal/registry" proto "git.fossy.my.id/bagas/tunnel-please-grpc/gen" @@ -754,6 +753,7 @@ type MockConfig struct { } func (m *MockConfig) Domain() string { return m.Called().String(0) } +func (m *MockConfig) FrontendURL() string { return m.Called().String(0) } func (m *MockConfig) SSHPort() string { return m.Called().String(0) } func (m *MockConfig) HTTPPort() string { return m.Called().String(0) } func (m *MockConfig) HTTPSPort() string { return m.Called().String(0) } @@ -885,16 +885,16 @@ func (m *mockLifecycle) Connection() ssh.Conn { return args.Get(0).(ssh.Conn) } func (m *mockLifecycle) User() string { return m.Called().String(0) } -func (m *mockLifecycle) SetChannel(channel ssh.Channel) { m.Called(channel) } +func (m *mockLifecycle) SetChannel(channel ssh.Channel) error { return m.Called(channel).Error(0) } func (m *mockLifecycle) SetStatus(status types.SessionStatus) { m.Called(status) } func (m *mockLifecycle) IsActive() bool { return m.Called().Bool(0) } func (m *mockLifecycle) StartedAt() time.Time { return m.Called().Get(0).(time.Time) } -func (m *mockLifecycle) PortRegistry() port.Port { +func (m *mockLifecycle) PortRegistry() lifecycle.PortRegistry { args := m.Called() if args.Get(0) == nil { return nil } - return args.Get(0).(port.Port) + return args.Get(0).(lifecycle.PortRegistry) } type mockEventServiceClient struct { diff --git a/internal/port/port.go b/internal/port/port.go index 6c60fbb..15f08a8 100644 --- a/internal/port/port.go +++ b/internal/port/port.go @@ -33,10 +33,15 @@ func (pm *port) AddRange(startPort, endPort uint16) error { if startPort > endPort { return fmt.Errorf("start port cannot be greater than end port") } - for index := startPort; index <= endPort; index++ { - if _, exists := pm.ports[index]; !exists { - pm.ports[index] = false - pm.sortedPorts = append(pm.sortedPorts, index) + for index := startPort; ; index++ { + if index != 0 { + if _, exists := pm.ports[index]; !exists { + pm.ports[index] = false + pm.sortedPorts = append(pm.sortedPorts, index) + } + } + if index == endPort { + break } } sort.Slice(pm.sortedPorts, func(i, j int) bool { @@ -51,6 +56,7 @@ func (pm *port) Unassigned() (uint16, bool) { for _, index := range pm.sortedPorts { if !pm.ports[index] { + pm.ports[index] = true return index, true } } @@ -61,6 +67,9 @@ func (pm *port) SetStatus(port uint16, assigned bool) error { pm.mu.Lock() defer pm.mu.Unlock() + if _, exists := pm.ports[port]; !exists { + return fmt.Errorf("port %d is not in the allowed range", port) + } pm.ports[port] = assigned return nil } @@ -70,16 +79,10 @@ func (pm *port) Claim(port uint16) (claimed bool) { defer pm.mu.Unlock() status, exists := pm.ports[port] - - if exists && status { + if !exists || status { return false } - if !exists { - pm.ports[port] = true - return true - } - pm.ports[port] = true return true } diff --git a/internal/port/port_test.go b/internal/port/port_test.go index fcc64d3..9f9a3d9 100644 --- a/internal/port/port_test.go +++ b/internal/port/port_test.go @@ -16,6 +16,8 @@ func TestAddRange(t *testing.T) { {"normal range", 1000, 1002, false}, {"invalid range", 2000, 1999, true}, {"single port range", 3000, 3000, false}, + {"range ending at max uint16", 65533, 65535, false}, + {"range including port zero", 0, 2, false}, } for _, tt := range tests { @@ -31,6 +33,22 @@ func TestAddRange(t *testing.T) { } } +func TestAddRangeBoundaries(t *testing.T) { + pm := New() + err := pm.AddRange(0, 3) + assert.NoError(t, err) + _, hasZero := pm.(*port).ports[0] + assert.False(t, hasZero, "port 0 must be skipped") + assert.Len(t, pm.(*port).ports, 3) + + pm2 := New() + err = pm2.AddRange(65533, 65535) + assert.NoError(t, err) + assert.Len(t, pm2.(*port).ports, 3) + _, hasMax := pm2.(*port).ports[65535] + assert.True(t, hasMax) +} + func TestUnassigned(t *testing.T) { pm := New() _ = pm.AddRange(1000, 1002) @@ -58,6 +76,21 @@ func TestUnassigned(t *testing.T) { } } +func TestUnassignedReservesPort(t *testing.T) { + pm := New() + _ = pm.AddRange(1000, 1002) + + p1, ok1 := pm.Unassigned() + assert.True(t, ok1) + assert.Equal(t, uint16(1000), p1) + + p2, ok2 := pm.Unassigned() + assert.True(t, ok2) + assert.Equal(t, uint16(1001), p2) + + assert.True(t, pm.(*port).ports[1000], "Unassigned must reserve the port") +} + func TestSetStatus(t *testing.T) { pm := New() _ = pm.AddRange(1000, 1002) @@ -83,6 +116,17 @@ func TestSetStatus(t *testing.T) { } } +func TestSetStatusUnknownPort(t *testing.T) { + pm := New() + _ = pm.AddRange(1000, 1002) + + err := pm.SetStatus(5000, true) + assert.Error(t, err) + + _, exists := pm.(*port).ports[5000] + assert.False(t, exists, "SetStatus must not create entries for unknown ports") +} + func TestClaim(t *testing.T) { pm := New() _ = pm.AddRange(1000, 1002) @@ -95,7 +139,7 @@ func TestClaim(t *testing.T) { }{ {"claim unassigned port", 1000, false, true}, {"claim already assigned port", 1001, true, false}, - {"claim non-existent port", 5000, false, true}, + {"claim non-existent port", 5000, false, false}, } for _, tt := range tests { @@ -107,8 +151,13 @@ func TestClaim(t *testing.T) { got := pm.Claim(tt.port) assert.Equal(t, tt.want, got) - finalState := pm.(*port).ports[tt.port] - assert.True(t, finalState) + finalState, exists := pm.(*port).ports[tt.port] + if !tt.want && tt.port == 5000 { + assert.False(t, exists, "out-of-range port must not be added to the registry") + } else { + assert.True(t, exists) + assert.True(t, finalState) + } }) } } diff --git a/internal/registry/registry.go b/internal/registry/registry.go index a4f64a8..e911a8a 100644 --- a/internal/registry/registry.go +++ b/internal/registry/registry.go @@ -94,13 +94,13 @@ func (r *registry) Update(user string, oldKey, newKey Key) error { return ErrInvalidSlug } + r.mu.Lock() + defer r.mu.Unlock() + if _, exists := r.slugIndex[newKey]; exists && newKey != oldKey { return ErrSlugInUse } - r.mu.Lock() - defer r.mu.Unlock() - client, ok := r.byUser[user][oldKey] if !ok { return ErrSessionNotFound diff --git a/internal/registry/registry_test.go b/internal/registry/registry_test.go index 9a80d47..489122c 100644 --- a/internal/registry/registry_test.go +++ b/internal/registry/registry_test.go @@ -4,7 +4,6 @@ import ( "sync" "testing" "time" - "tunnel_pls/internal/port" "tunnel_pls/internal/session/forwarder" "tunnel_pls/internal/session/interaction" "tunnel_pls/internal/session/lifecycle" @@ -78,15 +77,15 @@ func (ml *mockLifecycle) Connection() ssh.Conn { return args.Get(0).(ssh.Conn) } -func (ml *mockLifecycle) PortRegistry() port.Port { +func (ml *mockLifecycle) PortRegistry() lifecycle.PortRegistry { args := ml.Called() if args.Get(0) == nil { return nil } - return args.Get(0).(port.Port) + return args.Get(0).(lifecycle.PortRegistry) } -func (ml *mockLifecycle) SetChannel(channel ssh.Channel) { ml.Called(channel) } +func (ml *mockLifecycle) SetChannel(channel ssh.Channel) error { return ml.Called(channel).Error(0) } func (ml *mockLifecycle) SetStatus(status types.SessionStatus) { ml.Called(status) } func (ml *mockLifecycle) IsActive() bool { return ml.Called().Bool(0) } func (ml *mockLifecycle) StartedAt() time.Time { return ml.Called().Get(0).(time.Time) } diff --git a/internal/server/server_test.go b/internal/server/server_test.go index d4499a8..dbed841 100644 --- a/internal/server/server_test.go +++ b/internal/server/server_test.go @@ -33,6 +33,7 @@ type MockConfig struct { } func (m *MockConfig) Domain() string { return m.Called().String(0) } +func (m *MockConfig) FrontendURL() string { return m.Called().String(0) } func (m *MockConfig) SSHPort() string { return m.Called().String(0) } func (m *MockConfig) HTTPPort() string { return m.Called().String(0) } func (m *MockConfig) HTTPSPort() string { return m.Called().String(0) } diff --git a/internal/session/forwarder/forwarder.go b/internal/session/forwarder/forwarder.go index 6af3bae..f59d315 100644 --- a/internal/session/forwarder/forwarder.go +++ b/internal/session/forwarder/forwarder.go @@ -28,6 +28,7 @@ type Forwarder interface { Close() error } type forwarder struct { + mu sync.RWMutex listener net.Listener tunnelType types.TunnelType forwardedPort uint16 @@ -60,7 +61,7 @@ func (f *forwarder) copyWithBuffer(dst io.Writer, src io.Reader) (written int64, } func (f *forwarder) OpenForwardedChannel(ctx context.Context, origin net.Addr) (ssh.Channel, <-chan *ssh.Request, error) { - payload := createForwardedTCPIPPayload(origin, f.forwardedPort) + payload := createForwardedTCPIPPayload(origin, f.ForwardedPort()) type channelResult struct { channel ssh.Channel reqs <-chan *ssh.Request @@ -141,32 +142,44 @@ func (f *forwarder) HandleConnection(dst io.ReadWriter, src ssh.Channel) { } func (f *forwarder) SetType(tunnelType types.TunnelType) { + f.mu.Lock() + defer f.mu.Unlock() f.tunnelType = tunnelType } func (f *forwarder) TunnelType() types.TunnelType { + f.mu.RLock() + defer f.mu.RUnlock() return f.tunnelType } func (f *forwarder) ForwardedPort() uint16 { + f.mu.RLock() + defer f.mu.RUnlock() return f.forwardedPort } func (f *forwarder) SetForwardedPort(port uint16) { + f.mu.Lock() + defer f.mu.Unlock() f.forwardedPort = port } func (f *forwarder) SetListener(listener net.Listener) { + f.mu.Lock() + defer f.mu.Unlock() f.listener = listener } func (f *forwarder) Listener() net.Listener { + f.mu.RLock() + defer f.mu.RUnlock() return f.listener } func (f *forwarder) Close() error { - if f.Listener() != nil { - return f.listener.Close() + if listener := f.Listener(); listener != nil { + return listener.Close() } return nil } diff --git a/internal/session/forwarder/forwarder_test.go b/internal/session/forwarder/forwarder_test.go index 2a5c24d..37853de 100644 --- a/internal/session/forwarder/forwarder_test.go +++ b/internal/session/forwarder/forwarder_test.go @@ -24,6 +24,7 @@ type mockConfig struct { } func (m *mockConfig) Domain() string { return m.Called().String(0) } +func (m *mockConfig) FrontendURL() string { return m.Called().String(0) } func (m *mockConfig) SSHPort() string { return m.Called().String(0) } func (m *mockConfig) HTTPPort() string { return m.Called().String(0) } func (m *mockConfig) HTTPSPort() string { return m.Called().String(0) } diff --git a/internal/session/interaction/interaction_test.go b/internal/session/interaction/interaction_test.go index 679ec8a..1dc0408 100644 --- a/internal/session/interaction/interaction_test.go +++ b/internal/session/interaction/interaction_test.go @@ -32,6 +32,7 @@ type MockConfig struct { } func (m *MockConfig) Domain() string { return m.Called().String(0) } +func (m *MockConfig) FrontendURL() string { return m.Called().String(0) } func (m *MockConfig) SSHPort() string { return m.Called().String(0) } func (m *MockConfig) HTTPPort() string { return m.Called().String(0) } func (m *MockConfig) HTTPSPort() string { return m.Called().String(0) } @@ -1922,10 +1923,6 @@ func TestInteraction_Start_ProtocolSelection(t *testing.T) { time.Sleep(50 * time.Millisecond) i := mockInteraction.(*interaction) - if i.program != nil { - assert.NotNil(t, i.program, "program should be initialized") - } - i.Stop() mockConfig.AssertExpectations(t) diff --git a/internal/session/lifecycle/lifecycle.go b/internal/session/lifecycle/lifecycle.go index 286d8fb..0191e40 100644 --- a/internal/session/lifecycle/lifecycle.go +++ b/internal/session/lifecycle/lifecycle.go @@ -2,6 +2,7 @@ package lifecycle import ( "errors" + "fmt" "io" "net" "sync" @@ -9,8 +10,6 @@ import ( "tunnel_pls/internal/session/slug" "tunnel_pls/internal/types" - portUtil "tunnel_pls/internal/port" - "golang.org/x/crypto/ssh" ) @@ -24,6 +23,12 @@ type SessionRegistry interface { Remove(key types.SessionKey) } +type PortRegistry interface { + Unassigned() (uint16, bool) + Claim(port uint16) bool + SetStatus(port uint16, assigned bool) error +} + type lifecycle struct { mu sync.Mutex status types.SessionStatus @@ -34,18 +39,18 @@ type lifecycle struct { slug slug.Slug startedAt time.Time sessionRegistry SessionRegistry - portRegistry portUtil.Port + portRegistry PortRegistry user string } -func New(conn ssh.Conn, forwarder Forwarder, slugManager slug.Slug, port portUtil.Port, sessionRegistry SessionRegistry, user string) Lifecycle { +func New(conn ssh.Conn, forwarder Forwarder, slugManager slug.Slug, port PortRegistry, sessionRegistry SessionRegistry, user string) Lifecycle { return &lifecycle{ status: types.SessionStatusINITIALIZING, conn: conn, channel: nil, forwarder: forwarder, slug: slugManager, - startedAt: time.Now(), + startedAt: time.Time{}, sessionRegistry: sessionRegistry, portRegistry: port, user: user, @@ -55,16 +60,16 @@ func New(conn ssh.Conn, forwarder Forwarder, slugManager slug.Slug, port portUti type Lifecycle interface { Connection() ssh.Conn Channel() ssh.Channel - PortRegistry() portUtil.Port + PortRegistry() PortRegistry User() string - SetChannel(channel ssh.Channel) + SetChannel(channel ssh.Channel) error SetStatus(status types.SessionStatus) IsActive() bool StartedAt() time.Time Close() error } -func (l *lifecycle) PortRegistry() portUtil.Port { +func (l *lifecycle) PortRegistry() PortRegistry { return l.portRegistry } @@ -72,11 +77,25 @@ func (l *lifecycle) User() string { return l.user } -func (l *lifecycle) SetChannel(channel ssh.Channel) { +func (l *lifecycle) SetChannel(channel ssh.Channel) error { + l.mu.Lock() + defer l.mu.Unlock() + if l.status == types.SessionStatusCLOSED { + return fmt.Errorf("lifecycle is closed") + } + if channel == nil { + return fmt.Errorf("channel cannot be nil") + } + if l.channel != nil { + return fmt.Errorf("channel already set") + } l.channel = channel + return nil } func (l *lifecycle) Channel() ssh.Channel { + l.mu.Lock() + defer l.mu.Unlock() return l.channel } @@ -87,7 +106,13 @@ func (l *lifecycle) Connection() ssh.Conn { func (l *lifecycle) SetStatus(status types.SessionStatus) { l.mu.Lock() defer l.mu.Unlock() + if l.status == types.SessionStatusCLOSED { + return + } l.status = status + if status == types.SessionStatusRUNNING && l.startedAt.IsZero() { + l.startedAt = time.Now() + } } func (l *lifecycle) IsActive() bool { @@ -98,50 +123,74 @@ func (l *lifecycle) IsActive() bool { func (l *lifecycle) Close() error { l.mu.Lock() - defer l.mu.Unlock() if l.status == types.SessionStatusCLOSED { - return l.closeErr + closeErr := l.closeErr + l.mu.Unlock() + return closeErr } l.status = types.SessionStatusCLOSED + channel := l.channel + conn := l.conn + l.mu.Unlock() + var errs []error - tunnelType := l.forwarder.TunnelType() - - if l.channel != nil { - if err := l.channel.Close(); err != nil && !isClosedError(err) { + if channel != nil { + if err := channel.Close(); err != nil && !isClosedError(err) { + errs = append(errs, err) + } + } + if conn != nil { + if err := conn.Close(); err != nil && !isClosedError(err) { errs = append(errs, err) } } - if l.conn != nil { - if err := l.conn.Close(); err != nil && !isClosedError(err) { - errs = append(errs, err) - } + l.cleanupRegistry() + if err := l.cleanupForwarder(); err != nil { + errs = append(errs, err) } - clientSlug := l.slug.String() + closeErr := errors.Join(errs...) + + l.mu.Lock() + l.closeErr = closeErr + l.mu.Unlock() + + return closeErr +} + +func (l *lifecycle) cleanupRegistry() { + slugStr := l.slug.String() + if slugStr == "" { + return + } key := types.SessionKey{ - Id: clientSlug, - Type: tunnelType, + Id: slugStr, + Type: l.forwarder.TunnelType(), } l.sessionRegistry.Remove(key) +} - if tunnelType == types.TunnelTypeTCP { - errs = append(errs, l.PortRegistry().SetStatus(l.forwarder.ForwardedPort(), false)) - errs = append(errs, l.forwarder.Close()) +func (l *lifecycle) cleanupForwarder() error { + if l.forwarder.TunnelType() != types.TunnelTypeTCP { + return nil } - - l.closeErr = errors.Join(errs...) - return l.closeErr + var errs []error + errs = append(errs, l.portRegistry.SetStatus(l.forwarder.ForwardedPort(), false)) + errs = append(errs, l.forwarder.Close()) + return errors.Join(errs...) } func isClosedError(err error) bool { if err == nil { return false } - return errors.Is(err, io.EOF) || errors.Is(err, net.ErrClosed) || err.Error() == "EOF" + return errors.Is(err, io.EOF) || errors.Is(err, net.ErrClosed) } func (l *lifecycle) StartedAt() time.Time { + l.mu.Lock() + defer l.mu.Unlock() return l.startedAt } diff --git a/internal/session/lifecycle/lifecycle_test.go b/internal/session/lifecycle/lifecycle_test.go index b8fc7bd..b5b2ff6 100644 --- a/internal/session/lifecycle/lifecycle_test.go +++ b/internal/session/lifecycle/lifecycle_test.go @@ -5,6 +5,7 @@ import ( "errors" "io" "net" + "sync" "testing" "tunnel_pls/internal/types" @@ -177,8 +178,14 @@ func TestLifecycle_SetChannel(t *testing.T) { mockSSHChannel := &MockSSHChannel{} - mockLifecycle.SetChannel(mockSSHChannel) + err := mockLifecycle.SetChannel(mockSSHChannel) + assert.NoError(t, err) + assert.Equal(t, mockSSHChannel, mockLifecycle.Channel()) + anotherChannel := &MockSSHChannel{} + err = mockLifecycle.SetChannel(anotherChannel) + assert.Error(t, err) + assert.Contains(t, err.Error(), "channel already set") assert.Equal(t, mockSSHChannel, mockLifecycle.Channel()) } @@ -276,14 +283,15 @@ func TestLifecycle_Close(t *testing.T) { mockLifecycle := New(mockSSHConn, mockForwarder, mockSlug, mockPort, mockSessionRegistry, "mas-fuad") mockLifecycle.SetStatus(types.SessionStatusRUNNING) - mockLifecycle.SetChannel(mockSSHChannel) + err := mockLifecycle.SetChannel(mockSSHChannel) + assert.NoError(t, err) if tt.alreadyClosed { - err := mockLifecycle.Close() + err = mockLifecycle.Close() assert.NoError(t, err) } - err := mockLifecycle.Close() + err = mockLifecycle.Close() if tt.expectErr { assert.Error(t, err) @@ -301,3 +309,116 @@ func TestLifecycle_Close(t *testing.T) { }) } } + +func TestLifecycle_ConcurrentClose(t *testing.T) { + mockSSHConn := &MockSSHConn{} + mockSSHConn.On("Close").Return(nil) + + mockForwarder := &MockForwarder{} + mockForwarder.On("TunnelType").Return(types.TunnelTypeHTTP) + + mockSlug := &MockSlug{} + mockSlug.On("String").Return("test-slug") + + mockPort := &MockPort{} + + mockSessionRegistry := &MockSessionRegistry{} + mockSessionRegistry.On("Remove", mock.Anything).Return() + + mockSSHChannel := &MockSSHChannel{} + mockSSHChannel.On("Close").Return(nil) + + mockLifecycle := New(mockSSHConn, mockForwarder, mockSlug, mockPort, mockSessionRegistry, "mas-fuad") + mockLifecycle.SetStatus(types.SessionStatusRUNNING) + err := mockLifecycle.SetChannel(mockSSHChannel) + assert.NoError(t, err) + + const numGoroutines = 10 + var wg sync.WaitGroup + errChan := make(chan error, numGoroutines) + + for i := 0; i < numGoroutines; i++ { + wg.Add(1) + go func() { + defer wg.Done() + err := mockLifecycle.Close() + errChan <- err + }() + } + + wg.Wait() + close(errChan) + + for err := range errChan { + assert.NoError(t, err) + } + + assert.False(t, mockLifecycle.IsActive()) +} + +func TestLifecycle_SetChannel_AfterClose(t *testing.T) { + mockSSHConn := new(MockSSHConn) + mockSSHConn.On("Close").Return(nil) + mockForwarder := &MockForwarder{} + mockForwarder.On("TunnelType").Return(types.TunnelTypeHTTP) + mockSlug := &MockSlug{} + mockSlug.On("String").Return("test-slug") + mockPort := &MockPort{} + mockSessionRegistry := &MockSessionRegistry{} + mockSessionRegistry.On("Remove", mock.Anything).Return() + mockSSHChannel := &MockSSHChannel{} + mockSSHChannel.On("Close").Return(nil) + + mockLifecycle := New(mockSSHConn, mockForwarder, mockSlug, mockPort, mockSessionRegistry, "mas-fuad") + mockLifecycle.SetStatus(types.SessionStatusRUNNING) + err := mockLifecycle.SetChannel(mockSSHChannel) + assert.NoError(t, err) + + err = mockLifecycle.Close() + assert.NoError(t, err) + + anotherChannel := &MockSSHChannel{} + err = mockLifecycle.SetChannel(anotherChannel) + assert.Error(t, err) + assert.Contains(t, err.Error(), "lifecycle is closed") +} + +func TestLifecycle_SetChannel_Nil(t *testing.T) { + mockSSHConn := new(MockSSHConn) + mockForwarder := &MockForwarder{} + mockSlug := &MockSlug{} + mockPort := &MockPort{} + mockSessionRegistry := &MockSessionRegistry{} + + mockLifecycle := New(mockSSHConn, mockForwarder, mockSlug, mockPort, mockSessionRegistry, "mas-fuad") + + err := mockLifecycle.SetChannel(nil) + assert.Error(t, err) + assert.Contains(t, err.Error(), "channel cannot be nil") +} + +func TestLifecycle_SetStatus_AfterClose(t *testing.T) { + mockSSHConn := new(MockSSHConn) + mockSSHConn.On("Close").Return(nil) + mockForwarder := &MockForwarder{} + mockForwarder.On("TunnelType").Return(types.TunnelTypeHTTP) + mockSlug := &MockSlug{} + mockSlug.On("String").Return("test-slug") + mockPort := &MockPort{} + mockSessionRegistry := &MockSessionRegistry{} + mockSessionRegistry.On("Remove", mock.Anything).Return() + mockSSHChannel := &MockSSHChannel{} + mockSSHChannel.On("Close").Return(nil) + + mockLifecycle := New(mockSSHConn, mockForwarder, mockSlug, mockPort, mockSessionRegistry, "mas-fuad") + mockLifecycle.SetStatus(types.SessionStatusRUNNING) + err := mockLifecycle.SetChannel(mockSSHChannel) + assert.NoError(t, err) + + err = mockLifecycle.Close() + assert.NoError(t, err) + assert.False(t, mockLifecycle.IsActive()) + + mockLifecycle.SetStatus(types.SessionStatusRUNNING) + assert.False(t, mockLifecycle.IsActive(), "SetStatus should be ignored after Close") +} diff --git a/internal/session/session.go b/internal/session/session.go index 5978827..a2c3d81 100644 --- a/internal/session/session.go +++ b/internal/session/session.go @@ -26,7 +26,7 @@ type Session interface { HandleGlobalRequest(ch <-chan *ssh.Request) error HandleTCPIPForward(req *ssh.Request) error HandleHTTPForward(req *ssh.Request, port uint16) error - HandleTCPForward(req *ssh.Request, addr string, port uint16) error + HandleTCPForward(req *ssh.Request, addr string, port uint16, reserved bool) error Lifecycle() lifecycle.Lifecycle Interaction() interaction.Interaction Forwarder() forwarder.Forwarder @@ -158,13 +158,15 @@ func (s *session) setupInteractiveMode(channel ssh.NewChannel) error { } go func() { - err = s.HandleGlobalRequest(reqs) + err := s.HandleGlobalRequest(reqs) if err != nil { log.Printf("global request handler error: %v", err) } }() - s.lifecycle.SetChannel(ch) + if err = s.lifecycle.SetChannel(ch); err != nil { + return err + } s.interaction.SetChannel(ch) s.interaction.SetMode(types.InteractiveModeINTERACTIVE) @@ -198,23 +200,22 @@ func (s *session) waitForSessionEnd() error { } func (s *session) waitForTCPIPForward() *ssh.Request { - select { - case req, ok := <-s.initialReq: - if !ok { - log.Println("Forwarding request channel closed") + for { + select { + case req, ok := <-s.initialReq: + if !ok { + log.Println("Forwarding request channel closed") + return nil + } + if req.Type == "tcpip-forward" { + return req + } + log.Printf("Ignoring unexpected global request: %s", req.Type) + _ = req.Reply(false, nil) + case <-time.After(500 * time.Millisecond): + log.Println("No tcpip-forward request received within timeout") return nil } - if req.Type == "tcpip-forward" { - return req - } - if err := req.Reply(false, nil); err != nil { - log.Printf("Failed to reply to request: %v", err) - } - log.Printf("Expected tcpip-forward request, got: %s", req.Type) - return nil - case <-time.After(500 * time.Millisecond): - log.Println("No forwarding request received") - return nil } } @@ -253,35 +254,35 @@ func (s *session) HandleGlobalRequest(GlobalRequest <-chan *ssh.Request) error { return nil } -func (s *session) parseForwardPayload(payload []byte) (address string, port uint16, err error) { +func (s *session) parseForwardPayload(payload []byte) (address string, port uint16, reserved bool, err error) { var forwardPayload struct { BindAddr string BindPort uint32 } if err = ssh.Unmarshal(payload, &forwardPayload); err != nil { - return "", 0, fmt.Errorf("failed to unmarshal forward payload: %w", err) + return "", 0, false, fmt.Errorf("failed to unmarshal forward payload: %w", err) } if forwardPayload.BindPort > 65535 { - return "", 0, fmt.Errorf("port is larger than allowed port of 65535") + return "", 0, false, fmt.Errorf("port is larger than allowed port of 65535") } port = uint16(forwardPayload.BindPort) if isBlockedPort(port) { - return "", 0, fmt.Errorf("port is blocked") + return "", 0, false, fmt.Errorf("port is blocked") } if port == 0 { unassigned, ok := s.lifecycle.PortRegistry().Unassigned() if !ok { - return "", 0, fmt.Errorf("no available port") + return "", 0, false, fmt.Errorf("no available port") } - return forwardPayload.BindAddr, unassigned, nil + return forwardPayload.BindAddr, unassigned, true, nil } - return forwardPayload.BindAddr, port, nil + return forwardPayload.BindAddr, port, false, nil } func (s *session) denyForwardingRequest(req *ssh.Request, key *types.SessionKey, listener io.Closer, msg string) error { @@ -325,7 +326,7 @@ func (s *session) finalizeForwarding(req *ssh.Request, portToBind uint16, listen } func (s *session) HandleTCPIPForward(req *ssh.Request) error { - address, port, err := s.parseForwardPayload(req.Payload) + address, port, reserved, err := s.parseForwardPayload(req.Payload) if err != nil { return s.denyForwardingRequest(req, nil, nil, fmt.Sprintf("cannot parse forwarded payload: %s", err.Error())) } @@ -334,7 +335,7 @@ func (s *session) HandleTCPIPForward(req *ssh.Request) error { case 80, 443: return s.HandleHTTPForward(req, port) default: - return s.HandleTCPForward(req, address, port) + return s.HandleTCPForward(req, address, port, reserved) } } @@ -355,30 +356,40 @@ func (s *session) HandleHTTPForward(req *ssh.Request, portToBind uint16) error { return nil } -func (s *session) HandleTCPForward(req *ssh.Request, addr string, portToBind uint16) error { - if claimed := s.lifecycle.PortRegistry().Claim(portToBind); !claimed { - return s.denyForwardingRequest(req, nil, nil, fmt.Sprintf("Port %d is already in use or restricted", portToBind)) +func (s *session) HandleTCPForward(req *ssh.Request, addr string, portToBind uint16, reserved bool) error { + if !reserved { + if claimed := s.lifecycle.PortRegistry().Claim(portToBind); !claimed { + return s.denyForwardingRequest(req, nil, nil, fmt.Sprintf("Port %d is already in use or restricted", portToBind)) + } + } + + releasePort := func() { + if err := s.lifecycle.PortRegistry().SetStatus(portToBind, false); err != nil { + log.Printf("failed to release port %d: %v", portToBind, err) + } } tcpServer := transport.NewTCPServer(portToBind, s.forwarder) listener, err := tcpServer.Listen() if err != nil { + releasePort() return s.denyForwardingRequest(req, nil, listener, fmt.Sprintf("Port %d is already in use or restricted", portToBind)) } key := types.SessionKey{Id: fmt.Sprintf("%d", portToBind), Type: types.TunnelTypeTCP} if !s.registry.Register(key, s) { + releasePort() return s.denyForwardingRequest(req, nil, listener, fmt.Sprintf("Failed to register TunnelTypeTCP client with id: %s", key.Id)) } err = s.finalizeForwarding(req, portToBind, listener, types.TunnelTypeTCP, key.Id) if err != nil { + releasePort() return s.denyForwardingRequest(req, &key, listener, fmt.Sprintf("Failed to finalize forwarding: %s", err)) } go func() { - err = tcpServer.Serve(listener) - if err != nil { + if err := tcpServer.Serve(listener); err != nil { log.Printf("Failed serving tcp server: %s\n", err) } }() diff --git a/internal/session/session_test.go b/internal/session/session_test.go index 9cff8be..ad3cde3 100644 --- a/internal/session/session_test.go +++ b/internal/session/session_test.go @@ -38,8 +38,9 @@ type mockConfig struct { config.Config } -func (m *mockConfig) Domain() string { return m.Called().String(0) } -func (m *mockConfig) SSHPort() string { return m.Called().String(0) } +func (m *mockConfig) Domain() string { return m.Called().String(0) } +func (m *mockConfig) FrontendURL() string { return m.Called().String(0) } +func (m *mockConfig) SSHPort() string { return m.Called().String(0) } func (m *mockConfig) Mode() types.ServerMode { args := m.Called() if args.Get(0) == nil { @@ -396,6 +397,12 @@ func TestHandleTCPIPForward_Table(t *testing.T) { err := s.HandleTCPIPForward(req) assert.NoError(t, err) assert.Equal(t, uint16(12345), s.forwarder.ForwardedPort()) + + defer func() { + if l := s.forwarder.Listener(); l != nil { + _ = l.Close() + } + }() }) t.Run("Invalid Payload", func(t *testing.T) { @@ -699,6 +706,7 @@ func TestForwardingFailures(t *testing.T) { s, mRegistry, mPort, _, _, sReqs, cConn, cleanup := setup(t) defer cleanup() mPort.On("Claim", mock.Anything).Return(true) + mPort.On("SetStatus", uint16(1234), false).Return(nil) mRegistry.On("Register", mock.Anything, mock.Anything).Return(false) payload := make([]byte, 4+9+4) @@ -717,7 +725,7 @@ func TestForwardingFailures(t *testing.T) { }) t.Run("Finalize Forwarding Failure", func(t *testing.T) { - s, mRegistry, _, mRandom, _, sReqs, cConn, cleanup := setup(t) + s, mRegistry, _, mRandom, sConn, sReqs, cConn, cleanup := setup(t) defer cleanup() mRandom.On("String", 20).Return("test-slug", nil) mRegistry.On("Register", mock.Anything, mock.Anything).Return(true) @@ -736,7 +744,7 @@ func TestForwardingFailures(t *testing.T) { err := cConn.Close() assert.NoError(t, err) - time.Sleep(50 * time.Millisecond) + _ = sConn.Wait() err = s.HandleTCPIPForward(req) assert.Error(t, err) @@ -758,6 +766,7 @@ func TestForwardingFailures(t *testing.T) { }(l) _, portStr, _ := net.SplitHostPort(l.Addr().String()) port, _ := strconv.Atoi(portStr) + mPort.On("SetStatus", uint16(port), false).Return(nil) payload := make([]byte, 4+9+4) binary.BigEndian.PutUint32(payload[0:4], 9) @@ -807,7 +816,7 @@ func (m *mockNewChanFail) Accept() (ssh.Channel, <-chan *ssh.Request, error) { } func TestWaitForTCPIPForward_EdgeCases(t *testing.T) { - t.Run("Wrong Request Type", func(t *testing.T) { + t.Run("Wrong Request Type Then Timeout", func(t *testing.T) { _, sReqs, _, cConn, cleanup := setupSSH(t) defer cleanup() @@ -817,10 +826,65 @@ func TestWaitForTCPIPForward_EdgeCases(t *testing.T) { _, _, _ = cConn.SendRequest("not-tcpip-forward", true, nil) }() + start := time.Now() req := s.waitForTCPIPForward() + elapsed := time.Since(start) + if req != nil { t.Error("expected nil request") } + if elapsed < 400*time.Millisecond { + t.Errorf("expected timeout ~500ms, got %v", elapsed) + } + }) + + t.Run("Multiple Non-Forward Requests Then Success", func(t *testing.T) { + _, sReqs, _, cConn, cleanup := setupSSH(t) + defer cleanup() + + s := &session{initialReq: sReqs} + + go func() { + time.Sleep(100 * time.Millisecond) + _, _, _ = cConn.SendRequest("keepalive@openssh.com", false, nil) + time.Sleep(100 * time.Millisecond) + _, _, _ = cConn.SendRequest("hostkeys-00@openssh.com", false, nil) + time.Sleep(100 * time.Millisecond) + _, _, _ = cConn.SendRequest("tcpip-forward", true, nil) + }() + + req := s.waitForTCPIPForward() + if req == nil { + t.Error("expected tcpip-forward request, got nil") + } + if req != nil && req.Type != "tcpip-forward" { + t.Errorf("expected tcpip-forward, got %s", req.Type) + } + }) + + t.Run("Timeout After Non-Forward Requests", func(t *testing.T) { + _, sReqs, _, cConn, cleanup := setupSSH(t) + defer cleanup() + + s := &session{initialReq: sReqs} + + go func() { + time.Sleep(100 * time.Millisecond) + _, _, _ = cConn.SendRequest("keepalive@openssh.com", false, nil) + time.Sleep(100 * time.Millisecond) + _, _, _ = cConn.SendRequest("hostkeys-00@openssh.com", false, nil) + }() + + start := time.Now() + req := s.waitForTCPIPForward() + elapsed := time.Since(start) + + if req != nil { + t.Error("expected nil request after timeout") + } + if elapsed < 400*time.Millisecond { + t.Errorf("expected timeout ~500ms after last request, got %v", elapsed) + } }) t.Run("Channel Closed", func(t *testing.T) { @@ -982,7 +1046,7 @@ func TestParseForwardPayload_Errors(t *testing.T) { s := &session{} t.Run("Short Address", func(t *testing.T) { - _, _, err := s.parseForwardPayload([]byte{0, 0, 0, 4}) + _, _, _, err := s.parseForwardPayload([]byte{0, 0, 0, 4}) if err == nil { t.Error("expected error, got nil") } @@ -990,7 +1054,7 @@ func TestParseForwardPayload_Errors(t *testing.T) { t.Run("Short Port", func(t *testing.T) { payload := append([]byte{0, 0, 0, 4}, []byte("addr")...) - _, _, err := s.parseForwardPayload(payload) + _, _, _, err := s.parseForwardPayload(payload) if err == nil { t.Error("expected error, got nil") } @@ -1001,7 +1065,7 @@ func TestParseForwardPayload_Errors(t *testing.T) { portBuf := make([]byte, 4) binary.BigEndian.PutUint32(portBuf, 22) payload = append(payload, portBuf...) - _, _, err := s.parseForwardPayload(payload) + _, _, _, err := s.parseForwardPayload(payload) if err == nil { t.Error("expected error, got nil") } else if !strings.Contains(err.Error(), "port is block") { @@ -1165,7 +1229,7 @@ func TestHandleTCPForward_Failures(t *testing.T) { s, _, mPort, _, sReqs, cConn, cleanup := setup(t) defer cleanup() mPort.On("Claim", mock.Anything).Return(false) - err := s.HandleTCPForward(getReq(t, cConn, sReqs), "localhost", 1234) + err := s.HandleTCPForward(getReq(t, cConn, sReqs), "localhost", 1234, false) if err == nil { t.Error("expected error, got nil") } else if !strings.Contains(err.Error(), "already in use") { @@ -1186,44 +1250,50 @@ func TestHandleTCPForward_Failures(t *testing.T) { assert.NoError(t, err) }(l) port := uint16(l.Addr().(*net.TCPAddr).Port) + mPort.On("SetStatus", port, false).Return(nil) - err = s.HandleTCPForward(getReq(t, cConn, sReqs), "localhost", port) + err = s.HandleTCPForward(getReq(t, cConn, sReqs), "localhost", port, false) if err == nil { t.Error("expected error, got nil") } else if !strings.Contains(err.Error(), "already in use") { t.Errorf("expected error to contain %q, got %q", "already in use", err.Error()) } + mPort.AssertExpectations(t) }) t.Run("Registry Register fail", func(t *testing.T) { s, mRegistry, mPort, _, sReqs, cConn, cleanup := setup(t) defer cleanup() mPort.On("Claim", mock.Anything).Return(true) + mPort.On("SetStatus", mock.AnythingOfType("uint16"), false).Return(nil) mRegistry.On("Register", mock.Anything, mock.Anything).Return(false) - err := s.HandleTCPForward(getReq(t, cConn, sReqs), "localhost", 0) + err := s.HandleTCPForward(getReq(t, cConn, sReqs), "localhost", 0, false) if err == nil { t.Error("expected error, got nil") } else if !strings.Contains(err.Error(), "Failed to register") { t.Errorf("expected error to contain %q, got %q", "Failed to register", err.Error()) } + mPort.AssertExpectations(t) }) t.Run("Finalize fail (Reply fail)", func(t *testing.T) { - s, mRegistry, mPort, _, sReqs, cConn, cleanup := setup(t) + s, mRegistry, mPort, sConn, sReqs, cConn, cleanup := setup(t) defer cleanup() mPort.On("Claim", mock.Anything).Return(true) + mPort.On("SetStatus", mock.AnythingOfType("uint16"), false).Return(nil) mRegistry.On("Register", mock.Anything, mock.Anything).Return(true) req := getReq(t, cConn, sReqs) err := cConn.Close() assert.NoError(t, err) - time.Sleep(100 * time.Millisecond) + _ = sConn.Wait() - err = s.HandleTCPForward(req, "localhost", 0) + err = s.HandleTCPForward(req, "localhost", 0, false) if err == nil { t.Error("expected error, got nil") } else if !strings.Contains(err.Error(), "Failed to finalize forwarding") { t.Errorf("expected error to contain %q, got %q", "Failed to finalize forwarding", err.Error()) } + mPort.AssertExpectations(t) }) } diff --git a/internal/session/slug/slug.go b/internal/session/slug/slug.go index b9684d1..7fedc59 100644 --- a/internal/session/slug/slug.go +++ b/internal/session/slug/slug.go @@ -1,11 +1,14 @@ package slug +import "sync" + type Slug interface { String() string Set(slug string) } type slug struct { + mu sync.RWMutex slug string } @@ -16,9 +19,13 @@ func New() Slug { } func (s *slug) String() string { + s.mu.RLock() + defer s.mu.RUnlock() return s.slug } func (s *slug) Set(slug string) { + s.mu.Lock() + defer s.mu.Unlock() s.slug = slug } diff --git a/internal/transport/http_test.go b/internal/transport/http_test.go index cd3cf68..847931c 100644 --- a/internal/transport/http_test.go +++ b/internal/transport/http_test.go @@ -55,8 +55,7 @@ func TestHTTPServer_Serve(t *testing.T) { go func() { time.Sleep(100 * time.Millisecond) - err = listener.Close() - assert.NoError(t, err) + _ = listener.Close() }() err = srv.Serve(listener) diff --git a/internal/transport/httphandler.go b/internal/transport/httphandler.go index bcb74e8..23c3694 100644 --- a/internal/transport/httphandler.go +++ b/internal/transport/httphandler.go @@ -1,6 +1,7 @@ package transport import ( + "bufio" "bytes" "context" "errors" @@ -52,25 +53,43 @@ func (hh *httpHandler) badRequest(conn net.Conn) error { return nil } +func readHTTPHeader(br *bufio.Reader, limit int) ([]byte, error) { + var headerBuf []byte + for { + line, err := br.ReadSlice('\n') + headerBuf = append(headerBuf, line...) + if errors.Is(err, bufio.ErrBufferFull) { + if len(headerBuf) > limit { + return nil, fmt.Errorf("headers too large") + } + continue + } + if err != nil { + return nil, err + } + if bytes.HasSuffix(headerBuf, []byte("\r\n\r\n")) { + return headerBuf, nil + } + if len(headerBuf) > limit { + return nil, fmt.Errorf("headers too large") + } + } +} + func (hh *httpHandler) Handler(conn net.Conn, isTLS bool) { defer hh.closeConnection(conn) _ = conn.SetReadDeadline(time.Now().Add(10 * time.Second)) - buf := make([]byte, hh.config.HeaderSize()) - n, err := conn.Read(buf) + br := bufio.NewReaderSize(conn, hh.config.HeaderSize()) + headerBuf, err := readHTTPHeader(br, hh.config.HeaderSize()) if err != nil { _ = hh.badRequest(conn) return } - if idx := bytes.Index(buf[:n], []byte("\r\n\r\n")); idx == -1 { - _ = hh.badRequest(conn) - return - } - _ = conn.SetReadDeadline(time.Time{}) - reqhf, err := header.NewRequest(buf[:n]) + reqhf, err := header.NewRequest(headerBuf) if err != nil { log.Printf("Error creating request header: %v", err) _ = hh.badRequest(conn) @@ -97,11 +116,11 @@ func (hh *httpHandler) Handler(conn net.Conn, isTLS bool) { Type: types.TunnelTypeHTTP, }) if err != nil { - _ = hh.redirect(conn, http.StatusMovedPermanently, fmt.Sprintf("https://tunnl.live/tunnel-not-found?slug=%s\r\n", slug)) + _ = hh.redirect(conn, http.StatusMovedPermanently, fmt.Sprintf("%s/tunnel-not-found?slug=%s\r\n", hh.config.FrontendURL(), slug)) return } - hw := stream.New(conn, conn, conn.RemoteAddr()) + hw := stream.New(conn, br, conn.RemoteAddr()) defer func(hw stream.HTTP) { err = hw.Close() if err != nil { diff --git a/internal/transport/httphandler_test.go b/internal/transport/httphandler_test.go index 0be77dc..329e539 100644 --- a/internal/transport/httphandler_test.go +++ b/internal/transport/httphandler_test.go @@ -223,6 +223,7 @@ func TestNewHTTPHandler(t *testing.T) { msr := new(MockSessionRegistry) mockConfig := &MockConfig{} mockConfig.On("Domain").Return("domain") + mockConfig.On("FrontendURL").Return("https://domain") mockConfig.On("TLSRedirect").Return(false) hh := newHTTPHandler(mockConfig, msr) assert.NotNil(t, hh) @@ -290,7 +291,7 @@ func TestHandler(t *testing.T) { isTLS: true, redirectTLS: false, request: []byte("GET / HTTP/1.1\r\nHost: test.domain\r\n\r\n"), - expected: []byte("HTTP/1.1 301 Moved Permanently\r\nLocation: https://tunnl.live/tunnel-not-found?slug=test\r\nContent-Length: 0\r\nConnection: close\r\n\r\n"), + expected: []byte("HTTP/1.1 301 Moved Permanently\r\nLocation: https://example.com/tunnel-not-found?slug=test\r\nContent-Length: 0\r\nConnection: close\r\n\r\n"), setupMocks: func(msr *MockSessionRegistry) { msr.On("Get", types.SessionKey{ Id: "test", @@ -321,8 +322,14 @@ func TestHandler(t *testing.T) { isTLS: false, redirectTLS: false, request: []byte(""), - expected: []byte("HTTP/1.1 400 Bad Request\r\n\r\n"), - setupMocks: func(msr *MockSessionRegistry) { + expected: []byte(""), + setupConn: func() (net.Conn, net.Conn) { + mc := new(MockConn) + mc.ReadBuffer = bytes.NewBuffer(nil) + mc.On("SetReadDeadline", mock.Anything).Return(nil) + mc.On("Write", []byte("HTTP/1.1 400 Bad Request\r\n\r\n")).Return(0, nil) + mc.On("Close").Return(nil) + return mc, nil }, }, { @@ -610,6 +617,7 @@ func TestHandler(t *testing.T) { mockConfig := &MockConfig{} port := "0" mockConfig.On("Domain").Return("example.com") + mockConfig.On("FrontendURL").Return("https://example.com") mockConfig.On("HTTPPort").Return(port) mockConfig.On("HeaderSize").Return(4096) mockConfig.On("TLSRedirect").Return(true) @@ -715,3 +723,114 @@ func TestHandler(t *testing.T) { }) } } + +func TestHandlerForwardsPostBody(t *testing.T) { + mockSessionRegistry := new(MockSessionRegistry) + mockConfig := &MockConfig{} + mockConfig.On("Domain").Return("example.com") + mockConfig.On("FrontendURL").Return("https://example.com") + mockConfig.On("HTTPPort").Return("0") + mockConfig.On("HeaderSize").Return(4096) + mockConfig.On("TLSRedirect").Return(true) + hh := &httpHandler{ + sessionRegistry: mockSessionRegistry, + config: mockConfig, + } + + mockSession := new(MockSession) + mockForwarder := new(MockForwarder) + mockSSHChannel := new(MockSSHChannel) + + mockSessionRegistry.On("Get", types.SessionKey{ + Id: "test", + Type: types.TunnelTypeHTTP, + }).Return(mockSession, nil) + mockSession.On("Forwarder").Return(mockForwarder) + + reqCh := make(chan *ssh.Request) + mockForwarder.On("OpenForwardedChannel", mock.Anything, mock.Anything).Return(mockSSHChannel, (<-chan *ssh.Request)(reqCh), nil) + + var mu sync.Mutex + var capturedHeaders []byte + mockSSHChannel.On("Write", mock.Anything).Run(func(args mock.Arguments) { + mu.Lock() + capturedHeaders = append(capturedHeaders, args.Get(0).([]byte)...) + mu.Unlock() + }).Return(0, nil) + mockSSHChannel.On("Close").Return(nil) + + bodyChan := make(chan string, 1) + mockForwarder.On("HandleConnection", mock.Anything, mockSSHChannel).Run(func(args mock.Arguments) { + w := args.Get(0).(io.ReadWriter) + buf := make([]byte, len("hello=world")) + if _, err := io.ReadFull(w, buf); err != nil { + bodyChan <- "" + } else { + bodyChan <- string(buf) + } + _, _ = w.Write([]byte("HTTP/1.1 200 OK\r\nContent-Length: 2\r\n\r\nok")) + }) + + go func() { + for range reqCh { + } + }() + + serverConn, clientConn := net.Pipe() + defer func() { + _ = clientConn.Close() + }() + + remoteAddr, _ := net.ResolveTCPAddr("tcp", "127.0.0.1:12345") + wrappedServerConn := &wrappedConn{Conn: serverConn, remoteAddr: remoteAddr} + + go hh.Handler(wrappedServerConn, true) + + request := []byte("POST / HTTP/1.1\r\nHost: test.domain\r\nContent-Type: application/x-www-form-urlencoded\r\nContent-Length: 11\r\n\r\nhello=world") + go func() { + _, _ = clientConn.Write(request) + }() + + var response []byte + respDone := make(chan struct{}) + go func() { + defer close(respDone) + buf := make([]byte, 4096) + for { + n, err := clientConn.Read(buf) + if err != nil { + break + } + response = append(response, buf[:n]...) + if bytes.Contains(response, []byte("\r\n\r\nok")) { + break + } + } + }() + + select { + case body := <-bodyChan: + assert.Equal(t, "hello=world", body) + case <-time.After(5 * time.Second): + t.Fatal("timed out waiting for forwarded body") + } + + select { + case <-respDone: + resStr := string(response) + assert.True(t, strings.HasPrefix(resStr, "HTTP/1.1 200 OK\r\n")) + assert.Contains(t, resStr, "Server: Tunnel Please\r\n") + assert.True(t, strings.HasSuffix(resStr, "\r\n\r\nok")) + case <-time.After(5 * time.Second): + t.Fatal("timed out waiting for response") + } + + mu.Lock() + hdrStr := string(capturedHeaders) + mu.Unlock() + assert.Contains(t, hdrStr, "POST / HTTP/1.1\r\n") + assert.Contains(t, hdrStr, "Content-Length: 11\r\n") + assert.Contains(t, hdrStr, "X-Forwarded-For: 127.0.0.1\r\n") + + mockSessionRegistry.AssertExpectations(t) +} diff --git a/internal/transport/https_test.go b/internal/transport/https_test.go index 6081d97..42bbe72 100644 --- a/internal/transport/https_test.go +++ b/internal/transport/https_test.go @@ -63,8 +63,7 @@ func TestHTTPSServer_Serve(t *testing.T) { go func() { time.Sleep(100 * time.Millisecond) - err = listener.Close() - assert.NoError(t, err) + _ = listener.Close() }() err = srv.Serve(listener) diff --git a/internal/transport/tcp_test.go b/internal/transport/tcp_test.go index c4c4963..761b902 100644 --- a/internal/transport/tcp_test.go +++ b/internal/transport/tcp_test.go @@ -45,8 +45,7 @@ func TestTCPServer_Serve(t *testing.T) { go func() { time.Sleep(100 * time.Millisecond) - err = listener.Close() - assert.NoError(t, err) + _ = listener.Close() }() err = srv.Serve(listener) diff --git a/internal/transport/tls_test.go b/internal/transport/tls_test.go index 3de8066..eba738b 100644 --- a/internal/transport/tls_test.go +++ b/internal/transport/tls_test.go @@ -25,6 +25,7 @@ type MockConfig struct { } func (m *MockConfig) Domain() string { return m.Called().String(0) } +func (m *MockConfig) FrontendURL() string { return m.Called().String(0) } func (m *MockConfig) SSHPort() string { return m.Called().String(0) } func (m *MockConfig) HTTPPort() string { return m.Called().String(0) } func (m *MockConfig) HTTPSPort() string { return m.Called().String(0) }