refact portal and relay server
Kim committed
Mar 11, 2026 at 11:32 UTC
5581fdeb409ea7fd51df2681f789df7db456574e
12 files changed
+895
-493
cmd/relay-server/serve.go
+50
-106
@@ -14,7 +14,6 @@ import (
14
"github.com/gosuda/portal/v2/portal"
15
"github.com/gosuda/portal/v2/portal/acme"
16
"github.com/gosuda/portal/v2/portal/admin"
17
- "github.com/gosuda/portal/v2/portal/keyless"
17
"github.com/gosuda/portal/v2/portal/policy"
18
"github.com/gosuda/portal/v2/types"
19
)
@@ -37,70 +36,39 @@ func runServer(cfg relayServerConfig) error {
36
}
37
policy.SetTrustedProxyCIDRs(trustedProxyCIDRs)
38
40
- acmeManager, err := acme.NewManager(acme.Config{
41
- BaseDomain: rootHost,
42
- KeyDir: cfg.KeylessDir,
43
- DNSProvider: cfg.ACMEDNSProvider,
44
- CloudflareToken: cfg.CloudflareToken,
45
- AWSAccessKeyID: cfg.AWSAccessKeyID,
46
- AWSSecretAccessKey: cfg.AWSSecretAccessKey,
47
- AWSSessionToken: cfg.AWSSessionToken,
48
- AWSRegion: cfg.AWSRegion,
49
- AWSHostedZoneID: cfg.AWSHostedZoneID,
50
- })
51
- if err != nil {
52
- return fmt.Errorf("create acme manager: %w", err)
53
- }
54
-
55
- certFile, keyFile, err := acmeManager.EnsureCertificate(ctx)
56
- if err != nil {
57
- return fmt.Errorf("ensure relay certificate: %w", err)
58
- }
59
- signer, err := keyless.NewSigner(keyFile)
60
- if err != nil {
61
- return fmt.Errorf("create keyless signer: %w", err)
62
- }
63
-
64
- frontend := NewFrontend(cfg.PortalURL)
65
- adminHandler := admin.NewHandler(cfg.PortalURL, cfg.AdminSecretKey, "admin_settings.json", cfg.TrustProxyHeaders, func(w http.ResponseWriter, r *http.Request, appPath string) {
66
- frontend.ServeAppStatic(w, r, appPath)
67
- })
68
- if loadErr := adminHandler.LoadSettings(); loadErr != nil {
69
- logger.Warn().Err(loadErr).Msg("load admin settings")
70
- }
71
-
39
server, err := portal.NewServer(portal.ServerConfig{
73
- PortalURL: cfg.PortalURL,
40
+ PortalURL: cfg.PortalURL,
41
+ ACME: acme.Config{
42
+ KeyDir: cfg.KeylessDir,
43
+ DNSProvider: cfg.ACMEDNSProvider,
44
+ CloudflareToken: cfg.CloudflareToken,
45
+ AWSAccessKeyID: cfg.AWSAccessKeyID,
46
+ AWSSecretAccessKey: cfg.AWSSecretAccessKey,
47
+ AWSSessionToken: cfg.AWSSessionToken,
48
+ AWSRegion: cfg.AWSRegion,
49
+ AWSHostedZoneID: cfg.AWSHostedZoneID,
50
+ },
51
APIListenAddr: apiListenAddr,
52
SNIListenAddr: sniListenAddr,
76
- RootHost: rootHost,
77
- RootFallbackAddr: portal.HostPortOrLoopback(apiListenAddr),
78
- Policy: adminHandler.Runtime(),
53
TrustProxyHeaders: cfg.TrustProxyHeaders,
80
- KeylessSignerHandler: func() http.Handler {
81
- if signer == nil {
82
- return nil
83
- }
84
- return signer.Handler()
85
- }(),
86
- APITLS: keyless.TLSMaterialConfig{
87
- CertPEM: mustRead(certFile),
88
- KeyPEM: mustRead(keyFile),
89
- },
90
- APIHandlerWrapper: serveAPI(frontend, adminHandler, cfg),
54
})
55
if err != nil {
56
return fmt.Errorf("create relay server: %w", err)
57
}
58
59
+ frontend := NewFrontend(cfg.PortalURL)
60
+ adminHandler := admin.NewHandler(cfg.PortalURL, cfg.AdminSecretKey, "admin_settings.json", cfg.TrustProxyHeaders, func(w http.ResponseWriter, r *http.Request, appPath string) {
61
+ frontend.ServeAppStatic(w, r, appPath)
62
+ })
63
frontend.Bind(server)
64
adminHandler.Bind(server)
65
+ if loadErr := adminHandler.LoadSettings(); loadErr != nil {
66
+ logger.Warn().Err(loadErr).Msg("load admin settings")
67
+ }
68
99
- if err := server.Start(ctx); err != nil {
69
+ if err := server.Start(ctx, newAPIMux(frontend, adminHandler, cfg)); err != nil {
70
return fmt.Errorf("start relay server: %w", err)
71
}
102
- acmeManager.Start(ctx)
103
- defer acmeManager.Stop()
72
73
logger.Info().
74
Str("api_addr", portal.HostPortOrLoopback(server.APIAddr())).
@@ -113,63 +81,47 @@ func runServer(cfg relayServerConfig) error {
81
return server.Wait()
82
}
83
116
-func serveAPI(frontend *Frontend, adminHandler *admin.Handler, cfg relayServerConfig) func(http.Handler) http.Handler {
117
- return func(base http.Handler) http.Handler {
118
- return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
119
- switch {
120
- case isRelayControlPlanePath(r.URL.Path):
121
- base.ServeHTTP(w, r)
122
- case r.URL.Path == types.PathV1Sign:
123
- base.ServeHTTP(w, r)
124
- case r.URL.Path == types.PathHealthz:
125
- base.ServeHTTP(w, r)
126
- case isFrontendRootAssetPath(r.URL.Path):
127
- frontend.ServeAsset(w, r, strings.TrimPrefix(r.URL.Path, "/"), "")
128
- case strings.HasPrefix(strings.TrimSpace(r.URL.Path), types.PathAssetsPrefix):
129
- frontend.ServeAsset(w, r, strings.TrimPrefix(r.URL.Path, "/"), "")
130
- case r.URL.Path == types.PathRoot || r.URL.Path == types.PathApp || r.URL.Path == types.PathAppPrefix:
131
- frontend.ServeAppStatic(w, r, "")
132
- case strings.HasPrefix(strings.TrimSpace(r.URL.Path), types.PathAppPrefix):
133
- frontend.ServeAppStatic(w, r, strings.TrimPrefix(strings.TrimSpace(r.URL.Path), types.PathAppPrefix))
134
- case r.URL.Path == types.PathAdmin || r.URL.Path == types.PathAdminPrefix:
135
- adminHandler.HandleRequest(w, r)
136
- case strings.HasPrefix(strings.TrimSpace(r.URL.Path), types.PathAdminPrefix):
137
- adminHandler.HandleRequest(w, r)
138
- case r.URL.Path == types.PathTunnel:
139
- serveTunnelScript(w, r, cfg.PortalURL)
140
- case strings.HasPrefix(strings.TrimSpace(r.URL.Path), types.PathTunnelBinPrefix):
141
- serveTunnelBinary(w, r)
142
- default:
143
- base.ServeHTTP(w, r)
144
- }
84
+func newAPIMux(frontend *Frontend, adminHandler *admin.Handler, cfg relayServerConfig) *http.ServeMux {
85
+ mux := http.NewServeMux()
86
+
87
+ mux.HandleFunc("/{$}", func(w http.ResponseWriter, r *http.Request) {
88
+ frontend.ServeAppStatic(w, r, "")
89
+ })
90
+ mux.HandleFunc(types.PathApp, func(w http.ResponseWriter, r *http.Request) {
91
+ frontend.ServeAppStatic(w, r, "")
92
+ })
93
+ mux.HandleFunc(types.PathAppPrefix, func(w http.ResponseWriter, r *http.Request) {
94
+ frontend.ServeAppStatic(w, r, strings.TrimPrefix(strings.TrimSpace(r.URL.Path), types.PathAppPrefix))
95
+ })
96
+ mux.HandleFunc(types.PathAssetsPrefix, func(w http.ResponseWriter, r *http.Request) {
97
+ frontend.ServeAsset(w, r, strings.TrimPrefix(r.URL.Path, "/"), "")
98
+ })
99
+ for _, assetPath := range frontendRootAssetPaths() {
100
+ mux.HandleFunc(assetPath, func(w http.ResponseWriter, r *http.Request) {
101
+ frontend.ServeAsset(w, r, strings.TrimPrefix(assetPath, "/"), "")
102
})
103
}
104
+
105
+ mux.HandleFunc(types.PathAdmin, adminHandler.HandleRequest)
106
+ mux.HandleFunc(types.PathAdminPrefix, adminHandler.HandleRequest)
107
+ mux.HandleFunc(types.PathTunnel, func(w http.ResponseWriter, r *http.Request) {
108
+ serveTunnelScript(w, r, cfg.PortalURL)
109
+ })
110
+ mux.HandleFunc(types.PathTunnelBinPrefix, serveTunnelBinary)
111
+
112
+ return mux
113
}
114
149
-func isFrontendRootAssetPath(requestPath string) bool {
150
- switch requestPath {
151
- case "/favicon.ico",
115
+func frontendRootAssetPaths() []string {
116
+ return []string{
117
+ "/favicon.ico",
118
"/favicon.svg",
119
"/favicon-96x96.png",
120
"/apple-touch-icon.png",
121
"/web-app-manifest-192x192.png",
122
"/web-app-manifest-512x512.png",
157
- "/portal.jpg":
158
- return true
159
- default:
160
- return false
161
- }
162
-}
163
-
164
-func mustRead(path string) []byte {
165
- if path == "" {
166
- log.Fatal().Msg("missing required PEM path")
123
+ "/portal.jpg",
124
}
168
- data, err := os.ReadFile(path)
169
- if err != nil {
170
- log.Fatal().Err(err).Str("path", path).Msg("read pem file")
171
- }
172
- return data
125
}
126
127
func parseURLs(raw string) []string {
@@ -186,11 +138,3 @@ func parseURLs(raw string) []string {
138
}
139
return out
140
}
189
-
190
-func isRelayControlPlanePath(path string) bool {
191
- switch strings.TrimSpace(path) {
192
- case types.PathSDKRegister, types.PathSDKConnect, types.PathSDKRenew, types.PathSDKUnregister, types.PathSDKDomain:
193
- return true
194
- }
195
- return strings.HasPrefix(path, types.PathSDKPrefix)
196
-}
portal/acme/acme.go
+17
@@ -151,6 +151,23 @@ func (m *Manager) EnsureCertificate(ctx context.Context) (string, string, error)
151
return m.TLSFiles()
152
}
153
154
+func (m *Manager) EnsureTLSMaterial(ctx context.Context) ([]byte, []byte, error) {
155
+ certFile, keyFile, err := m.EnsureCertificate(ctx)
156
+ if err != nil {
157
+ return nil, nil, err
158
+ }
159
+
160
+ certPEM, err := os.ReadFile(certFile)
161
+ if err != nil {
162
+ return nil, nil, fmt.Errorf("read api tls certificate: %w", err)
163
+ }
164
+ keyPEM, err := os.ReadFile(keyFile)
165
+ if err != nil {
166
+ return nil, nil, fmt.Errorf("read api tls private key: %w", err)
167
+ }
168
+ return certPEM, keyPEM, nil
169
+}
170
+
171
func (m *Manager) Start(ctx context.Context) {
172
if m == nil || isLocalhost(m.cfg.BaseDomain) {
173
return
portal/acme/acme_test.go
+8
-4
@@ -17,14 +17,18 @@ func TestEnsureCertificateGeneratesLocalDevelopmentMaterial(t *testing.T) {
17
t.Fatalf("NewManager() error = %v", err)
18
}
19
20
- certFile, keyFile, err := manager.EnsureCertificate(context.Background())
20
+ certPEM, keyPEM, err := manager.EnsureTLSMaterial(context.Background())
21
if err != nil {
22
- t.Fatalf("EnsureCertificate() error = %v", err)
22
+ t.Fatalf("EnsureTLSMaterial() error = %v", err)
23
}
24
- if certFile == "" || keyFile == "" {
25
- t.Fatalf("EnsureCertificate() = %q, %q, want certificate paths", certFile, keyFile)
24
+ if len(certPEM) == 0 || len(keyPEM) == 0 {
25
+ t.Fatalf("EnsureTLSMaterial() returned empty PEM material")
26
}
27
28
+ certFile, _, err := manager.TLSFiles()
29
+ if err != nil {
30
+ t.Fatalf("TLSFiles() error = %v", err)
31
+ }
32
covered, err := certCoversDomains(certFile, []string{"localhost"})
33
if err != nil {
34
t.Fatalf("certCoversDomains() error = %v", err)
portal/admin/handler.go
+34
-24
@@ -16,7 +16,6 @@ const cookieName = "portal_admin"
16
17
type Handler struct {
18
auth *policy.Authenticator
19
- runtime *policy.Runtime
19
server *portal.Server
20
settings *stateStore
21
serveAppStatic func(http.ResponseWriter, *http.Request, string)
@@ -27,7 +26,6 @@ type Handler struct {
26
func NewHandler(portalURL, secret, settingsPath string, trustProxy bool, serveAppStatic func(http.ResponseWriter, *http.Request, string)) *Handler {
27
h := &Handler{
28
auth: policy.NewAuthenticator(strings.TrimSpace(secret)),
30
- runtime: policy.NewRuntime(),
29
settings: newStateStore(settingsPath),
30
buildLeaseRows: func(serv *portal.Server, includeAdmin bool) []LeaseRow {
31
return BuildLeaseRows(serv, includeAdmin, portalURL)
@@ -48,12 +46,8 @@ func (h *Handler) Bind(server *portal.Server) {
46
h.server = server
47
}
48
51
-func (h *Handler) Runtime() *policy.Runtime {
52
- return h.runtime
53
-}
54
-
49
func (h *Handler) LoadSettings() error {
56
- return h.settings.Load(h.runtime)
50
+ return h.settings.Load(h.policyRuntime())
51
}
52
53
func (h *Handler) HandleRequest(w http.ResponseWriter, r *http.Request) {
@@ -87,6 +81,10 @@ func (h *Handler) HandleRequest(w http.ResponseWriter, r *http.Request) {
81
writeAPIError(w, http.StatusUnauthorized, types.APIErrorCodeUnauthorized, "unauthorized")
82
return
83
}
84
+ if h.server == nil {
85
+ writeAPIError(w, http.StatusInternalServerError, types.APIErrorCodeFeatureUnavailable, "admin handler is not bound to a server")
86
+ return
87
+ }
88
89
switch path {
90
case types.PathAdminLeases:
@@ -207,7 +205,7 @@ func (h *Handler) handleBannedLeases(w http.ResponseWriter, r *http.Request) {
205
writeAPIError(w, http.StatusMethodNotAllowed, types.APIErrorCodeMethodNotAllowed, "method not allowed")
206
return
207
}
210
- writeAPIData(w, http.StatusOK, h.runtime.BannedLeases())
208
+ writeAPIData(w, http.StatusOK, h.policyRuntime().BannedLeases())
209
}
210
211
func (h *Handler) handleSettings(w http.ResponseWriter, r *http.Request) {
@@ -215,7 +213,7 @@ func (h *Handler) handleSettings(w http.ResponseWriter, r *http.Request) {
213
writeAPIError(w, http.StatusMethodNotAllowed, types.APIErrorCodeMethodNotAllowed, "method not allowed")
214
return
215
}
218
- approver := h.runtime.Approver()
216
+ approver := h.policyRuntime().Approver()
217
writeAPIData(w, http.StatusOK, types.AdminSettingsResponse{
218
ApprovalMode: string(approver.Mode()),
219
ApprovedLeases: approver.ApprovedLeases(),
@@ -224,7 +222,8 @@ func (h *Handler) handleSettings(w http.ResponseWriter, r *http.Request) {
222
}
223
224
func (h *Handler) handleApprovalMode(w http.ResponseWriter, r *http.Request) {
227
- approver := h.runtime.Approver()
225
+ runtime := h.policyRuntime()
226
+ approver := runtime.Approver()
227
switch r.Method {
228
case http.MethodGet:
229
writeAPIData(w, http.StatusOK, types.AdminApprovalModeResponse{ApprovalMode: string(approver.Mode())})
@@ -238,7 +237,7 @@ func (h *Handler) handleApprovalMode(w http.ResponseWriter, r *http.Request) {
237
writeAPIError(w, http.StatusBadRequest, types.APIErrorCodeInvalidMode, "invalid mode (must be 'auto' or 'manual')")
238
return
239
}
241
- _ = h.settings.Save(h.runtime)
240
+ _ = h.settings.Save(runtime)
241
writeAPIData(w, http.StatusOK, types.AdminApprovalModeResponse{ApprovalMode: string(approver.Mode())})
242
default:
243
writeAPIError(w, http.StatusMethodNotAllowed, types.APIErrorCodeMethodNotAllowed, "method not allowed")
@@ -274,14 +273,15 @@ func (h *Handler) handleLeaseAction(w http.ResponseWriter, r *http.Request, path
273
}
274
275
func (h *Handler) handleLeaseBan(w http.ResponseWriter, r *http.Request, leaseID string) {
276
+ runtime := h.policyRuntime()
277
switch r.Method {
278
case http.MethodPost:
279
- h.runtime.BanLease(leaseID)
280
- _ = h.settings.Save(h.runtime)
279
+ runtime.BanLease(leaseID)
280
+ _ = h.settings.Save(runtime)
281
writeAPIOK(w, http.StatusOK)
282
case http.MethodDelete:
283
- h.runtime.UnbanLease(leaseID)
284
- _ = h.settings.Save(h.runtime)
283
+ runtime.UnbanLease(leaseID)
284
+ _ = h.settings.Save(runtime)
285
writeAPIOK(w, http.StatusOK)
286
default:
287
writeAPIError(w, http.StatusMethodNotAllowed, types.APIErrorCodeMethodNotAllowed, "method not allowed")
@@ -289,16 +289,17 @@ func (h *Handler) handleLeaseBan(w http.ResponseWriter, r *http.Request, leaseID
289
}
290
291
func (h *Handler) handleLeaseApproval(w http.ResponseWriter, r *http.Request, leaseID string) {
292
- approver := h.runtime.Approver()
292
+ runtime := h.policyRuntime()
293
+ approver := runtime.Approver()
294
switch r.Method {
295
case http.MethodPost:
296
approver.Approve(leaseID)
297
approver.Undeny(leaseID)
297
- _ = h.settings.Save(h.runtime)
298
+ _ = h.settings.Save(runtime)
299
writeAPIOK(w, http.StatusOK)
300
case http.MethodDelete:
301
approver.Revoke(leaseID)
301
- _ = h.settings.Save(h.runtime)
302
+ _ = h.settings.Save(runtime)
303
writeAPIOK(w, http.StatusOK)
304
default:
305
writeAPIError(w, http.StatusMethodNotAllowed, types.APIErrorCodeMethodNotAllowed, "method not allowed")
@@ -306,15 +307,16 @@ func (h *Handler) handleLeaseApproval(w http.ResponseWriter, r *http.Request, le
307
}
308
309
func (h *Handler) handleLeaseDenial(w http.ResponseWriter, r *http.Request, leaseID string) {
309
- approver := h.runtime.Approver()
310
+ runtime := h.policyRuntime()
311
+ approver := runtime.Approver()
312
switch r.Method {
313
case http.MethodPost:
314
approver.Deny(leaseID)
313
- _ = h.settings.Save(h.runtime)
315
+ _ = h.settings.Save(runtime)
316
writeAPIOK(w, http.StatusOK)
317
case http.MethodDelete:
318
approver.Undeny(leaseID)
317
- _ = h.settings.Save(h.runtime)
319
+ _ = h.settings.Save(runtime)
320
writeAPIOK(w, http.StatusOK)
321
default:
322
writeAPIError(w, http.StatusMethodNotAllowed, types.APIErrorCodeMethodNotAllowed, "method not allowed")
@@ -333,15 +335,16 @@ func (h *Handler) handleIPBan(w http.ResponseWriter, r *http.Request, path strin
335
return
336
}
337
336
- ipFilter := h.runtime.IPFilter()
338
+ runtime := h.policyRuntime()
339
+ ipFilter := runtime.IPFilter()
340
switch r.Method {
341
case http.MethodPost:
342
ipFilter.BanIP(rawIP)
340
- _ = h.settings.Save(h.runtime)
343
+ _ = h.settings.Save(runtime)
344
writeAPIOK(w, http.StatusOK)
345
case http.MethodDelete:
346
ipFilter.UnbanIP(rawIP)
344
- _ = h.settings.Save(h.runtime)
347
+ _ = h.settings.Save(runtime)
348
writeAPIOK(w, http.StatusOK)
349
default:
350
writeAPIError(w, http.StatusMethodNotAllowed, types.APIErrorCodeMethodNotAllowed, "method not allowed")
@@ -404,3 +407,10 @@ func writeAPIErrorWithData(w http.ResponseWriter, status int, code, message stri
407
Error: &types.APIError{Code: code, Message: message},
408
})
409
}
410
+
411
+func (h *Handler) policyRuntime() *policy.Runtime {
412
+ if h.server == nil {
413
+ return nil
414
+ }
415
+ return h.server.PolicyRuntime()
416
+}
portal/admin/handler_test.go
+8
-2
@@ -8,6 +8,7 @@ import (
8
"path/filepath"
9
"testing"
10
11
+ "github.com/gosuda/portal/v2/portal"
12
"github.com/gosuda/portal/v2/portal/policy"
13
"github.com/gosuda/portal/v2/types"
14
)
@@ -18,6 +19,11 @@ func TestLoginAndProtectedActions(t *testing.T) {
19
handler := NewHandler("https://portal.example.com", "secret-key", filepath.Join(t.TempDir(), "admin_settings.json"), false, func(w http.ResponseWriter, _ *http.Request, _ string) {
20
w.WriteHeader(http.StatusOK)
21
})
22
+ server, err := portal.NewServer(portal.ServerConfig{PortalURL: "https://portal.example.com"})
23
+ if err != nil {
24
+ t.Fatalf("NewServer() error = %v", err)
25
+ }
26
+ handler.Bind(server)
27
28
loginRecorder := httptest.NewRecorder()
29
loginRequest := httptest.NewRequest(http.MethodPost, types.PathAdminLogin, bytes.NewBufferString(`{"key":"secret-key"}`))
@@ -54,7 +60,7 @@ func TestLoginAndProtectedActions(t *testing.T) {
60
if approvalRecorder.Code != http.StatusOK {
61
t.Fatalf("approval status = %d, want %d", approvalRecorder.Code, http.StatusOK)
62
}
57
- if got := handler.runtime.Approver().Mode(); got != policy.ModeManual {
63
+ if got := server.PolicyRuntime().Approver().Mode(); got != policy.ModeManual {
64
t.Fatalf("approval mode = %q, want %q", got, policy.ModeManual)
65
}
66
@@ -66,7 +72,7 @@ func TestLoginAndProtectedActions(t *testing.T) {
72
if ipBanRecorder.Code != http.StatusOK {
73
t.Fatalf("ip ban status = %d, want %d", ipBanRecorder.Code, http.StatusOK)
74
}
69
- if !handler.runtime.IPFilter().IsIPBanned("203.0.113.10") {
75
+ if !server.PolicyRuntime().IPFilter().IsIPBanned("203.0.113.10") {
76
t.Fatalf("IsIPBanned() = false, want true")
77
}
78
}
portal/api_server.go
+79
-107
@@ -20,20 +20,25 @@ import (
20
)
21
22
var (
23
- errLeaseNotFound = errors.New("lease not found")
24
- errIPBanned = errors.New("request denied because source IP is banned")
23
+ errLeaseNotFound = errors.New(types.APIErrorCodeLeaseNotFound)
24
+ errIPBanned = errors.New(types.APIErrorCodeIPBanned)
25
errUnauthorized = errors.New(types.APIErrorCodeUnauthorized)
26
- errHostnameConflict = errors.New("hostname already registered")
26
+ errHostnameConflict = errors.New(types.APIErrorCodeHostnameConflict)
27
)
28
29
-func (s *Server) newAPIServer(listener net.Listener) (net.Listener, *http.Server, io.Closer, error) {
29
+func (s *Server) newAPIServer(listener net.Listener, apiMux *http.ServeMux, apiTLS keyless.TLSMaterialConfig) (net.Listener, *http.Server, io.Closer, error) {
30
+ keylessSignerHandler, err := newKeylessSignerHandler(apiTLS)
31
+ if err != nil {
32
+ return nil, nil, nil, err
33
+ }
34
+
35
apiServer := &http.Server{
31
- Handler: s.wrapAPIHandler(s.apiHandler()),
36
+ Handler: s.apiHandler(apiMux, keylessSignerHandler),
37
ReadHeaderTimeout: 10 * time.Second,
38
TLSNextProto: make(map[string]func(*http.Server, *tls.Conn, http.Handler)),
39
}
40
36
- apiCloser, err := keyless.AttachToHTTPServer(apiServer, s.cfg.APITLS)
41
+ apiCloser, err := keyless.AttachToHTTPServer(apiServer, apiTLS)
42
if err != nil {
43
return nil, nil, nil, fmt.Errorf("configure api tls: %w", err)
44
}
@@ -41,25 +46,42 @@ func (s *Server) newAPIServer(listener net.Listener) (net.Listener, *http.Server
46
return tls.NewListener(listener, apiServer.TLSConfig), apiServer, apiCloser, nil
47
}
48
44
-func (s *Server) apiHandler() http.Handler {
45
- mux := http.NewServeMux()
46
- if s.cfg.KeylessSignerHandler != nil {
47
- mux.Handle(types.PathV1Sign, s.cfg.KeylessSignerHandler)
48
- }
49
- mux.HandleFunc(types.PathHealthz, s.handleHealthz)
50
- mux.HandleFunc(types.PathSDKDomain, s.handleDomain)
51
- mux.HandleFunc(types.PathSDKRegister, s.handleRegister)
52
- mux.HandleFunc(types.PathSDKRenew, s.handleRenew)
53
- mux.HandleFunc(types.PathSDKUnregister, s.handleUnregister)
54
- mux.HandleFunc(types.PathSDKConnect, s.handleConnect)
55
- mux.HandleFunc("/", s.handleRoot)
56
- return mux
49
+func (s *Server) apiHandler(base *http.ServeMux, keylessSignerHandler http.Handler) http.Handler {
50
+ if base == nil {
51
+ base = http.NewServeMux()
52
+ base.HandleFunc("/{$}", s.handleRoot)
53
+ }
54
+
55
+ return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
56
+ switch strings.TrimSpace(r.URL.Path) {
57
+ case types.PathHealthz:
58
+ s.handleHealthz(w, r)
59
+ case types.PathSDKDomain:
60
+ s.handleDomain(w, r)
61
+ case types.PathSDKRegister:
62
+ s.handleRegister(w, r)
63
+ case types.PathSDKRenew:
64
+ s.handleRenew(w, r)
65
+ case types.PathSDKUnregister:
66
+ s.handleUnregister(w, r)
67
+ case types.PathSDKConnect:
68
+ s.handleConnect(w, r)
69
+ case types.PathV1Sign:
70
+ if keylessSignerHandler == nil {
71
+ http.NotFound(w, r)
72
+ return
73
+ }
74
+ keylessSignerHandler.ServeHTTP(w, r)
75
+ default:
76
+ base.ServeHTTP(w, r)
77
+ }
78
+ })
79
}
80
81
func (s *Server) handleRoot(w http.ResponseWriter, _ *http.Request) {
82
writeAPIData(w, http.StatusOK, map[string]any{
83
"service": "portal-relay",
62
- "root": s.cfg.RootHost,
84
+ "root": s.rootHost,
85
})
86
}
87
@@ -75,8 +97,8 @@ func (s *Server) handleDomain(w http.ResponseWriter, r *http.Request) {
97
98
name := r.URL.Query().Get("name")
99
writeAPIData(w, http.StatusOK, types.DomainResponse{
78
- RootHost: s.cfg.RootHost,
79
- SuggestedHostname: suggestHostname(name, s.cfg.RootHost),
100
+ RootHost: s.rootHost,
101
+ SuggestedHostname: suggestHostname(name, s.rootHost),
102
Version: types.SDKProtocolVersion,
103
})
104
}
@@ -202,7 +224,7 @@ func (s *Server) handleConnect(w http.ResponseWriter, r *http.Request) {
224
writeAPIError(w, http.StatusNotFound, types.APIErrorCodeLeaseNotFound, err.Error())
225
return
226
}
205
- if !s.isLeaseRoutable(lease) {
227
+ if !s.registry.IsRoutable(lease) {
228
writeAPIError(w, http.StatusForbidden, types.APIErrorCodeLeaseRejected, "lease is not approved for routing")
229
return
230
}
@@ -245,7 +267,7 @@ func (s *Server) handleConnect(w http.ResponseWriter, r *http.Request) {
267
return
268
}
269
248
- s.touchLease(lease.ID, clientIP)
270
+ s.registry.Touch(lease.ID, clientIP, time.Now())
271
log.Info().
272
Str("component", "relay-server").
273
Str("lease_id", lease.ID).
@@ -268,16 +290,7 @@ func (s *Server) registerLease(req types.RegisterRequest, clientIP string) (type
290
291
hostnames := normalizeHostnames(req.Hostnames)
292
if len(hostnames) == 0 {
271
- hostnames = []string{suggestHostname(req.Name, s.cfg.RootHost)}
272
- }
273
-
274
- s.mu.Lock()
275
- defer s.mu.Unlock()
276
-
277
- for _, host := range hostnames {
278
- if owner := s.findLeaseByHostnameLocked(host); owner != nil {
279
- return types.RegisterResponse{}, fmt.Errorf("%w: %s", errHostnameConflict, host)
280
- }
293
+ hostnames = []string{suggestHostname(req.Name, s.rootHost)}
294
}
295
296
ttl := s.cfg.LeaseTTL
@@ -301,12 +314,8 @@ func (s *Server) registerLease(req types.RegisterRequest, clientIP string) (type
314
Broker: newLeaseBroker(leaseID, s.cfg.IdleKeepaliveInterval, s.cfg.ReadyQueueLimit),
315
}
316
304
- s.leases[leaseID] = record
305
- for _, host := range hostnames {
306
- s.routes.Set(host, leaseID)
307
- }
308
- if strings.TrimSpace(clientIP) != "" {
309
- s.cfg.Policy.IPFilter().RegisterLeaseIP(leaseID, clientIP)
317
+ if err := s.registry.Register(record); err != nil {
318
+ return types.RegisterResponse{}, err
319
}
320
321
return types.RegisterResponse{
@@ -323,62 +332,29 @@ func (s *Server) renewLease(req types.RenewRequest, clientIP string) (types.Rene
332
return types.RenewResponse{}, errIPBanned
333
}
334
326
- s.mu.Lock()
327
- defer s.mu.Unlock()
328
-
329
- record, ok := s.leases[strings.TrimSpace(req.LeaseID)]
330
- if !ok {
331
- return types.RenewResponse{}, errLeaseNotFound
332
- }
333
- if !tokenMatches(record.ReverseToken, req.ReverseToken) {
334
- return types.RenewResponse{}, errUnauthorized
335
- }
336
-
335
ttl := s.cfg.LeaseTTL
336
if req.TTL > 0 {
337
ttl = time.Duration(req.TTL) * time.Second
338
}
341
- record.ExpiresAt = time.Now().Add(ttl)
342
- record.LastSeenAt = time.Now()
343
- if strings.TrimSpace(clientIP) != "" {
344
- record.ClientIP = clientIP
345
- s.cfg.Policy.IPFilter().RegisterLeaseIP(record.ID, clientIP)
339
+ record, err := s.registry.Renew(req.LeaseID, req.ReverseToken, ttl, clientIP)
340
+ if err != nil {
341
+ return types.RenewResponse{}, err
342
}
343
344
return types.RenewResponse{LeaseID: record.ID, ExpiresAt: record.ExpiresAt}, nil
345
}
346
347
func (s *Server) unregisterLease(req types.UnregisterRequest) error {
352
- s.mu.Lock()
353
- record, ok := s.leases[strings.TrimSpace(req.LeaseID)]
354
- if !ok {
355
- s.mu.Unlock()
356
- return errLeaseNotFound
357
- }
358
- if !tokenMatches(record.ReverseToken, req.ReverseToken) {
359
- s.mu.Unlock()
360
- return errUnauthorized
348
+ record, err := s.registry.Unregister(req.LeaseID, req.ReverseToken)
349
+ if err != nil {
350
+ return err
351
}
362
- delete(s.leases, record.ID)
363
- s.mu.Unlock()
364
-
365
- s.routes.DeleteLease(record.Hostnames)
366
- s.cfg.Policy.ForgetLease(record.ID)
352
record.Broker.Close()
353
return nil
354
}
355
356
func (s *Server) findLeaseByID(leaseID string) (*leaseRecord, error) {
372
- s.mu.RLock()
373
- record, ok := s.leases[strings.TrimSpace(leaseID)]
374
- s.mu.RUnlock()
375
- if !ok {
376
- return nil, errLeaseNotFound
377
- }
378
- if time.Now().After(record.ExpiresAt) {
379
- return nil, errLeaseNotFound
380
- }
381
- return record, nil
357
+ return s.registry.FindByID(leaseID)
358
}
359
360
func (s *Server) authorizeLeaseToken(record *leaseRecord, token string) error {
@@ -391,22 +367,6 @@ func (s *Server) authorizeLeaseToken(record *leaseRecord, token string) error {
367
return nil
368
}
369
394
-func (s *Server) findLeaseByHostnameLocked(host string) *leaseRecord {
395
- host = normalizeHostname(host)
396
- now := time.Now()
397
- for _, lease := range s.leases {
398
- if now.After(lease.ExpiresAt) {
399
- continue
400
- }
401
- for _, candidate := range lease.Hostnames {
402
- if normalizeHostname(candidate) == host {
403
- return lease
404
- }
405
- }
406
- }
407
- return nil
408
-}
409
-
370
func (s *Server) runAPIServer() error {
371
err := s.apiServer.Serve(s.apiListener)
372
if err == nil || errors.Is(err, http.ErrServerClosed) || errors.Is(err, net.ErrClosed) {
@@ -417,19 +377,9 @@ func (s *Server) runAPIServer() error {
377
378
func (s *Server) connectURL() string {
379
base := strings.TrimRight(s.cfg.PortalURL, "/")
420
- if base == "" && s.apiListener != nil {
421
- return "https://" + HostPortOrLoopback(s.apiListener.Addr().String()) + types.PathSDKConnect
422
- }
380
return base + types.PathSDKConnect
381
}
382
426
-func (s *Server) wrapAPIHandler(base http.Handler) http.Handler {
427
- if s.cfg.APIHandlerWrapper == nil {
428
- return base
429
- }
430
- return s.cfg.APIHandlerWrapper(base)
431
-}
432
-
383
func decodeJSONBody(w http.ResponseWriter, r *http.Request, dst any) error {
384
r.Body = http.MaxBytesReader(w, r.Body, defaultControlBodyLimit)
385
defer r.Body.Close()
@@ -487,5 +437,27 @@ func (s *Server) clientIPFromRequest(r *http.Request) string {
437
}
438
439
func (s *Server) isClientIPBanned(clientIP string) bool {
490
- return s.cfg.Policy.IPFilter().IsIPBanned(clientIP)
440
+ return s.registry.IsClientIPBanned(clientIP)
441
+}
442
+
443
+func validateAPITLS(apiTLS keyless.TLSMaterialConfig) error {
444
+ if len(apiTLS.CertPEM) == 0 {
445
+ return errors.New("api tls certificate is required")
446
+ }
447
+ if len(apiTLS.KeyPEM) == 0 && apiTLS.Keyless == nil {
448
+ return errors.New("api tls key or keyless signer is required")
449
+ }
450
+ return nil
451
+}
452
+
453
+func newKeylessSignerHandler(apiTLS keyless.TLSMaterialConfig) (http.Handler, error) {
454
+ if len(apiTLS.KeyPEM) == 0 {
455
+ return nil, nil
456
+ }
457
+
458
+ signer, err := keyless.NewSigner(apiTLS.KeyPEM)
459
+ if err != nil {
460
+ return nil, fmt.Errorf("configure api signer: %w", err)
461
+ }
462
+ return signer.Handler(), nil
463
}
portal/keyless/signer.go
+1
-6
@@ -6,7 +6,6 @@ import (
6
"errors"
7
"fmt"
8
"net/http"
9
- "os"
9
"strings"
10
"time"
11
@@ -24,11 +23,7 @@ type Signer struct {
23
keyID string
24
}
25
27
-func NewSigner(keyFile string) (*Signer, error) {
28
- keyPEM, err := os.ReadFile(strings.TrimSpace(keyFile))
29
- if err != nil {
30
- return nil, fmt.Errorf("read keyless signing key: %w", err)
31
- }
26
+func NewSigner(keyPEM []byte) (*Signer, error) {
27
signingKey, err := ksigner.ParsePrivateKeyPEM(keyPEM)
28
if err != nil {
29
return nil, fmt.Errorf("parse keyless signing key: %w", err)
portal/lease.go
new
+346
@@ -0,0 +1,346 @@
1
+package portal
2
+
3
+import (
4
+ "context"
5
+ "errors"
6
+ "fmt"
7
+ "strings"
8
+ "sync"
9
+ "time"
10
+
11
+ "github.com/gosuda/portal/v2/portal/policy"
12
+ "github.com/gosuda/portal/v2/types"
13
+)
14
+
15
+type leaseRegistry struct {
16
+ routes *routeTable
17
+ leaseByID map[string]*leaseRecord
18
+ policy *policy.Runtime
19
+ mu sync.RWMutex
20
+}
21
+
22
+func newLeaseRegistry(runtime *policy.Runtime) *leaseRegistry {
23
+ if runtime == nil {
24
+ runtime = policy.NewRuntime()
25
+ }
26
+ return &leaseRegistry{
27
+ routes: newRouteTable(),
28
+ leaseByID: make(map[string]*leaseRecord),
29
+ policy: runtime,
30
+ }
31
+}
32
+
33
+func (r *leaseRegistry) PolicyRuntime() *policy.Runtime {
34
+ if r == nil {
35
+ return nil
36
+ }
37
+ return r.policy
38
+}
39
+
40
+func (r *leaseRegistry) CloseAll() []*leaseRecord {
41
+ r.mu.Lock()
42
+ defer r.mu.Unlock()
43
+
44
+ out := make([]*leaseRecord, 0, len(r.leaseByID))
45
+ for _, record := range r.leaseByID {
46
+ out = append(out, record)
47
+ r.policy.ForgetLease(record.ID)
48
+ }
49
+ r.routes = newRouteTable()
50
+ r.leaseByID = make(map[string]*leaseRecord)
51
+ return out
52
+}
53
+
54
+func (r *leaseRegistry) RunJanitor(ctx context.Context, interval time.Duration) error {
55
+ if interval <= 0 {
56
+ return errors.New("janitor interval must be positive")
57
+ }
58
+
59
+ ticker := time.NewTicker(interval)
60
+ defer ticker.Stop()
61
+
62
+ for {
63
+ select {
64
+ case <-ctx.Done():
65
+ return nil
66
+ case <-ticker.C:
67
+ r.cleanupExpired(time.Now())
68
+ }
69
+ }
70
+}
71
+
72
+func (r *leaseRegistry) List() []*leaseRecord {
73
+ r.mu.RLock()
74
+ defer r.mu.RUnlock()
75
+
76
+ out := make([]*leaseRecord, 0, len(r.leaseByID))
77
+ for _, record := range r.leaseByID {
78
+ out = append(out, record)
79
+ }
80
+ return out
81
+}
82
+
83
+func (r *leaseRegistry) Get(leaseID string) (*leaseRecord, bool) {
84
+ r.mu.RLock()
85
+ defer r.mu.RUnlock()
86
+
87
+ record, ok := r.leaseByID[strings.TrimSpace(leaseID)]
88
+ return record, ok
89
+}
90
+
91
+func (r *leaseRegistry) Lookup(host string) (*leaseRecord, bool) {
92
+ host = normalizeHostname(host)
93
+ if host == "" {
94
+ return nil, false
95
+ }
96
+
97
+ r.mu.RLock()
98
+ defer r.mu.RUnlock()
99
+
100
+ leaseID, ok := r.routes.Lookup(host)
101
+ if !ok {
102
+ return nil, false
103
+ }
104
+ record, ok := r.leaseByID[leaseID]
105
+ return record, ok && record != nil
106
+}
107
+
108
+func (r *leaseRegistry) Register(record *leaseRecord) error {
109
+ if record == nil {
110
+ return errors.New("lease record is required")
111
+ }
112
+
113
+ leaseID := strings.TrimSpace(record.ID)
114
+ if leaseID == "" {
115
+ return errors.New("lease id is required")
116
+ }
117
+ hostnames := normalizeHostnames(record.Hostnames)
118
+ if len(hostnames) == 0 {
119
+ return errors.New("lease hostnames are required")
120
+ }
121
+
122
+ r.mu.Lock()
123
+ defer r.mu.Unlock()
124
+
125
+ for _, host := range hostnames {
126
+ if ownerLeaseID, ok := r.routes.LookupExact(host); ok && ownerLeaseID != leaseID {
127
+ return fmt.Errorf("%w: %s", errHostnameConflict, host)
128
+ }
129
+ }
130
+
131
+ record.ID = leaseID
132
+ record.Hostnames = hostnames
133
+ r.leaseByID[leaseID] = record
134
+ for _, host := range hostnames {
135
+ r.routes.Set(host, leaseID)
136
+ }
137
+ if strings.TrimSpace(record.ClientIP) != "" {
138
+ r.policy.IPFilter().RegisterLeaseIP(leaseID, record.ClientIP)
139
+ }
140
+ return nil
141
+}
142
+
143
+func (r *leaseRegistry) Renew(leaseID, reverseToken string, ttl time.Duration, clientIP string) (*leaseRecord, error) {
144
+ r.mu.Lock()
145
+ defer r.mu.Unlock()
146
+
147
+ record, ok := r.leaseByID[strings.TrimSpace(leaseID)]
148
+ if !ok {
149
+ return nil, errLeaseNotFound
150
+ }
151
+ if !tokenMatches(record.ReverseToken, reverseToken) {
152
+ return nil, errUnauthorized
153
+ }
154
+
155
+ now := time.Now()
156
+ record.ExpiresAt = now.Add(ttl)
157
+ record.LastSeenAt = now
158
+ if strings.TrimSpace(clientIP) != "" {
159
+ record.ClientIP = clientIP
160
+ r.policy.IPFilter().RegisterLeaseIP(record.ID, clientIP)
161
+ }
162
+ return record, nil
163
+}
164
+
165
+func (r *leaseRegistry) Unregister(leaseID, reverseToken string) (*leaseRecord, error) {
166
+ r.mu.Lock()
167
+ defer r.mu.Unlock()
168
+
169
+ record, ok := r.leaseByID[strings.TrimSpace(leaseID)]
170
+ if !ok {
171
+ return nil, errLeaseNotFound
172
+ }
173
+ if !tokenMatches(record.ReverseToken, reverseToken) {
174
+ return nil, errUnauthorized
175
+ }
176
+
177
+ delete(r.leaseByID, record.ID)
178
+ r.routes.DeleteLease(record.Hostnames)
179
+ r.policy.ForgetLease(record.ID)
180
+ return record, nil
181
+}
182
+
183
+func (r *leaseRegistry) FindByID(leaseID string) (*leaseRecord, error) {
184
+ r.mu.RLock()
185
+ defer r.mu.RUnlock()
186
+
187
+ record, ok := r.leaseByID[strings.TrimSpace(leaseID)]
188
+ if !ok || time.Now().After(record.ExpiresAt) {
189
+ return nil, errLeaseNotFound
190
+ }
191
+ return record, nil
192
+}
193
+
194
+func (r *leaseRegistry) Touch(leaseID, clientIP string, now time.Time) *leaseRecord {
195
+ r.mu.Lock()
196
+ defer r.mu.Unlock()
197
+
198
+ record := r.leaseByID[strings.TrimSpace(leaseID)]
199
+ if record == nil {
200
+ return nil
201
+ }
202
+ record.LastSeenAt = now
203
+ if strings.TrimSpace(clientIP) != "" {
204
+ record.ClientIP = clientIP
205
+ r.policy.IPFilter().RegisterLeaseIP(record.ID, clientIP)
206
+ }
207
+ return record
208
+}
209
+
210
+func (r *leaseRegistry) cleanupExpired(now time.Time) {
211
+ for _, lease := range r.removeExpired(now) {
212
+ lease.Broker.Close()
213
+ }
214
+}
215
+
216
+func (r *leaseRegistry) removeExpired(now time.Time) []*leaseRecord {
217
+ r.mu.Lock()
218
+ defer r.mu.Unlock()
219
+
220
+ expired := make([]*leaseRecord, 0)
221
+ for leaseID, record := range r.leaseByID {
222
+ if now.After(record.ExpiresAt) {
223
+ expired = append(expired, record)
224
+ delete(r.leaseByID, leaseID)
225
+ r.routes.DeleteLease(record.Hostnames)
226
+ r.policy.ForgetLease(record.ID)
227
+ }
228
+ }
229
+ return expired
230
+}
231
+
232
+func (r *leaseRegistry) IsClientIPBanned(clientIP string) bool {
233
+ return r.policy.IPFilter().IsIPBanned(clientIP)
234
+}
235
+
236
+func (r *leaseRegistry) IsRoutable(record *leaseRecord) bool {
237
+ if record == nil {
238
+ return false
239
+ }
240
+ return r.policy.IsLeaseRoutable(record.ID)
241
+}
242
+
243
+func (r *leaseRegistry) Snapshot(record *leaseRecord) LeaseSnapshot {
244
+ if record == nil {
245
+ return LeaseSnapshot{}
246
+ }
247
+
248
+ clientIP := record.ClientIP
249
+ return LeaseSnapshot{
250
+ ID: record.ID,
251
+ Name: record.Name,
252
+ ClientIP: clientIP,
253
+ Hostnames: append([]string(nil), record.Hostnames...),
254
+ Metadata: record.Metadata,
255
+ ExpiresAt: record.ExpiresAt,
256
+ FirstSeenAt: record.FirstSeenAt,
257
+ LastSeenAt: record.LastSeenAt,
258
+ Ready: record.Broker.ReadyCount(),
259
+ IsApproved: r.policy.EffectiveApproval(record.ID),
260
+ IsBanned: r.policy.IsLeaseBanned(record.ID),
261
+ IsDenied: r.policy.IsLeaseDenied(record.ID),
262
+ IsIPBanned: r.policy.IPFilter().IsIPBanned(clientIP),
263
+ }
264
+}
265
+
266
+type leaseRecord struct {
267
+ ExpiresAt time.Time
268
+ FirstSeenAt time.Time
269
+ LastSeenAt time.Time
270
+ Broker *leaseBroker
271
+ ID string
272
+ Name string
273
+ ReverseToken string
274
+ ClientIP string
275
+ Hostnames []string
276
+ Metadata types.LeaseMetadata
277
+}
278
+
279
+type LeaseSnapshot struct {
280
+ ExpiresAt time.Time
281
+ FirstSeenAt time.Time
282
+ LastSeenAt time.Time
283
+ ID string
284
+ Name string
285
+ ClientIP string
286
+ Hostnames []string
287
+ Metadata types.LeaseMetadata
288
+ Ready int
289
+ IsApproved bool
290
+ IsBanned bool
291
+ IsDenied bool
292
+ IsIPBanned bool
293
+}
294
+
295
+type routeTable struct {
296
+ exact map[string]string
297
+}
298
+
299
+func newRouteTable() *routeTable {
300
+ return &routeTable{exact: make(map[string]string)}
301
+}
302
+
303
+func (t *routeTable) Set(host, leaseID string) {
304
+ host = normalizeHostname(host)
305
+ if host == "" {
306
+ return
307
+ }
308
+ t.exact[host] = leaseID
309
+}
310
+
311
+func (t *routeTable) DeleteLease(hosts []string) {
312
+ for _, host := range hosts {
313
+ delete(t.exact, normalizeHostname(host))
314
+ }
315
+}
316
+
317
+func (t *routeTable) LookupExact(host string) (string, bool) {
318
+ host = normalizeHostname(host)
319
+ if host == "" {
320
+ return "", false
321
+ }
322
+ leaseID, ok := t.exact[host]
323
+ return leaseID, ok
324
+}
325
+
326
+func (t *routeTable) Lookup(host string) (string, bool) {
327
+ host = normalizeHostname(host)
328
+ if host == "" {
329
+ return "", false
330
+ }
331
+
332
+ if leaseID, ok := t.exact[host]; ok {
333
+ return leaseID, true
334
+ }
335
+
336
+ parts := stringsSplit(host, ".")
337
+ if len(parts) < 3 {
338
+ return "", false
339
+ }
340
+ wildcard := "*." + stringsJoin(parts[1:], ".")
341
+ leaseID, ok := t.exact[wildcard]
342
+ return leaseID, ok
343
+}
344
+
345
+func stringsSplit(s, sep string) []string { return strings.Split(s, sep) }
346
+func stringsJoin(parts []string, sep string) string { return strings.Join(parts, sep) }
portal/lease_test.go
new
+171
@@ -0,0 +1,171 @@
1
+package portal
2
+
3
+import (
4
+ "context"
5
+ "errors"
6
+ "testing"
7
+ "time"
8
+
9
+ "github.com/gosuda/portal/v2/portal/policy"
10
+)
11
+
12
+func TestLeaseRegistryLifecycle(t *testing.T) {
13
+ t.Parallel()
14
+
15
+ runtime := policy.NewRuntime()
16
+ registry := newLeaseRegistry(runtime)
17
+ record := &leaseRecord{
18
+ ID: "lease_1",
19
+ Hostnames: []string{"demo.example.com"},
20
+ ReverseToken: "tok_1",
21
+ ExpiresAt: time.Now().Add(30 * time.Second),
22
+ }
23
+
24
+ if err := registry.Register(record); err != nil {
25
+ t.Fatalf("Register() error = %v", err)
26
+ }
27
+
28
+ lookedUp, ok := registry.Lookup("demo.example.com")
29
+ if !ok || lookedUp != record {
30
+ t.Fatalf("Lookup() = %v, %v, want registered lease", lookedUp, ok)
31
+ }
32
+
33
+ renewed, err := registry.Renew(record.ID, record.ReverseToken, time.Minute, "203.0.113.10")
34
+ if err != nil {
35
+ t.Fatalf("Renew() error = %v", err)
36
+ }
37
+ if renewed.ClientIP != "203.0.113.10" {
38
+ t.Fatalf("Renew() client ip = %q, want %q", renewed.ClientIP, "203.0.113.10")
39
+ }
40
+ if got := runtime.IPFilter().LeaseIP(record.ID); got != "203.0.113.10" {
41
+ t.Fatalf("Renew() did not register client IP for lease")
42
+ }
43
+
44
+ removed, err := registry.Unregister(record.ID, record.ReverseToken)
45
+ if err != nil {
46
+ t.Fatalf("Unregister() error = %v", err)
47
+ }
48
+ if removed != record {
49
+ t.Fatalf("Unregister() record = %v, want original record", removed)
50
+ }
51
+
52
+ if _, ok := registry.Lookup("demo.example.com"); ok {
53
+ t.Fatal("Lookup() after Unregister() = true, want false")
54
+ }
55
+ if got := runtime.IPFilter().LeaseIP(record.ID); got != "" {
56
+ t.Fatalf("Unregister() lease IP = %q, want empty", got)
57
+ }
58
+}
59
+
60
+func TestLeaseRegistryWildcardAndConflict(t *testing.T) {
61
+ t.Parallel()
62
+
63
+ registry := newLeaseRegistry(policy.NewRuntime())
64
+ wildcardLease := &leaseRecord{
65
+ ID: "lease_wildcard",
66
+ Hostnames: []string{"*.example.com"},
67
+ ReverseToken: "tok_wildcard",
68
+ ExpiresAt: time.Now().Add(30 * time.Second),
69
+ }
70
+ if err := registry.Register(wildcardLease); err != nil {
71
+ t.Fatalf("Register(wildcard) error = %v", err)
72
+ }
73
+
74
+ if _, ok := registry.Lookup("app.example.com"); !ok {
75
+ t.Fatal("Lookup(one-level wildcard) = false, want true")
76
+ }
77
+ if _, ok := registry.Lookup("deep.app.example.com"); ok {
78
+ t.Fatal("Lookup(multi-level wildcard) = true, want false")
79
+ }
80
+
81
+ conflict := &leaseRecord{
82
+ ID: "lease_conflict",
83
+ Hostnames: []string{"*.example.com"},
84
+ ReverseToken: "tok_conflict",
85
+ ExpiresAt: time.Now().Add(30 * time.Second),
86
+ }
87
+ err := registry.Register(conflict)
88
+ if !errors.Is(err, errHostnameConflict) {
89
+ t.Fatalf("Register(conflict) error = %v, want hostname conflict", err)
90
+ }
91
+}
92
+
93
+func TestLeaseRegistrySnapshotAndRoutableUsePolicy(t *testing.T) {
94
+ t.Parallel()
95
+
96
+ runtime := policy.NewRuntime()
97
+ if err := runtime.Approver().SetMode(policy.ModeManual); err != nil {
98
+ t.Fatalf("SetMode() error = %v", err)
99
+ }
100
+
101
+ registry := newLeaseRegistry(runtime)
102
+ record := &leaseRecord{
103
+ ID: "lease_policy",
104
+ Name: "demo",
105
+ Hostnames: []string{"demo.example.com"},
106
+ ReverseToken: "tok_policy",
107
+ ExpiresAt: time.Now().Add(30 * time.Second),
108
+ ClientIP: "203.0.113.20",
109
+ Broker: newLeaseBroker("lease_policy", time.Minute, 1),
110
+ }
111
+ if err := registry.Register(record); err != nil {
112
+ t.Fatalf("Register() error = %v", err)
113
+ }
114
+
115
+ if registry.IsRoutable(record) {
116
+ t.Fatal("IsRoutable() = true, want false before approval")
117
+ }
118
+
119
+ snapshot := registry.Snapshot(record)
120
+ if snapshot.IsApproved {
121
+ t.Fatal("Snapshot().IsApproved = true, want false before approval")
122
+ }
123
+ if got := runtime.IPFilter().LeaseIP(record.ID); got != "203.0.113.20" {
124
+ t.Fatalf("Register() lease IP = %q, want %q", got, "203.0.113.20")
125
+ }
126
+
127
+ runtime.Approver().Approve(record.ID)
128
+ if !registry.IsRoutable(record) {
129
+ t.Fatal("IsRoutable() = false, want true after approval")
130
+ }
131
+
132
+ snapshot = registry.Snapshot(record)
133
+ if !snapshot.IsApproved {
134
+ t.Fatal("Snapshot().IsApproved = false, want true after approval")
135
+ }
136
+}
137
+
138
+func TestLeaseRegistryCleanupExpiredClosesBroker(t *testing.T) {
139
+ t.Parallel()
140
+
141
+ registry := newLeaseRegistry(policy.NewRuntime())
142
+ record := &leaseRecord{
143
+ ID: "lease_expired",
144
+ Hostnames: []string{"expired.example.com"},
145
+ ReverseToken: "tok_expired",
146
+ ExpiresAt: time.Now().Add(-time.Second),
147
+ Broker: newLeaseBroker("lease_expired", time.Minute, 1),
148
+ }
149
+ if err := registry.Register(record); err != nil {
150
+ t.Fatalf("Register() error = %v", err)
151
+ }
152
+
153
+ registry.cleanupExpired(time.Now())
154
+
155
+ if _, ok := registry.Lookup("expired.example.com"); ok {
156
+ t.Fatal("Lookup() after cleanupExpired() = true, want false")
157
+ }
158
+ if _, err := record.Broker.Claim(context.Background()); !errors.Is(err, errBrokerClosed) {
159
+ t.Fatalf("Claim() after cleanupExpired() error = %v, want %v", err, errBrokerClosed)
160
+ }
161
+}
162
+
163
+func TestLeaseRegistryRunJanitorRejectsNonPositiveInterval(t *testing.T) {
164
+ t.Parallel()
165
+
166
+ registry := newLeaseRegistry(policy.NewRuntime())
167
+ err := registry.RunJanitor(context.Background(), 0)
168
+ if err == nil {
169
+ t.Fatal("RunJanitor() error = nil, want validation error")
170
+ }
171
+}
portal/routing.go
deleted
-68
@@ -1,68 +0,0 @@
1
-package portal
2
-
3
-import (
4
- "strings"
5
- "sync"
6
-)
7
-
8
-type routeTable struct {
9
- exact map[string]string
10
- mu sync.RWMutex
11
-}
12
-
13
-func newRouteTable() *routeTable {
14
- return &routeTable{exact: make(map[string]string)}
15
-}
16
-
17
-func (t *routeTable) Set(host, leaseID string) {
18
- host = normalizeHostname(host)
19
- if host == "" {
20
- return
21
- }
22
- t.mu.Lock()
23
- defer t.mu.Unlock()
24
- t.exact[host] = leaseID
25
-}
26
-
27
-func (t *routeTable) Delete(host string) {
28
- host = normalizeHostname(host)
29
- if host == "" {
30
- return
31
- }
32
- t.mu.Lock()
33
- defer t.mu.Unlock()
34
- delete(t.exact, host)
35
-}
36
-
37
-func (t *routeTable) DeleteLease(hosts []string) {
38
- t.mu.Lock()
39
- defer t.mu.Unlock()
40
- for _, host := range hosts {
41
- delete(t.exact, normalizeHostname(host))
42
- }
43
-}
44
-
45
-func (t *routeTable) Lookup(host string) (string, bool) {
46
- host = normalizeHostname(host)
47
- if host == "" {
48
- return "", false
49
- }
50
-
51
- t.mu.RLock()
52
- defer t.mu.RUnlock()
53
-
54
- if leaseID, ok := t.exact[host]; ok {
55
- return leaseID, true
56
- }
57
-
58
- parts := stringsSplit(host, ".")
59
- if len(parts) < 3 {
60
- return "", false
61
- }
62
- wildcard := "*." + stringsJoin(parts[1:], ".")
63
- leaseID, ok := t.exact[wildcard]
64
- return leaseID, ok
65
-}
66
-
67
-func stringsSplit(s, sep string) []string { return strings.Split(s, sep) }
68
-func stringsJoin(parts []string, sep string) string { return strings.Join(parts, sep) }
portal/server.go
+84
-176
@@ -7,16 +7,15 @@ import (
7
"io"
8
"net"
9
"net/http"
10
- "strings"
10
"sync"
11
"time"
12
13
"github.com/gosuda/keyless_tls/relay/l4"
14
"golang.org/x/sync/errgroup"
15
16
+ "github.com/gosuda/portal/v2/portal/acme"
17
"github.com/gosuda/portal/v2/portal/keyless"
18
"github.com/gosuda/portal/v2/portal/policy"
19
- "github.com/gosuda/portal/v2/types"
19
)
20
21
const (
@@ -30,15 +29,10 @@ const (
29
)
30
31
type ServerConfig struct {
33
- APIHandlerWrapper func(http.Handler) http.Handler
34
- KeylessSignerHandler http.Handler
35
- Policy *policy.Runtime
32
PortalURL string
33
+ ACME acme.Config
34
APIListenAddr string
35
SNIListenAddr string
39
- RootHost string
40
- RootFallbackAddr string
41
- APITLS keyless.TLSMaterialConfig
36
LeaseTTL time.Duration
37
ClaimTimeout time.Duration
38
IdleKeepaliveInterval time.Duration
@@ -49,47 +43,18 @@ type ServerConfig struct {
43
44
type Server struct {
45
sniListener net.Listener
52
- apiTLSClose io.Closer
46
apiListener net.Listener
47
apiServer *http.Server
48
+ apiTLSClose io.Closer
49
+ acmeManager *acme.Manager
50
cancel context.CancelFunc
51
group *errgroup.Group
57
- routes *routeTable
58
- leases map[string]*leaseRecord
52
+ registry *leaseRegistry
53
cfg ServerConfig
60
- mu sync.RWMutex
54
+ rootHost string
55
shutdownOnce sync.Once
56
}
57
64
-type leaseRecord struct {
65
- ExpiresAt time.Time
66
- FirstSeenAt time.Time
67
- LastSeenAt time.Time
68
- Broker *leaseBroker
69
- ID string
70
- Name string
71
- ReverseToken string
72
- ClientIP string
73
- Hostnames []string
74
- Metadata types.LeaseMetadata
75
-}
76
-
77
-type LeaseSnapshot struct {
78
- ExpiresAt time.Time
79
- FirstSeenAt time.Time
80
- LastSeenAt time.Time
81
- ID string
82
- Name string
83
- ClientIP string
84
- Hostnames []string
85
- Metadata types.LeaseMetadata
86
- Ready int
87
- IsApproved bool
88
- IsBanned bool
89
- IsDenied bool
90
- IsIPBanned bool
91
-}
92
-
58
func NewServer(cfg ServerConfig) (*Server, error) {
59
if cfg.APIListenAddr == "" {
60
cfg.APIListenAddr = ":4017"
@@ -102,53 +67,48 @@ func NewServer(cfg ServerConfig) (*Server, error) {
67
cfg.IdleKeepaliveInterval = durationOrDefault(cfg.IdleKeepaliveInterval, defaultIdleKeepalive)
68
cfg.ReadyQueueLimit = intOrDefault(cfg.ReadyQueueLimit, defaultReadyQueueLimit)
69
cfg.ClientHelloTimeout = durationOrDefault(cfg.ClientHelloTimeout, defaultClientHelloWait)
105
- if cfg.RootHost == "" {
106
- cfg.RootHost = PortalRootHost(cfg.PortalURL)
107
- }
108
- if cfg.Policy == nil {
109
- cfg.Policy = policy.NewRuntime()
110
- }
111
- if cfg.RootHost == "" {
70
+ rootHost := PortalRootHost(cfg.PortalURL)
71
+ if rootHost == "" {
72
return nil, errors.New("root host is required")
73
}
114
- if len(cfg.APITLS.CertPEM) == 0 {
115
- return nil, errors.New("api tls certificate is required")
116
- }
117
- if len(cfg.APITLS.KeyPEM) == 0 && cfg.APITLS.Keyless == nil {
118
- return nil, errors.New("api tls key or keyless signer is required")
119
- }
74
75
return &Server{
122
- cfg: cfg,
123
- routes: newRouteTable(),
124
- leases: make(map[string]*leaseRecord),
76
+ cfg: cfg,
77
+ rootHost: rootHost,
78
+ registry: newLeaseRegistry(policy.NewRuntime()),
79
}, nil
80
}
81
128
-func (s *Server) Start(ctx context.Context) error {
82
+func (s *Server) Start(ctx context.Context, apiMux *http.ServeMux) error {
83
if s.group != nil {
84
return errors.New("server already started")
85
}
86
+ apiTLS, acmeManager, err := s.prepareAPITLS(ctx)
87
+ if err != nil {
88
+ return err
89
+ }
90
91
serverCtx, cancel := context.WithCancel(ctx)
92
var listenConfig net.ListenConfig
93
94
apiListener, err := listenConfig.Listen(serverCtx, "tcp", s.cfg.APIListenAddr)
95
if err != nil {
96
+ acmeManager.Stop()
97
cancel()
98
return fmt.Errorf("listen api: %w", err)
99
}
100
sniListener, err := listenConfig.Listen(serverCtx, "tcp", s.cfg.SNIListenAddr)
101
if err != nil {
102
+ acmeManager.Stop()
103
_ = apiListener.Close()
104
cancel()
105
return fmt.Errorf("listen sni: %w", err)
106
}
107
108
group, groupCtx := errgroup.WithContext(serverCtx)
149
-
150
- wrappedAPIListener, apiServer, apiCloser, err := s.newAPIServer(apiListener)
109
+ wrappedAPIListener, apiServer, apiCloser, err := s.newAPIServer(apiListener, apiMux, apiTLS)
110
if err != nil {
111
+ acmeManager.Stop()
112
_ = apiListener.Close()
113
_ = sniListener.Close()
114
cancel()
@@ -159,13 +119,15 @@ func (s *Server) Start(ctx context.Context) error {
119
s.sniListener = sniListener
120
s.apiServer = apiServer
121
s.apiTLSClose = apiCloser
122
+ s.acmeManager = acmeManager
123
s.cancel = cancel
124
s.group = group
125
126
group.Go(s.runAPIServer)
127
group.Go(func() error { return s.runSNIListener(groupCtx) })
167
- group.Go(func() error { return s.runLeaseJanitor(groupCtx) })
128
+ group.Go(func() error { return s.registry.RunJanitor(groupCtx, 5*time.Second) })
129
group.Go(func() error { return s.watchContext(groupCtx) })
130
+ s.acmeManager.Start(serverCtx)
131
return nil
132
}
133
@@ -183,11 +145,9 @@ func (s *Server) Shutdown(ctx context.Context) error {
145
s.cancel()
146
}
147
186
- s.mu.Lock()
187
- for _, lease := range s.leases {
148
+ for _, lease := range s.registry.CloseAll() {
149
lease.Broker.Close()
150
}
190
- s.mu.Unlock()
151
152
if s.sniListener != nil {
153
if err := s.sniListener.Close(); err != nil && !errors.Is(err, net.ErrClosed) {
@@ -202,10 +162,17 @@ func (s *Server) Shutdown(ctx context.Context) error {
162
if s.apiTLSClose != nil {
163
_ = s.apiTLSClose.Close()
164
}
165
+ if s.acmeManager != nil {
166
+ s.acmeManager.Stop()
167
+ }
168
})
169
return shutdownErr
170
}
171
172
+func (s *Server) PolicyRuntime() *policy.Runtime {
173
+ return s.registry.PolicyRuntime()
174
+}
175
+
176
func (s *Server) APIAddr() string {
177
if s.apiListener == nil {
178
return ""
@@ -221,36 +188,63 @@ func (s *Server) SNIAddr() string {
188
}
189
190
func (s *Server) GetLease(leaseID string) (LeaseSnapshot, bool) {
224
- s.mu.RLock()
225
- record, ok := s.leases[strings.TrimSpace(leaseID)]
226
- s.mu.RUnlock()
191
+ record, ok := s.registry.Get(leaseID)
192
if !ok {
193
return LeaseSnapshot{}, false
194
}
230
- return s.snapshotForLease(record), true
195
+ return s.registry.Snapshot(record), true
196
}
197
198
func (s *Server) ListLeases() []LeaseSnapshot {
234
- s.mu.RLock()
235
- defer s.mu.RUnlock()
236
-
237
- out := make([]LeaseSnapshot, 0, len(s.leases))
238
- for _, record := range s.leases {
239
- out = append(out, s.snapshotForLease(record))
199
+ records := s.registry.List()
200
+ out := make([]LeaseSnapshot, 0, len(records))
201
+ for _, record := range records {
202
+ out = append(out, s.registry.Snapshot(record))
203
}
204
return out
205
}
206
207
+func (s *Server) prepareAPITLS(ctx context.Context) (keyless.TLSMaterialConfig, *acme.Manager, error) {
208
+ acmeCfg := s.cfg.ACME
209
+ if baseDomain := normalizeHostname(acmeCfg.BaseDomain); baseDomain != "" && baseDomain != s.rootHost {
210
+ return keyless.TLSMaterialConfig{}, nil, fmt.Errorf("acme base domain %q does not match portal root host %q", acmeCfg.BaseDomain, s.rootHost)
211
+ }
212
+ acmeCfg.BaseDomain = s.rootHost
213
+
214
+ manager, err := acme.NewManager(acmeCfg)
215
+ if err != nil {
216
+ return keyless.TLSMaterialConfig{}, nil, fmt.Errorf("create acme manager: %w", err)
217
+ }
218
+
219
+ certPEM, keyPEM, err := manager.EnsureTLSMaterial(ctx)
220
+ if err != nil {
221
+ manager.Stop()
222
+ return keyless.TLSMaterialConfig{}, nil, fmt.Errorf("ensure relay certificate: %w", err)
223
+ }
224
+
225
+ apiTLS := keyless.TLSMaterialConfig{
226
+ CertPEM: certPEM,
227
+ KeyPEM: keyPEM,
228
+ }
229
+ if err := validateAPITLS(apiTLS); err != nil {
230
+ manager.Stop()
231
+ return keyless.TLSMaterialConfig{}, nil, err
232
+ }
233
+
234
+ return apiTLS, manager, nil
235
+}
236
+
237
func (s *Server) runSNIListener(ctx context.Context) error {
238
for {
239
conn, err := s.sniListener.Accept()
247
- if err != nil {
248
- if errors.Is(err, net.ErrClosed) || ctx.Err() != nil {
249
- return nil
250
- }
240
+ switch {
241
+ case err == nil:
242
+ go s.handleSNIConn(ctx, conn)
243
+ case errors.Is(err, net.ErrClosed):
244
+ return nil
245
+ default:
246
return err
247
}
253
- go s.handleSNIConn(ctx, conn)
248
}
249
}
250
@@ -267,25 +261,17 @@ func (s *Server) handleSNIConn(ctx context.Context, conn net.Conn) {
261
return
262
}
263
270
- if serverName == s.cfg.RootHost && s.cfg.RootFallbackAddr != "" {
271
- s.bridgeToFallback(ctx, wrappedConn)
272
- return
273
- }
274
-
275
- leaseID, ok := s.routes.Lookup(serverName)
276
- if !ok {
277
- _ = wrappedConn.Close()
264
+ if serverName == s.rootHost {
265
+ s.bridgeToAPI(ctx, wrappedConn)
266
return
267
}
268
281
- s.mu.RLock()
282
- record := s.leases[leaseID]
283
- s.mu.RUnlock()
284
- if record == nil || time.Now().After(record.ExpiresAt) {
269
+ record, ok := s.registry.Lookup(serverName)
270
+ if !ok || time.Now().After(record.ExpiresAt) {
271
_ = wrappedConn.Close()
272
return
273
}
288
- if !s.isLeaseRoutable(record) {
274
+ if !s.registry.IsRoutable(record) {
275
_ = wrappedConn.Close()
276
return
277
}
@@ -303,9 +289,13 @@ func (s *Server) handleSNIConn(ctx context.Context, conn net.Conn) {
289
_ = session.Close()
290
}
291
306
-func (s *Server) bridgeToFallback(ctx context.Context, conn net.Conn) {
292
+func (s *Server) bridgeToAPI(ctx context.Context, conn net.Conn) {
293
+ if s.apiListener == nil {
294
+ _ = conn.Close()
295
+ return
296
+ }
297
dialer := &net.Dialer{Timeout: 5 * time.Second}
308
- upstream, err := dialer.DialContext(ctx, "tcp", HostPortOrLoopback(s.cfg.RootFallbackAddr))
298
+ upstream, err := dialer.DialContext(ctx, "tcp", HostPortOrLoopback(s.apiListener.Addr().String()))
299
if err != nil {
300
_ = conn.Close()
301
return
@@ -313,40 +303,6 @@ func (s *Server) bridgeToFallback(ctx context.Context, conn net.Conn) {
303
bridgeConns(conn, upstream)
304
}
305
316
-func (s *Server) runLeaseJanitor(ctx context.Context) error {
317
- ticker := time.NewTicker(5 * time.Second)
318
- defer ticker.Stop()
319
-
320
- for {
321
- select {
322
- case <-ctx.Done():
323
- return nil
324
- case <-ticker.C:
325
- s.cleanupExpiredLeases()
326
- }
327
- }
328
-}
329
-
330
-func (s *Server) cleanupExpiredLeases() {
331
- now := time.Now()
332
-
333
- s.mu.Lock()
334
- expired := make([]*leaseRecord, 0)
335
- for leaseID, lease := range s.leases {
336
- if now.After(lease.ExpiresAt) {
337
- expired = append(expired, lease)
338
- delete(s.leases, leaseID)
339
- }
340
- }
341
- s.mu.Unlock()
342
-
343
- for _, lease := range expired {
344
- s.routes.DeleteLease(lease.Hostnames)
345
- s.cfg.Policy.ForgetLease(lease.ID)
346
- lease.Broker.Close()
347
- }
348
-}
349
-
306
func (s *Server) watchContext(ctx context.Context) error {
307
<-ctx.Done()
308
shutdownCtx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
@@ -380,51 +336,3 @@ func closeWrite(conn net.Conn) {
336
_ = cw.CloseWrite()
337
}
338
}
383
-
384
-func (s *Server) snapshotForLease(record *leaseRecord) LeaseSnapshot {
385
- if record == nil {
386
- return LeaseSnapshot{}
387
- }
388
- clientIP := record.ClientIP
389
- runtime := s.cfg.Policy
390
- return LeaseSnapshot{
391
- ID: record.ID,
392
- Name: record.Name,
393
- ClientIP: clientIP,
394
- Hostnames: append([]string(nil), record.Hostnames...),
395
- Metadata: record.Metadata,
396
- ExpiresAt: record.ExpiresAt,
397
- FirstSeenAt: record.FirstSeenAt,
398
- LastSeenAt: record.LastSeenAt,
399
- Ready: record.Broker.ReadyCount(),
400
- IsApproved: runtime.EffectiveApproval(record.ID),
401
- IsBanned: runtime.IsLeaseBanned(record.ID),
402
- IsDenied: runtime.IsLeaseDenied(record.ID),
403
- IsIPBanned: runtime.IPFilter().IsIPBanned(clientIP),
404
- }
405
-}
406
-
407
-func (s *Server) isLeaseRoutable(record *leaseRecord) bool {
408
- if record == nil {
409
- return false
410
- }
411
- return s.cfg.Policy.IsLeaseRoutable(record.ID)
412
-}
413
-
414
-func (s *Server) touchLease(leaseID, clientIP string) {
415
- now := time.Now()
416
-
417
- s.mu.Lock()
418
- record := s.leases[strings.TrimSpace(leaseID)]
419
- if record != nil {
420
- record.LastSeenAt = now
421
- if strings.TrimSpace(clientIP) != "" {
422
- record.ClientIP = clientIP
423
- }
424
- }
425
- s.mu.Unlock()
426
-
427
- if record != nil && strings.TrimSpace(clientIP) != "" {
428
- s.cfg.Policy.IPFilter().RegisterLeaseIP(record.ID, clientIP)
429
- }
430
-}
portal/server_test.go
new
+97
@@ -0,0 +1,97 @@
1
+package portal
2
+
3
+import (
4
+ "context"
5
+ "crypto/tls"
6
+ "encoding/json"
7
+ "net/http"
8
+ "strings"
9
+ "testing"
10
+
11
+ "github.com/gosuda/portal/v2/portal/acme"
12
+ "github.com/gosuda/portal/v2/types"
13
+)
14
+
15
+func TestServerStartInitializesLocalACMEAndSigner(t *testing.T) {
16
+ t.Parallel()
17
+
18
+ server, err := NewServer(ServerConfig{
19
+ PortalURL: "https://localhost:4017",
20
+ ACME: acme.Config{KeyDir: t.TempDir()},
21
+ APIListenAddr: "127.0.0.1:0",
22
+ SNIListenAddr: "127.0.0.1:0",
23
+ })
24
+ if err != nil {
25
+ t.Fatalf("NewServer() error = %v", err)
26
+ }
27
+
28
+ ctx, cancel := context.WithCancel(context.Background())
29
+ defer cancel()
30
+
31
+ if err := server.Start(ctx, nil); err != nil {
32
+ t.Fatalf("Start() error = %v", err)
33
+ }
34
+
35
+ client := &http.Client{
36
+ Transport: &http.Transport{
37
+ TLSClientConfig: &tls.Config{InsecureSkipVerify: true},
38
+ },
39
+ }
40
+ t.Cleanup(func() {
41
+ client.CloseIdleConnections()
42
+ cancel()
43
+ if err := server.Wait(); err != nil {
44
+ t.Fatalf("Wait() error = %v", err)
45
+ }
46
+ })
47
+
48
+ healthResp, err := client.Get("https://" + HostPortOrLoopback(server.APIAddr()) + types.PathHealthz)
49
+ if err != nil {
50
+ t.Fatalf("GET /healthz error = %v", err)
51
+ }
52
+ defer healthResp.Body.Close()
53
+
54
+ if healthResp.StatusCode != http.StatusOK {
55
+ t.Fatalf("GET /healthz status = %d, want %d", healthResp.StatusCode, http.StatusOK)
56
+ }
57
+
58
+ var healthEnvelope types.APIEnvelope[map[string]string]
59
+ if err := json.NewDecoder(healthResp.Body).Decode(&healthEnvelope); err != nil {
60
+ t.Fatalf("decode /healthz response: %v", err)
61
+ }
62
+ if !healthEnvelope.OK || healthEnvelope.Data["status"] != "ok" {
63
+ t.Fatalf("GET /healthz response = %+v, want ok status", healthEnvelope)
64
+ }
65
+
66
+ signResp, err := client.Get("https://" + HostPortOrLoopback(server.APIAddr()) + types.PathV1Sign)
67
+ if err != nil {
68
+ t.Fatalf("GET /v1/sign error = %v", err)
69
+ }
70
+ defer signResp.Body.Close()
71
+
72
+ if signResp.StatusCode != http.StatusMethodNotAllowed {
73
+ t.Fatalf("GET /v1/sign status = %d, want %d", signResp.StatusCode, http.StatusMethodNotAllowed)
74
+ }
75
+}
76
+
77
+func TestServerStartRejectsMismatchedACMEBaseDomain(t *testing.T) {
78
+ t.Parallel()
79
+
80
+ server, err := NewServer(ServerConfig{
81
+ PortalURL: "https://portal.example.com",
82
+ ACME: acme.Config{BaseDomain: "other.example.com", KeyDir: t.TempDir()},
83
+ APIListenAddr: "127.0.0.1:0",
84
+ SNIListenAddr: "127.0.0.1:0",
85
+ })
86
+ if err != nil {
87
+ t.Fatalf("NewServer() error = %v", err)
88
+ }
89
+
90
+ err = server.Start(context.Background(), nil)
91
+ if err == nil {
92
+ t.Fatal("Start() error = nil, want mismatch error")
93
+ }
94
+ if !strings.Contains(err.Error(), "does not match portal root host") {
95
+ t.Fatalf("Start() error = %v, want base domain mismatch", err)
96
+ }
97
+}