Compare commits

..

3 Commits

Author SHA1 Message Date
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
14 changed files with 11 additions and 44 deletions
-1
View File
@@ -36,7 +36,6 @@ The following environment variables can be configured in the `.env` file:
| Variable | Description | Default | Required | | Variable | Description | Default | Required |
|---------------------|-----------------------------------------------------------------------------|-------------------------|---------------------| |---------------------|-----------------------------------------------------------------------------|-------------------------|---------------------|
| `DOMAIN` | Domain name for subdomain routing | `localhost` | No | | `DOMAIN` | Domain name for subdomain routing | `localhost` | No |
| `FRONTEND_URL` | URL for the frontend dashboard/landing page | `https://<DOMAIN>` | No |
| `PORT` | SSH server port | `2200` | No | | `PORT` | SSH server port | `2200` | No |
| `HTTP_PORT` | HTTP server port | `8080` | No | | `HTTP_PORT` | HTTP server port | `8080` | No |
| `HTTPS_PORT` | HTTPS server port | `8443` | No | | `HTTPS_PORT` | HTTPS server port | `8443` | No |
-1
View File
@@ -80,7 +80,6 @@ type MockConfig struct {
} }
func (m *MockConfig) Domain() string { return m.Called().String(0) } func (m *MockConfig) Domain() string { return m.Called().String(0) }
func (m *MockConfig) FrontendURL() string { return m.Called().String(0) }
func (m *MockConfig) SSHPort() string { return m.Called().String(0) } func (m *MockConfig) SSHPort() string { return m.Called().String(0) }
func (m *MockConfig) HTTPPort() string { return m.Called().String(0) } func (m *MockConfig) HTTPPort() string { return m.Called().String(0) }
func (m *MockConfig) HTTPSPort() string { return m.Called().String(0) } func (m *MockConfig) HTTPSPort() string { return m.Called().String(0) }
-2
View File
@@ -6,7 +6,6 @@ import (
type Config interface { type Config interface {
Domain() string Domain() string
FrontendURL() string
SSHPort() string SSHPort() string
HTTPPort() string HTTPPort() string
@@ -51,7 +50,6 @@ func MustLoad() (Config, error) {
} }
func (c *config) Domain() string { return c.domain } func (c *config) Domain() string { return c.domain }
func (c *config) FrontendURL() string { return c.frontendURL }
func (c *config) SSHPort() string { return c.sshPort } func (c *config) SSHPort() string { return c.sshPort }
func (c *config) HTTPPort() string { return c.httpPort } func (c *config) HTTPPort() string { return c.httpPort }
func (c *config) HTTPSPort() string { return c.httpsPort } func (c *config) HTTPSPort() string { return c.httpsPort }
+2 -5
View File
@@ -12,9 +12,8 @@ import (
) )
type config struct { type config struct {
domain string domain string
frontendURL string sshPort string
sshPort string
httpPort string httpPort string
httpsPort string httpsPort string
@@ -50,7 +49,6 @@ func parse() (*config, error) {
} }
domain := getenv("DOMAIN", "localhost") domain := getenv("DOMAIN", "localhost")
frontendURL := getenv("FRONTEND_URL", "https://"+domain)
sshPort := getenv("PORT", "2200") sshPort := getenv("PORT", "2200")
httpPort := getenv("HTTP_PORT", "8080") httpPort := getenv("HTTP_PORT", "8080")
@@ -91,7 +89,6 @@ func parse() (*config, error) {
return &config{ return &config{
domain: domain, domain: domain,
frontendURL: frontendURL,
sshPort: sshPort, sshPort: sshPort,
httpPort: httpPort, httpPort: httpPort,
httpsPort: httpsPort, httpsPort: httpsPort,
-1
View File
@@ -753,7 +753,6 @@ type MockConfig struct {
} }
func (m *MockConfig) Domain() string { return m.Called().String(0) } func (m *MockConfig) Domain() string { return m.Called().String(0) }
func (m *MockConfig) FrontendURL() string { return m.Called().String(0) }
func (m *MockConfig) SSHPort() string { return m.Called().String(0) } func (m *MockConfig) SSHPort() string { return m.Called().String(0) }
func (m *MockConfig) HTTPPort() string { return m.Called().String(0) } func (m *MockConfig) HTTPPort() string { return m.Called().String(0) }
func (m *MockConfig) HTTPSPort() string { return m.Called().String(0) } func (m *MockConfig) HTTPSPort() string { return m.Called().String(0) }
+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
-1
View File
@@ -33,7 +33,6 @@ type MockConfig struct {
} }
func (m *MockConfig) Domain() string { return m.Called().String(0) } func (m *MockConfig) Domain() string { return m.Called().String(0) }
func (m *MockConfig) FrontendURL() string { return m.Called().String(0) }
func (m *MockConfig) SSHPort() string { return m.Called().String(0) } func (m *MockConfig) SSHPort() string { return m.Called().String(0) }
func (m *MockConfig) HTTPPort() string { return m.Called().String(0) } func (m *MockConfig) HTTPPort() string { return m.Called().String(0) }
func (m *MockConfig) HTTPSPort() string { return m.Called().String(0) } func (m *MockConfig) HTTPSPort() string { return m.Called().String(0) }
@@ -24,7 +24,6 @@ type mockConfig struct {
} }
func (m *mockConfig) Domain() string { return m.Called().String(0) } func (m *mockConfig) Domain() string { return m.Called().String(0) }
func (m *mockConfig) FrontendURL() string { return m.Called().String(0) }
func (m *mockConfig) SSHPort() string { return m.Called().String(0) } func (m *mockConfig) SSHPort() string { return m.Called().String(0) }
func (m *mockConfig) HTTPPort() string { return m.Called().String(0) } func (m *mockConfig) HTTPPort() string { return m.Called().String(0) }
func (m *mockConfig) HTTPSPort() string { return m.Called().String(0) } func (m *mockConfig) HTTPSPort() string { return m.Called().String(0) }
@@ -32,7 +32,6 @@ type MockConfig struct {
} }
func (m *MockConfig) Domain() string { return m.Called().String(0) } func (m *MockConfig) Domain() string { return m.Called().String(0) }
func (m *MockConfig) FrontendURL() string { return m.Called().String(0) }
func (m *MockConfig) SSHPort() string { return m.Called().String(0) } func (m *MockConfig) SSHPort() string { return m.Called().String(0) }
func (m *MockConfig) HTTPPort() string { return m.Called().String(0) } func (m *MockConfig) HTTPPort() string { return m.Called().String(0) }
func (m *MockConfig) HTTPSPort() string { return m.Called().String(0) } func (m *MockConfig) HTTPSPort() string { return m.Called().String(0) }
-9
View File
@@ -363,28 +363,19 @@ func (s *session) HandleTCPForward(req *ssh.Request, addr string, portToBind uin
} }
} }
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))
} }
+4 -13
View File
@@ -38,9 +38,8 @@ type mockConfig struct {
config.Config config.Config
} }
func (m *mockConfig) Domain() string { return m.Called().String(0) } func (m *mockConfig) Domain() string { return m.Called().String(0) }
func (m *mockConfig) FrontendURL() string { return m.Called().String(0) } func (m *mockConfig) SSHPort() string { return m.Called().String(0) }
func (m *mockConfig) SSHPort() string { return m.Called().String(0) }
func (m *mockConfig) Mode() types.ServerMode { func (m *mockConfig) Mode() types.ServerMode {
args := m.Called() args := m.Called()
if args.Get(0) == nil { if args.Get(0) == nil {
@@ -706,7 +705,6 @@ 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)
@@ -725,7 +723,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, sConn, sReqs, cConn, cleanup := setup(t) s, mRegistry, _, mRandom, _, 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)
@@ -744,7 +742,7 @@ func TestForwardingFailures(t *testing.T) {
err := cConn.Close() err := cConn.Close()
assert.NoError(t, err) assert.NoError(t, err)
_ = sConn.Wait() time.Sleep(50 * time.Millisecond)
err = s.HandleTCPIPForward(req) err = s.HandleTCPIPForward(req)
assert.Error(t, err) assert.Error(t, err)
@@ -766,7 +764,6 @@ 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)
@@ -1250,7 +1247,6 @@ 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, false) err = s.HandleTCPForward(getReq(t, cConn, sReqs), "localhost", port, false)
if err == nil { if err == nil {
@@ -1258,14 +1254,12 @@ func TestHandleTCPForward_Failures(t *testing.T) {
} 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, false) err := s.HandleTCPForward(getReq(t, cConn, sReqs), "localhost", 0, false)
if err == nil { if err == nil {
@@ -1273,14 +1267,12 @@ func TestHandleTCPForward_Failures(t *testing.T) {
} 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, sConn, 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()
@@ -1293,7 +1285,6 @@ func TestHandleTCPForward_Failures(t *testing.T) {
} 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)
}) })
} }
+1 -1
View File
@@ -116,7 +116,7 @@ func (hh *httpHandler) Handler(conn net.Conn, isTLS bool) {
Type: types.TunnelTypeHTTP, Type: types.TunnelTypeHTTP,
}) })
if err != nil { if err != nil {
_ = hh.redirect(conn, http.StatusMovedPermanently, fmt.Sprintf("%s/tunnel-not-found?slug=%s\r\n", hh.config.FrontendURL(), slug)) _ = hh.redirect(conn, http.StatusMovedPermanently, fmt.Sprintf("https://tunnl.live/tunnel-not-found?slug=%s\r\n", slug))
return return
} }
+1 -4
View File
@@ -223,7 +223,6 @@ func TestNewHTTPHandler(t *testing.T) {
msr := new(MockSessionRegistry) msr := new(MockSessionRegistry)
mockConfig := &MockConfig{} mockConfig := &MockConfig{}
mockConfig.On("Domain").Return("domain") mockConfig.On("Domain").Return("domain")
mockConfig.On("FrontendURL").Return("https://domain")
mockConfig.On("TLSRedirect").Return(false) mockConfig.On("TLSRedirect").Return(false)
hh := newHTTPHandler(mockConfig, msr) hh := newHTTPHandler(mockConfig, msr)
assert.NotNil(t, hh) assert.NotNil(t, hh)
@@ -291,7 +290,7 @@ func TestHandler(t *testing.T) {
isTLS: true, isTLS: true,
redirectTLS: false, redirectTLS: false,
request: []byte("GET / HTTP/1.1\r\nHost: test.domain\r\n\r\n"), request: []byte("GET / HTTP/1.1\r\nHost: test.domain\r\n\r\n"),
expected: []byte("HTTP/1.1 301 Moved Permanently\r\nLocation: https://example.com/tunnel-not-found?slug=test\r\nContent-Length: 0\r\nConnection: close\r\n\r\n"), expected: []byte("HTTP/1.1 301 Moved Permanently\r\nLocation: https://tunnl.live/tunnel-not-found?slug=test\r\nContent-Length: 0\r\nConnection: close\r\n\r\n"),
setupMocks: func(msr *MockSessionRegistry) { setupMocks: func(msr *MockSessionRegistry) {
msr.On("Get", types.SessionKey{ msr.On("Get", types.SessionKey{
Id: "test", Id: "test",
@@ -617,7 +616,6 @@ func TestHandler(t *testing.T) {
mockConfig := &MockConfig{} mockConfig := &MockConfig{}
port := "0" port := "0"
mockConfig.On("Domain").Return("example.com") mockConfig.On("Domain").Return("example.com")
mockConfig.On("FrontendURL").Return("https://example.com")
mockConfig.On("HTTPPort").Return(port) mockConfig.On("HTTPPort").Return(port)
mockConfig.On("HeaderSize").Return(4096) mockConfig.On("HeaderSize").Return(4096)
mockConfig.On("TLSRedirect").Return(true) mockConfig.On("TLSRedirect").Return(true)
@@ -728,7 +726,6 @@ func TestHandlerForwardsPostBody(t *testing.T) {
mockSessionRegistry := new(MockSessionRegistry) mockSessionRegistry := new(MockSessionRegistry)
mockConfig := &MockConfig{} mockConfig := &MockConfig{}
mockConfig.On("Domain").Return("example.com") mockConfig.On("Domain").Return("example.com")
mockConfig.On("FrontendURL").Return("https://example.com")
mockConfig.On("HTTPPort").Return("0") mockConfig.On("HTTPPort").Return("0")
mockConfig.On("HeaderSize").Return(4096) mockConfig.On("HeaderSize").Return(4096)
mockConfig.On("TLSRedirect").Return(true) mockConfig.On("TLSRedirect").Return(true)
-1
View File
@@ -25,7 +25,6 @@ type MockConfig struct {
} }
func (m *MockConfig) Domain() string { return m.Called().String(0) } func (m *MockConfig) Domain() string { return m.Called().String(0) }
func (m *MockConfig) FrontendURL() string { return m.Called().String(0) }
func (m *MockConfig) SSHPort() string { return m.Called().String(0) } func (m *MockConfig) SSHPort() string { return m.Called().String(0) }
func (m *MockConfig) HTTPPort() string { return m.Called().String(0) } func (m *MockConfig) HTTPPort() string { return m.Called().String(0) }
func (m *MockConfig) HTTPSPort() string { return m.Called().String(0) } func (m *MockConfig) HTTPSPort() string { return m.Called().String(0) }