Compare commits

...

13 Commits

Author SHA1 Message Date
bagas 53cc5c2328 fix(registry): data race when updating slug
Docker Build and Push / Run Tests (push) Successful in 2m33s
Tests / Run Tests (pull_request) Successful in 2m27s
Docker Build and Push / Build and Push Docker Image (push) Successful in 14m51s
2026-07-19 10:54:26 +07:00
bagas bd8ab00fa5 fix(session): release port when error happen
Docker Build and Push / Run Tests (push) Successful in 2m33s
Docker Build and Push / Build and Push Docker Image (push) Successful in 14m44s
Tests / Run Tests (pull_request) Successful in 2m28s
2026-07-18 17:43:17 +07:00
bagas 4fcc41eb8a fix(port,session): harden port registry and fix concurrent allocation
Docker Build and Push / Run Tests (push) Successful in 2m35s
Docker Build and Push / Build and Push Docker Image (push) Successful in 14m55s
2026-07-18 17:10:35 +07:00
bagas 34000962ef fix(httphandler): use readSlice for reading http header
Docker Build and Push / Run Tests (push) Successful in 2m30s
Docker Build and Push / Build and Push Docker Image (push) Successful in 14m44s
Tests / Run Tests (pull_request) Successful in 2m27s
2026-07-18 15:29:01 +07:00
bagas 9d623ae3f9 fix(httphandler): post/put request hang
Docker Build and Push / Run Tests (push) Successful in 2m34s
Docker Build and Push / Build and Push Docker Image (push) Successful in 14m49s
2026-07-18 13:19:51 +07:00
bagas 6fd25387f7 fix(session): resolve race conditions (#151)
SonarQube Scan / SonarQube Trigger (push) Successful in 4m1s
Docker Build and Push / Run Tests (push) Successful in 2m28s
Docker Build and Push / Build and Push Docker Image (push) Successful in 14m52s
## Summary

Comprehensive fixes for race conditions, concurrency bugs, and a critical session handling bug across the session package and its child packages.

## Changes

### Bug Fix: waitForTCPIPForward fails on non-tcpip-forward global requests

- Changed `waitForTCPIPForward()` from a single select to a for loop that drains non-tcpip-forward requests gracefully
- Non-tcpip-forward requests now receive a false reply and the function continues waiting for the actual forward request
- The 500ms timeout resets after each request to handle slow clients with multiple early requests
- Updated and added tests to cover the new behavior: Wrong_Request_Type_Then_Timeout, Multiple_Non-Forward_Requests_Then_Success, Timeout_After_Non-Forward_Requests

### Race Condition Fixes

- Added sync.RWMutex to the forwarder struct to protect concurrent field access
- Fixed data race in forwarder setters (`SetType()`, `SetForwardedPort()`, `SetListener()`) by adding `Lock()`/`Unlock()` protection
- Fixed data race in forwarder getters (`TunnelType()`, `ForwardedPort()`, `Listener()`) by adding `RLock()`/`RUnlock()` protection
- Added sync.RWMutex to the slug struct to protect concurrent field access
- Fixed data race in `slug.Set()` by adding `Lock()`/`Unlock()` protection
- Fixed data race in `slug.String()` by adding `RLock()`/`RUnlock()` protection

### Concurrency Safety Improvements

- Changed `OpenForwardedChannel()` to use `f.ForwardedPort()` getter instead of direct field access
- Fixed `Close()` in forwarder to use `f.Listener()` getter consistently instead of accessing `f.listener` directly
- Fixed data race in `HandleTCPForward()` where the goroutine running `tcpServer.Serve(listener)` wrote to the outer err variable from the enclosing function scope
- Changed err = `tcpServer.Serve(listener)` to `if err := tcpServer.Serve(listener)` so the goroutine uses a local variable

Reviewed-on: #151
Co-authored-by: Bagas <bagas@fossy.my.id>
Co-committed-by: Bagas <bagas@fossy.my.id>
2026-07-17 22:22:25 +07:00
bagas e8be839dac Fix race conditions and improve lifecycle safety (#150)
SonarQube Scan / SonarQube Trigger (push) Successful in 4m1s
Docker Build and Push / Run Tests (push) Successful in 2m28s
Docker Build and Push / Build and Push Docker Image (push) Successful in 18m27s
Tests / Run Tests (pull_request) Successful in 2m29s
## Summary
Comprehensive fixes for race conditions, concurrency bugs, and code quality improvements in the lifecycle module and related test files.

## Changes

### Race Condition Fixes
- Fixed data race in `SetChannel()` and `Channel()` by adding mutex protection
- Fixed data race in `StartedAt()` by adding mutex protection
- Fixed test race conditions in `transport` and `interaction` packages where variables were shared between goroutines

### Concurrency Safety Improvements
- Refactored `Close()` to release mutex before calling external code, preventing potential deadlocks
- Made `SetChannel()` reject calls after `Close()` to prevent setting channel on torn-down lifecycle
- Made `SetStatus()` ignore calls after `Close()` to ensure CLOSED is a terminal state
- Added nil check to `SetChannel()` to prevent masking upstream bugs

### Code Quality Improvements
- Refactored `Close()` to extract helper methods: `closeChannel()`, `closeConnection()`, `cleanupRegistry()`, `cleanupForwarder()`
- Fixed `isClosedError()` to remove redundant string comparison
- Added `PortRegistry` interface for better interface segregation (removed dependency on full `port.Port`)
- Made `cleanupRegistry()` conditional on non-empty slug to avoid unnecessary registry calls
- Fixed `startedAt` to be set when session starts running instead of at object creation

Reviewed-on: #150
Co-authored-by: Bagas <bagas@fossy.my.id>
Co-committed-by: Bagas <bagas@fossy.my.id>
2026-07-16 17:31:14 +07:00
bagas d7ec5e852b Merge pull request 'refactor: move types package to internal' (#122) from refactor/move-types-to-internal into staging
SonarQube Scan / SonarQube Trigger (push) Successful in 3m32s
Tests / Run Tests (pull_request) Successful in 2m6s
Reviewed-on: #122
2026-03-30 11:53:16 +07:00
bagas 4af011b91d refactor: move types package to internal
Tests / Run Tests (pull_request) Successful in 2m7s
2026-03-30 11:50:24 +07:00
bagas dc0a4af27f Merge pull request 'refactor: move session package to internal' (#121) from refactor/move-session-to-internal into staging
SonarQube Scan / SonarQube Trigger (push) Successful in 3m48s
Reviewed-on: #121
2026-03-30 11:47:35 +07:00
bagas 032b6e5096 Merge pull request 'refactor: move server package to internal' (#120) from refactor/move-server-to-internal into staging
SonarQube Scan / SonarQube Trigger (push) Successful in 3m36s
Reviewed-on: #120
2026-03-30 11:38:52 +07:00
bagas ae71b46482 refactor: move session package to internal
Tests / Run Tests (pull_request) Successful in 2m4s
2026-03-30 11:37:59 +07:00
bagas 7ecc56aa3c refactor: move server package to internal
Tests / Run Tests (pull_request) Successful in 2m37s
2026-03-30 11:18:44 +07:00
35 changed files with 629 additions and 179 deletions
+2 -2
View File
@@ -15,10 +15,10 @@ import (
"tunnel_pls/internal/port"
"tunnel_pls/internal/random"
"tunnel_pls/internal/registry"
"tunnel_pls/internal/server"
"tunnel_pls/internal/transport"
"tunnel_pls/internal/types"
"tunnel_pls/internal/version"
"tunnel_pls/server"
"tunnel_pls/types"
"golang.org/x/crypto/ssh"
)
+2 -2
View File
@@ -14,8 +14,8 @@ import (
"tunnel_pls/internal/config"
"tunnel_pls/internal/port"
"tunnel_pls/internal/registry"
"tunnel_pls/session/slug"
"tunnel_pls/types"
"tunnel_pls/internal/session/slug"
"tunnel_pls/internal/types"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/mock"
+3 -1
View File
@@ -1,6 +1,8 @@
package config
import "tunnel_pls/types"
import (
"tunnel_pls/internal/types"
)
type Config interface {
Domain() string
+1 -1
View File
@@ -3,7 +3,7 @@ package config
import (
"os"
"testing"
"tunnel_pls/types"
"tunnel_pls/internal/types"
"github.com/stretchr/testify/assert"
)
+2 -2
View File
@@ -6,7 +6,7 @@ import (
"os"
"strconv"
"strings"
"tunnel_pls/types"
"tunnel_pls/internal/types"
"github.com/joho/godotenv"
)
@@ -32,7 +32,7 @@ type config struct {
bufferSize int
headerSize int
pprofEnabled bool
pprofPort string
+1 -1
View File
@@ -9,7 +9,7 @@ import (
"time"
"tunnel_pls/internal/config"
"tunnel_pls/internal/registry"
"tunnel_pls/types"
"tunnel_pls/internal/types"
proto "git.fossy.my.id/bagas/tunnel-please-grpc/gen"
"google.golang.org/grpc"
+8 -9
View File
@@ -7,14 +7,13 @@ import (
"io"
"testing"
"time"
"tunnel_pls/internal/session/forwarder"
"tunnel_pls/internal/session/interaction"
"tunnel_pls/internal/session/lifecycle"
"tunnel_pls/internal/session/slug"
"tunnel_pls/internal/types"
"tunnel_pls/internal/port"
"tunnel_pls/internal/registry"
"tunnel_pls/session/forwarder"
"tunnel_pls/session/interaction"
"tunnel_pls/session/lifecycle"
"tunnel_pls/session/slug"
"tunnel_pls/types"
proto "git.fossy.my.id/bagas/tunnel-please-grpc/gen"
"github.com/stretchr/testify/assert"
@@ -885,16 +884,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 {
+14 -11
View File
@@ -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
}
+52 -3
View File
@@ -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)
}
})
}
}
+8 -8
View File
@@ -3,11 +3,11 @@ package registry
import (
"fmt"
"sync"
"tunnel_pls/session/forwarder"
"tunnel_pls/session/interaction"
"tunnel_pls/session/lifecycle"
"tunnel_pls/session/slug"
"tunnel_pls/types"
"tunnel_pls/internal/session/forwarder"
"tunnel_pls/internal/session/interaction"
"tunnel_pls/internal/session/lifecycle"
"tunnel_pls/internal/session/slug"
"tunnel_pls/internal/types"
)
type Key = types.SessionKey
@@ -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
+8 -9
View File
@@ -4,12 +4,11 @@ import (
"sync"
"testing"
"time"
"tunnel_pls/internal/port"
"tunnel_pls/session/forwarder"
"tunnel_pls/session/interaction"
"tunnel_pls/session/lifecycle"
"tunnel_pls/session/slug"
"tunnel_pls/types"
"tunnel_pls/internal/session/forwarder"
"tunnel_pls/internal/session/interaction"
"tunnel_pls/internal/session/lifecycle"
"tunnel_pls/internal/session/slug"
"tunnel_pls/internal/types"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/mock"
@@ -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) }
@@ -13,7 +13,7 @@ import (
"tunnel_pls/internal/port"
"tunnel_pls/internal/random"
"tunnel_pls/internal/registry"
"tunnel_pls/session"
"tunnel_pls/internal/session"
"golang.org/x/crypto/ssh"
)
@@ -10,8 +10,8 @@ import (
"testing"
"time"
"tunnel_pls/internal/registry"
"tunnel_pls/session/slug"
"tunnel_pls/types"
"tunnel_pls/internal/session/slug"
"tunnel_pls/internal/types"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/mock"
@@ -10,8 +10,8 @@ import (
"strconv"
"sync"
"tunnel_pls/internal/config"
"tunnel_pls/session/slug"
"tunnel_pls/types"
"tunnel_pls/internal/session/slug"
"tunnel_pls/internal/types"
"golang.org/x/crypto/ssh"
)
@@ -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
}
@@ -10,8 +10,8 @@ import (
"sync/atomic"
"testing"
"time"
"tunnel_pls/session/slug"
"tunnel_pls/types"
"tunnel_pls/internal/session/slug"
"tunnel_pls/internal/types"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/mock"
@@ -6,8 +6,8 @@ import (
"sync"
"tunnel_pls/internal/config"
"tunnel_pls/internal/random"
"tunnel_pls/session/slug"
"tunnel_pls/types"
"tunnel_pls/internal/session/slug"
"tunnel_pls/internal/types"
"github.com/charmbracelet/bubbles/help"
"github.com/charmbracelet/bubbles/key"
@@ -7,7 +7,7 @@ import (
"net"
"testing"
"time"
"tunnel_pls/types"
"tunnel_pls/internal/types"
"github.com/charmbracelet/bubbles/key"
"github.com/charmbracelet/bubbles/list"
@@ -1922,10 +1922,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)
@@ -4,7 +4,7 @@ import (
"fmt"
"time"
"tunnel_pls/internal/random"
"tunnel_pls/types"
"tunnel_pls/internal/types"
"github.com/charmbracelet/bubbles/help"
"github.com/charmbracelet/bubbles/key"
@@ -3,7 +3,7 @@ package interaction
import (
"fmt"
"strings"
"tunnel_pls/types"
"tunnel_pls/internal/types"
"github.com/charmbracelet/bubbles/key"
"github.com/charmbracelet/bubbles/textinput"
@@ -2,14 +2,13 @@ package lifecycle
import (
"errors"
"fmt"
"io"
"net"
"sync"
"time"
portUtil "tunnel_pls/internal/port"
"tunnel_pls/session/slug"
"tunnel_pls/types"
"tunnel_pls/internal/session/slug"
"tunnel_pls/internal/types"
"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
}
@@ -5,8 +5,9 @@ import (
"errors"
"io"
"net"
"sync"
"testing"
"tunnel_pls/types"
"tunnel_pls/internal/types"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/mock"
@@ -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")
}
@@ -12,12 +12,12 @@ import (
portUtil "tunnel_pls/internal/port"
"tunnel_pls/internal/random"
"tunnel_pls/internal/registry"
"tunnel_pls/internal/session/forwarder"
"tunnel_pls/internal/session/interaction"
"tunnel_pls/internal/session/lifecycle"
"tunnel_pls/internal/session/slug"
"tunnel_pls/internal/transport"
"tunnel_pls/session/forwarder"
"tunnel_pls/session/interaction"
"tunnel_pls/session/lifecycle"
"tunnel_pls/session/slug"
"tunnel_pls/types"
"tunnel_pls/internal/types"
"golang.org/x/crypto/ssh"
)
@@ -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)
}
}()
@@ -15,8 +15,8 @@ import (
"time"
"tunnel_pls/internal/config"
"tunnel_pls/internal/registry"
"tunnel_pls/session/lifecycle"
"tunnel_pls/types"
"tunnel_pls/internal/session/lifecycle"
"tunnel_pls/internal/types"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/mock"
@@ -396,6 +396,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 +705,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 +724,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 +743,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 +765,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 +815,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 +825,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 +1045,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 +1053,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 +1064,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 +1228,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 +1249,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)
})
}
@@ -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
}
+1 -2
View File
@@ -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)
+29 -10
View File
@@ -1,6 +1,7 @@
package transport
import (
"bufio"
"bytes"
"context"
"errors"
@@ -16,7 +17,7 @@ import (
"tunnel_pls/internal/http/stream"
"tunnel_pls/internal/middleware"
"tunnel_pls/internal/registry"
"tunnel_pls/types"
"tunnel_pls/internal/types"
"golang.org/x/crypto/ssh"
)
@@ -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)
@@ -101,7 +120,7 @@ func (hh *httpHandler) Handler(conn net.Conn, isTLS bool) {
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 {
+123 -7
View File
@@ -11,11 +11,11 @@ import (
"testing"
"time"
"tunnel_pls/internal/registry"
"tunnel_pls/session/forwarder"
"tunnel_pls/session/interaction"
"tunnel_pls/session/lifecycle"
"tunnel_pls/session/slug"
"tunnel_pls/types"
"tunnel_pls/internal/session/forwarder"
"tunnel_pls/internal/session/interaction"
"tunnel_pls/internal/session/lifecycle"
"tunnel_pls/internal/session/slug"
"tunnel_pls/internal/types"
"golang.org/x/crypto/ssh"
@@ -321,8 +321,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
},
},
{
@@ -715,3 +721,113 @@ func TestHandler(t *testing.T) {
})
}
}
func TestHandlerForwardsPostBody(t *testing.T) {
mockSessionRegistry := new(MockSessionRegistry)
mockConfig := &MockConfig{}
mockConfig.On("Domain").Return("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)
}
+1 -2
View File
@@ -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)
+1 -2
View File
@@ -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)
+1 -1
View File
@@ -14,7 +14,7 @@ import (
"testing"
"time"
"tunnel_pls/internal/config"
"tunnel_pls/types"
"tunnel_pls/internal/types"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/mock"