Compare commits

..

7 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
23 changed files with 596 additions and 197 deletions
+4 -4
View File
@@ -40,10 +40,10 @@ jobs:
uses: actions/checkout@v6 uses: actions/checkout@v6
- name: Set up Docker Buildx - name: Set up Docker Buildx
uses: docker/setup-buildx-action@v4 uses: docker/setup-buildx-action@v3
- name: Log in to Docker Registry - name: Log in to Docker Registry
uses: docker/login-action@v4 uses: docker/login-action@v3
with: with:
registry: git.fossy.my.id registry: git.fossy.my.id
username: ${{ secrets.DOCKER_USERNAME }} username: ${{ secrets.DOCKER_USERNAME }}
@@ -79,7 +79,7 @@ jobs:
fi fi
- name: Build and push Docker image (release) - name: Build and push Docker image (release)
uses: docker/build-push-action@v7 uses: docker/build-push-action@v6
with: with:
context: . context: .
push: true push: true
@@ -97,7 +97,7 @@ jobs:
if: steps.version.outputs.IS_PRERELEASE == 'false' if: steps.version.outputs.IS_PRERELEASE == 'false'
- name: Build and push Docker image (pre-release) - name: Build and push Docker image (pre-release)
uses: docker/build-push-action@v7 uses: docker/build-push-action@v6
with: with:
context: . context: .
push: true push: true
+1 -1
View File
@@ -43,7 +43,7 @@ jobs:
--output.checkstyle.path=golangci-lint-report.xml --output.checkstyle.path=golangci-lint-report.xml
- name: SonarQube Scan - name: SonarQube Scan
uses: SonarSource/sonarqube-scan-action@v8.2.1 uses: SonarSource/sonarqube-scan-action@v7.0.0
env: env:
SONAR_HOST_URL: ${{ secrets.SONARQUBE_HOST }} SONAR_HOST_URL: ${{ secrets.SONARQUBE_HOST }}
SONAR_TOKEN: ${{ secrets.SONARQUBE_TOKEN }} SONAR_TOKEN: ${{ secrets.SONARQUBE_TOKEN }}
+1 -1
View File
@@ -1,4 +1,4 @@
FROM golang:1.26.2-alpine AS go_builder FROM golang:1.26.0-alpine AS go_builder
ARG VERSION=dev ARG VERSION=dev
ARG BUILD_DATE=unknown ARG BUILD_DATE=unknown
+13 -13
View File
@@ -4,7 +4,7 @@ go 1.26.0
require ( require (
git.fossy.my.id/bagas/tunnel-please-grpc v1.5.0 git.fossy.my.id/bagas/tunnel-please-grpc v1.5.0
github.com/caddyserver/certmagic v0.25.3 github.com/caddyserver/certmagic v0.25.1
github.com/charmbracelet/bubbles v1.0.0 github.com/charmbracelet/bubbles v1.0.0
github.com/charmbracelet/bubbletea v1.3.10 github.com/charmbracelet/bubbletea v1.3.10
github.com/charmbracelet/lipgloss v1.1.0 github.com/charmbracelet/lipgloss v1.1.0
@@ -12,15 +12,15 @@ require (
github.com/libdns/cloudflare v0.2.2 github.com/libdns/cloudflare v0.2.2
github.com/muesli/termenv v0.16.0 github.com/muesli/termenv v0.16.0
github.com/stretchr/testify v1.11.1 github.com/stretchr/testify v1.11.1
golang.org/x/crypto v0.50.0 golang.org/x/crypto v0.48.0
google.golang.org/grpc v1.80.0 google.golang.org/grpc v1.78.0
google.golang.org/protobuf v1.36.11 google.golang.org/protobuf v1.36.11
) )
require ( require (
github.com/atotto/clipboard v0.1.4 // indirect github.com/atotto/clipboard v0.1.4 // indirect
github.com/aymanbagabas/go-osc52/v2 v2.0.1 // indirect github.com/aymanbagabas/go-osc52/v2 v2.0.1 // indirect
github.com/caddyserver/zerossl v0.1.5 // indirect github.com/caddyserver/zerossl v0.1.4 // indirect
github.com/charmbracelet/colorprofile v0.4.1 // indirect github.com/charmbracelet/colorprofile v0.4.1 // indirect
github.com/charmbracelet/x/ansi v0.11.6 // indirect github.com/charmbracelet/x/ansi v0.11.6 // indirect
github.com/charmbracelet/x/cellbuf v0.0.15 // indirect github.com/charmbracelet/x/cellbuf v0.0.15 // indirect
@@ -36,8 +36,8 @@ require (
github.com/mattn/go-isatty v0.0.20 // indirect github.com/mattn/go-isatty v0.0.20 // indirect
github.com/mattn/go-localereader v0.0.1 // indirect github.com/mattn/go-localereader v0.0.1 // indirect
github.com/mattn/go-runewidth v0.0.19 // indirect github.com/mattn/go-runewidth v0.0.19 // indirect
github.com/mholt/acmez/v3 v3.1.6 // indirect github.com/mholt/acmez/v3 v3.1.4 // indirect
github.com/miekg/dns v1.1.72 // indirect github.com/miekg/dns v1.1.69 // indirect
github.com/muesli/ansi v0.0.0-20230316100256-276c6243b2f6 // indirect github.com/muesli/ansi v0.0.0-20230316100256-276c6243b2f6 // indirect
github.com/muesli/cancelreader v0.2.2 // indirect github.com/muesli/cancelreader v0.2.2 // indirect
github.com/pmezard/go-difflib v1.0.0 // indirect github.com/pmezard/go-difflib v1.0.0 // indirect
@@ -49,12 +49,12 @@ require (
go.uber.org/multierr v1.11.0 // indirect go.uber.org/multierr v1.11.0 // indirect
go.uber.org/zap v1.27.1 // indirect go.uber.org/zap v1.27.1 // indirect
go.uber.org/zap/exp v0.3.0 // indirect go.uber.org/zap/exp v0.3.0 // indirect
golang.org/x/mod v0.35.0 // indirect golang.org/x/mod v0.32.0 // indirect
golang.org/x/net v0.53.0 // indirect golang.org/x/net v0.49.0 // indirect
golang.org/x/sync v0.20.0 // indirect golang.org/x/sync v0.19.0 // indirect
golang.org/x/sys v0.43.0 // indirect golang.org/x/sys v0.41.0 // indirect
golang.org/x/text v0.36.0 // indirect golang.org/x/text v0.34.0 // indirect
golang.org/x/tools v0.44.0 // indirect golang.org/x/tools v0.41.0 // indirect
google.golang.org/genproto/googleapis/rpc v0.0.0-20260120221211-b8f7ae30c516 // indirect google.golang.org/genproto/googleapis/rpc v0.0.0-20251222181119-0a764e51fe1b // indirect
gopkg.in/yaml.v3 v3.0.1 // indirect gopkg.in/yaml.v3 v3.0.1 // indirect
) )
-46
View File
@@ -8,14 +8,8 @@ github.com/aymanbagabas/go-udiff v0.3.1 h1:LV+qyBQ2pqe0u42ZsUEtPiCaUoqgA9gYRDs3v
github.com/aymanbagabas/go-udiff v0.3.1/go.mod h1:G0fsKmG+P6ylD0r6N/KgQD/nWzgfnl8ZBcNLgcbrw8E= github.com/aymanbagabas/go-udiff v0.3.1/go.mod h1:G0fsKmG+P6ylD0r6N/KgQD/nWzgfnl8ZBcNLgcbrw8E=
github.com/caddyserver/certmagic v0.25.1 h1:4sIKKbOt5pg6+sL7tEwymE1x2bj6CHr80da1CRRIPbY= github.com/caddyserver/certmagic v0.25.1 h1:4sIKKbOt5pg6+sL7tEwymE1x2bj6CHr80da1CRRIPbY=
github.com/caddyserver/certmagic v0.25.1/go.mod h1:VhyvndxtVton/Fo/wKhRoC46Rbw1fmjvQ3GjHYSQTEY= github.com/caddyserver/certmagic v0.25.1/go.mod h1:VhyvndxtVton/Fo/wKhRoC46Rbw1fmjvQ3GjHYSQTEY=
github.com/caddyserver/certmagic v0.25.2 h1:D7xcS7ggX/WEY54x0czj7ioTkmDWKIgxtIi2OcQclUc=
github.com/caddyserver/certmagic v0.25.2/go.mod h1:llW/CvsNmza8S6hmsuggsZeiX+uS27dkqY27wDIuBWg=
github.com/caddyserver/certmagic v0.25.3 h1:mGf5ba8F7xA4c5jfDZZbK2buY1VEkbnwpMDixaju94A=
github.com/caddyserver/certmagic v0.25.3/go.mod h1:YVs43D5+H/Dckt4bTga1KSO/xYfFBfVZainGDywYPAA=
github.com/caddyserver/zerossl v0.1.4 h1:CVJOE3MZeFisCERZjkxIcsqIH4fnFdlYWnPYeFtBHRw= github.com/caddyserver/zerossl v0.1.4 h1:CVJOE3MZeFisCERZjkxIcsqIH4fnFdlYWnPYeFtBHRw=
github.com/caddyserver/zerossl v0.1.4/go.mod h1:CxA0acn7oEGO6//4rtrRjYgEoa4MFw/XofZnrYwGqG4= github.com/caddyserver/zerossl v0.1.4/go.mod h1:CxA0acn7oEGO6//4rtrRjYgEoa4MFw/XofZnrYwGqG4=
github.com/caddyserver/zerossl v0.1.5 h1:dkvOjBAEEtY6LIGAHei7sw2UgqSD6TrWweXpV7lvEvE=
github.com/caddyserver/zerossl v0.1.5/go.mod h1:CxA0acn7oEGO6//4rtrRjYgEoa4MFw/XofZnrYwGqG4=
github.com/charmbracelet/bubbles v1.0.0 h1:12J8/ak/uCZEMQ6KU7pcfwceyjLlWsDLAxB5fXonfvc= github.com/charmbracelet/bubbles v1.0.0 h1:12J8/ak/uCZEMQ6KU7pcfwceyjLlWsDLAxB5fXonfvc=
github.com/charmbracelet/bubbles v1.0.0/go.mod h1:9d/Zd5GdnauMI5ivUIVisuEm3ave1XwXtD1ckyV6r3E= github.com/charmbracelet/bubbles v1.0.0/go.mod h1:9d/Zd5GdnauMI5ivUIVisuEm3ave1XwXtD1ckyV6r3E=
github.com/charmbracelet/bubbletea v1.3.10 h1:otUDHWMMzQSB0Pkc87rm691KZ3SWa4KUlvF9nRvCICw= github.com/charmbracelet/bubbletea v1.3.10 h1:otUDHWMMzQSB0Pkc87rm691KZ3SWa4KUlvF9nRvCICw=
@@ -72,12 +66,8 @@ github.com/mattn/go-runewidth v0.0.19 h1:v++JhqYnZuu5jSKrk9RbgF5v4CGUjqRfBm05byF
github.com/mattn/go-runewidth v0.0.19/go.mod h1:XBkDxAl56ILZc9knddidhrOlY5R/pDhgLpndooCuJAs= github.com/mattn/go-runewidth v0.0.19/go.mod h1:XBkDxAl56ILZc9knddidhrOlY5R/pDhgLpndooCuJAs=
github.com/mholt/acmez/v3 v3.1.4 h1:DyzZe/RnAzT3rpZj/2Ii5xZpiEvvYk3cQEN/RmqxwFQ= github.com/mholt/acmez/v3 v3.1.4 h1:DyzZe/RnAzT3rpZj/2Ii5xZpiEvvYk3cQEN/RmqxwFQ=
github.com/mholt/acmez/v3 v3.1.4/go.mod h1:L1wOU06KKvq7tswuMDwKdcHeKpFFgkppZy/y0DFxagQ= github.com/mholt/acmez/v3 v3.1.4/go.mod h1:L1wOU06KKvq7tswuMDwKdcHeKpFFgkppZy/y0DFxagQ=
github.com/mholt/acmez/v3 v3.1.6 h1:eGVQNObP0pBN4sxqrXeg7MYqTOWyoiYpQqITVWlrevk=
github.com/mholt/acmez/v3 v3.1.6/go.mod h1:5nTPosTGosLxF3+LU4ygbgMRFDhbAVpqMI4+a4aHLBY=
github.com/miekg/dns v1.1.69 h1:Kb7Y/1Jo+SG+a2GtfoFUfDkG//csdRPwRLkCsxDG9Sc= github.com/miekg/dns v1.1.69 h1:Kb7Y/1Jo+SG+a2GtfoFUfDkG//csdRPwRLkCsxDG9Sc=
github.com/miekg/dns v1.1.69/go.mod h1:7OyjD9nEba5OkqQ/hB4fy3PIoxafSZJtducccIelz3g= github.com/miekg/dns v1.1.69/go.mod h1:7OyjD9nEba5OkqQ/hB4fy3PIoxafSZJtducccIelz3g=
github.com/miekg/dns v1.1.72 h1:vhmr+TF2A3tuoGNkLDFK9zi36F2LS+hKTRW0Uf8kbzI=
github.com/miekg/dns v1.1.72/go.mod h1:+EuEPhdHOsfk6Wk5TT2CzssZdqkmFhf8r+aVyDEToIs=
github.com/muesli/ansi v0.0.0-20230316100256-276c6243b2f6 h1:ZK8zHtRHOkbHy6Mmr5D264iyp3TiX5OmNcI5cIARiQI= github.com/muesli/ansi v0.0.0-20230316100256-276c6243b2f6 h1:ZK8zHtRHOkbHy6Mmr5D264iyp3TiX5OmNcI5cIARiQI=
github.com/muesli/ansi v0.0.0-20230316100256-276c6243b2f6/go.mod h1:CJlz5H+gyd6CUWT45Oy4q24RdLyn7Md9Vj2/ldJBSIo= github.com/muesli/ansi v0.0.0-20230316100256-276c6243b2f6/go.mod h1:CJlz5H+gyd6CUWT45Oy4q24RdLyn7Md9Vj2/ldJBSIo=
github.com/muesli/cancelreader v0.2.2 h1:3I4Kt4BQjOR54NavqnDogx/MIoWBFa0StPA8ELUXHmA= github.com/muesli/cancelreader v0.2.2 h1:3I4Kt4BQjOR54NavqnDogx/MIoWBFa0StPA8ELUXHmA=
@@ -124,66 +114,30 @@ go.uber.org/zap/exp v0.3.0 h1:6JYzdifzYkGmTdRR59oYH+Ng7k49H9qVpWwNSsGJj3U=
go.uber.org/zap/exp v0.3.0/go.mod h1:5I384qq7XGxYyByIhHm6jg5CHkGY0nsTfbDLgDDlgJQ= go.uber.org/zap/exp v0.3.0/go.mod h1:5I384qq7XGxYyByIhHm6jg5CHkGY0nsTfbDLgDDlgJQ=
golang.org/x/crypto v0.48.0 h1:/VRzVqiRSggnhY7gNRxPauEQ5Drw9haKdM0jqfcCFts= golang.org/x/crypto v0.48.0 h1:/VRzVqiRSggnhY7gNRxPauEQ5Drw9haKdM0jqfcCFts=
golang.org/x/crypto v0.48.0/go.mod h1:r0kV5h3qnFPlQnBSrULhlsRfryS2pmewsg+XfMgkVos= golang.org/x/crypto v0.48.0/go.mod h1:r0kV5h3qnFPlQnBSrULhlsRfryS2pmewsg+XfMgkVos=
golang.org/x/crypto v0.49.0 h1:+Ng2ULVvLHnJ/ZFEq4KdcDd/cfjrrjjNSXNzxg0Y4U4=
golang.org/x/crypto v0.49.0/go.mod h1:ErX4dUh2UM+CFYiXZRTcMpEcN8b/1gxEuv3nODoYtCA=
golang.org/x/crypto v0.50.0 h1:zO47/JPrL6vsNkINmLoo/PH1gcxpls50DNogFvB5ZGI=
golang.org/x/crypto v0.50.0/go.mod h1:3muZ7vA7PBCE6xgPX7nkzzjiUq87kRItoJQM1Yo8S+Q=
golang.org/x/exp v0.0.0-20231006140011-7918f672742d h1:jtJma62tbqLibJ5sFQz8bKtEM8rJBtfilJ2qTU199MI= golang.org/x/exp v0.0.0-20231006140011-7918f672742d h1:jtJma62tbqLibJ5sFQz8bKtEM8rJBtfilJ2qTU199MI=
golang.org/x/exp v0.0.0-20231006140011-7918f672742d/go.mod h1:ldy0pHrwJyGW56pPQzzkH36rKxoZW1tw7ZJpeKx+hdo= golang.org/x/exp v0.0.0-20231006140011-7918f672742d/go.mod h1:ldy0pHrwJyGW56pPQzzkH36rKxoZW1tw7ZJpeKx+hdo=
golang.org/x/mod v0.32.0 h1:9F4d3PHLljb6x//jOyokMv3eX+YDeepZSEo3mFJy93c= golang.org/x/mod v0.32.0 h1:9F4d3PHLljb6x//jOyokMv3eX+YDeepZSEo3mFJy93c=
golang.org/x/mod v0.32.0/go.mod h1:SgipZ/3h2Ci89DlEtEXWUk/HteuRin+HHhN+WbNhguU= golang.org/x/mod v0.32.0/go.mod h1:SgipZ/3h2Ci89DlEtEXWUk/HteuRin+HHhN+WbNhguU=
golang.org/x/mod v0.33.0 h1:tHFzIWbBifEmbwtGz65eaWyGiGZatSrT9prnU8DbVL8=
golang.org/x/mod v0.33.0/go.mod h1:swjeQEj+6r7fODbD2cqrnje9PnziFuw4bmLbBZFrQ5w=
golang.org/x/mod v0.34.0 h1:xIHgNUUnW6sYkcM5Jleh05DvLOtwc6RitGHbDk4akRI=
golang.org/x/mod v0.34.0/go.mod h1:ykgH52iCZe79kzLLMhyCUzhMci+nQj+0XkbXpNYtVjY=
golang.org/x/mod v0.35.0 h1:Ww1D637e6Pg+Zb2KrWfHQUnH2dQRLBQyAtpr/haaJeM=
golang.org/x/mod v0.35.0/go.mod h1:+GwiRhIInF8wPm+4AoT6L0FA1QWAad3OMdTRx4tFYlU=
golang.org/x/net v0.49.0 h1:eeHFmOGUTtaaPSGNmjBKpbng9MulQsJURQUAfUwY++o= golang.org/x/net v0.49.0 h1:eeHFmOGUTtaaPSGNmjBKpbng9MulQsJURQUAfUwY++o=
golang.org/x/net v0.49.0/go.mod h1:/ysNB2EvaqvesRkuLAyjI1ycPZlQHM3q01F02UY/MV8= golang.org/x/net v0.49.0/go.mod h1:/ysNB2EvaqvesRkuLAyjI1ycPZlQHM3q01F02UY/MV8=
golang.org/x/net v0.51.0 h1:94R/GTO7mt3/4wIKpcR5gkGmRLOuE/2hNGeWq/GBIFo=
golang.org/x/net v0.51.0/go.mod h1:aamm+2QF5ogm02fjy5Bb7CQ0WMt1/WVM7FtyaTLlA9Y=
golang.org/x/net v0.52.0 h1:He/TN1l0e4mmR3QqHMT2Xab3Aj3L9qjbhRm78/6jrW0=
golang.org/x/net v0.52.0/go.mod h1:R1MAz7uMZxVMualyPXb+VaqGSa3LIaUqk0eEt3w36Sw=
golang.org/x/net v0.53.0 h1:d+qAbo5L0orcWAr0a9JweQpjXF19LMXJE8Ey7hwOdUA=
golang.org/x/net v0.53.0/go.mod h1:JvMuJH7rrdiCfbeHoo3fCQU24Lf5JJwT9W3sJFulfgs=
golang.org/x/sync v0.19.0 h1:vV+1eWNmZ5geRlYjzm2adRgW2/mcpevXNg50YZtPCE4= golang.org/x/sync v0.19.0 h1:vV+1eWNmZ5geRlYjzm2adRgW2/mcpevXNg50YZtPCE4=
golang.org/x/sync v0.19.0/go.mod h1:9KTHXmSnoGruLpwFjVSX0lNNA75CykiMECbovNTZqGI= golang.org/x/sync v0.19.0/go.mod h1:9KTHXmSnoGruLpwFjVSX0lNNA75CykiMECbovNTZqGI=
golang.org/x/sync v0.20.0 h1:e0PTpb7pjO8GAtTs2dQ6jYa5BWYlMuX047Dco/pItO4=
golang.org/x/sync v0.20.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0=
golang.org/x/sys v0.0.0-20210809222454-d867a43fc93e/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.0.0-20210809222454-d867a43fc93e/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.6.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.6.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
golang.org/x/sys v0.41.0 h1:Ivj+2Cp/ylzLiEU89QhWblYnOE9zerudt9Ftecq2C6k= golang.org/x/sys v0.41.0 h1:Ivj+2Cp/ylzLiEU89QhWblYnOE9zerudt9Ftecq2C6k=
golang.org/x/sys v0.41.0/go.mod h1:OgkHotnGiDImocRcuBABYBEXf8A9a87e/uXjp9XT3ks= golang.org/x/sys v0.41.0/go.mod h1:OgkHotnGiDImocRcuBABYBEXf8A9a87e/uXjp9XT3ks=
golang.org/x/sys v0.42.0 h1:omrd2nAlyT5ESRdCLYdm3+fMfNFE/+Rf4bDIQImRJeo=
golang.org/x/sys v0.42.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
golang.org/x/sys v0.43.0 h1:Rlag2XtaFTxp19wS8MXlJwTvoh8ArU6ezoyFsMyCTNI=
golang.org/x/sys v0.43.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
golang.org/x/term v0.40.0 h1:36e4zGLqU4yhjlmxEaagx2KuYbJq3EwY8K943ZsHcvg= golang.org/x/term v0.40.0 h1:36e4zGLqU4yhjlmxEaagx2KuYbJq3EwY8K943ZsHcvg=
golang.org/x/term v0.40.0/go.mod h1:w2P8uVp06p2iyKKuvXIm7N/y0UCRt3UfJTfZ7oOpglM= golang.org/x/term v0.40.0/go.mod h1:w2P8uVp06p2iyKKuvXIm7N/y0UCRt3UfJTfZ7oOpglM=
golang.org/x/text v0.34.0 h1:oL/Qq0Kdaqxa1KbNeMKwQq0reLCCaFtqu2eNuSeNHbk= golang.org/x/text v0.34.0 h1:oL/Qq0Kdaqxa1KbNeMKwQq0reLCCaFtqu2eNuSeNHbk=
golang.org/x/text v0.34.0/go.mod h1:homfLqTYRFyVYemLBFl5GgL/DWEiH5wcsQ5gSh1yziA= golang.org/x/text v0.34.0/go.mod h1:homfLqTYRFyVYemLBFl5GgL/DWEiH5wcsQ5gSh1yziA=
golang.org/x/text v0.35.0 h1:JOVx6vVDFokkpaq1AEptVzLTpDe9KGpj5tR4/X+ybL8=
golang.org/x/text v0.35.0/go.mod h1:khi/HExzZJ2pGnjenulevKNX1W67CUy0AsXcNubPGCA=
golang.org/x/text v0.36.0 h1:JfKh3XmcRPqZPKevfXVpI1wXPTqbkE5f7JA92a55Yxg=
golang.org/x/text v0.36.0/go.mod h1:NIdBknypM8iqVmPiuco0Dh6P5Jcdk8lJL0CUebqK164=
golang.org/x/tools v0.41.0 h1:a9b8iMweWG+S0OBnlU36rzLp20z1Rp10w+IY2czHTQc= golang.org/x/tools v0.41.0 h1:a9b8iMweWG+S0OBnlU36rzLp20z1Rp10w+IY2czHTQc=
golang.org/x/tools v0.41.0/go.mod h1:XSY6eDqxVNiYgezAVqqCeihT4j1U2CCsqvH3WhQpnlg= golang.org/x/tools v0.41.0/go.mod h1:XSY6eDqxVNiYgezAVqqCeihT4j1U2CCsqvH3WhQpnlg=
golang.org/x/tools v0.42.0 h1:uNgphsn75Tdz5Ji2q36v/nsFSfR/9BRFvqhGBaJGd5k=
golang.org/x/tools v0.42.0/go.mod h1:Ma6lCIwGZvHK6XtgbswSoWroEkhugApmsXyrUmBhfr0=
golang.org/x/tools v0.43.0 h1:12BdW9CeB3Z+J/I/wj34VMl8X+fEXBxVR90JeMX5E7s=
golang.org/x/tools v0.43.0/go.mod h1:uHkMso649BX2cZK6+RpuIPXS3ho2hZo4FVwfoy1vIk0=
golang.org/x/tools v0.44.0 h1:UP4ajHPIcuMjT1GqzDWRlalUEoY+uzoZKnhOjbIPD2c=
golang.org/x/tools v0.44.0/go.mod h1:KA0AfVErSdxRZIsOVipbv3rQhVXTnlU6UhKxHd1seDI=
gonum.org/v1/gonum v0.16.0 h1:5+ul4Swaf3ESvrOnidPp4GZbzf0mxVQpDCYUQE7OJfk= gonum.org/v1/gonum v0.16.0 h1:5+ul4Swaf3ESvrOnidPp4GZbzf0mxVQpDCYUQE7OJfk=
gonum.org/v1/gonum v0.16.0/go.mod h1:fef3am4MQ93R2HHpKnLk4/Tbh/s0+wqD5nfa6Pnwy4E= gonum.org/v1/gonum v0.16.0/go.mod h1:fef3am4MQ93R2HHpKnLk4/Tbh/s0+wqD5nfa6Pnwy4E=
google.golang.org/genproto/googleapis/rpc v0.0.0-20251222181119-0a764e51fe1b h1:Mv8VFug0MP9e5vUxfBcE3vUkV6CImK3cMNMIDFjmzxU= google.golang.org/genproto/googleapis/rpc v0.0.0-20251222181119-0a764e51fe1b h1:Mv8VFug0MP9e5vUxfBcE3vUkV6CImK3cMNMIDFjmzxU=
google.golang.org/genproto/googleapis/rpc v0.0.0-20251222181119-0a764e51fe1b/go.mod h1:j9x/tPzZkyxcgEFkiKEEGxfvyumM01BEtsW8xzOahRQ= google.golang.org/genproto/googleapis/rpc v0.0.0-20251222181119-0a764e51fe1b/go.mod h1:j9x/tPzZkyxcgEFkiKEEGxfvyumM01BEtsW8xzOahRQ=
google.golang.org/genproto/googleapis/rpc v0.0.0-20260120221211-b8f7ae30c516 h1:sNrWoksmOyF5bvJUcnmbeAmQi8baNhqg5IWaI3llQqU=
google.golang.org/genproto/googleapis/rpc v0.0.0-20260120221211-b8f7ae30c516/go.mod h1:j9x/tPzZkyxcgEFkiKEEGxfvyumM01BEtsW8xzOahRQ=
google.golang.org/grpc v1.78.0 h1:K1XZG/yGDJnzMdd/uZHAkVqJE+xIDOcmdSFZkBUicNc= google.golang.org/grpc v1.78.0 h1:K1XZG/yGDJnzMdd/uZHAkVqJE+xIDOcmdSFZkBUicNc=
google.golang.org/grpc v1.78.0/go.mod h1:I47qjTo4OKbMkjA/aOOwxDIiPSBofUtQUI5EfpWvW7U= google.golang.org/grpc v1.78.0/go.mod h1:I47qjTo4OKbMkjA/aOOwxDIiPSBofUtQUI5EfpWvW7U=
google.golang.org/grpc v1.80.0 h1:Xr6m2WmWZLETvUNvIUmeD5OAagMw3FiKmMlTdViWsHM=
google.golang.org/grpc v1.80.0/go.mod h1:ho/dLnxwi3EDJA4Zghp7k2Ec1+c2jqup0bFkw07bwF4=
google.golang.org/protobuf v1.36.11 h1:fV6ZwhNocDyBLK0dj+fg8ektcVegBBuEolpbTQyBNVE= google.golang.org/protobuf v1.36.11 h1:fV6ZwhNocDyBLK0dj+fg8ektcVegBBuEolpbTQyBNVE=
google.golang.org/protobuf v1.36.11/go.mod h1:HTf+CrKn2C3g5S8VImy6tdcUvCska2kB7j23XfzDpco= google.golang.org/protobuf v1.36.11/go.mod h1:HTf+CrKn2C3g5S8VImy6tdcUvCska2kB7j23XfzDpco=
gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405 h1:yhCVgyC4o1eVCa2tZl7eS0r+SDo693bJlVdllGtEeKM= gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405 h1:yhCVgyC4o1eVCa2tZl7eS0r+SDo693bJlVdllGtEeKM=
+3 -4
View File
@@ -13,7 +13,6 @@ 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"
@@ -885,16 +884,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) { 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) 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() port.Port { func (m *mockLifecycle) PortRegistry() lifecycle.PortRegistry {
args := m.Called() args := m.Called()
if args.Get(0) == nil { if args.Get(0) == nil {
return nil return nil
} }
return args.Get(0).(port.Port) return args.Get(0).(lifecycle.PortRegistry)
} }
type mockEventServiceClient struct { type mockEventServiceClient struct {
+14 -11
View File
@@ -33,10 +33,15 @@ 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 <= endPort; index++ { for index := startPort; ; index++ {
if _, exists := pm.ports[index]; !exists { if index != 0 {
pm.ports[index] = false if _, exists := pm.ports[index]; !exists {
pm.sortedPorts = append(pm.sortedPorts, index) pm.ports[index] = false
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 {
@@ -51,6 +56,7 @@ 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
} }
} }
@@ -61,6 +67,9 @@ 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
} }
@@ -70,16 +79,10 @@ 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
} }
+52 -3
View File
@@ -16,6 +16,8 @@ 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 {
@@ -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) { func TestUnassigned(t *testing.T) {
pm := New() pm := New()
_ = pm.AddRange(1000, 1002) _ = 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) { func TestSetStatus(t *testing.T) {
pm := New() pm := New()
_ = pm.AddRange(1000, 1002) _ = 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) { func TestClaim(t *testing.T) {
pm := New() pm := New()
_ = pm.AddRange(1000, 1002) _ = pm.AddRange(1000, 1002)
@@ -95,7 +139,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, true}, {"claim non-existent port", 5000, false, false},
} }
for _, tt := range tests { for _, tt := range tests {
@@ -107,8 +151,13 @@ 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 := pm.(*port).ports[tt.port] finalState, exists := pm.(*port).ports[tt.port]
assert.True(t, finalState) 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)
}
}) })
} }
} }
+3 -3
View File
@@ -94,13 +94,13 @@ func (r *registry) Update(user string, oldKey, newKey Key) error {
return ErrInvalidSlug return ErrInvalidSlug
} }
r.mu.Lock()
defer r.mu.Unlock()
if _, exists := r.slugIndex[newKey]; exists && newKey != oldKey { if _, exists := r.slugIndex[newKey]; exists && newKey != oldKey {
return ErrSlugInUse return ErrSlugInUse
} }
r.mu.Lock()
defer r.mu.Unlock()
client, ok := r.byUser[user][oldKey] client, ok := r.byUser[user][oldKey]
if !ok { if !ok {
return ErrSessionNotFound return ErrSessionNotFound
+3 -4
View File
@@ -4,7 +4,6 @@ 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"
@@ -78,15 +77,15 @@ func (ml *mockLifecycle) Connection() ssh.Conn {
return args.Get(0).(ssh.Conn) return args.Get(0).(ssh.Conn)
} }
func (ml *mockLifecycle) PortRegistry() port.Port { func (ml *mockLifecycle) PortRegistry() lifecycle.PortRegistry {
args := ml.Called() args := ml.Called()
if args.Get(0) == nil { if args.Get(0) == nil {
return 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) 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) }
+16 -3
View File
@@ -28,6 +28,7 @@ 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
@@ -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) { 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
@@ -141,32 +142,44 @@ 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 f.Listener() != nil { if listener := f.Listener(); listener != nil {
return f.listener.Close() return listener.Close()
} }
return nil return nil
} }
@@ -1922,10 +1922,6 @@ 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)
+78 -29
View File
@@ -2,6 +2,7 @@ package lifecycle
import ( import (
"errors" "errors"
"fmt"
"io" "io"
"net" "net"
"sync" "sync"
@@ -9,8 +10,6 @@ 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"
) )
@@ -24,6 +23,12 @@ 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
@@ -34,18 +39,18 @@ type lifecycle struct {
slug slug.Slug slug slug.Slug
startedAt time.Time startedAt time.Time
sessionRegistry SessionRegistry sessionRegistry SessionRegistry
portRegistry portUtil.Port portRegistry PortRegistry
user string 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{ 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.Now(), startedAt: time.Time{},
sessionRegistry: sessionRegistry, sessionRegistry: sessionRegistry,
portRegistry: port, portRegistry: port,
user: user, user: user,
@@ -55,16 +60,16 @@ func New(conn ssh.Conn, forwarder Forwarder, slugManager slug.Slug, port portUti
type Lifecycle interface { type Lifecycle interface {
Connection() ssh.Conn Connection() ssh.Conn
Channel() ssh.Channel Channel() ssh.Channel
PortRegistry() portUtil.Port PortRegistry() PortRegistry
User() string User() string
SetChannel(channel ssh.Channel) SetChannel(channel ssh.Channel) error
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() portUtil.Port { func (l *lifecycle) PortRegistry() PortRegistry {
return l.portRegistry return l.portRegistry
} }
@@ -72,11 +77,25 @@ func (l *lifecycle) User() string {
return l.user 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 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
} }
@@ -87,7 +106,13 @@ 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 {
@@ -98,50 +123,74 @@ 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 {
return l.closeErr closeErr := 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
tunnelType := l.forwarder.TunnelType() if channel != nil {
if err := channel.Close(); err != nil && !isClosedError(err) {
if l.channel != nil { errs = append(errs, err)
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)
} }
} }
if l.conn != nil { l.cleanupRegistry()
if err := l.conn.Close(); err != nil && !isClosedError(err) { if err := l.cleanupForwarder(); err != nil {
errs = append(errs, err) 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{ key := types.SessionKey{
Id: clientSlug, Id: slugStr,
Type: tunnelType, Type: l.forwarder.TunnelType(),
} }
l.sessionRegistry.Remove(key) l.sessionRegistry.Remove(key)
}
if tunnelType == types.TunnelTypeTCP { func (l *lifecycle) cleanupForwarder() error {
errs = append(errs, l.PortRegistry().SetStatus(l.forwarder.ForwardedPort(), false)) if l.forwarder.TunnelType() != types.TunnelTypeTCP {
errs = append(errs, l.forwarder.Close()) return nil
} }
var errs []error
l.closeErr = errors.Join(errs...) errs = append(errs, l.portRegistry.SetStatus(l.forwarder.ForwardedPort(), false))
return l.closeErr 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) || err.Error() == "EOF" return errors.Is(err, io.EOF) || errors.Is(err, net.ErrClosed)
} }
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
} }
+125 -4
View File
@@ -5,6 +5,7 @@ import (
"errors" "errors"
"io" "io"
"net" "net"
"sync"
"testing" "testing"
"tunnel_pls/internal/types" "tunnel_pls/internal/types"
@@ -177,8 +178,14 @@ func TestLifecycle_SetChannel(t *testing.T) {
mockSSHChannel := &MockSSHChannel{} 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()) 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 := New(mockSSHConn, mockForwarder, mockSlug, mockPort, mockSessionRegistry, "mas-fuad")
mockLifecycle.SetStatus(types.SessionStatusRUNNING) mockLifecycle.SetStatus(types.SessionStatusRUNNING)
mockLifecycle.SetChannel(mockSSHChannel) err := 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)
@@ -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")
}
+43 -32
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) error HandleTCPForward(req *ssh.Request, addr string, port uint16, reserved bool) error
Lifecycle() lifecycle.Lifecycle Lifecycle() lifecycle.Lifecycle
Interaction() interaction.Interaction Interaction() interaction.Interaction
Forwarder() forwarder.Forwarder Forwarder() forwarder.Forwarder
@@ -158,13 +158,15 @@ 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)
} }
}() }()
s.lifecycle.SetChannel(ch) if err = s.lifecycle.SetChannel(ch); err != nil {
return err
}
s.interaction.SetChannel(ch) s.interaction.SetChannel(ch)
s.interaction.SetMode(types.InteractiveModeINTERACTIVE) s.interaction.SetMode(types.InteractiveModeINTERACTIVE)
@@ -198,23 +200,22 @@ func (s *session) waitForSessionEnd() error {
} }
func (s *session) waitForTCPIPForward() *ssh.Request { func (s *session) waitForTCPIPForward() *ssh.Request {
select { for {
case req, ok := <-s.initialReq: select {
if !ok { case req, ok := <-s.initialReq:
log.Println("Forwarding request channel closed") 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 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 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 { 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, 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 { 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) port = uint16(forwardPayload.BindPort)
if isBlockedPort(port) { if isBlockedPort(port) {
return "", 0, fmt.Errorf("port is blocked") return "", 0, false, 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, 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 { 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 { 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 { 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()))
} }
@@ -334,7 +335,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) return s.HandleTCPForward(req, address, port, reserved)
} }
} }
@@ -355,30 +356,40 @@ 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) error { func (s *session) HandleTCPForward(req *ssh.Request, addr string, portToBind uint16, reserved bool) error {
if claimed := s.lifecycle.PortRegistry().Claim(portToBind); !claimed { if !reserved {
return s.denyForwardingRequest(req, nil, nil, fmt.Sprintf("Port %d is already in use or restricted", portToBind)) 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) tcpServer := transport.NewTCPServer(portToBind, s.forwarder)
listener, err := tcpServer.Listen() listener, err := tcpServer.Listen()
if err != nil { if err != nil {
releasePort()
return s.denyForwardingRequest(req, nil, listener, fmt.Sprintf("Port %d is already in use or restricted", portToBind)) 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} key := types.SessionKey{Id: fmt.Sprintf("%d", portToBind), Type: types.TunnelTypeTCP}
if !s.registry.Register(key, s) { 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)) 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) err = s.finalizeForwarding(req, portToBind, listener, types.TunnelTypeTCP, key.Id)
if err != nil { if err != nil {
releasePort()
return s.denyForwardingRequest(req, &key, listener, fmt.Sprintf("Failed to finalize forwarding: %s", err)) return s.denyForwardingRequest(req, &key, listener, fmt.Sprintf("Failed to finalize forwarding: %s", err))
} }
go func() { go func() {
err = tcpServer.Serve(listener) if err := tcpServer.Serve(listener); err != nil {
if err != nil {
log.Printf("Failed serving tcp server: %s\n", err) log.Printf("Failed serving tcp server: %s\n", err)
} }
}() }()
+81 -12
View File
@@ -396,6 +396,12 @@ 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) {
@@ -699,6 +705,7 @@ func TestForwardingFailures(t *testing.T) {
s, mRegistry, mPort, _, _, 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)
mPort.On("SetStatus", uint16(1234), false).Return(nil)
mRegistry.On("Register", mock.Anything, mock.Anything).Return(false) mRegistry.On("Register", mock.Anything, mock.Anything).Return(false)
payload := make([]byte, 4+9+4) payload := make([]byte, 4+9+4)
@@ -717,7 +724,7 @@ func TestForwardingFailures(t *testing.T) {
}) })
t.Run("Finalize Forwarding Failure", func(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() defer cleanup()
mRandom.On("String", 20).Return("test-slug", nil) mRandom.On("String", 20).Return("test-slug", nil)
mRegistry.On("Register", mock.Anything, mock.Anything).Return(true) mRegistry.On("Register", mock.Anything, mock.Anything).Return(true)
@@ -736,7 +743,7 @@ func TestForwardingFailures(t *testing.T) {
err := cConn.Close() err := cConn.Close()
assert.NoError(t, err) assert.NoError(t, err)
time.Sleep(50 * time.Millisecond) _ = sConn.Wait()
err = s.HandleTCPIPForward(req) err = s.HandleTCPIPForward(req)
assert.Error(t, err) assert.Error(t, err)
@@ -758,6 +765,7 @@ func TestForwardingFailures(t *testing.T) {
}(l) }(l)
_, portStr, _ := net.SplitHostPort(l.Addr().String()) _, portStr, _ := net.SplitHostPort(l.Addr().String())
port, _ := strconv.Atoi(portStr) port, _ := strconv.Atoi(portStr)
mPort.On("SetStatus", uint16(port), false).Return(nil)
payload := make([]byte, 4+9+4) payload := make([]byte, 4+9+4)
binary.BigEndian.PutUint32(payload[0:4], 9) 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) { 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) _, sReqs, _, cConn, cleanup := setupSSH(t)
defer cleanup() defer cleanup()
@@ -817,10 +825,65 @@ 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) {
@@ -982,7 +1045,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")
} }
@@ -990,7 +1053,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")
} }
@@ -1001,7 +1064,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") {
@@ -1165,7 +1228,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) err := s.HandleTCPForward(getReq(t, cConn, sReqs), "localhost", 1234, false)
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") {
@@ -1186,44 +1249,50 @@ func TestHandleTCPForward_Failures(t *testing.T) {
assert.NoError(t, err) assert.NoError(t, err)
}(l) }(l)
port := uint16(l.Addr().(*net.TCPAddr).Port) 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 { 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") {
t.Errorf("expected error to contain %q, got %q", "already in use", err.Error()) 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) { t.Run("Registry Register fail", func(t *testing.T) {
s, mRegistry, mPort, _, 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)
mPort.On("SetStatus", mock.AnythingOfType("uint16"), false).Return(nil)
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) err := s.HandleTCPForward(getReq(t, cConn, sReqs), "localhost", 0, false)
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") {
t.Errorf("expected error to contain %q, got %q", "Failed to register", err.Error()) 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) { 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() defer cleanup()
mPort.On("Claim", mock.Anything).Return(true) 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) 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)
time.Sleep(100 * time.Millisecond) _ = sConn.Wait()
err = s.HandleTCPForward(req, "localhost", 0) err = s.HandleTCPForward(req, "localhost", 0, false)
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") {
t.Errorf("expected error to contain %q, got %q", "Failed to finalize forwarding", err.Error()) t.Errorf("expected error to contain %q, got %q", "Failed to finalize forwarding", err.Error())
} }
mPort.AssertExpectations(t)
}) })
} }
+7
View File
@@ -1,11 +1,14 @@
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
} }
@@ -16,9 +19,13 @@ 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
} }
+1 -2
View File
@@ -55,8 +55,7 @@ func TestHTTPServer_Serve(t *testing.T) {
go func() { go func() {
time.Sleep(100 * time.Millisecond) time.Sleep(100 * time.Millisecond)
err = listener.Close() _ = listener.Close()
assert.NoError(t, err)
}() }()
err = srv.Serve(listener) err = srv.Serve(listener)
+28 -9
View File
@@ -1,6 +1,7 @@
package transport package transport
import ( import (
"bufio"
"bytes" "bytes"
"context" "context"
"errors" "errors"
@@ -52,25 +53,43 @@ 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))
buf := make([]byte, hh.config.HeaderSize()) br := bufio.NewReaderSize(conn, hh.config.HeaderSize())
n, err := conn.Read(buf) headerBuf, err := readHTTPHeader(br, hh.config.HeaderSize())
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(buf[:n]) reqhf, err := header.NewRequest(headerBuf)
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)
@@ -101,7 +120,7 @@ func (hh *httpHandler) Handler(conn net.Conn, isTLS bool) {
return return
} }
hw := stream.New(conn, conn, conn.RemoteAddr()) hw := stream.New(conn, br, conn.RemoteAddr())
defer func(hw stream.HTTP) { defer func(hw stream.HTTP) {
err = hw.Close() err = hw.Close()
if err != nil { if err != nil {
+118 -2
View File
@@ -321,8 +321,14 @@ func TestHandler(t *testing.T) {
isTLS: false, isTLS: false,
redirectTLS: false, redirectTLS: false,
request: []byte(""), request: []byte(""),
expected: []byte("HTTP/1.1 400 Bad Request\r\n\r\n"), expected: []byte(""),
setupMocks: func(msr *MockSessionRegistry) { 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() { go func() {
time.Sleep(100 * time.Millisecond) time.Sleep(100 * time.Millisecond)
err = listener.Close() _ = listener.Close()
assert.NoError(t, err)
}() }()
err = srv.Serve(listener) err = srv.Serve(listener)
+1 -2
View File
@@ -45,8 +45,7 @@ func TestTCPServer_Serve(t *testing.T) {
go func() { go func() {
time.Sleep(100 * time.Millisecond) time.Sleep(100 * time.Millisecond)
err = listener.Close() _ = listener.Close()
assert.NoError(t, err)
}() }()
err = srv.Serve(listener) err = srv.Serve(listener)
+3 -6
View File
@@ -2,8 +2,6 @@
"extends": [ "extends": [
"config:recommended" "config:recommended"
], ],
"prConcurrentLimit": 1,
"prHourlyLimit": 1,
"packageRules": [ "packageRules": [
{ {
"matchUpdateTypes": [ "matchUpdateTypes": [
@@ -12,11 +10,10 @@
"pin", "pin",
"digest" "digest"
], ],
"groupName": "all-dependencies",
"automerge": true, "automerge": true,
"matchPackageNames": [ "baseBranchPatterns": [
"*" "staging"
] ]
} }
] ]
} }