Compare commits

..

2 Commits

Author SHA1 Message Date
bagas 1b369be5f8 ci: migrate docker build to multi-machine multi-arch matrix
Docker Build and Push / Run Tests (push) Successful in 2m9s
Docker Build and Push / Build (linux/arm64) (push) Failing after 1m18s
Docker Build and Push / Build (linux/amd64) (push) Failing after 2m8s
Docker Build and Push / Merge Multi-Arch Manifest (push) Has been skipped
2026-04-07 23:53:00 +07:00
bagas d029f9e93d chore: configure renovate to group and automerge dependencies
Tests / Run Tests (pull_request) Successful in 2m8s
2026-04-01 21:11:45 +07:00
18 changed files with 280 additions and 597 deletions
+151 -40
View File
@@ -9,7 +9,6 @@ jobs:
test: test:
name: Run Tests name: Run Tests
runs-on: ubuntu-latest runs-on: ubuntu-latest
steps: steps:
- name: Checkout repository - name: Checkout repository
uses: actions/checkout@v6 uses: actions/checkout@v6
@@ -29,16 +28,35 @@ jobs:
- name: Run tests - name: Run tests
run: go test -v -p 4 ./... run: go test -v -p 4 ./...
build-amd64:
build-and-push: name: Build (linux/amd64)
name: Build and Push Docker Image runs-on: [ubuntu-latest, amd64]
runs-on: ubuntu-latest
needs: test needs: test
steps: steps:
- name: Checkout repository - name: Checkout repository
uses: actions/checkout@v6 uses: actions/checkout@v6
- name: Extract version
id: version
run: |
VERSION=${GITHUB_REF#refs/tags/v}
echo "VERSION=$VERSION" >> $GITHUB_OUTPUT
echo "BUILD_DATE=$(date -u +'%Y-%m-%dT%H:%M:%SZ')" >> $GITHUB_OUTPUT
echo "COMMIT=${{ github.sha }}" >> $GITHUB_OUTPUT
if echo "$VERSION" | grep -qE '^[0-9]+\.[0-9]+\.[0-9]+(-[a-zA-Z0-9.]+)?$'; then
echo "MAJOR=$(echo "$VERSION" | cut -d. -f1)" >> $GITHUB_OUTPUT
echo "MINOR=$(echo "$VERSION" | cut -d. -f2)" >> $GITHUB_OUTPUT
if echo "$VERSION" | grep -q '-'; then
echo "IS_PRERELEASE=true" >> $GITHUB_OUTPUT
else
echo "IS_PRERELEASE=false" >> $GITHUB_OUTPUT
fi
else
echo "Invalid version format: $VERSION"
exit 1
fi
- name: Set up Docker Buildx - name: Set up Docker Buildx
uses: docker/setup-buildx-action@v3 uses: docker/setup-buildx-action@v3
@@ -49,7 +67,42 @@ jobs:
username: ${{ secrets.DOCKER_USERNAME }} username: ${{ secrets.DOCKER_USERNAME }}
password: ${{ secrets.DOCKER_PASSWORD }} password: ${{ secrets.DOCKER_PASSWORD }}
- name: Extract version and determine release type - name: Build and push by digest
id: build
uses: docker/build-push-action@v6
with:
context: .
platforms: linux/amd64
outputs: >-
type=image,name=git.fossy.my.id/${{ secrets.DOCKER_USERNAME }}/tunnel-please,push-by-digest=true,name-canonical=true,push=true
build-args: |
VERSION=${{ steps.version.outputs.VERSION }}
BUILD_DATE=${{ steps.version.outputs.BUILD_DATE }}
COMMIT=${{ steps.version.outputs.COMMIT }}
- name: Export digest
run: |
mkdir -p /tmp/digests
digest="${{ steps.build.outputs.digest }}"
touch "/tmp/digests/${digest#sha256:}"
- name: Upload digest
uses: actions/upload-artifact@v7
with:
name: digests-linux-amd64
path: /tmp/digests/*
if-no-files-found: error
retention-days: 1
build-arm64:
name: Build (linux/arm64)
runs-on: [ubuntu-latest, arm64]
needs: test
steps:
- name: Checkout repository
uses: actions/checkout@v6
- name: Extract version
id: version id: version
run: | run: |
VERSION=${GITHUB_REF#refs/tags/v} VERSION=${GITHUB_REF#refs/tags/v}
@@ -58,18 +111,10 @@ jobs:
echo "COMMIT=${{ github.sha }}" >> $GITHUB_OUTPUT echo "COMMIT=${{ github.sha }}" >> $GITHUB_OUTPUT
if echo "$VERSION" | grep -qE '^[0-9]+\.[0-9]+\.[0-9]+(-[a-zA-Z0-9.]+)?$'; then if echo "$VERSION" | grep -qE '^[0-9]+\.[0-9]+\.[0-9]+(-[a-zA-Z0-9.]+)?$'; then
MAJOR=$(echo "$VERSION" | cut -d. -f1) echo "MAJOR=$(echo "$VERSION" | cut -d. -f1)" >> $GITHUB_OUTPUT
MINOR=$(echo "$VERSION" | cut -d. -f2) echo "MINOR=$(echo "$VERSION" | cut -d. -f2)" >> $GITHUB_OUTPUT
PATCH=$(echo "$VERSION" | cut -d. -f3 | cut -d- -f1)
echo "MAJOR=$MAJOR" >> $GITHUB_OUTPUT
echo "MINOR=$MINOR" >> $GITHUB_OUTPUT
echo "PATCH=$PATCH" >> $GITHUB_OUTPUT
if echo "$VERSION" | grep -q '-'; then if echo "$VERSION" | grep -q '-'; then
PRERELEASE_TAG=$(echo "$VERSION" | cut -d- -f2 | cut -d. -f1)
echo "IS_PRERELEASE=true" >> $GITHUB_OUTPUT echo "IS_PRERELEASE=true" >> $GITHUB_OUTPUT
echo "PRERELEASE_TAG=$PRERELEASE_TAG" >> $GITHUB_OUTPUT
else else
echo "IS_PRERELEASE=false" >> $GITHUB_OUTPUT echo "IS_PRERELEASE=false" >> $GITHUB_OUTPUT
fi fi
@@ -78,35 +123,101 @@ jobs:
exit 1 exit 1
fi fi
- name: Build and push Docker image (release) - name: Set up Docker Buildx
uses: docker/build-push-action@v6 uses: docker/setup-buildx-action@v3
with:
context: .
push: true
tags: |
git.fossy.my.id/${{ secrets.DOCKER_USERNAME }}/tunnel-please:v${{ steps.version.outputs.VERSION }}
git.fossy.my.id/${{ secrets.DOCKER_USERNAME }}/tunnel-please:release
git.fossy.my.id/${{ secrets.DOCKER_USERNAME }}/tunnel-please:v${{ steps.version.outputs.MAJOR }}.${{ steps.version.outputs.MINOR }}
git.fossy.my.id/${{ secrets.DOCKER_USERNAME }}/tunnel-please:v${{ steps.version.outputs.MAJOR }}
git.fossy.my.id/${{ secrets.DOCKER_USERNAME }}/tunnel-please:latest
platforms: linux/amd64,linux/arm64
build-args: |
VERSION=${{ steps.version.outputs.VERSION }}
BUILD_DATE=${{ steps.version.outputs.BUILD_DATE }}
COMMIT=${{ steps.version.outputs.COMMIT }}
if: steps.version.outputs.IS_PRERELEASE == 'false'
- name: Build and push Docker image (pre-release) - name: Log in to Docker Registry
uses: docker/login-action@v3
with:
registry: git.fossy.my.id
username: ${{ secrets.DOCKER_USERNAME }}
password: ${{ secrets.DOCKER_PASSWORD }}
- name: Build and push by digest
id: build
uses: docker/build-push-action@v6 uses: docker/build-push-action@v6
with: with:
context: . context: .
push: true platforms: linux/arm64
tags: | outputs: >-
git.fossy.my.id/${{ secrets.DOCKER_USERNAME }}/tunnel-please:v${{ steps.version.outputs.VERSION }} type=image,name=git.fossy.my.id/${{ secrets.DOCKER_USERNAME }}/tunnel-please,push-by-digest=true,name-canonical=true,push=true
git.fossy.my.id/${{ secrets.DOCKER_USERNAME }}/tunnel-please:staging
platforms: linux/amd64,linux/arm64
build-args: | build-args: |
VERSION=${{ steps.version.outputs.VERSION }} VERSION=${{ steps.version.outputs.VERSION }}
BUILD_DATE=${{ steps.version.outputs.BUILD_DATE }} BUILD_DATE=${{ steps.version.outputs.BUILD_DATE }}
COMMIT=${{ steps.version.outputs.COMMIT }} COMMIT=${{ steps.version.outputs.COMMIT }}
- name: Export digest
run: |
mkdir -p /tmp/digests
digest="${{ steps.build.outputs.digest }}"
touch "/tmp/digests/${digest#sha256:}"
- name: Upload digest
uses: actions/upload-artifact@v7
with:
name: digests-linux-arm64
path: /tmp/digests/*
if-no-files-found: error
retention-days: 1
merge:
name: Merge Multi-Arch Manifest
runs-on: ubuntu-latest
needs: [build-amd64, build-arm64]
steps:
- name: Download all digests
uses: actions/download-artifact@v8.0.1
with:
path: /tmp/digests
pattern: digests-*
merge-multiple: true
- name: Extract version
id: version
run: |
VERSION=${GITHUB_REF#refs/tags/v}
echo "VERSION=$VERSION" >> $GITHUB_OUTPUT
if echo "$VERSION" | grep -qE '^[0-9]+\.[0-9]+\.[0-9]+(-[a-zA-Z0-9.]+)?$'; then
echo "MAJOR=$(echo "$VERSION" | cut -d. -f1)" >> $GITHUB_OUTPUT
echo "MINOR=$(echo "$VERSION" | cut -d. -f2)" >> $GITHUB_OUTPUT
if echo "$VERSION" | grep -q '-'; then
echo "IS_PRERELEASE=true" >> $GITHUB_OUTPUT
else
echo "IS_PRERELEASE=false" >> $GITHUB_OUTPUT
fi
else
echo "Invalid version format: $VERSION"
exit 1
fi
- name: Set up Docker Buildx
uses: docker/setup-buildx-action@v3
- name: Log in to Docker Registry
uses: docker/login-action@v3
with:
registry: git.fossy.my.id
username: ${{ secrets.DOCKER_USERNAME }}
password: ${{ secrets.DOCKER_PASSWORD }}
- name: Create and push manifest (release)
working-directory: /tmp/digests
if: steps.version.outputs.IS_PRERELEASE == 'false'
run: |
docker buildx imagetools create \
-t git.fossy.my.id/${{ secrets.DOCKER_USERNAME }}/tunnel-please:v${{ steps.version.outputs.VERSION }} \
-t git.fossy.my.id/${{ secrets.DOCKER_USERNAME }}/tunnel-please:release \
-t git.fossy.my.id/${{ secrets.DOCKER_USERNAME }}/tunnel-please:v${{ steps.version.outputs.MAJOR }}.${{ steps.version.outputs.MINOR }} \
-t git.fossy.my.id/${{ secrets.DOCKER_USERNAME }}/tunnel-please:v${{ steps.version.outputs.MAJOR }} \
-t git.fossy.my.id/${{ secrets.DOCKER_USERNAME }}/tunnel-please:latest \
$(printf 'git.fossy.my.id/${{ secrets.DOCKER_USERNAME }}/tunnel-please@sha256:%s ' *)
- name: Create and push manifest (pre-release)
working-directory: /tmp/digests
if: steps.version.outputs.IS_PRERELEASE == 'true' if: steps.version.outputs.IS_PRERELEASE == 'true'
run: |
docker buildx imagetools create \
-t git.fossy.my.id/${{ secrets.DOCKER_USERNAME }}/tunnel-please:v${{ steps.version.outputs.VERSION }} \
-t git.fossy.my.id/${{ secrets.DOCKER_USERNAME }}/tunnel-please:staging \
$(printf 'git.fossy.my.id/${{ secrets.DOCKER_USERNAME }}/tunnel-please@sha256:%s ' *)
+4 -3
View File
@@ -13,6 +13,7 @@ import (
"tunnel_pls/internal/session/slug" "tunnel_pls/internal/session/slug"
"tunnel_pls/internal/types" "tunnel_pls/internal/types"
"tunnel_pls/internal/port"
"tunnel_pls/internal/registry" "tunnel_pls/internal/registry"
proto "git.fossy.my.id/bagas/tunnel-please-grpc/gen" proto "git.fossy.my.id/bagas/tunnel-please-grpc/gen"
@@ -884,16 +885,16 @@ func (m *mockLifecycle) Connection() ssh.Conn {
return args.Get(0).(ssh.Conn) return args.Get(0).(ssh.Conn)
} }
func (m *mockLifecycle) User() string { return m.Called().String(0) } func (m *mockLifecycle) User() string { return m.Called().String(0) }
func (m *mockLifecycle) SetChannel(channel ssh.Channel) error { return m.Called(channel).Error(0) } func (m *mockLifecycle) SetChannel(channel ssh.Channel) { m.Called(channel) }
func (m *mockLifecycle) SetStatus(status types.SessionStatus) { m.Called(status) } func (m *mockLifecycle) SetStatus(status types.SessionStatus) { m.Called(status) }
func (m *mockLifecycle) IsActive() bool { return m.Called().Bool(0) } 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) StartedAt() time.Time { return m.Called().Get(0).(time.Time) }
func (m *mockLifecycle) PortRegistry() lifecycle.PortRegistry { func (m *mockLifecycle) PortRegistry() port.Port {
args := m.Called() args := m.Called()
if args.Get(0) == nil { if args.Get(0) == nil {
return nil return nil
} }
return args.Get(0).(lifecycle.PortRegistry) return args.Get(0).(port.Port)
} }
type mockEventServiceClient struct { type mockEventServiceClient struct {
+8 -11
View File
@@ -33,17 +33,12 @@ func (pm *port) AddRange(startPort, endPort uint16) error {
if startPort > endPort { if startPort > endPort {
return fmt.Errorf("start port cannot be greater than end port") return fmt.Errorf("start port cannot be greater than end port")
} }
for index := startPort; ; index++ { for index := startPort; index <= endPort; index++ {
if index != 0 {
if _, exists := pm.ports[index]; !exists { if _, exists := pm.ports[index]; !exists {
pm.ports[index] = false pm.ports[index] = false
pm.sortedPorts = append(pm.sortedPorts, index) pm.sortedPorts = append(pm.sortedPorts, index)
} }
} }
if index == endPort {
break
}
}
sort.Slice(pm.sortedPorts, func(i, j int) bool { sort.Slice(pm.sortedPorts, func(i, j int) bool {
return pm.sortedPorts[i] < pm.sortedPorts[j] return pm.sortedPorts[i] < pm.sortedPorts[j]
}) })
@@ -56,7 +51,6 @@ func (pm *port) Unassigned() (uint16, bool) {
for _, index := range pm.sortedPorts { for _, index := range pm.sortedPorts {
if !pm.ports[index] { if !pm.ports[index] {
pm.ports[index] = true
return index, true return index, true
} }
} }
@@ -67,9 +61,6 @@ func (pm *port) SetStatus(port uint16, assigned bool) error {
pm.mu.Lock() pm.mu.Lock()
defer pm.mu.Unlock() 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 pm.ports[port] = assigned
return nil return nil
} }
@@ -79,10 +70,16 @@ func (pm *port) Claim(port uint16) (claimed bool) {
defer pm.mu.Unlock() defer pm.mu.Unlock()
status, exists := pm.ports[port] status, exists := pm.ports[port]
if !exists || status {
if exists && status {
return false return false
} }
if !exists {
pm.ports[port] = true
return true
}
pm.ports[port] = true pm.ports[port] = true
return true return true
} }
+2 -51
View File
@@ -16,8 +16,6 @@ func TestAddRange(t *testing.T) {
{"normal range", 1000, 1002, false}, {"normal range", 1000, 1002, false},
{"invalid range", 2000, 1999, true}, {"invalid range", 2000, 1999, true},
{"single port range", 3000, 3000, false}, {"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 { for _, tt := range tests {
@@ -33,22 +31,6 @@ 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) { func TestUnassigned(t *testing.T) {
pm := New() pm := New()
_ = pm.AddRange(1000, 1002) _ = pm.AddRange(1000, 1002)
@@ -76,21 +58,6 @@ 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) { func TestSetStatus(t *testing.T) {
pm := New() pm := New()
_ = pm.AddRange(1000, 1002) _ = pm.AddRange(1000, 1002)
@@ -116,17 +83,6 @@ 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) { func TestClaim(t *testing.T) {
pm := New() pm := New()
_ = pm.AddRange(1000, 1002) _ = pm.AddRange(1000, 1002)
@@ -139,7 +95,7 @@ func TestClaim(t *testing.T) {
}{ }{
{"claim unassigned port", 1000, false, true}, {"claim unassigned port", 1000, false, true},
{"claim already assigned port", 1001, true, false}, {"claim already assigned port", 1001, true, false},
{"claim non-existent port", 5000, false, false}, {"claim non-existent port", 5000, false, true},
} }
for _, tt := range tests { for _, tt := range tests {
@@ -151,13 +107,8 @@ func TestClaim(t *testing.T) {
got := pm.Claim(tt.port) got := pm.Claim(tt.port)
assert.Equal(t, tt.want, got) assert.Equal(t, tt.want, got)
finalState, exists := pm.(*port).ports[tt.port] finalState := 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) assert.True(t, finalState)
}
}) })
} }
} }
+4 -3
View File
@@ -4,6 +4,7 @@ import (
"sync" "sync"
"testing" "testing"
"time" "time"
"tunnel_pls/internal/port"
"tunnel_pls/internal/session/forwarder" "tunnel_pls/internal/session/forwarder"
"tunnel_pls/internal/session/interaction" "tunnel_pls/internal/session/interaction"
"tunnel_pls/internal/session/lifecycle" "tunnel_pls/internal/session/lifecycle"
@@ -77,15 +78,15 @@ func (ml *mockLifecycle) Connection() ssh.Conn {
return args.Get(0).(ssh.Conn) return args.Get(0).(ssh.Conn)
} }
func (ml *mockLifecycle) PortRegistry() lifecycle.PortRegistry { func (ml *mockLifecycle) PortRegistry() port.Port {
args := ml.Called() args := ml.Called()
if args.Get(0) == nil { if args.Get(0) == nil {
return nil return nil
} }
return args.Get(0).(lifecycle.PortRegistry) return args.Get(0).(port.Port)
} }
func (ml *mockLifecycle) SetChannel(channel ssh.Channel) error { return ml.Called(channel).Error(0) } func (ml *mockLifecycle) SetChannel(channel ssh.Channel) { ml.Called(channel) }
func (ml *mockLifecycle) SetStatus(status types.SessionStatus) { ml.Called(status) } func (ml *mockLifecycle) SetStatus(status types.SessionStatus) { ml.Called(status) }
func (ml *mockLifecycle) IsActive() bool { return ml.Called().Bool(0) } func (ml *mockLifecycle) IsActive() bool { return ml.Called().Bool(0) }
func (ml *mockLifecycle) StartedAt() time.Time { return ml.Called().Get(0).(time.Time) } func (ml *mockLifecycle) StartedAt() time.Time { return ml.Called().Get(0).(time.Time) }
+3 -16
View File
@@ -28,7 +28,6 @@ type Forwarder interface {
Close() error Close() error
} }
type forwarder struct { type forwarder struct {
mu sync.RWMutex
listener net.Listener listener net.Listener
tunnelType types.TunnelType tunnelType types.TunnelType
forwardedPort uint16 forwardedPort uint16
@@ -61,7 +60,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) { 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 { type channelResult struct {
channel ssh.Channel channel ssh.Channel
reqs <-chan *ssh.Request reqs <-chan *ssh.Request
@@ -142,44 +141,32 @@ func (f *forwarder) HandleConnection(dst io.ReadWriter, src ssh.Channel) {
} }
func (f *forwarder) SetType(tunnelType types.TunnelType) { func (f *forwarder) SetType(tunnelType types.TunnelType) {
f.mu.Lock()
defer f.mu.Unlock()
f.tunnelType = tunnelType f.tunnelType = tunnelType
} }
func (f *forwarder) TunnelType() types.TunnelType { func (f *forwarder) TunnelType() types.TunnelType {
f.mu.RLock()
defer f.mu.RUnlock()
return f.tunnelType return f.tunnelType
} }
func (f *forwarder) ForwardedPort() uint16 { func (f *forwarder) ForwardedPort() uint16 {
f.mu.RLock()
defer f.mu.RUnlock()
return f.forwardedPort return f.forwardedPort
} }
func (f *forwarder) SetForwardedPort(port uint16) { func (f *forwarder) SetForwardedPort(port uint16) {
f.mu.Lock()
defer f.mu.Unlock()
f.forwardedPort = port f.forwardedPort = port
} }
func (f *forwarder) SetListener(listener net.Listener) { func (f *forwarder) SetListener(listener net.Listener) {
f.mu.Lock()
defer f.mu.Unlock()
f.listener = listener f.listener = listener
} }
func (f *forwarder) Listener() net.Listener { func (f *forwarder) Listener() net.Listener {
f.mu.RLock()
defer f.mu.RUnlock()
return f.listener return f.listener
} }
func (f *forwarder) Close() error { func (f *forwarder) Close() error {
if listener := f.Listener(); listener != nil { if f.Listener() != nil {
return listener.Close() return f.listener.Close()
} }
return nil return nil
} }
@@ -1922,6 +1922,10 @@ func TestInteraction_Start_ProtocolSelection(t *testing.T) {
time.Sleep(50 * time.Millisecond) time.Sleep(50 * time.Millisecond)
i := mockInteraction.(*interaction) i := mockInteraction.(*interaction)
if i.program != nil {
assert.NotNil(t, i.program, "program should be initialized")
}
i.Stop() i.Stop()
mockConfig.AssertExpectations(t) mockConfig.AssertExpectations(t)
+27 -76
View File
@@ -2,7 +2,6 @@ package lifecycle
import ( import (
"errors" "errors"
"fmt"
"io" "io"
"net" "net"
"sync" "sync"
@@ -10,6 +9,8 @@ import (
"tunnel_pls/internal/session/slug" "tunnel_pls/internal/session/slug"
"tunnel_pls/internal/types" "tunnel_pls/internal/types"
portUtil "tunnel_pls/internal/port"
"golang.org/x/crypto/ssh" "golang.org/x/crypto/ssh"
) )
@@ -23,12 +24,6 @@ type SessionRegistry interface {
Remove(key types.SessionKey) Remove(key types.SessionKey)
} }
type PortRegistry interface {
Unassigned() (uint16, bool)
Claim(port uint16) bool
SetStatus(port uint16, assigned bool) error
}
type lifecycle struct { type lifecycle struct {
mu sync.Mutex mu sync.Mutex
status types.SessionStatus status types.SessionStatus
@@ -39,18 +34,18 @@ type lifecycle struct {
slug slug.Slug slug slug.Slug
startedAt time.Time startedAt time.Time
sessionRegistry SessionRegistry sessionRegistry SessionRegistry
portRegistry PortRegistry portRegistry portUtil.Port
user string user string
} }
func New(conn ssh.Conn, forwarder Forwarder, slugManager slug.Slug, port PortRegistry, sessionRegistry SessionRegistry, user string) Lifecycle { func New(conn ssh.Conn, forwarder Forwarder, slugManager slug.Slug, port portUtil.Port, sessionRegistry SessionRegistry, user string) Lifecycle {
return &lifecycle{ return &lifecycle{
status: types.SessionStatusINITIALIZING, status: types.SessionStatusINITIALIZING,
conn: conn, conn: conn,
channel: nil, channel: nil,
forwarder: forwarder, forwarder: forwarder,
slug: slugManager, slug: slugManager,
startedAt: time.Time{}, startedAt: time.Now(),
sessionRegistry: sessionRegistry, sessionRegistry: sessionRegistry,
portRegistry: port, portRegistry: port,
user: user, user: user,
@@ -60,16 +55,16 @@ func New(conn ssh.Conn, forwarder Forwarder, slugManager slug.Slug, port PortReg
type Lifecycle interface { type Lifecycle interface {
Connection() ssh.Conn Connection() ssh.Conn
Channel() ssh.Channel Channel() ssh.Channel
PortRegistry() PortRegistry PortRegistry() portUtil.Port
User() string User() string
SetChannel(channel ssh.Channel) error SetChannel(channel ssh.Channel)
SetStatus(status types.SessionStatus) SetStatus(status types.SessionStatus)
IsActive() bool IsActive() bool
StartedAt() time.Time StartedAt() time.Time
Close() error Close() error
} }
func (l *lifecycle) PortRegistry() PortRegistry { func (l *lifecycle) PortRegistry() portUtil.Port {
return l.portRegistry return l.portRegistry
} }
@@ -77,25 +72,11 @@ func (l *lifecycle) User() string {
return l.user return l.user
} }
func (l *lifecycle) SetChannel(channel ssh.Channel) error { func (l *lifecycle) SetChannel(channel ssh.Channel) {
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 l.channel = channel
return nil
} }
func (l *lifecycle) Channel() ssh.Channel { func (l *lifecycle) Channel() ssh.Channel {
l.mu.Lock()
defer l.mu.Unlock()
return l.channel return l.channel
} }
@@ -106,13 +87,7 @@ func (l *lifecycle) Connection() ssh.Conn {
func (l *lifecycle) SetStatus(status types.SessionStatus) { func (l *lifecycle) SetStatus(status types.SessionStatus) {
l.mu.Lock() l.mu.Lock()
defer l.mu.Unlock() defer l.mu.Unlock()
if l.status == types.SessionStatusCLOSED {
return
}
l.status = status l.status = status
if status == types.SessionStatusRUNNING && l.startedAt.IsZero() {
l.startedAt = time.Now()
}
} }
func (l *lifecycle) IsActive() bool { func (l *lifecycle) IsActive() bool {
@@ -123,74 +98,50 @@ func (l *lifecycle) IsActive() bool {
func (l *lifecycle) Close() error { func (l *lifecycle) Close() error {
l.mu.Lock() l.mu.Lock()
defer l.mu.Unlock()
if l.status == types.SessionStatusCLOSED { if l.status == types.SessionStatusCLOSED {
closeErr := l.closeErr return l.closeErr
l.mu.Unlock()
return closeErr
} }
l.status = types.SessionStatusCLOSED l.status = types.SessionStatusCLOSED
channel := l.channel
conn := l.conn
l.mu.Unlock()
var errs []error var errs []error
if channel != nil { tunnelType := l.forwarder.TunnelType()
if err := channel.Close(); err != nil && !isClosedError(err) {
errs = append(errs, err) if l.channel != nil {
} if err := l.channel.Close(); err != nil && !isClosedError(err) {
}
if conn != nil {
if err := conn.Close(); err != nil && !isClosedError(err) {
errs = append(errs, err) errs = append(errs, err)
} }
} }
l.cleanupRegistry() if l.conn != nil {
if err := l.cleanupForwarder(); err != nil { if err := l.conn.Close(); err != nil && !isClosedError(err) {
errs = append(errs, err) errs = append(errs, err)
} }
closeErr := errors.Join(errs...)
l.mu.Lock()
l.closeErr = closeErr
l.mu.Unlock()
return closeErr
} }
func (l *lifecycle) cleanupRegistry() { clientSlug := l.slug.String()
slugStr := l.slug.String()
if slugStr == "" {
return
}
key := types.SessionKey{ key := types.SessionKey{
Id: slugStr, Id: clientSlug,
Type: l.forwarder.TunnelType(), Type: tunnelType,
} }
l.sessionRegistry.Remove(key) 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 { l.closeErr = errors.Join(errs...)
if l.forwarder.TunnelType() != types.TunnelTypeTCP { return l.closeErr
return nil
}
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 { func isClosedError(err error) bool {
if err == nil { if err == nil {
return false return false
} }
return errors.Is(err, io.EOF) || errors.Is(err, net.ErrClosed) return errors.Is(err, io.EOF) || errors.Is(err, net.ErrClosed) || err.Error() == "EOF"
} }
func (l *lifecycle) StartedAt() time.Time { func (l *lifecycle) StartedAt() time.Time {
l.mu.Lock()
defer l.mu.Unlock()
return l.startedAt return l.startedAt
} }
+4 -125
View File
@@ -5,7 +5,6 @@ import (
"errors" "errors"
"io" "io"
"net" "net"
"sync"
"testing" "testing"
"tunnel_pls/internal/types" "tunnel_pls/internal/types"
@@ -178,14 +177,8 @@ func TestLifecycle_SetChannel(t *testing.T) {
mockSSHChannel := &MockSSHChannel{} mockSSHChannel := &MockSSHChannel{}
err := mockLifecycle.SetChannel(mockSSHChannel) 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()) assert.Equal(t, mockSSHChannel, mockLifecycle.Channel())
} }
@@ -283,15 +276,14 @@ func TestLifecycle_Close(t *testing.T) {
mockLifecycle := New(mockSSHConn, mockForwarder, mockSlug, mockPort, mockSessionRegistry, "mas-fuad") mockLifecycle := New(mockSSHConn, mockForwarder, mockSlug, mockPort, mockSessionRegistry, "mas-fuad")
mockLifecycle.SetStatus(types.SessionStatusRUNNING) mockLifecycle.SetStatus(types.SessionStatusRUNNING)
err := mockLifecycle.SetChannel(mockSSHChannel) mockLifecycle.SetChannel(mockSSHChannel)
assert.NoError(t, err)
if tt.alreadyClosed { if tt.alreadyClosed {
err = mockLifecycle.Close() err := mockLifecycle.Close()
assert.NoError(t, err) assert.NoError(t, err)
} }
err = mockLifecycle.Close() err := mockLifecycle.Close()
if tt.expectErr { if tt.expectErr {
assert.Error(t, err) assert.Error(t, err)
@@ -309,116 +301,3 @@ 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")
}
+22 -24
View File
@@ -26,7 +26,7 @@ type Session interface {
HandleGlobalRequest(ch <-chan *ssh.Request) error HandleGlobalRequest(ch <-chan *ssh.Request) error
HandleTCPIPForward(req *ssh.Request) error HandleTCPIPForward(req *ssh.Request) error
HandleHTTPForward(req *ssh.Request, port uint16) error HandleHTTPForward(req *ssh.Request, port uint16) error
HandleTCPForward(req *ssh.Request, addr string, port uint16, reserved bool) error HandleTCPForward(req *ssh.Request, addr string, port uint16) error
Lifecycle() lifecycle.Lifecycle Lifecycle() lifecycle.Lifecycle
Interaction() interaction.Interaction Interaction() interaction.Interaction
Forwarder() forwarder.Forwarder Forwarder() forwarder.Forwarder
@@ -158,15 +158,13 @@ func (s *session) setupInteractiveMode(channel ssh.NewChannel) error {
} }
go func() { go func() {
err := s.HandleGlobalRequest(reqs) err = s.HandleGlobalRequest(reqs)
if err != nil { if err != nil {
log.Printf("global request handler error: %v", err) log.Printf("global request handler error: %v", err)
} }
}() }()
if err = s.lifecycle.SetChannel(ch); err != nil { s.lifecycle.SetChannel(ch)
return err
}
s.interaction.SetChannel(ch) s.interaction.SetChannel(ch)
s.interaction.SetMode(types.InteractiveModeINTERACTIVE) s.interaction.SetMode(types.InteractiveModeINTERACTIVE)
@@ -200,7 +198,6 @@ func (s *session) waitForSessionEnd() error {
} }
func (s *session) waitForTCPIPForward() *ssh.Request { func (s *session) waitForTCPIPForward() *ssh.Request {
for {
select { select {
case req, ok := <-s.initialReq: case req, ok := <-s.initialReq:
if !ok { if !ok {
@@ -210,12 +207,14 @@ func (s *session) waitForTCPIPForward() *ssh.Request {
if req.Type == "tcpip-forward" { if req.Type == "tcpip-forward" {
return req return req
} }
log.Printf("Ignoring unexpected global request: %s", req.Type) if err := req.Reply(false, nil); err != nil {
_ = req.Reply(false, nil) log.Printf("Failed to reply to request: %v", err)
case <-time.After(500 * time.Millisecond):
log.Println("No tcpip-forward request received within timeout")
return nil
} }
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
} }
} }
@@ -254,35 +253,35 @@ func (s *session) HandleGlobalRequest(GlobalRequest <-chan *ssh.Request) error {
return nil return nil
} }
func (s *session) parseForwardPayload(payload []byte) (address string, port uint16, reserved bool, err error) { func (s *session) parseForwardPayload(payload []byte) (address string, port uint16, err error) {
var forwardPayload struct { var forwardPayload struct {
BindAddr string BindAddr string
BindPort uint32 BindPort uint32
} }
if err = ssh.Unmarshal(payload, &forwardPayload); err != nil { if err = ssh.Unmarshal(payload, &forwardPayload); err != nil {
return "", 0, false, fmt.Errorf("failed to unmarshal forward payload: %w", err) return "", 0, fmt.Errorf("failed to unmarshal forward payload: %w", err)
} }
if forwardPayload.BindPort > 65535 { if forwardPayload.BindPort > 65535 {
return "", 0, false, fmt.Errorf("port is larger than allowed port of 65535") return "", 0, fmt.Errorf("port is larger than allowed port of 65535")
} }
port = uint16(forwardPayload.BindPort) port = uint16(forwardPayload.BindPort)
if isBlockedPort(port) { if isBlockedPort(port) {
return "", 0, false, fmt.Errorf("port is blocked") return "", 0, fmt.Errorf("port is blocked")
} }
if port == 0 { if port == 0 {
unassigned, ok := s.lifecycle.PortRegistry().Unassigned() unassigned, ok := s.lifecycle.PortRegistry().Unassigned()
if !ok { if !ok {
return "", 0, false, fmt.Errorf("no available port") return "", 0, fmt.Errorf("no available port")
} }
return forwardPayload.BindAddr, unassigned, true, nil return forwardPayload.BindAddr, unassigned, nil
} }
return forwardPayload.BindAddr, port, false, nil return forwardPayload.BindAddr, port, nil
} }
func (s *session) denyForwardingRequest(req *ssh.Request, key *types.SessionKey, listener io.Closer, msg string) error { func (s *session) denyForwardingRequest(req *ssh.Request, key *types.SessionKey, listener io.Closer, msg string) error {
@@ -326,7 +325,7 @@ func (s *session) finalizeForwarding(req *ssh.Request, portToBind uint16, listen
} }
func (s *session) HandleTCPIPForward(req *ssh.Request) error { func (s *session) HandleTCPIPForward(req *ssh.Request) error {
address, port, reserved, err := s.parseForwardPayload(req.Payload) address, port, err := s.parseForwardPayload(req.Payload)
if err != nil { if err != nil {
return s.denyForwardingRequest(req, nil, nil, fmt.Sprintf("cannot parse forwarded payload: %s", err.Error())) return s.denyForwardingRequest(req, nil, nil, fmt.Sprintf("cannot parse forwarded payload: %s", err.Error()))
} }
@@ -335,7 +334,7 @@ func (s *session) HandleTCPIPForward(req *ssh.Request) error {
case 80, 443: case 80, 443:
return s.HandleHTTPForward(req, port) return s.HandleHTTPForward(req, port)
default: default:
return s.HandleTCPForward(req, address, port, reserved) return s.HandleTCPForward(req, address, port)
} }
} }
@@ -356,12 +355,10 @@ func (s *session) HandleHTTPForward(req *ssh.Request, portToBind uint16) error {
return nil return nil
} }
func (s *session) HandleTCPForward(req *ssh.Request, addr string, portToBind uint16, reserved bool) error { func (s *session) HandleTCPForward(req *ssh.Request, addr string, portToBind uint16) error {
if !reserved {
if claimed := s.lifecycle.PortRegistry().Claim(portToBind); !claimed { 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)) return s.denyForwardingRequest(req, nil, nil, fmt.Sprintf("Port %d is already in use or restricted", portToBind))
} }
}
tcpServer := transport.NewTCPServer(portToBind, s.forwarder) tcpServer := transport.NewTCPServer(portToBind, s.forwarder)
listener, err := tcpServer.Listen() listener, err := tcpServer.Listen()
@@ -380,7 +377,8 @@ func (s *session) HandleTCPForward(req *ssh.Request, addr string, portToBind uin
} }
go func() { go func() {
if err := tcpServer.Serve(listener); err != nil { err = tcpServer.Serve(listener)
if err != nil {
log.Printf("Failed serving tcp server: %s\n", err) log.Printf("Failed serving tcp server: %s\n", err)
} }
}() }()
+10 -71
View File
@@ -396,12 +396,6 @@ func TestHandleTCPIPForward_Table(t *testing.T) {
err := s.HandleTCPIPForward(req) err := s.HandleTCPIPForward(req)
assert.NoError(t, err) assert.NoError(t, err)
assert.Equal(t, uint16(12345), s.forwarder.ForwardedPort()) 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) { t.Run("Invalid Payload", func(t *testing.T) {
@@ -813,7 +807,7 @@ func (m *mockNewChanFail) Accept() (ssh.Channel, <-chan *ssh.Request, error) {
} }
func TestWaitForTCPIPForward_EdgeCases(t *testing.T) { func TestWaitForTCPIPForward_EdgeCases(t *testing.T) {
t.Run("Wrong Request Type Then Timeout", func(t *testing.T) { t.Run("Wrong Request Type", func(t *testing.T) {
_, sReqs, _, cConn, cleanup := setupSSH(t) _, sReqs, _, cConn, cleanup := setupSSH(t)
defer cleanup() defer cleanup()
@@ -823,65 +817,10 @@ func TestWaitForTCPIPForward_EdgeCases(t *testing.T) {
_, _, _ = cConn.SendRequest("not-tcpip-forward", true, nil) _, _, _ = cConn.SendRequest("not-tcpip-forward", true, nil)
}() }()
start := time.Now()
req := s.waitForTCPIPForward() req := s.waitForTCPIPForward()
elapsed := time.Since(start)
if req != nil { if req != nil {
t.Error("expected nil request") 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) { t.Run("Channel Closed", func(t *testing.T) {
@@ -1043,7 +982,7 @@ func TestParseForwardPayload_Errors(t *testing.T) {
s := &session{} s := &session{}
t.Run("Short Address", func(t *testing.T) { 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 { if err == nil {
t.Error("expected error, got nil") t.Error("expected error, got nil")
} }
@@ -1051,7 +990,7 @@ func TestParseForwardPayload_Errors(t *testing.T) {
t.Run("Short Port", func(t *testing.T) { t.Run("Short Port", func(t *testing.T) {
payload := append([]byte{0, 0, 0, 4}, []byte("addr")...) payload := append([]byte{0, 0, 0, 4}, []byte("addr")...)
_, _, _, err := s.parseForwardPayload(payload) _, _, err := s.parseForwardPayload(payload)
if err == nil { if err == nil {
t.Error("expected error, got nil") t.Error("expected error, got nil")
} }
@@ -1062,7 +1001,7 @@ func TestParseForwardPayload_Errors(t *testing.T) {
portBuf := make([]byte, 4) portBuf := make([]byte, 4)
binary.BigEndian.PutUint32(portBuf, 22) binary.BigEndian.PutUint32(portBuf, 22)
payload = append(payload, portBuf...) payload = append(payload, portBuf...)
_, _, _, err := s.parseForwardPayload(payload) _, _, err := s.parseForwardPayload(payload)
if err == nil { if err == nil {
t.Error("expected error, got nil") t.Error("expected error, got nil")
} else if !strings.Contains(err.Error(), "port is block") { } else if !strings.Contains(err.Error(), "port is block") {
@@ -1226,7 +1165,7 @@ func TestHandleTCPForward_Failures(t *testing.T) {
s, _, mPort, _, sReqs, cConn, cleanup := setup(t) s, _, mPort, _, sReqs, cConn, cleanup := setup(t)
defer cleanup() defer cleanup()
mPort.On("Claim", mock.Anything).Return(false) mPort.On("Claim", mock.Anything).Return(false)
err := s.HandleTCPForward(getReq(t, cConn, sReqs), "localhost", 1234, false) err := s.HandleTCPForward(getReq(t, cConn, sReqs), "localhost", 1234)
if err == nil { if err == nil {
t.Error("expected error, got nil") t.Error("expected error, got nil")
} else if !strings.Contains(err.Error(), "already in use") { } else if !strings.Contains(err.Error(), "already in use") {
@@ -1248,7 +1187,7 @@ func TestHandleTCPForward_Failures(t *testing.T) {
}(l) }(l)
port := uint16(l.Addr().(*net.TCPAddr).Port) port := uint16(l.Addr().(*net.TCPAddr).Port)
err = s.HandleTCPForward(getReq(t, cConn, sReqs), "localhost", port, false) err = s.HandleTCPForward(getReq(t, cConn, sReqs), "localhost", port)
if err == nil { if err == nil {
t.Error("expected error, got nil") t.Error("expected error, got nil")
} else if !strings.Contains(err.Error(), "already in use") { } else if !strings.Contains(err.Error(), "already in use") {
@@ -1261,7 +1200,7 @@ func TestHandleTCPForward_Failures(t *testing.T) {
defer cleanup() defer cleanup()
mPort.On("Claim", mock.Anything).Return(true) mPort.On("Claim", mock.Anything).Return(true)
mRegistry.On("Register", mock.Anything, mock.Anything).Return(false) mRegistry.On("Register", mock.Anything, mock.Anything).Return(false)
err := s.HandleTCPForward(getReq(t, cConn, sReqs), "localhost", 0, false) err := s.HandleTCPForward(getReq(t, cConn, sReqs), "localhost", 0)
if err == nil { if err == nil {
t.Error("expected error, got nil") t.Error("expected error, got nil")
} else if !strings.Contains(err.Error(), "Failed to register") { } else if !strings.Contains(err.Error(), "Failed to register") {
@@ -1270,16 +1209,16 @@ func TestHandleTCPForward_Failures(t *testing.T) {
}) })
t.Run("Finalize fail (Reply fail)", func(t *testing.T) { t.Run("Finalize fail (Reply fail)", func(t *testing.T) {
s, mRegistry, mPort, sConn, sReqs, cConn, cleanup := setup(t) s, mRegistry, mPort, _, sReqs, cConn, cleanup := setup(t)
defer cleanup() defer cleanup()
mPort.On("Claim", mock.Anything).Return(true) mPort.On("Claim", mock.Anything).Return(true)
mRegistry.On("Register", mock.Anything, mock.Anything).Return(true) mRegistry.On("Register", mock.Anything, mock.Anything).Return(true)
req := getReq(t, cConn, sReqs) req := getReq(t, cConn, sReqs)
err := cConn.Close() err := cConn.Close()
assert.NoError(t, err) assert.NoError(t, err)
_ = sConn.Wait() time.Sleep(100 * time.Millisecond)
err = s.HandleTCPForward(req, "localhost", 0, false) err = s.HandleTCPForward(req, "localhost", 0)
if err == nil { if err == nil {
t.Error("expected error, got nil") t.Error("expected error, got nil")
} else if !strings.Contains(err.Error(), "Failed to finalize forwarding") { } else if !strings.Contains(err.Error(), "Failed to finalize forwarding") {
-7
View File
@@ -1,14 +1,11 @@
package slug package slug
import "sync"
type Slug interface { type Slug interface {
String() string String() string
Set(slug string) Set(slug string)
} }
type slug struct { type slug struct {
mu sync.RWMutex
slug string slug string
} }
@@ -19,13 +16,9 @@ func New() Slug {
} }
func (s *slug) String() string { func (s *slug) String() string {
s.mu.RLock()
defer s.mu.RUnlock()
return s.slug return s.slug
} }
func (s *slug) Set(slug string) { func (s *slug) Set(slug string) {
s.mu.Lock()
defer s.mu.Unlock()
s.slug = slug s.slug = slug
} }
+2 -1
View File
@@ -55,7 +55,8 @@ func TestHTTPServer_Serve(t *testing.T) {
go func() { go func() {
time.Sleep(100 * time.Millisecond) time.Sleep(100 * time.Millisecond)
_ = listener.Close() err = listener.Close()
assert.NoError(t, err)
}() }()
err = srv.Serve(listener) err = srv.Serve(listener)
+9 -28
View File
@@ -1,7 +1,6 @@
package transport package transport
import ( import (
"bufio"
"bytes" "bytes"
"context" "context"
"errors" "errors"
@@ -53,43 +52,25 @@ func (hh *httpHandler) badRequest(conn net.Conn) error {
return nil 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) { func (hh *httpHandler) Handler(conn net.Conn, isTLS bool) {
defer hh.closeConnection(conn) defer hh.closeConnection(conn)
_ = conn.SetReadDeadline(time.Now().Add(10 * time.Second)) _ = conn.SetReadDeadline(time.Now().Add(10 * time.Second))
br := bufio.NewReaderSize(conn, hh.config.HeaderSize()) buf := make([]byte, hh.config.HeaderSize())
headerBuf, err := readHTTPHeader(br, hh.config.HeaderSize()) n, err := conn.Read(buf)
if err != nil { if err != nil {
_ = hh.badRequest(conn) _ = hh.badRequest(conn)
return return
} }
if idx := bytes.Index(buf[:n], []byte("\r\n\r\n")); idx == -1 {
_ = hh.badRequest(conn)
return
}
_ = conn.SetReadDeadline(time.Time{}) _ = conn.SetReadDeadline(time.Time{})
reqhf, err := header.NewRequest(headerBuf) reqhf, err := header.NewRequest(buf[:n])
if err != nil { if err != nil {
log.Printf("Error creating request header: %v", err) log.Printf("Error creating request header: %v", err)
_ = hh.badRequest(conn) _ = hh.badRequest(conn)
@@ -120,7 +101,7 @@ func (hh *httpHandler) Handler(conn net.Conn, isTLS bool) {
return return
} }
hw := stream.New(conn, br, conn.RemoteAddr()) hw := stream.New(conn, conn, conn.RemoteAddr())
defer func(hw stream.HTTP) { defer func(hw stream.HTTP) {
err = hw.Close() err = hw.Close()
if err != nil { if err != nil {
+2 -118
View File
@@ -321,14 +321,8 @@ func TestHandler(t *testing.T) {
isTLS: false, isTLS: false,
redirectTLS: false, redirectTLS: false,
request: []byte(""), request: []byte(""),
expected: []byte(""), expected: []byte("HTTP/1.1 400 Bad Request\r\n\r\n"),
setupConn: func() (net.Conn, net.Conn) { setupMocks: func(msr *MockSessionRegistry) {
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
}, },
}, },
{ {
@@ -721,113 +715,3 @@ 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)
}
+2 -1
View File
@@ -63,7 +63,8 @@ func TestHTTPSServer_Serve(t *testing.T) {
go func() { go func() {
time.Sleep(100 * time.Millisecond) time.Sleep(100 * time.Millisecond)
_ = listener.Close() err = listener.Close()
assert.NoError(t, err)
}() }()
err = srv.Serve(listener) err = srv.Serve(listener)
+2 -1
View File
@@ -45,7 +45,8 @@ func TestTCPServer_Serve(t *testing.T) {
go func() { go func() {
time.Sleep(100 * time.Millisecond) time.Sleep(100 * time.Millisecond)
_ = listener.Close() err = listener.Close()
assert.NoError(t, err)
}() }()
err = srv.Serve(listener) err = srv.Serve(listener)
+5 -2
View File
@@ -2,6 +2,8 @@
"extends": [ "extends": [
"config:recommended" "config:recommended"
], ],
"prConcurrentLimit": 1,
"prHourlyLimit": 1,
"packageRules": [ "packageRules": [
{ {
"matchUpdateTypes": [ "matchUpdateTypes": [
@@ -10,9 +12,10 @@
"pin", "pin",
"digest" "digest"
], ],
"groupName": "all-dependencies",
"automerge": true, "automerge": true,
"baseBranchPatterns": [ "matchPackageNames": [
"staging" "*"
] ]
} }
] ]