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"`