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 +}