refactor: tidy domain handling
rabbitprincess committed
Mar 2, 2026 at 13:06 UTC
5f808bea13e2278f2b735e15ed64564c96d27661
7 files changed
+144
-179
cmd/relay-server/main.go
+6
-20
@@ -7,7 +7,6 @@ import (
7
"net"
8
"os"
9
"os/signal"
10
- "strconv"
10
"strings"
11
"syscall"
12
"time"
@@ -90,7 +89,10 @@ func runServer(cfg relayServerConfig) error {
89
Strs("bootstrap_uris", cfg.Bootstraps).
90
Msg("[server] frontend configuration")
91
93
- serv, err := portal.NewRelayServer(ctx, cfg.Bootstraps, sniListenAddr, cfg.PortalURL, cfg.KeylessKeyFile, cfg.CloudflareToken)
92
+ baseHost := extractBaseDomain(cfg.PortalURL)
93
+ rootSNI := portalRootHost(cfg.PortalURL)
94
+ apiUpstreamAddr := loopbackForwardAddr(fmt.Sprintf(":%d", cfg.AdminPort))
95
+ serv, err := portal.NewRelayServer(ctx, cfg.Bootstraps, sniListenAddr, baseHost, cfg.KeylessKeyFile, cfg.CloudflareToken)
96
if err != nil {
97
return fmt.Errorf("create relay server: %w", err)
98
}
@@ -134,6 +136,8 @@ func runServer(cfg relayServerConfig) error {
136
reverseConn.Close()
137
})
138
139
+ serv.ConfigurePortalRootFallback(rootSNI, apiUpstreamAddr)
140
+
141
if err := serv.Start(); err != nil {
142
return fmt.Errorf("start relay server: %w", err)
143
}
@@ -155,21 +159,3 @@ func runServer(cfg relayServerConfig) error {
159
log.Info().Msg("[server] shutdown complete")
160
return nil
161
}
158
-
159
-func parsePortNumber(raw string, fallback int, source string) int {
160
- value := strings.TrimSpace(raw)
161
- if value == "" {
162
- return fallback
163
- }
164
- value = strings.TrimPrefix(value, ":")
165
- port, err := strconv.Atoi(value)
166
- if err != nil || port < 1 || port > 65535 {
167
- log.Warn().
168
- Str("source", source).
169
- Str("value", raw).
170
- Int("fallback_port", fallback).
171
- Msg("[server] invalid port value; using fallback")
172
- return fallback
173
- }
174
- return port
175
-}
cmd/relay-server/serve.go
-42
@@ -12,13 +12,11 @@ import (
12
"net/http"
13
"strconv"
14
"strings"
15
- "time"
15
16
"github.com/rs/zerolog/log"
17
18
"gosuda.org/portal/portal"
19
"gosuda.org/portal/portal/keyless"
21
- "gosuda.org/portal/portal/sni"
20
"gosuda.org/portal/sdk"
21
)
22
@@ -135,7 +133,6 @@ func serveAPI(addr string, serv *portal.RelayServer, admin *Admin, frontend *Fro
133
}
134
acmeManager := serv.GetACMEManager()
135
tlsCertFile, tlsKeyFile := acmeManager.TLSFiles()
138
- configurePortalRootFallback(addr, serv)
136
137
go func() {
138
var err error
@@ -327,42 +324,3 @@ func handleKeylessSign(w http.ResponseWriter, r *http.Request, signer *keyless.S
324
writeSignError(w, http.StatusInternalServerError, "failed to encode response")
325
}
326
}
330
-
331
-func configurePortalRootFallback(adminListenAddr string, serv *portal.RelayServer) {
332
- portalRootSNI := portalRootHost(flagPortalURL)
333
- if portalRootSNI == "" {
334
- return
335
- }
336
-
337
- apiAddr, ok := loopbackForwardAddr(adminListenAddr)
338
- if !ok {
339
- log.Warn().
340
- Str("listen_addr", adminListenAddr).
341
- Msg("[SNI] invalid admin listen address; root-domain fallback disabled")
342
- return
343
- }
344
-
345
- serv.GetSNIRouter().SetNoRouteHandler(func(clientConn net.Conn, serverName string) bool {
346
- if !strings.EqualFold(strings.TrimSpace(serverName), portalRootSNI) {
347
- return false
348
- }
349
-
350
- upstreamConn, err := net.DialTimeout("tcp", apiAddr, 5*time.Second)
351
- if err != nil {
352
- log.Warn().
353
- Err(err).
354
- Str("sni", serverName).
355
- Str("upstream", apiAddr).
356
- Msg("[SNI] failed to forward root domain to admin/API listener")
357
- clientConn.Close()
358
- return true
359
- }
360
-
361
- log.Debug().
362
- Str("sni", serverName).
363
- Str("upstream", apiAddr).
364
- Msg("[SNI] forwarding root domain to admin/API listener")
365
- sni.BridgeConnections(clientConn, upstreamConn)
366
- return true
367
- })
368
-}
cmd/relay-server/utils.go
+84
-45
@@ -49,6 +49,52 @@ func parseURLs(raw string) []string {
49
return out
50
}
51
52
+func parsePortNumber(raw string, fallback int, source string) int {
53
+ value := strings.TrimSpace(raw)
54
+ if value == "" {
55
+ return fallback
56
+ }
57
+ value = strings.TrimPrefix(value, ":")
58
+ port, err := strconv.Atoi(value)
59
+ if err != nil || port < 1 || port > 65535 {
60
+ log.Warn().
61
+ Str("source", source).
62
+ Str("value", raw).
63
+ Int("fallback_port", fallback).
64
+ Msg("[server] invalid port value; using fallback")
65
+ return fallback
66
+ }
67
+ return port
68
+}
69
+
70
+func loopbackForwardAddr(listenAddr string) string {
71
+ raw := strings.TrimSpace(listenAddr)
72
+ if raw == "" {
73
+ return ""
74
+ }
75
+
76
+ port := ""
77
+ switch {
78
+ case strings.HasPrefix(raw, ":"):
79
+ port = strings.TrimPrefix(raw, ":")
80
+ case strings.Count(raw, ":") == 0:
81
+ port = raw
82
+ default:
83
+ _, p, err := net.SplitHostPort(raw)
84
+ if err != nil {
85
+ return ""
86
+ }
87
+ port = p
88
+ }
89
+
90
+ portNum, err := strconv.Atoi(port)
91
+ if err != nil || portNum < 1 || portNum > 65535 {
92
+ return ""
93
+ }
94
+
95
+ return net.JoinHostPort("127.0.0.1", strconv.Itoa(portNum))
96
+}
97
+
98
// isSubdomain reports whether host matches the given domain pattern.
99
func isSubdomain(domain, host string) bool {
100
if host == "" || domain == "" {
@@ -155,51 +201,6 @@ func portalHostPort(portalURL string) string {
201
))
202
}
203
158
-// loopbackForwardAddr resolves a listen address into 127.0.0.1:<port>.
159
-func loopbackForwardAddr(listenAddr string) (string, bool) {
160
- raw := strings.TrimSpace(listenAddr)
161
- if raw == "" {
162
- return "", false
163
- }
164
-
165
- port := ""
166
- switch {
167
- case strings.HasPrefix(raw, ":"):
168
- port = strings.TrimPrefix(raw, ":")
169
- case strings.Count(raw, ":") == 0:
170
- port = raw
171
- default:
172
- _, p, err := net.SplitHostPort(raw)
173
- if err != nil {
174
- return "", false
175
- }
176
- port = p
177
- }
178
-
179
- portNum, err := strconv.Atoi(port)
180
- if err != nil || portNum < 1 || portNum > 65535 {
181
- return "", false
182
- }
183
-
184
- return net.JoinHostPort("127.0.0.1", strconv.Itoa(portNum)), true
185
-}
186
-
187
-func portalRootHost(portalURL string) string {
188
- raw := strings.TrimSpace(portalURL)
189
- if raw == "" {
190
- return ""
191
- }
192
- if !strings.Contains(raw, "://") {
193
- raw = "https://" + raw
194
- }
195
-
196
- parsed, err := url.Parse(raw)
197
- if err != nil || parsed.Hostname() == "" {
198
- return ""
199
- }
200
- return strings.TrimPrefix(strings.ToLower(strings.TrimSpace(parsed.Hostname())), "*.")
201
-}
202
-
204
// servicePublicURL returns a service URL derived from portalURL and service name.
205
func servicePublicURL(portalURL, serviceName string) string {
206
serviceName = strings.TrimSpace(serviceName)
@@ -233,6 +234,44 @@ func servicePublicURL(portalURL, serviceName string) string {
234
return fmt.Sprintf("%s://%s.%s", scheme, serviceName, host)
235
}
236
237
+func extractBaseDomain(rawURL string) string {
238
+ trimmed := strings.TrimSpace(rawURL)
239
+ if trimmed == "" {
240
+ return ""
241
+ }
242
+ if !strings.Contains(trimmed, "://") {
243
+ trimmed = "https://" + trimmed
244
+ }
245
+
246
+ u, err := url.Parse(trimmed)
247
+ if err != nil || u.Hostname() == "" {
248
+ return ""
249
+ }
250
+
251
+ host := strings.TrimPrefix(strings.ToLower(strings.TrimSpace(u.Hostname())), "*.")
252
+ parts := strings.Split(host, ".")
253
+ if len(parts) < 2 {
254
+ return ""
255
+ }
256
+ return parts[len(parts)-2] + "." + parts[len(parts)-1]
257
+}
258
+
259
+func portalRootHost(portalURL string) string {
260
+ raw := strings.TrimSpace(portalURL)
261
+ if raw == "" {
262
+ return ""
263
+ }
264
+ if !strings.Contains(raw, "://") {
265
+ raw = "https://" + raw
266
+ }
267
+
268
+ parsed, err := url.Parse(raw)
269
+ if err != nil || parsed.Hostname() == "" {
270
+ return ""
271
+ }
272
+ return strings.TrimPrefix(strings.ToLower(strings.TrimSpace(parsed.Hostname())), "*.")
273
+}
274
+
275
// getContentType returns the MIME type for a file extension
276
func getContentType(ext string) string {
277
switch ext {
portal/acme/acme.go
+4
-30
@@ -12,7 +12,6 @@ import (
12
"encoding/pem"
13
"errors"
14
"fmt"
15
- "net/url"
15
"os"
16
"path/filepath"
17
"strings"
@@ -43,7 +42,7 @@ type provisionConfig struct {
42
}
43
44
type Config struct {
46
- PortalURL string
45
+ BaseDomain string
46
KeyFile string
47
CloudflareToken string
48
}
@@ -55,7 +54,7 @@ type Manager struct {
54
func NewManager(cfg Config) *Manager {
55
return &Manager{
56
cfg: Config{
58
- PortalURL: strings.TrimSpace(cfg.PortalURL),
57
+ BaseDomain: strings.ToLower(strings.TrimSpace(cfg.BaseDomain)),
58
KeyFile: strings.TrimSpace(cfg.KeyFile),
59
CloudflareToken: strings.TrimSpace(cfg.CloudflareToken),
60
},
@@ -119,9 +118,9 @@ func (m *Manager) EnsureSigningKey(ctx context.Context) (string, error) {
118
return keyFile, nil
119
}
120
122
- baseDomain := extractBaseDomain(m.cfg.PortalURL)
121
+ baseDomain := m.cfg.BaseDomain
122
if baseDomain == "" {
124
- return "", fmt.Errorf("derive base domain from PORTAL_URL for ACME provisioning")
123
+ return "", fmt.Errorf("base domain is required for ACME provisioning")
124
}
125
126
cfg, err := buildProvisionConfig(baseDomain, keyFile, m.cfg.CloudflareToken)
@@ -409,31 +408,6 @@ func hasCloudflareToken(cloudflareToken string) bool {
408
return strings.TrimSpace(cloudflareToken) != ""
409
}
410
412
-func extractBaseDomain(portalURL string) string {
413
- raw := strings.TrimSpace(portalURL)
414
- if raw == "" {
415
- return ""
416
- }
417
- if !strings.Contains(raw, "://") {
418
- raw = "https://" + raw
419
- }
420
-
421
- parsed, err := url.Parse(raw)
422
- if err != nil {
423
- return ""
424
- }
425
- if parsed.Hostname() == "" {
426
- return ""
427
- }
428
-
429
- host := strings.TrimPrefix(strings.ToLower(strings.TrimSpace(parsed.Hostname())), "*.")
430
- parts := strings.Split(host, ".")
431
- if len(parts) < 2 {
432
- return ""
433
- }
434
- return parts[len(parts)-2] + "." + parts[len(parts)-1]
435
-}
436
-
411
func fullChainPath(keyFile string) string {
412
return filepath.Join(filepath.Dir(keyFile), fullChainFileName)
413
}
portal/relay.go
+47
-31
@@ -4,7 +4,7 @@ import (
4
"context"
5
"crypto/subtle"
6
"fmt"
7
- "net/url"
7
+ "net"
8
"strings"
9
"sync"
10
"time"
@@ -35,23 +35,18 @@ func NewRelayServer(
35
ctx context.Context,
36
address []string,
37
sniPort string,
38
- portalURL string,
38
+ baseHost string,
39
keylessKey string,
40
cloudflareToken string,
41
) (*RelayServer, error) {
42
- baseDomain := extractBaseDomain(portalURL)
43
- if baseDomain == "" {
44
- log.Warn().Msg("[RelayServer] Could not extract base domain from portal URL")
45
- }
46
-
42
server := &RelayServer{
48
- BaseHost: baseDomain,
43
+ BaseHost: strings.ToLower(strings.TrimSpace(baseHost)),
44
address: address,
45
leaseManager: NewLeaseManager(30 * time.Second),
46
reverseHub: NewReverseHub(),
47
sniRouter: sni.NewRouter(sniPort),
48
acmeManager: acme.NewManager(acme.Config{
54
- PortalURL: portalURL,
49
+ BaseDomain: strings.ToLower(strings.TrimSpace(baseHost)),
50
KeyFile: keylessKey,
51
CloudflareToken: cloudflareToken,
52
}),
@@ -91,28 +86,6 @@ func NewRelayServer(
86
return server, nil
87
}
88
94
-func extractBaseDomain(rawURL string) string {
95
- trimmed := strings.TrimSpace(rawURL)
96
- if trimmed == "" {
97
- return ""
98
- }
99
- if !strings.Contains(trimmed, "://") {
100
- trimmed = "https://" + trimmed
101
- }
102
-
103
- u, err := url.Parse(trimmed)
104
- if err != nil || u.Hostname() == "" {
105
- return ""
106
- }
107
-
108
- host := strings.TrimPrefix(strings.ToLower(strings.TrimSpace(u.Hostname())), "*.")
109
- parts := strings.Split(host, ".")
110
- if len(parts) < 2 {
111
- return ""
112
- }
113
- return parts[len(parts)-2] + "." + parts[len(parts)-1]
114
-}
115
-
89
// GetLeaseManager returns the lease manager instance.
90
func (g *RelayServer) GetLeaseManager() *LeaseManager {
91
return g.leaseManager
@@ -138,6 +111,49 @@ func (g *RelayServer) GetACMEManager() *acme.Manager {
111
return g.acmeManager
112
}
113
114
+// ConfigurePortalRootFallback forwards unmatched root-domain SNI traffic to the provided upstream listener.
115
+func (g *RelayServer) ConfigurePortalRootFallback(rootSNI, upstreamAddr string) {
116
+ if g == nil || g.sniRouter == nil {
117
+ return
118
+ }
119
+
120
+ rootSNI = strings.ToLower(strings.TrimSpace(rootSNI))
121
+ if rootSNI == "" {
122
+ return
123
+ }
124
+
125
+ upstreamAddr = strings.TrimSpace(upstreamAddr)
126
+ if upstreamAddr == "" {
127
+ log.Warn().
128
+ Msg("[RelayServer] root-domain SNI fallback upstream is empty; fallback disabled")
129
+ return
130
+ }
131
+
132
+ g.sniRouter.SetNoRouteHandler(func(clientConn net.Conn, serverName string) bool {
133
+ if !strings.EqualFold(strings.TrimSpace(serverName), rootSNI) {
134
+ return false
135
+ }
136
+
137
+ upstreamConn, err := net.DialTimeout("tcp", upstreamAddr, 5*time.Second)
138
+ if err != nil {
139
+ log.Warn().
140
+ Err(err).
141
+ Str("sni", serverName).
142
+ Str("upstream", upstreamAddr).
143
+ Msg("[SNI] failed to forward root domain to admin/API listener")
144
+ clientConn.Close()
145
+ return true
146
+ }
147
+
148
+ log.Debug().
149
+ Str("sni", serverName).
150
+ Str("upstream", upstreamAddr).
151
+ Msg("[SNI] forwarding root domain to admin/API listener")
152
+ sni.BridgeConnections(clientConn, upstreamConn)
153
+ return true
154
+ })
155
+}
156
+
157
// Start starts the relay server.
158
func (g *RelayServer) Start() error {
159
g.leaseManager.Start()
sdk/client.go
+2
-10
@@ -101,7 +101,7 @@ func (c *Client) Listen(name string, options ...MetadataOption) (net.Listener, e
101
lease := &portal.Lease{
102
ID: generateID(),
103
Name: name,
104
- TLSMode: string(normalizeTLSMode(c.config.TLSMode)),
104
+ TLSMode: string(c.config.TLSMode),
105
ReverseToken: reverseToken,
106
Metadata: portal.Metadata{
107
Description: metadata.Description,
@@ -116,7 +116,7 @@ func (c *Client) Listen(name string, options ...MetadataOption) (net.Listener, e
116
// Build TLS config if enabled
117
var tlsConfig *tls.Config
118
var listenerCloseFns []func()
119
- tlsMode := normalizeTLSMode(c.config.TLSMode)
119
+ tlsMode := c.config.TLSMode
120
tlsEnabled := tlsMode != TLSModeNoTLS
121
if tlsEnabled {
122
switch tlsMode {
@@ -442,11 +442,3 @@ func fetchEndpointCertificateChain(ctx context.Context, endpoint string, serverN
442
}
443
return chainPEM, nil
444
}
445
-
446
-func normalizeTLSMode(mode TLSMode) TLSMode {
447
- normalized := TLSMode(strings.ToLower(strings.TrimSpace(string(mode))))
448
- if normalized == "" {
449
- return TLSModeNoTLS
450
- }
451
- return normalized
452
-}
sdk/listener.go
+1
-1
@@ -340,7 +340,7 @@ func (l *Listener) registerWithRelay() error {
340
LeaseID: l.lease.ID,
341
Name: l.lease.Name,
342
Metadata: l.lease.Metadata,
343
- TLSMode: normalizeTLSMode(TLSMode(l.lease.TLSMode)),
343
+ TLSMode: TLSMode(l.lease.TLSMode),
344
ReverseToken: l.lease.ReverseToken,
345
}
346