Refactor relay server
gosunuts committed
Feb 26, 2026 at 19:00 UTC
1ab6cab5a6451d707b844be3a3be180f46b117ed
16 files changed
+632
-331
cmd/relay-server/main.go
+18
-70
@@ -2,8 +2,6 @@ package main
2
3
import (
4
"context"
5
- "crypto/rand"
6
- "encoding/hex"
5
"flag"
6
"fmt"
7
"net"
@@ -18,7 +16,6 @@ import (
16
17
"gosuda.org/portal/cmd/relay-server/manager"
18
"gosuda.org/portal/portal"
21
- "gosuda.org/portal/portal/utils/cert"
19
"gosuda.org/portal/portal/utils/sni"
20
)
21
@@ -36,6 +33,9 @@ var (
33
flagACMEDNSProvider string
34
flagACMEEmail string
35
flagACMEDirectory string
36
+
37
+ // SNI router
38
+ flagSNIPort string
39
)
40
41
func main() {
@@ -61,15 +61,15 @@ func main() {
61
62
defaultNoIndex := os.Getenv("NOINDEX") == "true"
63
flag.BoolVar(&flagNoIndex, "noindex", defaultNoIndex, "disallow all crawlers via robots.txt (env: NOINDEX)")
64
-
65
- defaultAdminSecretKey := os.Getenv("ADMIN_SECRET_KEY")
66
- flag.StringVar(&flagAdminSecretKey, "admin-secret-key", defaultAdminSecretKey, "secret key for admin authentication (env: ADMIN_SECRET_KEY)")
64
+ flag.StringVar(&flagAdminSecretKey, "admin-secret-key", os.Getenv("ADMIN_SECRET_KEY"), "secret key for admin authentication (env: ADMIN_SECRET_KEY)")
65
66
// ACME DNS-01 flags
67
flag.StringVar(&flagACMEDNSProvider, "acme-dns-provider", os.Getenv("ACME_DNS_PROVIDER"), "DNS provider for ACME DNS-01 challenge (cloudflare, route53)")
68
flag.StringVar(&flagACMEEmail, "acme-email", os.Getenv("ACME_EMAIL"), "email for ACME account registration")
69
flag.StringVar(&flagACMEDirectory, "acme-directory", os.Getenv("ACME_DIRECTORY"), "ACME directory URL (default: Let's Encrypt production)")
70
71
+ flag.StringVar(&flagSNIPort, "sni-port", os.Getenv("SNI_PORT"), "SNI router port for TLS passthrough (env: SNI_PORT)")
72
+
73
flag.Parse()
74
75
flagBootstraps = parseURLs(flagBootstrapsCSV)
@@ -87,67 +87,27 @@ func runServer() error {
87
Str("bootstrap_uris", strings.Join(flagBootstraps, ",")).
88
Msg("[server] frontend configuration")
89
90
- serv := portal.NewRelayServer(flagBootstraps)
91
-
92
- // Create AuthManager for admin authentication
93
- // Auto-generate secret key if not provided
94
- if flagAdminSecretKey == "" {
95
- randomBytes := make([]byte, 16)
96
- if _, err := rand.Read(randomBytes); err != nil {
97
- log.Fatal().Err(err).Msg("[server] failed to generate random admin secret key")
98
- }
99
- flagAdminSecretKey = hex.EncodeToString(randomBytes)
100
- log.Warn().Str("key", flagAdminSecretKey).Msg("[server] auto-generated ADMIN_SECRET_KEY (set ADMIN_SECRET_KEY env to use your own)")
101
- } else {
102
- log.Info().Str("key", flagAdminSecretKey).Msg("[server] admin authentication enabled")
103
- }
104
- authManager := manager.NewAuthManager(flagAdminSecretKey)
105
-
106
- // Create certificate manager if ACME DNS provider is configured
107
- var certManager cert.Manager
108
- if flagACMEDNSProvider != "" && flagACMEEmail != "" {
109
- baseDomain := extractBaseDomain(flagPortalURL)
110
- if baseDomain == "" {
111
- log.Warn().Msg("[server] could not extract base domain from PORTAL_URL, ACME disabled")
112
- } else {
113
- acmeCfg := &cert.ACMEConfig{
114
- BaseDomain: baseDomain,
115
- DNSProviderType: flagACMEDNSProvider,
116
- Email: flagACMEEmail,
117
- DirectoryURL: flagACMEDirectory,
118
- }
119
- var err error
120
- certManager, err = cert.NewACMEManager(ctx, acmeCfg)
121
- if err != nil {
122
- log.Error().Err(err).Msg("[server] failed to create ACME manager, TLSAuto disabled")
123
- } else {
124
- log.Info().
125
- Str("dns_provider", flagACMEDNSProvider).
126
- Str("base_domain", baseDomain).
127
- Msg("[server] ACME certificate manager initialized")
128
- }
129
- }
90
+ if flagSNIPort == "" {
91
+ flagSNIPort = ":443"
92
}
93
+ serv := portal.NewRelayServer(ctx, flagBootstraps, flagSNIPort, flagPortalURL, flagACMEDNSProvider, flagACMEEmail, flagACMEDirectory)
94
132
- // Create Frontend first, then Admin, then attach Admin back to Frontend.
95
frontend := NewFrontend()
96
+ authManager := manager.NewAuthManager(flagAdminSecretKey)
97
admin := NewAdmin(int64(flagLeaseBPS), frontend, authManager)
98
frontend.SetAdmin(admin)
99
100
// Load persisted admin settings (ban list, BPS limits, IP bans)
101
admin.LoadSettings(serv)
102
140
- // Start SNI-based TCP router for TLS passthrough
141
- sniRouter := sni.NewRouter()
142
-
143
- // Set up connection callback to route to tunnel backends
144
- sniRouter.SetConnectionCallback(func(clientConn net.Conn, route *sni.Route) {
103
+ // Set up SNI connection callback to route to tunnel backends
104
+ serv.GetSNIRouter().SetConnectionCallback(func(clientConn net.Conn, route *sni.Route) {
105
if _, ok := serv.GetLeaseManager().GetLeaseByID(route.LeaseID); !ok {
106
log.Warn().
107
Str("lease_id", route.LeaseID).
108
Str("sni", route.SNI).
109
Msg("[SNI] Lease not active; dropping connection and unregistering route")
150
- sniRouter.UnregisterRouteByLeaseID(route.LeaseID)
110
+ serv.GetSNIRouter().UnregisterRouteByLeaseID(route.LeaseID)
111
clientConn.Close()
112
return
113
}
@@ -155,7 +115,7 @@ func runServer() error {
115
// Get BPS manager for rate limiting
116
bpsManager := admin.GetBPSManager()
117
158
- reverseConn, err := serv.GetReverseHub().AcquireStarted(route.LeaseID, portal.ReverseSNIAcquireWait)
118
+ reverseConn, err := serv.GetReverseHub().AcquireForTLS(route.LeaseID, portal.TLSAcquireWait)
119
if err != nil {
120
log.Warn().
121
Err(err).
@@ -165,30 +125,18 @@ func runServer() error {
125
clientConn.Close()
126
return
127
}
168
- defer reverseConn.Close()
128
129
// SNI path is reverse-only (NAT-friendly): relay never dials app directly.
130
manager.EstablishRelayWithBPS(clientConn, reverseConn.Conn, route.LeaseID, bpsManager)
131
+ reverseConn.Close()
132
})
133
174
- // Start SNI router on port 443 (or configurable port)
175
- sniPort := ":443"
176
- if envPort := os.Getenv("SNI_PORT"); envPort != "" {
177
- sniPort = envPort
134
+ if err := serv.Start(); err != nil {
135
+ log.Fatal().Err(err).Msg("[server] Failed to start relay server")
136
}
179
-
180
- if err := sniRouter.Start(sniPort); err != nil {
181
- log.Error().Err(err).Str("port", sniPort).Msg("[server] Failed to start SNI router")
182
- // Continue without SNI router - HTTP proxy still works
183
- } else {
184
- log.Info().Str("port", sniPort).Msg("[server] SNI router started")
185
- defer sniRouter.Stop()
186
- }
187
-
188
- serv.Start()
137
defer serv.Stop()
138
191
- httpSrv := serveHTTP(fmt.Sprintf(":%d", flagPort), sniPort, serv, sniRouter, admin, frontend, flagNoIndex, certManager, stop)
139
+ httpSrv := serveHTTP(fmt.Sprintf(":%d", flagPort), serv, admin, frontend, flagNoIndex, stop)
140
141
<-ctx.Done()
142
log.Info().Msg("[server] shutting down...")
cmd/relay-server/manager/auth_manager.go
+15
@@ -6,6 +6,8 @@ import (
6
"encoding/hex"
7
"sync"
8
"time"
9
+
10
+ "github.com/rs/zerolog/log"
11
)
12
13
const (
@@ -29,6 +31,19 @@ type loginAttempt struct {
31
32
// NewAuthManager creates a new AuthManager with the given secret key
33
func NewAuthManager(secretKey string) *AuthManager {
34
+ // Create AuthManager for admin authentication
35
+ // Auto-generate secret key if not provided
36
+ if secretKey == "" {
37
+ randomBytes := make([]byte, 16)
38
+ if _, err := rand.Read(randomBytes); err != nil {
39
+ log.Fatal().Err(err).Msg("[server] failed to generate random admin secret key")
40
+ }
41
+ secretKey = hex.EncodeToString(randomBytes)
42
+ log.Warn().Str("key", secretKey).Msg("[server] auto-generated ADMIN_SECRET_KEY (set ADMIN_SECRET_KEY env to use your own)")
43
+ } else {
44
+ log.Info().Str("key", secretKey).Msg("[server] admin authentication enabled")
45
+ }
46
+
47
return &AuthManager{
48
secretKey: secretKey,
49
failedLogins: make(map[string]*loginAttempt),
cmd/relay-server/registry.go
+85
-65
@@ -9,33 +9,48 @@ import (
9
"time"
10
11
"github.com/rs/zerolog/log"
12
+ "golang.org/x/net/websocket"
13
+
14
"gosuda.org/portal/portal"
15
"gosuda.org/portal/portal/utils/cert"
14
- "gosuda.org/portal/portal/utils/sni"
16
"gosuda.org/portal/sdk"
17
)
18
19
// SDKRegistry handles HTTP API for SDK lease registration
19
-// Used by both tunnel clients and native applications
20
-type SDKRegistry struct {
21
- server *portal.RelayServer
22
- sniRouter *sni.Router
23
- baseHost string
24
- certManager cert.Manager
20
+type SDKRegistry struct{}
21
+
22
+// HandleSDKRequest routes /sdk/* requests.
23
+func (r *SDKRegistry) HandleSDKRequest(w http.ResponseWriter, req *http.Request, serv *portal.RelayServer) {
24
+ route := strings.Trim(strings.TrimPrefix(req.URL.Path, "/sdk"), "/")
25
+
26
+ switch route {
27
+ case "register":
28
+ r.handleRegister(w, req, serv)
29
+ case "unregister":
30
+ r.handleUnregister(w, req, serv)
31
+ case "renew":
32
+ r.handleRenew(w, req, serv)
33
+ case "csr":
34
+ r.handleCSR(w, req, serv)
35
+ case "domain":
36
+ r.handleDomain(w, req, serv)
37
+ case "connect":
38
+ r.handleConnect(w, req, serv)
39
+ default:
40
+ http.NotFound(w, req)
41
+ }
42
}
43
27
-// NewSDKRegistry creates a new SDK registry
28
-func NewSDKRegistry(server *portal.RelayServer, sniRouter *sni.Router, certManager cert.Manager) *SDKRegistry {
29
- return &SDKRegistry{
30
- server: server,
31
- sniRouter: sniRouter,
32
- baseHost: portalBaseHostNoPort(flagPortalURL),
33
- certManager: certManager,
44
+func (r *SDKRegistry) handleConnect(w http.ResponseWriter, req *http.Request, serv *portal.RelayServer) {
45
+ wsHandler := websocket.Server{
46
+ Handshake: func(*websocket.Config, *http.Request) error { return nil },
47
+ Handler: websocket.Handler(serv.GetReverseHub().HandleConnect),
48
}
49
+ wsHandler.ServeHTTP(w, req)
50
}
51
37
-// HandleRegister handles SDK lease registration requests
38
-func (r *SDKRegistry) HandleRegister(w http.ResponseWriter, req *http.Request) {
52
+// handleRegister handles SDK lease registration requests
53
+func (r *SDKRegistry) handleRegister(w http.ResponseWriter, req *http.Request, serv *portal.RelayServer) {
54
if req.Method != http.MethodPost {
55
w.Header().Set("Allow", http.MethodPost)
56
http.Error(w, "method not allowed", http.StatusMethodNotAllowed)
@@ -88,7 +103,7 @@ func (r *SDKRegistry) HandleRegister(w http.ResponseWriter, req *http.Request) {
103
}
104
105
// Register with lease manager
91
- if !r.server.GetLeaseManager().UpdateLease(lease) {
106
+ if !serv.GetLeaseManager().UpdateLease(lease) {
107
writeJSON(w, sdk.RegisterResponse{
108
Success: false,
109
Message: "failed to register lease (name conflict or policy violation)",
@@ -97,16 +112,19 @@ func (r *SDKRegistry) HandleRegister(w http.ResponseWriter, req *http.Request) {
112
}
113
114
// Clear dropped state in case this is a re-registration after disconnect
100
- r.server.GetReverseHub().ClearDropped(registerReq.LeaseID)
115
+ serv.GetReverseHub().ClearDropped(registerReq.LeaseID)
116
102
- if err := r.registerSNIRoute(registerReq.LeaseID, registerReq.Name); err != nil {
103
- // Keep lease and route state consistent on partial failure.
104
- r.server.GetLeaseManager().DeleteLease(registerReq.LeaseID)
105
- writeJSON(w, sdk.RegisterResponse{
106
- Success: false,
107
- Message: fmt.Sprintf("failed to register SNI route: %v", err),
108
- })
109
- return
117
+ // Only register SNI route for TLS-enabled leases
118
+ if registerReq.TLSEnabled {
119
+ if err := registerSNIRoute(serv, registerReq.LeaseID, registerReq.Name); err != nil {
120
+ // Keep lease and route state consistent on partial failure.
121
+ serv.GetLeaseManager().DeleteLease(registerReq.LeaseID)
122
+ writeJSON(w, sdk.RegisterResponse{
123
+ Success: false,
124
+ Message: fmt.Sprintf("failed to register SNI route: %v", err),
125
+ })
126
+ return
127
+ }
128
}
129
130
log.Info().
@@ -125,8 +143,8 @@ func (r *SDKRegistry) HandleRegister(w http.ResponseWriter, req *http.Request) {
143
})
144
}
145
128
-// HandleUnregister handles SDK lease unregistration requests
129
-func (r *SDKRegistry) HandleUnregister(w http.ResponseWriter, req *http.Request) {
146
+// handleUnregister handles SDK lease unregistration requests
147
+func (r *SDKRegistry) handleUnregister(w http.ResponseWriter, req *http.Request, serv *portal.RelayServer) {
148
if req.Method != http.MethodPost {
149
w.Header().Set("Allow", http.MethodPost)
150
http.Error(w, "method not allowed", http.StatusMethodNotAllowed)
@@ -139,7 +157,7 @@ func (r *SDKRegistry) HandleUnregister(w http.ResponseWriter, req *http.Request)
157
158
if err := json.NewDecoder(req.Body).Decode(&unregisterReq); err != nil {
159
log.Error().Err(err).Msg("[Registry] Failed to decode unregistration request")
142
- writeJSON(w, map[string]interface{}{
160
+ writeJSON(w, map[string]any{
161
"success": false,
162
"message": "invalid request body",
163
})
@@ -147,7 +165,7 @@ func (r *SDKRegistry) HandleUnregister(w http.ResponseWriter, req *http.Request)
165
}
166
167
if unregisterReq.LeaseID == "" {
150
- writeJSON(w, map[string]interface{}{
168
+ writeJSON(w, map[string]any{
169
"success": false,
170
"message": "lease_id is required",
171
})
@@ -155,21 +173,21 @@ func (r *SDKRegistry) HandleUnregister(w http.ResponseWriter, req *http.Request)
173
}
174
175
// Delete from lease manager
158
- if r.server.GetLeaseManager().DeleteLease(unregisterReq.LeaseID) {
176
+ if serv.GetLeaseManager().DeleteLease(unregisterReq.LeaseID) {
177
log.Info().
178
Str("lease_id", unregisterReq.LeaseID).
179
Msg("[Registry] Lease unregistered")
180
}
163
- r.unregisterSNIRoute(unregisterReq.LeaseID)
164
- r.server.GetReverseHub().DropLease(unregisterReq.LeaseID)
181
+ unregisterSNIRoute(serv, unregisterReq.LeaseID)
182
+ serv.GetReverseHub().DropLease(unregisterReq.LeaseID)
183
166
- writeJSON(w, map[string]interface{}{
184
+ writeJSON(w, map[string]any{
185
"success": true,
186
})
187
}
188
171
-// HandleRenew handles SDK lease renewal requests (keepalive)
172
-func (r *SDKRegistry) HandleRenew(w http.ResponseWriter, req *http.Request) {
189
+// handleRenew handles SDK lease renewal requests (keepalive)
190
+func (r *SDKRegistry) handleRenew(w http.ResponseWriter, req *http.Request, serv *portal.RelayServer) {
191
if req.Method != http.MethodPost {
192
w.Header().Set("Allow", http.MethodPost)
193
http.Error(w, "method not allowed", http.StatusMethodNotAllowed)
@@ -183,7 +201,7 @@ func (r *SDKRegistry) HandleRenew(w http.ResponseWriter, req *http.Request) {
201
202
if err := json.NewDecoder(req.Body).Decode(&renewReq); err != nil {
203
log.Error().Err(err).Msg("[Registry] Failed to decode renewal request")
186
- writeJSON(w, map[string]interface{}{
204
+ writeJSON(w, map[string]any{
205
"success": false,
206
"message": "invalid request body",
207
})
@@ -191,14 +209,14 @@ func (r *SDKRegistry) HandleRenew(w http.ResponseWriter, req *http.Request) {
209
}
210
211
if renewReq.LeaseID == "" {
194
- writeJSON(w, map[string]interface{}{
212
+ writeJSON(w, map[string]any{
213
"success": false,
214
"message": "lease_id is required",
215
})
216
return
217
}
218
if strings.TrimSpace(renewReq.ReverseToken) == "" {
201
- writeJSON(w, map[string]interface{}{
219
+ writeJSON(w, map[string]any{
220
"success": false,
221
"message": "reverse_token is required",
222
})
@@ -206,16 +224,16 @@ func (r *SDKRegistry) HandleRenew(w http.ResponseWriter, req *http.Request) {
224
}
225
226
// Get existing lease
209
- entry, ok := r.server.GetLeaseManager().GetLeaseByID(renewReq.LeaseID)
227
+ entry, ok := serv.GetLeaseManager().GetLeaseByID(renewReq.LeaseID)
228
if !ok {
211
- writeJSON(w, map[string]interface{}{
229
+ writeJSON(w, map[string]any{
230
"success": false,
231
"message": "lease not found",
232
})
233
return
234
}
235
if subtle.ConstantTimeCompare([]byte(strings.TrimSpace(entry.Lease.ReverseToken)), []byte(strings.TrimSpace(renewReq.ReverseToken))) != 1 {
218
- writeJSON(w, map[string]interface{}{
236
+ writeJSON(w, map[string]any{
237
"success": false,
238
"message": "unauthorized lease renewal",
239
})
@@ -224,8 +242,8 @@ func (r *SDKRegistry) HandleRenew(w http.ResponseWriter, req *http.Request) {
242
243
// Update expiration
244
entry.Lease.Expires = time.Now().Add(30 * time.Second)
227
- if !r.server.GetLeaseManager().UpdateLease(entry.Lease) {
228
- writeJSON(w, map[string]interface{}{
245
+ if !serv.GetLeaseManager().UpdateLease(entry.Lease) {
246
+ writeJSON(w, map[string]any{
247
"success": false,
248
"message": "failed to renew lease",
249
})
@@ -233,7 +251,7 @@ func (r *SDKRegistry) HandleRenew(w http.ResponseWriter, req *http.Request) {
251
}
252
253
// Re-register route if needed (e.g., router restarted while lease remained active).
236
- if err := r.registerSNIRoute(entry.Lease.ID, entry.Lease.Name); err != nil {
254
+ if err := registerSNIRoute(serv, entry.Lease.ID, entry.Lease.Name); err != nil {
255
log.Warn().
256
Err(err).
257
Str("lease_id", entry.Lease.ID).
@@ -241,32 +259,34 @@ func (r *SDKRegistry) HandleRenew(w http.ResponseWriter, req *http.Request) {
259
Msg("[Registry] Failed to refresh SNI route on renew")
260
}
261
244
- writeJSON(w, map[string]interface{}{
262
+ writeJSON(w, map[string]any{
263
"success": true,
264
})
265
}
266
249
-func (r *SDKRegistry) registerSNIRoute(leaseID, name string) error {
250
- if r.sniRouter == nil {
267
+func registerSNIRoute(serv *portal.RelayServer, leaseID, name string) error {
268
+ sniRouter := serv.GetSNIRouter()
269
+ if sniRouter == nil {
270
return nil
271
}
253
- if r.baseHost == "" {
254
- return fmt.Errorf("invalid app domain configuration")
272
+ if serv.BaseHost == "" {
273
+ return nil
274
}
256
- sniName := strings.ToLower(strings.TrimSpace(name)) + "." + r.baseHost
257
- return r.sniRouter.RegisterRoute(sniName, leaseID, name)
275
+ sniName := strings.ToLower(strings.TrimSpace(name)) + "." + serv.BaseHost
276
+ return sniRouter.RegisterRoute(sniName, leaseID, name)
277
}
278
260
-func (r *SDKRegistry) unregisterSNIRoute(leaseID string) {
261
- if r.sniRouter == nil {
279
+func unregisterSNIRoute(serv *portal.RelayServer, leaseID string) {
280
+ sniRouter := serv.GetSNIRouter()
281
+ if sniRouter == nil {
282
return
283
}
264
- r.sniRouter.UnregisterRouteByLeaseID(leaseID)
284
+ sniRouter.UnregisterRouteByLeaseID(leaseID)
285
}
286
267
-// HandleCSR handles Certificate Signing Request submissions
287
+// handleCSR handles Certificate Signing Request submissions
288
// The tunnel client submits a CSR, and the relay issues a certificate via ACME DNS-01
269
-func (r *SDKRegistry) HandleCSR(w http.ResponseWriter, req *http.Request) {
289
+func (r *SDKRegistry) handleCSR(w http.ResponseWriter, req *http.Request, serv *portal.RelayServer) {
290
if req.Method != http.MethodPost {
291
w.Header().Set("Allow", http.MethodPost)
292
http.Error(w, "method not allowed", http.StatusMethodNotAllowed)
@@ -274,7 +294,7 @@ func (r *SDKRegistry) HandleCSR(w http.ResponseWriter, req *http.Request) {
294
}
295
296
// Check if certificate manager is available
277
- if r.certManager == nil {
297
+ if serv.GetCertManager() == nil {
298
writeJSON(w, sdk.CSRResponse{
299
Success: false,
300
Message: "certificate issuance not configured on this relay",
@@ -316,7 +336,7 @@ func (r *SDKRegistry) HandleCSR(w http.ResponseWriter, req *http.Request) {
336
}
337
338
// Authenticate via lease
319
- entry, ok := r.server.GetLeaseManager().GetLeaseByID(csrReq.LeaseID)
339
+ entry, ok := serv.GetLeaseManager().GetLeaseByID(csrReq.LeaseID)
340
if !ok {
341
writeJSON(w, sdk.CSRResponse{
342
Success: false,
@@ -343,7 +363,7 @@ func (r *SDKRegistry) HandleCSR(w http.ResponseWriter, req *http.Request) {
363
}
364
365
// Validate domain matches lease name + base host
346
- expectedDomain := strings.ToLower(entry.Lease.Name) + "." + r.baseHost
366
+ expectedDomain := strings.ToLower(entry.Lease.Name) + "." + serv.BaseHost
367
if strings.ToLower(csrDomain) != expectedDomain {
368
writeJSON(w, sdk.CSRResponse{
369
Success: false,
@@ -358,7 +378,7 @@ func (r *SDKRegistry) HandleCSR(w http.ResponseWriter, req *http.Request) {
378
CSR: csrReq.CSR,
379
}
380
361
- cert, err := r.certManager.IssueCertificate(req.Context(), certReq)
381
+ cert, err := serv.GetCertManager().IssueCertificate(req.Context(), certReq)
382
if err != nil {
383
log.Error().Err(err).
384
Str("lease_id", csrReq.LeaseID).
@@ -384,9 +404,9 @@ func (r *SDKRegistry) HandleCSR(w http.ResponseWriter, req *http.Request) {
404
})
405
}
406
387
-// HandleDomain returns the relay's base domain for TLS certificate construction.
388
-func (r *SDKRegistry) HandleDomain(w http.ResponseWriter, req *http.Request) {
389
- if r.baseHost == "" {
407
+// handleDomain returns the relay's base domain for TLS certificate construction.
408
+func (r *SDKRegistry) handleDomain(w http.ResponseWriter, req *http.Request, serv *portal.RelayServer) {
409
+ if serv.BaseHost == "" {
410
writeJSON(w, map[string]any{
411
"success": false,
412
"message": "base domain not configured",
@@ -395,6 +415,6 @@ func (r *SDKRegistry) HandleDomain(w http.ResponseWriter, req *http.Request) {
415
}
416
writeJSON(w, map[string]any{
417
"success": true,
398
- "base_domain": r.baseHost,
418
+ "base_domain": serv.BaseHost,
419
})
420
}
cmd/relay-server/serve.go
+17
-21
@@ -11,18 +11,15 @@ import (
11
"strings"
12
13
"github.com/rs/zerolog/log"
14
- "golang.org/x/net/websocket"
14
15
"gosuda.org/portal/portal"
17
- "gosuda.org/portal/portal/utils/cert"
18
- "gosuda.org/portal/portal/utils/sni"
16
)
17
18
//go:embed dist/*
19
var distFS embed.FS
20
21
// serveHTTP builds the HTTP mux and returns the server.
25
-func serveHTTP(addr, sniListenAddr string, serv *portal.RelayServer, sniRouter *sni.Router, admin *Admin, frontend *Frontend, noIndex bool, certManager cert.Manager, cancel context.CancelFunc) *http.Server {
22
+func serveHTTP(addr string, serv *portal.RelayServer, admin *Admin, frontend *Frontend, noIndex bool, cancel context.CancelFunc) *http.Server {
23
if addr == "" {
24
addr = ":0"
25
}
@@ -44,10 +41,15 @@ func serveHTTP(addr, sniListenAddr string, serv *portal.RelayServer, sniRouter *
41
}
42
43
// Portal app assets (JS, CSS, etc.) - served from /app/
47
- appMux.HandleFunc("/app/", withCORSMiddleware(func(w http.ResponseWriter, r *http.Request) {
44
+ appMux.HandleFunc("/app/", func(w http.ResponseWriter, r *http.Request) {
45
+ setCORSHeaders(w)
46
+ if r.Method == http.MethodOptions {
47
+ w.WriteHeader(http.StatusOK)
48
+ return
49
+ }
50
p := strings.TrimPrefix(r.URL.Path, "/app/")
51
frontend.ServeAppStatic(w, r, p, serv)
50
- }))
52
+ })
53
54
// Tunnel installer script and binaries
55
appMux.HandleFunc("/tunnel", func(w http.ResponseWriter, r *http.Request) {
@@ -57,16 +59,10 @@ func serveHTTP(addr, sniListenAddr string, serv *portal.RelayServer, sniRouter *
59
serveTunnelBinary(w, r)
60
})
61
60
- // SDK Registry API for lease registration (used by SDK and tunnel clients)
61
- registry := NewSDKRegistry(serv, sniRouter, certManager)
62
- appMux.HandleFunc("/api/register", registry.HandleRegister)
63
- appMux.HandleFunc("/api/unregister", registry.HandleUnregister)
64
- appMux.HandleFunc("/api/renew", registry.HandleRenew)
65
- appMux.HandleFunc("/api/csr", registry.HandleCSR)
66
- appMux.HandleFunc("/api/domain", registry.HandleDomain)
67
- appMux.Handle("/api/connect", websocket.Server{
68
- Handshake: func(*websocket.Config, *http.Request) error { return nil },
69
- Handler: websocket.Handler(serv.GetReverseHub().HandleConnect),
62
+ // SDK Registry API for lease registration
63
+ registry := &SDKRegistry{}
64
+ appMux.HandleFunc("/sdk/", func(w http.ResponseWriter, r *http.Request) {
65
+ registry.HandleSDKRequest(w, r, serv)
66
})
67
68
// App UI index page - serve React frontend with SSR (delegates to serveAppStatic)
@@ -115,7 +111,7 @@ func serveHTTP(addr, sniListenAddr string, serv *portal.RelayServer, sniRouter *
111
}
112
// TLS is enabled, redirect to HTTPS.
113
log.Debug().Str("host", r.Host).Msg("[server] redirecting to HTTPS")
118
- redirectToHTTPS(w, r, sniListenAddr)
114
+ redirectToHTTPS(w, r, serv.GetSNIRouter().GetAddr())
115
return
116
}
117
appMux.ServeHTTP(w, r)
@@ -158,8 +154,7 @@ func shouldProxyHTTP(host string, serv *portal.RelayServer) bool {
154
log.Debug().
155
Str("lease_name", leaseName).
156
Bool("tls_enabled", entry.Lease.TLSEnabled).
161
- Bool("should_proxy_http", shouldProxy).
162
- Msg("[proxy] shouldProxyHTTP check")
157
+ Msg("[proxy] shouldProxyHTTP")
158
return shouldProxy
159
}
160
@@ -181,7 +176,7 @@ func proxyToHTTP(w http.ResponseWriter, r *http.Request, serv *portal.RelayServe
176
return
177
}
178
184
- targetConn, releaseConn, err := openLeaseConnection(entry.Lease.ID, serv)
179
+ reverseConn, err := serv.GetReverseHub().AcquireForHTTP(entry.Lease.ID, portal.HTTPProxyWait)
180
if err != nil {
181
log.Error().
182
Err(err).
@@ -191,7 +186,8 @@ func proxyToHTTP(w http.ResponseWriter, r *http.Request, serv *portal.RelayServe
186
http.Error(w, "service unavailable", http.StatusServiceUnavailable)
187
return
188
}
194
- defer releaseConn()
189
+ defer reverseConn.Close()
190
+ targetConn := reverseConn.Conn
191
192
// Write the HTTP request to the tunnel
193
if err := r.Write(targetConn); err != nil {
cmd/relay-server/utils.go
-60
@@ -4,7 +4,6 @@ import (
4
"encoding/base64"
5
"encoding/json"
6
"fmt"
7
- "net"
7
"net/http"
8
"net/url"
9
"strings"
@@ -139,11 +138,6 @@ func portalHostPort(portalURL string) string {
138
))
139
}
140
142
-// portalBaseHostNoPort returns host without port from a portal URL-like input.
143
-func portalBaseHostNoPort(portalURL string) string {
144
- return strings.ToLower(strings.TrimSpace(stripPort(portalHostPort(portalURL))))
145
-}
146
-
141
// servicePublicURL returns a service URL derived from portalURL and service name.
142
func servicePublicURL(portalURL, serviceName string) string {
143
serviceName = strings.TrimSpace(serviceName)
@@ -210,39 +204,6 @@ func setCORSHeaders(w http.ResponseWriter) {
204
w.Header().Set("Access-Control-Allow-Headers", "Content-Type, Accept, Accept-Encoding")
205
}
206
213
-// extractBaseDomain extracts the base domain from a URL.
214
-// For example, "https://app.portal.com" -> "portal.com"
215
-func extractBaseDomain(portalURL string) string {
216
- portalURL = strings.TrimSpace(portalURL)
217
- if portalURL == "" {
218
- return ""
219
- }
220
-
221
- for _, prefix := range []string{"https://", "http://"} {
222
- if strings.HasPrefix(strings.ToLower(portalURL), prefix) {
223
- portalURL = portalURL[len(prefix):]
224
- break
225
- }
226
- }
227
-
228
- if idx := strings.Index(portalURL, ":"); idx > 0 {
229
- portalURL = portalURL[:idx]
230
- }
231
-
232
- if idx := strings.Index(portalURL, "/"); idx > 0 {
233
- portalURL = portalURL[:idx]
234
- }
235
-
236
- portalURL = strings.TrimPrefix(portalURL, "*.")
237
-
238
- parts := strings.Split(portalURL, ".")
239
- if len(parts) < 2 {
240
- return ""
241
- }
242
-
243
- return parts[len(parts)-2] + "." + parts[len(parts)-1]
244
-}
245
-
207
// leaseNameFromHost extracts the lease name from a subdomain host.
208
func leaseNameFromHost(host, appURL string) (string, bool) {
209
if !isSubdomain(appURL, host) {
@@ -271,27 +232,6 @@ func leaseNameFromHost(host, appURL string) (string, bool) {
232
return leaseName, true
233
}
234
274
-// openLeaseConnection acquires a reverse connection for the given lease ID.
275
-func openLeaseConnection(leaseID string, serv *portal.RelayServer) (net.Conn, func(), error) {
276
- reverseConn, err := serv.GetReverseHub().AcquireStarted(leaseID, portal.ReverseHTTPWait)
277
- if err != nil {
278
- return nil, nil, fmt.Errorf("no reverse connection available for lease %s: %w", leaseID, err)
279
- }
280
- return reverseConn.Conn, reverseConn.Close, nil
281
-}
282
-
283
-// withCORSMiddleware wraps a handler with CORS headers.
284
-func withCORSMiddleware(h http.HandlerFunc) http.HandlerFunc {
285
- return func(w http.ResponseWriter, r *http.Request) {
286
- setCORSHeaders(w)
287
- if r.Method == http.MethodOptions {
288
- w.WriteHeader(http.StatusOK)
289
- return
290
- }
291
- h(w, r)
292
- }
293
-}
294
-
235
// leaseRow represents a lease entry for display in admin UI and frontend.
236
type leaseRow struct {
237
Peer string
docs/architecture.md
+1
-1
@@ -39,7 +39,7 @@ Portal connects local applications to web users through a secure relay layer wit
39
40
## Connection Flow
41
42
-1. **Register**: App/Tunnel → Relay (`POST /api/register`)
42
+1. **Register**: App/Tunnel → Relay (`POST /sdk/register`)
43
2. **Reverse Connect**: App/Tunnel ← Relay (`TCP reverse tunnel`)
44
3. **Client Request**: Browser → Relay (`GET *.localhost:4017`)
45
4. **Proxy**: Relay ↔ App/Tunnel ↔ Local Service
docs/portal-deploy-guide.md
+3
-3
@@ -134,7 +134,7 @@ portal-tunnel --host localhost:3000 --name myapp --relay https://example.com --t
134
135
How it works:
136
1. Tunnel generates private key and CSR locally
137
-2. Sends CSR to relay via `/api/csr`
137
+2. Sends CSR to relay via `/sdk/csr`
138
3. Relay issues certificate via ACME DNS-01
139
4. TLS connections go directly to tunnel on port 443
140
@@ -166,7 +166,7 @@ curl https://example.com/healthz
166
# Expected: {"status":"ok"}
167
168
# Domain API
169
-curl https://example.com/api/domain
169
+curl https://example.com/sdk/domain
170
# Expected: {"success":true,"base_domain":"example.com"}
171
172
# Tunnel script
@@ -194,7 +194,7 @@ curl -fsSL https://example.com/tunnel | HOST=localhost:3000 NAME=test sh
194
│ ▼ │
195
│ ┌─────────────┐ ┌─────────────┐ │
196
│ │ Admin UI │ │ API │ │
197
- │ │ /admin │ │ /api/* │ │
197
+ │ │ /admin │ │ /sdk/* │ │
198
│ └─────────────┘ └─────────────┘ │
199
│ │
200
└─────────────────────────────────────────────────────┘
portal/relay.go
+98
-3
@@ -1,30 +1,69 @@
1
package portal
2
3
import (
4
+ "context"
5
"crypto/subtle"
6
+ "net/url"
7
"strings"
8
"sync"
9
"time"
10
11
"github.com/rs/zerolog/log"
12
+
13
+ "gosuda.org/portal/portal/utils/cert"
14
+ "gosuda.org/portal/portal/utils/sni"
15
)
16
17
type RelayServer struct {
13
- address []string
18
+ address []string
19
+ BaseHost string
20
21
leaseManager *LeaseManager
22
reverseHub *ReverseHub
23
+ sniRouter *sni.Router
24
+ certManager cert.Manager
25
26
stopch chan struct{}
27
waitgroup sync.WaitGroup
28
}
29
30
// NewRelayServer creates a new relay server.
23
-func NewRelayServer(address []string) *RelayServer {
31
+func NewRelayServer(
32
+ ctx context.Context,
33
+ address []string,
34
+ sniPort string,
35
+ portalURL string,
36
+ acmeDNSProvider string,
37
+ acmeEmail string,
38
+ acmeDirectory string,
39
+) *RelayServer {
40
+ var certManager cert.Manager
41
+ baseDomain := extractBaseDomain(portalURL)
42
+ if baseDomain == "" {
43
+ log.Warn().Msg("[RelayServer] Could not extract base domain from portal URL")
44
+ }
45
+
46
+ acmeCfg := buildACMEConfig(baseDomain, acmeDNSProvider, acmeEmail, acmeDirectory)
47
+ if acmeCfg != nil {
48
+ var err error
49
+ certManager, err = cert.NewACMEManager(ctx, acmeCfg)
50
+ if err != nil {
51
+ log.Error().Err(err).Msg("[RelayServer] Failed to create ACME manager, TLSAuto disabled")
52
+ } else {
53
+ log.Info().
54
+ Str("dns_provider", acmeCfg.DNSProviderType).
55
+ Str("base_domain", acmeCfg.BaseDomain).
56
+ Msg("[RelayServer] ACME certificate manager initialized")
57
+ }
58
+ }
59
+
60
server := &RelayServer{
61
+ BaseHost: baseDomain,
62
address: address,
63
leaseManager: NewLeaseManager(30 * time.Second),
64
reverseHub: NewReverseHub(),
65
+ sniRouter: sni.NewRouter(sniPort),
66
+ certManager: certManager,
67
stopch: make(chan struct{}),
68
}
69
server.leaseManager.SetOnLeaseDeleted(server.reverseHub.DropLease)
@@ -42,6 +81,43 @@ func NewRelayServer(address []string) *RelayServer {
81
return server
82
}
83
84
+func buildACMEConfig(baseDomain, dnsProvider, email, directory string) *cert.ACMEConfig {
85
+ dnsProvider = strings.TrimSpace(dnsProvider)
86
+ email = strings.TrimSpace(email)
87
+ if dnsProvider == "" || email == "" {
88
+ return nil
89
+ }
90
+
91
+ return &cert.ACMEConfig{
92
+ BaseDomain: baseDomain,
93
+ DNSProviderType: dnsProvider,
94
+ Email: email,
95
+ DirectoryURL: strings.TrimSpace(directory),
96
+ }
97
+}
98
+
99
+func extractBaseDomain(rawURL string) string {
100
+ trimmed := strings.TrimSpace(rawURL)
101
+ if trimmed == "" {
102
+ return ""
103
+ }
104
+ if !strings.Contains(trimmed, "://") {
105
+ trimmed = "https://" + trimmed
106
+ }
107
+
108
+ u, err := url.Parse(trimmed)
109
+ if err != nil || u.Hostname() == "" {
110
+ return ""
111
+ }
112
+
113
+ host := strings.TrimPrefix(strings.ToLower(strings.TrimSpace(u.Hostname())), "*.")
114
+ parts := strings.Split(host, ".")
115
+ if len(parts) < 2 {
116
+ return ""
117
+ }
118
+ return parts[len(parts)-2] + "." + parts[len(parts)-1]
119
+}
120
+
121
// GetLeaseManager returns the lease manager instance.
122
func (g *RelayServer) GetLeaseManager() *LeaseManager {
123
return g.leaseManager
@@ -52,16 +128,35 @@ func (g *RelayServer) GetReverseHub() *ReverseHub {
128
return g.reverseHub
129
}
130
131
+// GetSNIRouter returns the SNI router instance.
132
+func (g *RelayServer) GetSNIRouter() *sni.Router {
133
+ return g.sniRouter
134
+}
135
+
136
+// GetCertManager returns the certificate manager instance.
137
+func (g *RelayServer) GetCertManager() cert.Manager {
138
+ return g.certManager
139
+}
140
+
141
// Start starts the relay server.
56
-func (g *RelayServer) Start() {
142
+func (g *RelayServer) Start() error {
143
g.leaseManager.Start()
144
+
145
+ if err := g.sniRouter.Start(); err != nil {
146
+ log.Error().Err(err).Str("addr", g.sniRouter.GetAddr()).Msg("[RelayServer] Failed to start SNI router")
147
+ return err
148
+ }
149
+ log.Info().Str("addr", g.sniRouter.GetAddr()).Msg("[RelayServer] SNI router started")
150
+
151
log.Info().Msg("[RelayServer] Started")
152
+ return nil
153
}
154
155
// Stop stops the relay server.
156
func (g *RelayServer) Stop() {
157
close(g.stopch)
158
g.leaseManager.Stop()
159
+ g.sniRouter.Stop()
160
g.waitgroup.Wait()
161
log.Info().Msg("[RelayServer] Stopped")
162
}
portal/reverse_hub.go
+123
-74
@@ -5,6 +5,7 @@ import (
5
"net"
6
"strings"
7
"sync"
8
+ "sync/atomic"
9
"time"
10
11
"github.com/rs/zerolog/log"
@@ -12,20 +13,40 @@ import (
13
)
14
15
const (
15
- ReverseStartMarker = byte(0x01)
16
- ReverseQueueSize = 64
17
- ReverseAcquireWait = 2 * time.Second
18
- ReverseHTTPWait = 1500 * time.Millisecond
19
- ReverseSNIAcquireWait = 2 * time.Second
20
- ReverseHandleConnectDelay = 2 * time.Second
16
+ // HTTPStartMarker is sent by the relay to activate a reverse connection
17
+ // for HTTP proxy mode.
18
+ HTTPStartMarker = byte(0x01)
19
+
20
+ // TLSStartMarker is sent by the relay to activate a reverse connection
21
+ // for TLS passthrough mode.
22
+ TLSStartMarker = byte(0x02)
23
+
24
+ // QueueSize is the maximum number of pending reverse connections per lease.
25
+ QueueSize = 64
26
+
27
+ // DefaultAcquireTimeout is the default timeout for acquiring a reverse connection.
28
+ DefaultAcquireTimeout = 2 * time.Second
29
+
30
+ // HTTPProxyWait is the timeout for HTTP proxy connections (shorter for better UX).
31
+ HTTPProxyWait = 1500 * time.Millisecond
32
+
33
+ // TLSAcquireWait is the timeout for TLS passthrough connections.
34
+ TLSAcquireWait = 2 * time.Second
35
+
36
+ // AuthFailureDelay is the delay before closing unauthorized connections (rate limiting).
37
+ AuthFailureDelay = 2 * time.Second
38
)
39
40
+// ReverseConn wraps a net.Conn with lifecycle management for the connection pool.
41
type ReverseConn struct {
42
Conn net.Conn
43
done chan struct{}
44
once sync.Once
45
+ // closed tracks local close to help queue consumers skip stale entries.
46
+ closed atomic.Bool
47
}
48
49
+// NewReverseConn creates a new pooled connection.
50
func NewReverseConn(conn net.Conn) *ReverseConn {
51
return &ReverseConn{
52
Conn: conn,
@@ -33,57 +54,64 @@ func NewReverseConn(conn net.Conn) *ReverseConn {
54
}
55
}
56
57
+// Close closes the connection and signals completion.
58
func (c *ReverseConn) Close() {
59
+ c.closed.Store(true)
60
c.Conn.Close()
61
c.once.Do(func() {
62
close(c.done)
63
})
64
}
65
66
+// Wait blocks until the connection is closed.
67
func (c *ReverseConn) Wait() {
68
<-c.done
69
}
70
71
+func (c *ReverseConn) IsClosed() bool {
72
+ return c == nil || c.closed.Load()
73
+}
74
+
75
type ReverseHub struct {
76
mu sync.RWMutex
49
- pending map[string]chan *ReverseConn
50
- dropped map[string]struct{} // leases that have been dropped and should reject offers
51
- authorizer func(string, string) bool
77
+ pools map[string]chan *ReverseConn
78
+ dropped map[string]struct{}
79
+ authorizer func(leaseID, token string) bool
80
}
81
82
+// NewReverseHub creates a new reverse connection hub.
83
func NewReverseHub() *ReverseHub {
84
return &ReverseHub{
56
- pending: make(map[string]chan *ReverseConn),
85
+ pools: make(map[string]chan *ReverseConn),
86
dropped: make(map[string]struct{}),
87
}
88
}
89
61
-func (h *ReverseHub) getOrCreate(leaseID string) chan *ReverseConn {
90
+func (h *ReverseHub) getOrCreatePool(leaseID string) chan *ReverseConn {
91
h.mu.Lock()
92
defer h.mu.Unlock()
93
65
- // Don't create channels for dropped leases
94
if _, dropped := h.dropped[leaseID]; dropped {
95
return nil
96
}
97
70
- ch, ok := h.pending[leaseID]
71
- if ok {
72
- return ch
98
+ if pool, ok := h.pools[leaseID]; ok {
99
+ return pool
100
}
74
- ch = make(chan *ReverseConn, ReverseQueueSize)
75
- h.pending[leaseID] = ch
76
- return ch
101
+ pool := make(chan *ReverseConn, QueueSize)
102
+ h.pools[leaseID] = pool
103
+ return pool
104
}
105
79
-func (h *ReverseHub) get(leaseID string) (chan *ReverseConn, bool) {
106
+func (h *ReverseHub) getPool(leaseID string) (chan *ReverseConn, bool) {
107
h.mu.RLock()
108
defer h.mu.RUnlock()
82
- ch, ok := h.pending[leaseID]
83
- return ch, ok
109
+ pool, ok := h.pools[leaseID]
110
+ return pool, ok
111
}
112
86
-func (h *ReverseHub) SetAuthorizer(authorizer func(string, string) bool) {
113
+// SetAuthorizer sets the authentication function for new connections.
114
+func (h *ReverseHub) SetAuthorizer(authorizer func(leaseID, token string) bool) {
115
h.mu.Lock()
116
defer h.mu.Unlock()
117
h.authorizer = authorizer
@@ -100,82 +128,100 @@ func (h *ReverseHub) isAuthorized(leaseID, token string) bool {
128
}
129
130
func (h *ReverseHub) Offer(leaseID string, conn *ReverseConn) bool {
103
- ch := h.getOrCreate(leaseID)
104
- if ch == nil {
105
- // Lease was dropped, reject the offer
131
+ pool := h.getOrCreatePool(leaseID)
132
+ if pool == nil {
133
return false
134
}
108
- select {
109
- case ch <- conn:
110
- return true
111
- default:
112
- return false
135
+
136
+ for i := 0; i < QueueSize+1; i++ {
137
+ select {
138
+ case pool <- conn:
139
+ return true
140
+ default:
141
+ }
142
+
143
+ // Pool full: evict one oldest entry and retry.
144
+ select {
145
+ case old := <-pool:
146
+ if old != nil {
147
+ old.Close()
148
+ }
149
+ default:
150
+ }
151
}
152
+
153
+ return false
154
}
155
116
-func (h *ReverseHub) Acquire(leaseID string, timeout time.Duration) (*ReverseConn, error) {
117
- ch, ok := h.get(leaseID)
156
+func (h *ReverseHub) AcquireForTLS(leaseID string, timeout time.Duration) (*ReverseConn, error) {
157
+ return h.acquireWithStartMarker(leaseID, timeout, TLSStartMarker, "TLS")
158
+}
159
+
160
+// AcquireForHTTP retrieves a connection for HTTP proxy mode.
161
+// A mode-specific start marker is sent before returning the connection.
162
+func (h *ReverseHub) AcquireForHTTP(leaseID string, timeout time.Duration) (*ReverseConn, error) {
163
+ return h.acquireWithStartMarker(leaseID, timeout, HTTPStartMarker, "HTTP")
164
+}
165
+
166
+func (h *ReverseHub) acquireWithStartMarker(leaseID string, timeout time.Duration, marker byte, mode string) (*ReverseConn, error) {
167
+ pool, ok := h.getPool(leaseID)
168
if !ok {
119
- return nil, fmt.Errorf("no reverse tunnel for lease %s", leaseID)
169
+ return nil, fmt.Errorf("no tunnel available for lease %s", leaseID)
170
}
171
172
if timeout <= 0 {
123
- timeout = ReverseAcquireWait
173
+ timeout = DefaultAcquireTimeout
174
}
175
176
timer := time.NewTimer(timeout)
177
defer timer.Stop()
178
129
- select {
130
- case conn := <-ch:
131
- if conn == nil {
132
- return nil, fmt.Errorf("reverse tunnel unavailable for lease %s", leaseID)
133
- }
134
- return conn, nil
135
- case <-timer.C:
136
- return nil, fmt.Errorf("reverse tunnel timeout for lease %s", leaseID)
137
- }
138
-}
139
-
140
-func (h *ReverseHub) AcquireStarted(leaseID string, timeout time.Duration) (*ReverseConn, error) {
141
- if timeout <= 0 {
142
- timeout = ReverseAcquireWait
143
- }
144
-
179
deadline := time.Now().Add(timeout)
180
for {
181
remaining := time.Until(deadline)
182
if remaining <= 0 {
149
- return nil, fmt.Errorf("reverse tunnel timeout for lease %s", leaseID)
183
+ return nil, fmt.Errorf("tunnel acquisition timeout for lease %s", leaseID)
184
}
151
-
152
- conn, err := h.Acquire(leaseID, remaining)
153
- if err != nil {
154
- return nil, err
185
+ if !timer.Stop() {
186
+ select {
187
+ case <-timer.C:
188
+ default:
189
+ }
190
}
191
+ timer.Reset(remaining)
192
157
- _ = conn.Conn.SetWriteDeadline(time.Now().Add(2 * time.Second))
158
- _, err = conn.Conn.Write([]byte{ReverseStartMarker})
159
- _ = conn.Conn.SetWriteDeadline(time.Time{})
160
- if err == nil {
161
- return conn, nil
162
- }
193
+ select {
194
+ case conn := <-pool:
195
+ if conn == nil || conn.IsClosed() {
196
+ continue
197
+ }
198
+ // Signal tunnel worker to release this connection to application Accept().
199
+ _ = conn.Conn.SetWriteDeadline(time.Now().Add(2 * time.Second))
200
+ _, err := conn.Conn.Write([]byte{marker})
201
+ _ = conn.Conn.SetWriteDeadline(time.Time{})
202
+ if err == nil {
203
+ return conn, nil
204
+ }
205
164
- log.Warn().
165
- Err(err).
166
- Str("lease_id", leaseID).
167
- Msg("[ReverseHub] Failed to start reverse stream; retrying")
168
- conn.Close()
206
+ log.Warn().
207
+ Err(err).
208
+ Str("lease_id", leaseID).
209
+ Str("mode", mode).
210
+ Msg("[ReverseHub] Failed to send start marker; retrying with new connection")
211
+ conn.Close()
212
+ continue
213
+ case <-timer.C:
214
+ return nil, fmt.Errorf("tunnel acquisition timeout for lease %s", leaseID)
215
+ }
216
}
217
}
218
219
func (h *ReverseHub) DropLease(leaseID string) {
220
h.mu.Lock()
174
- ch, ok := h.pending[leaseID]
221
+ pool, ok := h.pools[leaseID]
222
if ok {
176
- delete(h.pending, leaseID)
223
+ delete(h.pools, leaseID)
224
}
178
- // Mark as dropped to prevent new offers from creating channels
225
h.dropped[leaseID] = struct{}{}
226
h.mu.Unlock()
227
@@ -183,10 +229,10 @@ func (h *ReverseHub) DropLease(leaseID string) {
229
return
230
}
231
186
- // Drain and close any pending connections
232
+ // Drain and close pending connections
233
for {
234
select {
189
- case conn := <-ch:
235
+ case conn := <-pool:
236
if conn != nil {
237
conn.Close()
238
}
@@ -214,25 +260,28 @@ func (h *ReverseHub) HandleConnect(ws *websocket.Conn) {
260
leaseID = strings.TrimSpace(req.URL.Query().Get("lease_id"))
261
token = strings.TrimSpace(req.URL.Query().Get("token"))
262
}
263
+
264
if leaseID == "" {
265
log.Warn().Msg("[ReverseHub] Missing lease_id on reverse connect")
219
- time.Sleep(ReverseHandleConnectDelay)
266
+ time.Sleep(AuthFailureDelay)
267
ws.Close()
268
return
269
}
270
+
271
if !h.isAuthorized(leaseID, token) {
272
log.Warn().Str("lease_id", leaseID).Msg("[ReverseHub] Unauthorized reverse connect")
225
- time.Sleep(ReverseHandleConnectDelay)
273
+ time.Sleep(AuthFailureDelay)
274
ws.Close()
275
return
276
}
277
278
conn := NewReverseConn(ws)
279
if !h.Offer(leaseID, conn) {
232
- log.Warn().Str("lease_id", leaseID).Msg("[ReverseHub] Reverse queue full")
280
+ log.Warn().Str("lease_id", leaseID).Msg("[ReverseHub] Connection pool full for lease")
281
conn.Close()
282
return
283
}
284
285
+ // Wait until the connection is used and closed
286
conn.Wait()
287
}
portal/reverse_hub_test.go
+100
-1
@@ -1,6 +1,11 @@
1
package portal
2
3
-import "testing"
3
+import (
4
+ "io"
5
+ "net"
6
+ "testing"
7
+ "time"
8
+)
9
10
func TestReverseHubAuthorization(t *testing.T) {
11
hub := NewReverseHub()
@@ -20,3 +25,97 @@ func TestReverseHubAuthorization(t *testing.T) {
25
t.Fatal("expected unauthorized for wrong token")
26
}
27
}
28
+
29
+func TestAcquireForTLSSendsStartMarker(t *testing.T) {
30
+ hub := NewReverseHub()
31
+ leaseID := "lease-tls-marker"
32
+
33
+ local, peer := net.Pipe()
34
+ defer func() {
35
+ _ = peer.Close()
36
+ }()
37
+ conn := NewReverseConn(local)
38
+ defer conn.Close()
39
+
40
+ if ok := hub.Offer(leaseID, conn); !ok {
41
+ t.Fatal("offer failed")
42
+ }
43
+
44
+ markerRead := make(chan byte, 1)
45
+ readErr := make(chan error, 1)
46
+ go func() {
47
+ var b [1]byte
48
+ _, err := io.ReadFull(peer, b[:])
49
+ if err != nil {
50
+ readErr <- err
51
+ return
52
+ }
53
+ markerRead <- b[0]
54
+ }()
55
+
56
+ got, err := hub.AcquireForTLS(leaseID, 500*time.Millisecond)
57
+ if err != nil {
58
+ t.Fatalf("AcquireForTLS failed: %v", err)
59
+ }
60
+ if got != conn {
61
+ t.Fatal("AcquireForTLS returned unexpected connection")
62
+ }
63
+
64
+ select {
65
+ case err := <-readErr:
66
+ t.Fatalf("failed to read marker: %v", err)
67
+ case b := <-markerRead:
68
+ if b != TLSStartMarker {
69
+ t.Fatalf("unexpected marker: %d", b)
70
+ }
71
+ case <-time.After(500 * time.Millisecond):
72
+ t.Fatal("timed out waiting for start marker")
73
+ }
74
+}
75
+
76
+func TestAcquireForHTTPSendsStartMarker(t *testing.T) {
77
+ hub := NewReverseHub()
78
+ leaseID := "lease-http-marker"
79
+
80
+ local, peer := net.Pipe()
81
+ defer func() {
82
+ _ = peer.Close()
83
+ }()
84
+ conn := NewReverseConn(local)
85
+ defer conn.Close()
86
+
87
+ if ok := hub.Offer(leaseID, conn); !ok {
88
+ t.Fatal("offer failed")
89
+ }
90
+
91
+ markerRead := make(chan byte, 1)
92
+ readErr := make(chan error, 1)
93
+ go func() {
94
+ var b [1]byte
95
+ _, err := io.ReadFull(peer, b[:])
96
+ if err != nil {
97
+ readErr <- err
98
+ return
99
+ }
100
+ markerRead <- b[0]
101
+ }()
102
+
103
+ got, err := hub.AcquireForHTTP(leaseID, 500*time.Millisecond)
104
+ if err != nil {
105
+ t.Fatalf("AcquireForHTTP failed: %v", err)
106
+ }
107
+ if got != conn {
108
+ t.Fatal("AcquireForHTTP returned unexpected connection")
109
+ }
110
+
111
+ select {
112
+ case err := <-readErr:
113
+ t.Fatalf("failed to read marker: %v", err)
114
+ case b := <-markerRead:
115
+ if b != HTTPStartMarker {
116
+ t.Fatalf("unexpected marker: %d", b)
117
+ }
118
+ case <-time.After(500 * time.Millisecond):
119
+ t.Fatal("timed out waiting for start marker")
120
+ }
121
+}
portal/utils/sni/router.go
+13
-6
@@ -39,6 +39,7 @@ type Router struct {
39
routes map[string]*Route // SNI -> Route
40
leases map[string]*Route // LeaseID -> Route
41
listener net.Listener
42
+ addr string
43
44
// Callback for new connections
45
onConnection func(conn net.Conn, route *Route)
@@ -49,14 +50,20 @@ type Router struct {
50
}
51
52
// NewRouter creates a new SNI router
52
-func NewRouter() *Router {
53
+func NewRouter(addr string) *Router {
54
return &Router{
55
+ addr: addr,
56
routes: make(map[string]*Route),
57
leases: make(map[string]*Route),
58
stopCh: make(chan struct{}),
59
}
60
}
61
62
+// GetAddr returns the listen address.
63
+func (r *Router) GetAddr() string {
64
+ return r.addr
65
+}
66
+
67
// SetConnectionCallback sets the callback for new connections
68
func (r *Router) SetConnectionCallback(cb func(conn net.Conn, route *Route)) {
69
r.mu.Lock()
@@ -192,11 +199,11 @@ func (r *Router) GetAllRoutes() []*Route {
199
return routes
200
}
201
195
-// Start starts the SNI router on the given address
196
-func (r *Router) Start(addr string) error {
197
- listener, err := net.Listen("tcp", addr)
202
+// Start starts the SNI router on the configured address.
203
+func (r *Router) Start() error {
204
+ listener, err := net.Listen("tcp", r.addr)
205
if err != nil {
199
- return fmt.Errorf("failed to listen on %s: %w", addr, err)
206
+ return fmt.Errorf("failed to listen on %s: %w", r.addr, err)
207
}
208
209
r.mu.Lock()
@@ -204,7 +211,7 @@ func (r *Router) Start(addr string) error {
211
r.mu.Unlock()
212
213
log.Info().
207
- Str("addr", addr).
214
+ Str("addr", r.addr).
215
Msg("[SNI] Router started")
216
217
r.wg.Add(1)
portal/utils/sni/router_test.go
+10
-10
@@ -5,7 +5,7 @@ import (
5
)
6
7
func TestRouter_RegisterRoute(t *testing.T) {
8
- router := NewRouter()
8
+ router := NewRouter("")
9
10
// Test basic registration
11
err := router.RegisterRoute("example.com", "lease-1", "test")
@@ -29,7 +29,7 @@ func TestRouter_RegisterRoute(t *testing.T) {
29
}
30
31
func TestRouter_UnregisterRoute(t *testing.T) {
32
- router := NewRouter()
32
+ router := NewRouter("")
33
34
err := router.RegisterRoute("example.com", "lease-1", "test")
35
if err != nil {
@@ -45,7 +45,7 @@ func TestRouter_UnregisterRoute(t *testing.T) {
45
}
46
47
func TestRouter_UnregisterRouteByLeaseID(t *testing.T) {
48
- router := NewRouter()
48
+ router := NewRouter("")
49
50
err := router.RegisterRoute("example.com", "lease-1", "test")
51
if err != nil {
@@ -61,7 +61,7 @@ func TestRouter_UnregisterRouteByLeaseID(t *testing.T) {
61
}
62
63
func TestRouter_GetRoute_Wildcard(t *testing.T) {
64
- router := NewRouter()
64
+ router := NewRouter("")
65
66
// Register wildcard route
67
err := router.RegisterRoute("*.example.com", "lease-1", "test")
@@ -96,7 +96,7 @@ func TestRouter_GetRoute_Wildcard(t *testing.T) {
96
}
97
98
func TestRouter_GetRoute_ExactBeforeWildcard(t *testing.T) {
99
- router := NewRouter()
99
+ router := NewRouter("")
100
101
// Register both exact and wildcard routes
102
err := router.RegisterRoute("*.example.com", "lease-1", "wildcard")
@@ -129,7 +129,7 @@ func TestRouter_GetRoute_ExactBeforeWildcard(t *testing.T) {
129
}
130
131
func TestRouter_GetRouteByLeaseID(t *testing.T) {
132
- router := NewRouter()
132
+ router := NewRouter("")
133
134
err := router.RegisterRoute("example.com", "lease-1", "test")
135
if err != nil {
@@ -151,7 +151,7 @@ func TestRouter_GetRouteByLeaseID(t *testing.T) {
151
}
152
153
func TestRouter_GetAllRoutes(t *testing.T) {
154
- router := NewRouter()
154
+ router := NewRouter("")
155
156
_ = router.RegisterRoute("example.com", "lease-1", "test1")
157
_ = router.RegisterRoute("other.com", "lease-2", "test2")
@@ -163,7 +163,7 @@ func TestRouter_GetAllRoutes(t *testing.T) {
163
}
164
165
func TestRouter_CaseInsensitive(t *testing.T) {
166
- router := NewRouter()
166
+ router := NewRouter("")
167
168
err := router.RegisterRoute("Example.COM", "lease-1", "test")
169
if err != nil {
@@ -181,7 +181,7 @@ func TestRouter_CaseInsensitive(t *testing.T) {
181
}
182
183
func TestRouter_LeaseRename(t *testing.T) {
184
- router := NewRouter()
184
+ router := NewRouter("")
185
186
// Register with name1
187
err := router.RegisterRoute("name1.example.com", "lease-1", "name1")
@@ -212,7 +212,7 @@ func TestRouter_LeaseRename(t *testing.T) {
212
}
213
214
func TestRouter_Stop(t *testing.T) {
215
- router := NewRouter()
215
+ router := NewRouter("")
216
217
err := router.RegisterRoute("example.com", "lease-1", "test")
218
if err != nil {
sdk/cert.go
+10
-6
@@ -85,7 +85,7 @@ func (c *CertificateClient) RequestCertificate(ctx context.Context, leaseID, rev
85
return nil, fmt.Errorf("marshal request: %w", err)
86
}
87
88
- url := c.relayAPIURL + "/api/csr"
88
+ url := c.relayAPIURL + "/sdk/csr"
89
req, err := http.NewRequestWithContext(ctx, http.MethodPost, url, bytes.NewReader(body))
90
if err != nil {
91
return nil, fmt.Errorf("create request: %w", err)
@@ -117,7 +117,7 @@ func (c *CertificateClient) RequestCertificate(ctx context.Context, leaseID, rev
117
118
// GetBaseDomain fetches the relay's base domain for TLS certificate construction.
119
func (c *CertificateClient) GetBaseDomain(ctx context.Context) (string, error) {
120
- url := c.relayAPIURL + "/api/domain"
120
+ url := c.relayAPIURL + "/sdk/domain"
121
req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil)
122
if err != nil {
123
return "", fmt.Errorf("create request: %w", err)
@@ -135,16 +135,20 @@ func (c *CertificateClient) GetBaseDomain(ctx context.Context) (string, error) {
135
}
136
137
var domainResp struct {
138
+ Success bool `json:"success"`
139
BaseDomain string `json:"base_domain"`
140
+ Message string `json:"message"`
141
}
142
if err := json.Unmarshal(respBody, &domainResp); err != nil {
143
return "", fmt.Errorf("parse response: %w", err)
144
}
143
-
144
- if domainResp.BaseDomain == "" {
145
- return "", fmt.Errorf("relay did not return base domain")
145
+ if !domainResp.Success {
146
+ msg := domainResp.Message
147
+ if msg == "" {
148
+ msg = "base domain not configured"
149
+ }
150
+ return "", fmt.Errorf("get base domain: %s", msg)
151
}
147
-
152
return domainResp.BaseDomain, nil
153
}
154
sdk/listener.go
+14
-9
@@ -260,7 +260,12 @@ func (l *Listener) reverseAcceptWorker(workerID int) {
260
continue
261
}
262
263
- if err := l.waitForReverseStart(conn); err != nil {
263
+ expectedMarker := portal.HTTPStartMarker
264
+ if l.tlsConfig != nil {
265
+ expectedMarker = portal.TLSStartMarker
266
+ }
267
+
268
+ if err := l.waitForReverseStart(conn, expectedMarker); err != nil {
269
conn.Close()
270
if errors.Is(err, net.ErrClosed) {
271
return
@@ -306,17 +311,17 @@ func (l *Listener) openReverseConnection() (net.Conn, error) {
311
return conn, nil
312
}
313
309
-func (l *Listener) waitForReverseStart(conn net.Conn) error {
314
+func (l *Listener) waitForReverseStart(conn net.Conn, expectedMarker byte) error {
315
var marker [1]byte
316
for {
317
_ = conn.SetReadDeadline(time.Now().Add(reverseReadTimeout))
318
_, err := io.ReadFull(conn, marker[:])
319
if err == nil {
320
_ = conn.SetReadDeadline(time.Time{})
316
- if marker[0] != portal.ReverseStartMarker {
317
- return fmt.Errorf("invalid reverse marker: %d", marker[0])
321
+ if marker[0] == expectedMarker {
322
+ return nil
323
}
319
- return nil
324
+ return fmt.Errorf("invalid reverse marker: %d", marker[0])
325
}
326
327
if netErr, ok := err.(net.Error); ok && netErr.Timeout() {
@@ -346,14 +351,14 @@ func (l *Listener) registerWithRelay() error {
351
ReverseToken: l.lease.ReverseToken,
352
}
353
349
- return l.postJSON("/api/register", reqBody)
354
+ return l.postJSON("/sdk/register", reqBody)
355
}
356
357
func (l *Listener) unregisterFromRelay() error {
358
reqBody := UnregisterRequest{
359
LeaseID: l.lease.ID,
360
}
356
- return l.postJSON("/api/unregister", reqBody)
361
+ return l.postJSON("/sdk/unregister", reqBody)
362
}
363
364
func (l *Listener) sendKeepalive() error {
@@ -361,7 +366,7 @@ func (l *Listener) sendKeepalive() error {
366
LeaseID: l.lease.ID,
367
ReverseToken: l.lease.ReverseToken,
368
}
364
- return l.postJSON("/api/renew", reqBody)
369
+ return l.postJSON("/sdk/renew", reqBody)
370
}
371
372
func (l *Listener) reRegisterLease() error {
@@ -482,7 +487,7 @@ func relayConnectURL(relayAddr, leaseID, token string) (string, error) {
487
default:
488
return "", fmt.Errorf("unsupported relay URL scheme: %q", u.Scheme)
489
}
485
- u.Path = "/api/connect"
490
+ u.Path = "/sdk/connect"
491
q := u.Query()
492
q.Set("lease_id", leaseID)
493
q.Set("token", token)
sdk/listener_test.go
+124
-1
@@ -1,8 +1,13 @@
1
package sdk
2
3
import (
4
+ "crypto/tls"
5
+ "net"
6
"strings"
7
"testing"
8
+ "time"
9
+
10
+ "gosuda.org/portal/portal"
11
)
12
13
func TestNormalizeRelayAPIURL(t *testing.T) {
@@ -70,7 +75,7 @@ func TestRelayConnectURL(t *testing.T) {
75
if err != nil {
76
t.Fatalf("unexpected error: %v", err)
77
}
73
- if !strings.HasPrefix(got, "ws://localhost:4017/api/connect?") {
78
+ if !strings.HasPrefix(got, "ws://localhost:4017/sdk/connect?") {
79
t.Fatalf("unexpected URL prefix: %q", got)
80
}
81
if !strings.Contains(got, "lease_id=lease-1") {
@@ -87,3 +92,121 @@ func TestRelayConnectURL(t *testing.T) {
92
t.Fatal("expected error for empty token")
93
}
94
}
95
+
96
+func TestWaitForReverseStart_HTTPMode(t *testing.T) {
97
+ t.Parallel()
98
+
99
+ l := &Listener{stopCh: make(chan struct{})}
100
+ local, peer := net.Pipe()
101
+ defer local.Close()
102
+ defer peer.Close()
103
+
104
+ done := make(chan error, 1)
105
+ go func() {
106
+ done <- l.waitForReverseStart(local, portal.HTTPStartMarker)
107
+ }()
108
+
109
+ _, err := peer.Write([]byte{portal.HTTPStartMarker})
110
+ if err != nil {
111
+ t.Fatalf("write marker: %v", err)
112
+ }
113
+
114
+ select {
115
+ case err := <-done:
116
+ if err != nil {
117
+ t.Fatalf("waitForReverseStart failed: %v", err)
118
+ }
119
+ case <-time.After(500 * time.Millisecond):
120
+ t.Fatal("timed out waiting for marker")
121
+ }
122
+}
123
+
124
+func TestWaitForReverseStart_TLSMode(t *testing.T) {
125
+ t.Parallel()
126
+
127
+ l := &Listener{
128
+ stopCh: make(chan struct{}),
129
+ tlsConfig: &tls.Config{MinVersion: tls.VersionTLS12},
130
+ }
131
+ local, peer := net.Pipe()
132
+ defer local.Close()
133
+ defer peer.Close()
134
+
135
+ done := make(chan error, 1)
136
+ go func() {
137
+ done <- l.waitForReverseStart(local, portal.TLSStartMarker)
138
+ }()
139
+
140
+ _, err := peer.Write([]byte{portal.TLSStartMarker})
141
+ if err != nil {
142
+ t.Fatalf("write marker: %v", err)
143
+ }
144
+
145
+ select {
146
+ case err := <-done:
147
+ if err != nil {
148
+ t.Fatalf("waitForReverseStart failed: %v", err)
149
+ }
150
+ case <-time.After(500 * time.Millisecond):
151
+ t.Fatal("timed out waiting for marker")
152
+ }
153
+}
154
+
155
+func TestWaitForReverseStart_TLSRejectsHTTPMarker(t *testing.T) {
156
+ t.Parallel()
157
+
158
+ l := &Listener{
159
+ stopCh: make(chan struct{}),
160
+ tlsConfig: &tls.Config{MinVersion: tls.VersionTLS12},
161
+ }
162
+ local, peer := net.Pipe()
163
+ defer local.Close()
164
+ defer peer.Close()
165
+
166
+ done := make(chan error, 1)
167
+ go func() {
168
+ done <- l.waitForReverseStart(local, portal.TLSStartMarker)
169
+ }()
170
+
171
+ _, err := peer.Write([]byte{portal.HTTPStartMarker})
172
+ if err != nil {
173
+ t.Fatalf("write marker: %v", err)
174
+ }
175
+
176
+ select {
177
+ case err := <-done:
178
+ if err == nil {
179
+ t.Fatal("expected invalid marker error")
180
+ }
181
+ case <-time.After(500 * time.Millisecond):
182
+ t.Fatal("timed out waiting for marker")
183
+ }
184
+}
185
+
186
+func TestWaitForReverseStart_HTTPRejectsTLSMarker(t *testing.T) {
187
+ t.Parallel()
188
+
189
+ l := &Listener{stopCh: make(chan struct{})}
190
+ local, peer := net.Pipe()
191
+ defer local.Close()
192
+ defer peer.Close()
193
+
194
+ done := make(chan error, 1)
195
+ go func() {
196
+ done <- l.waitForReverseStart(local, portal.HTTPStartMarker)
197
+ }()
198
+
199
+ _, err := peer.Write([]byte{portal.TLSStartMarker})
200
+ if err != nil {
201
+ t.Fatalf("write marker: %v", err)
202
+ }
203
+
204
+ select {
205
+ case err := <-done:
206
+ if err == nil {
207
+ t.Fatal("expected invalid marker error")
208
+ }
209
+ case <-time.After(500 * time.Millisecond):
210
+ t.Fatal("timed out waiting for marker")
211
+ }
212
+}
sdk/types.go
+1
-1
@@ -119,7 +119,7 @@ func WithHide(hide bool) MetadataOption {
119
}
120
}
121
122
-// API Types for /api/ endpoints
122
+// API Types for /sdk/ endpoints
123
// These types are shared between SDK and relay server
124
type RegisterRequest struct {
125
LeaseID string `json:"lease_id"`