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