refactor codes

rabbitprincess committed Mar 2, 2026 at 22:49 UTC d788782848e01d0570e67dd11b0f8cde724ba181
10 files changed +167 -371
cmd/demo-app/main.go
+10 -11
@@ -27,17 +27,16 @@ var staticFiles embed.FS
27 var thumbnailPNG []byte
28
29 var (
30 - flagServerURL string
31 - flagPort int
32 - flagName string
33 - flagDesc string
34 - flagTags string
35 - flagOwner string
36 - flagHide bool
37 - flagTLSMode string
38 - flagTLSCert string
39 - flagTLSKey string
40 - flagTLSBaseDomain string
30 + flagServerURL string
31 + flagPort int
32 + flagName string
33 + flagDesc string
34 + flagTags string
35 + flagOwner string
36 + flagHide bool
37 + flagTLSMode string
38 + flagTLSCert string
39 + flagTLSKey string
40 )
41
42 func main() {
cmd/portal-tunnel/main.go
+66 -97
@@ -20,17 +20,17 @@ import (
20 )
21
22 var (
23 - flagRelayURLs string
24 - flagHost string
25 - flagName string
26 - flagTLSMode string
27 - flagTLSCertFile string
28 - flagTLSKeyFile string
29 - flagDescription string
30 - flagTags string
31 - flagThumbnail string
32 - flagOwner string
33 - flagHide bool
23 + flagRelayURLs string
24 + flagHost string
25 + flagName string
26 + flagDesc string
27 + flagTags string
28 + flagThumbnail string
29 + flagOwner string
30 + flagHide bool
31 + flagTLSMode string
32 + flagTLSCert string
33 + flagTLSKey string
34 )
35
36 func main() {
@@ -50,58 +50,29 @@ func main() {
50 defaultTLSMode = string(sdk.TLSModeNoTLS)
51 }
52 flag.StringVar(&flagTLSMode, "tls-mode", defaultTLSMode, "TLS mode: no-tls, self, or keyless [env: TLS_MODE]")
53 - flag.StringVar(&flagTLSCertFile, "tls-cert-file", os.Getenv("TLS_CERT_FILE"), "PEM certificate chain for --tls-mode self [env: TLS_CERT_FILE]")
54 - flag.StringVar(&flagTLSKeyFile, "tls-key-file", os.Getenv("TLS_KEY_FILE"), "PEM private key for --tls-mode self [env: TLS_KEY_FILE]")
53 + flag.StringVar(&flagTLSCert, "tls-cert-file", os.Getenv("TLS_CERT_FILE"), "PEM certificate chain for --tls-mode self [env: TLS_CERT_FILE]")
54 + flag.StringVar(&flagTLSKey, "tls-key-file", os.Getenv("TLS_KEY_FILE"), "PEM private key for --tls-mode self [env: TLS_KEY_FILE]")
55
56 - flag.StringVar(&flagDescription, "description", os.Getenv("APP_DESCRIPTION"), "Service description metadata [env: APP_DESCRIPTION]")
56 + flag.StringVar(&flagDesc, "description", os.Getenv("APP_DESCRIPTION"), "Service description metadata [env: APP_DESCRIPTION]")
57 flag.StringVar(&flagTags, "tags", os.Getenv("APP_TAGS"), "Service tags metadata (comma-separated) [env: APP_TAGS]")
58 flag.StringVar(&flagThumbnail, "thumbnail", os.Getenv("APP_THUMBNAIL"), "Service thumbnail URL metadata [env: APP_THUMBNAIL]")
59 flag.StringVar(&flagOwner, "owner", os.Getenv("APP_OWNER"), "Service owner metadata [env: APP_OWNER]")
60
61 defaultHide := os.Getenv("APP_HIDE") == "true"
62 flag.BoolVar(&flagHide, "hide", defaultHide, "Hide service from discovery (metadata) [env: APP_HIDE]")
63 -
63 flag.Parse()
64
66 - if flagHost == "" || flagName == "" {
67 - flag.Usage()
68 - os.Exit(1)
69 - }
70 - if flagTLSMode != string(sdk.TLSModeNoTLS) &&
71 - flagTLSMode != string(sdk.TLSModeSelf) &&
72 - flagTLSMode != string(sdk.TLSModeKeyless) {
73 - log.Error().Str("tls_mode", flagTLSMode).Msg("--tls-mode must be one of: no-tls, self, keyless")
74 - os.Exit(1)
75 - }
76 -
77 - relayURLs := parseURLs(flagRelayURLs)
78 - if len(relayURLs) == 0 {
79 - log.Error().Msg("--relay must include at least one non-empty URL")
80 - os.Exit(1)
81 - }
82 -
83 - ctx, cancel := context.WithCancel(context.Background())
84 - defer cancel()
85 -
86 - sigCh := make(chan os.Signal, 1)
87 - signal.Notify(sigCh, syscall.SIGINT, syscall.SIGTERM)
88 - defer signal.Stop(sigCh)
89 -
90 - go func() {
91 - <-sigCh
92 - log.Info().Msg("Shutting down tunnel...")
93 - cancel()
94 - }()
95 -
96 - if err := runServiceTunnel(ctx, relayURLs); err != nil {
65 + if err := runTunnel(); err != nil {
66 log.Error().Err(err).Msg("Exited with error")
67 os.Exit(1)
68 }
100 -
101 - log.Info().Msg("Tunnel stopped")
69 }
70
104 -func runServiceTunnel(ctx context.Context, relayURLs []string) error {
71 +func runTunnel() error {
72 + ctx, stop := signal.NotifyContext(context.Background(), os.Interrupt, syscall.SIGTERM)
73 + defer stop()
74 +
75 + relayURLs := splitCSV(flagRelayURLs)
76 if len(relayURLs) == 0 {
77 return fmt.Errorf("no relay URLs provided")
78 }
@@ -112,35 +83,32 @@ func runServiceTunnel(ctx context.Context, relayURLs []string) error {
83 log.Info().Msgf(" Relays: %s", strings.Join(relayURLs, ", "))
84 log.Info().Msgf(" TLS Mode: %s", flagTLSMode)
85
115 - var clientOpts []sdk.ClientOption
116 - clientOpts = append(clientOpts, sdk.WithBootstrapServers(relayURLs))
117 -
118 - if flagTLSMode == string(sdk.TLSModeSelf) {
119 - clientOpts = append(clientOpts, sdk.WithTLSSelfCertificateFiles(flagTLSCertFile, flagTLSKeyFile))
120 - log.Info().
121 - Str("cert_file", flagTLSCertFile).
122 - Str("key_file", flagTLSKeyFile).
123 - Msg("TLS: Using self-managed local certificate")
124 - } else if flagTLSMode == string(sdk.TLSModeKeyless) {
125 - clientOpts = append(clientOpts, sdk.WithTLSKeylessDefaults())
126 - log.Info().Msg("TLS: Using keyless remote signer (SDK auto configuration)")
86 + opts := []sdk.ClientOption{sdk.WithBootstrapServers(relayURLs)}
87 + mode := sdk.TLSMode(flagTLSMode)
88 + switch mode {
89 + case sdk.TLSModeNoTLS:
90 + case sdk.TLSModeSelf:
91 + opts = append(opts, sdk.WithTLSSelfCertificateFiles(flagTLSCert, flagTLSKey))
92 + case sdk.TLSModeKeyless:
93 + opts = append(opts, sdk.WithTLSKeylessDefaults())
94 + default:
95 + return fmt.Errorf("unsupported tls mode: %s", flagTLSMode)
96 }
97
129 - client, err := sdk.NewClient(clientOpts...)
98 + client, err := sdk.NewClient(opts...)
99 if err != nil {
100 return fmt.Errorf("service %s: failed to create client: %w", flagName, err)
101 }
102 defer client.Close()
103
135 - metadataOptions := []sdk.MetadataOption{
136 - sdk.WithDescription(flagDescription),
104 + listener, err := client.Listen(
105 + flagName,
106 + sdk.WithDescription(flagDesc),
107 sdk.WithTags(splitCSV(flagTags)),
108 sdk.WithOwner(flagOwner),
109 sdk.WithThumbnail(flagThumbnail),
110 sdk.WithHide(flagHide),
141 - }
142 -
143 - listener, err := client.Listen(flagName, metadataOptions...)
111 + )
112 if err != nil {
113 return fmt.Errorf("service %s: failed to register service: %w", flagName, err)
114 }
@@ -165,11 +133,13 @@ func runServiceTunnel(ctx context.Context, relayURLs []string) error {
133
134 connCount := 0
135 var connWG sync.WaitGroup
168 - defer connWG.Wait()
136 +
137 +loop:
138 for {
139 select {
140 case <-ctx.Done():
172 - return nil
141 + log.Info().Msg("[tunnel] shutting down...")
142 + break loop
143 default:
144 }
145
@@ -177,7 +147,7 @@ func runServiceTunnel(ctx context.Context, relayURLs []string) error {
147 if err != nil {
148 select {
149 case <-ctx.Done():
180 - return nil
150 + break loop
151 default:
152 log.Error().Err(err).Msg("Failed to accept connection")
153 continue
@@ -201,34 +171,21 @@ func runServiceTunnel(ctx context.Context, relayURLs []string) error {
171 log.Info().Str("proxy", proxyType).Msg("Connection closed")
172 }(relayConn)
173 }
204 -}
174
206 -func parseURLs(raw string) []string {
207 - raw = strings.TrimSpace(raw)
208 - if raw == "" {
209 - return nil
210 - }
211 - parts := strings.Split(raw, ",")
212 - out := make([]string, 0, len(parts))
213 - for _, p := range parts {
214 - p = strings.TrimSpace(p)
215 - if p != "" {
216 - out = append(out, p)
217 - }
218 - }
219 - return out
220 -}
175 + done := make(chan struct{})
176 + go func() {
177 + connWG.Wait()
178 + close(done)
179 + }()
180
222 -func splitCSV(raw string) []string {
223 - parts := strings.Split(raw, ",")
224 - out := make([]string, 0, len(parts))
225 - for _, part := range parts {
226 - part = strings.TrimSpace(part)
227 - if part != "" {
228 - out = append(out, part)
229 - }
181 + select {
182 + case <-done:
183 + case <-time.After(5 * time.Second):
184 + log.Warn().Msg("[tunnel] shutdown timeout, some connections still active")
185 }
231 - return out
186 +
187 + log.Info().Msg("[tunnel] shutdown complete")
188 + return nil
189 }
190
191 var bufferPool = sync.Pool{
@@ -245,7 +202,7 @@ func proxyConnection(ctx context.Context, localAddr string, relayConn net.Conn,
202 localConn, err := dialer.DialContext(ctx, "tcp", localAddr)
203 if err != nil {
204 log.Debug().
248 - Str("local_addr", localAddr).
205 + Str("addr", localAddr).
206 Err(err).
207 Msg("Local service unavailable")
208 if tlsEnabled {
@@ -255,7 +212,7 @@ func proxyConnection(ctx context.Context, localAddr string, relayConn net.Conn,
212 }
213 defer localConn.Close()
214
258 - log.Info().Str("local_addr", localAddr).Msg("Connected to local service")
215 + log.Info().Str("addr", localAddr).Msg("Connected to local service")
216
217 errCh := make(chan error, 2)
218 stopCh := make(chan struct{})
@@ -321,3 +278,15 @@ func writeEmptyHTTPResponse(conn net.Conn) error {
278 _, err := conn.Write([]byte(response))
279 return err
280 }
281 +
282 +func splitCSV(raw string) []string {
283 + parts := strings.Split(raw, ",")
284 + out := make([]string, 0, len(parts))
285 + for _, part := range parts {
286 + part = strings.TrimSpace(part)
287 + if part != "" {
288 + out = append(out, part)
289 + }
290 + }
291 + return out
292 +}
cmd/relay-server/frontend/src/components/TunnelCommandModal.tsx
+19 -1
@@ -1,4 +1,4 @@
1 -import { useState, useMemo, useRef } from "react";
1 +import { useState, useMemo, useRef, useEffect } from "react";
2 import { Copy, Check, Terminal, X } from "lucide-react";
3 import { cn } from "@/lib/utils";
4 import {
@@ -40,6 +40,24 @@ export function TunnelCommandModal({ trigger }: TunnelCommandModalProps) {
40 const [tlsCertFile, setTlsCertFile] = useState("");
41 const [tlsKeyFile, setTlsKeyFile] = useState("");
42 const urlInputRef = useRef<HTMLInputElement>(null);
43 + const keylessAvailable = useMemo(() => {
44 + if (relayUrls.length === 0) {
45 + return false;
46 + }
47 + return relayUrls.every((raw) => {
48 + try {
49 + return new URL(raw).protocol === "https:";
50 + } catch {
51 + return false;
52 + }
53 + });
54 + }, [relayUrls]);
55 +
56 + useEffect(() => {
57 + if (!keylessAvailable && tlsMode === "keyless") {
58 + setTlsMode("no-tls");
59 + }
60 + }, [keylessAvailable, tlsMode]);
61
62 const addRelayUrl = (url: string) => {
63 const trimmed = url.trim();
cmd/relay-server/registry.go
+30 -121
@@ -64,48 +64,13 @@ func (r *SDKRegistry) handleRegister(w http.ResponseWriter, req *http.Request, s
64 return
65 }
66
67 - // Validate request
68 - if registerReq.LeaseID == "" {
69 - writeJSON(w, sdk.RegisterResponse{
70 - Success: false,
71 - Message: "lease_id is required",
72 - })
73 - return
74 - }
75 -
76 - if registerReq.Name == "" {
77 - writeJSON(w, sdk.RegisterResponse{
78 - Success: false,
79 - Message: "name is required",
80 - })
81 - return
82 - }
83 -
84 - if strings.TrimSpace(registerReq.ReverseToken) == "" {
85 - writeJSON(w, sdk.RegisterResponse{
86 - Success: false,
87 - Message: "reverse_token is required",
88 - })
89 - return
90 - }
91 - tlsMode := normalizeTLSMode(registerReq.TLSMode)
92 - switch tlsMode {
93 - case sdk.TLSModeNoTLS, sdk.TLSModeSelf, sdk.TLSModeKeyless:
94 - default:
95 - writeJSON(w, sdk.RegisterResponse{
96 - Success: false,
97 - Message: "tls_mode must be one of: no-tls, self, keyless",
98 - })
99 - return
100 - }
101 -
67 // Create lease
68 lease := &portal.Lease{
69 ID: registerReq.LeaseID,
70 Name: registerReq.Name,
71 Metadata: registerReq.Metadata,
72 Expires: time.Now().Add(30 * time.Second),
108 - TLSMode: string(tlsMode),
73 + TLSMode: string(registerReq.TLSMode),
74 ReverseToken: strings.TrimSpace(registerReq.ReverseToken),
75 }
76
@@ -122,8 +87,9 @@ func (r *SDKRegistry) handleRegister(w http.ResponseWriter, req *http.Request, s
87 serv.GetReverseHub().ClearDropped(registerReq.LeaseID)
88
89 // Only register SNI route for TLS leases.
125 - if normalizeTLSMode(tlsMode) != sdk.TLSModeNoTLS {
126 - if err := registerSNIRoute(serv, registerReq.LeaseID, registerReq.Name); err != nil {
90 + if registerReq.TLSMode != sdk.TLSModeNoTLS {
91 + sniName := strings.ToLower(strings.TrimSpace(registerReq.Name)) + "." + serv.BaseHost
92 + if err := serv.GetSNIRouter().RegisterRoute(sniName, registerReq.LeaseID, registerReq.Name); err != nil {
93 // Keep lease and route state consistent on partial failure.
94 serv.GetLeaseManager().DeleteLease(registerReq.LeaseID)
95 writeJSON(w, sdk.RegisterResponse{
@@ -137,7 +103,7 @@ func (r *SDKRegistry) handleRegister(w http.ResponseWriter, req *http.Request, s
103 log.Info().
104 Str("lease_id", registerReq.LeaseID).
105 Str("name", registerReq.Name).
140 - Str("tls_mode", string(tlsMode)).
106 + Str("tls_mode", string(registerReq.TLSMode)).
107 Msg("[Registry] Lease registered")
108
109 // Build public URL
@@ -158,23 +124,12 @@ func (r *SDKRegistry) handleUnregister(w http.ResponseWriter, req *http.Request,
124 return
125 }
126
161 - var unregisterReq struct {
162 - LeaseID string `json:"lease_id"`
163 - }
164 -
127 + var unregisterReq sdk.UnregisterRequest
128 if err := json.NewDecoder(req.Body).Decode(&unregisterReq); err != nil {
129 log.Error().Err(err).Msg("[Registry] Failed to decode unregistration request")
167 - writeJSON(w, map[string]any{
168 - "success": false,
169 - "message": "invalid request body",
170 - })
171 - return
172 - }
173 -
174 - if unregisterReq.LeaseID == "" {
175 - writeJSON(w, map[string]any{
176 - "success": false,
177 - "message": "lease_id is required",
130 + writeJSON(w, sdk.APIResponse{
131 + Success: false,
132 + Message: "invalid request body",
133 })
134 return
135 }
@@ -185,11 +140,11 @@ func (r *SDKRegistry) handleUnregister(w http.ResponseWriter, req *http.Request,
140 Str("lease_id", unregisterReq.LeaseID).
141 Msg("[Registry] Lease unregistered")
142 }
188 - unregisterSNIRoute(serv, unregisterReq.LeaseID)
143 + serv.GetSNIRouter().UnregisterRouteByLeaseID(unregisterReq.LeaseID)
144 serv.GetReverseHub().DropLease(unregisterReq.LeaseID)
145
191 - writeJSON(w, map[string]any{
192 - "success": true,
146 + writeJSON(w, sdk.APIResponse{
147 + Success: true,
148 })
149 }
150
@@ -201,31 +156,12 @@ func (r *SDKRegistry) handleRenew(w http.ResponseWriter, req *http.Request, serv
156 return
157 }
158
204 - var renewReq struct {
205 - LeaseID string `json:"lease_id"`
206 - ReverseToken string `json:"reverse_token"`
207 - }
208 -
159 + var renewReq sdk.RenewRequest
160 if err := json.NewDecoder(req.Body).Decode(&renewReq); err != nil {
161 log.Error().Err(err).Msg("[Registry] Failed to decode renewal request")
211 - writeJSON(w, map[string]any{
212 - "success": false,
213 - "message": "invalid request body",
214 - })
215 - return
216 - }
217 -
218 - if renewReq.LeaseID == "" {
219 - writeJSON(w, map[string]any{
220 - "success": false,
221 - "message": "lease_id is required",
222 - })
223 - return
224 - }
225 - if strings.TrimSpace(renewReq.ReverseToken) == "" {
226 - writeJSON(w, map[string]any{
227 - "success": false,
228 - "message": "reverse_token is required",
162 + writeJSON(w, sdk.APIResponse{
163 + Success: false,
164 + Message: "invalid request body",
165 })
166 return
167 }
@@ -233,16 +169,16 @@ func (r *SDKRegistry) handleRenew(w http.ResponseWriter, req *http.Request, serv
169 // Get existing lease
170 entry, ok := serv.GetLeaseManager().GetLeaseByID(renewReq.LeaseID)
171 if !ok {
236 - writeJSON(w, map[string]any{
237 - "success": false,
238 - "message": "lease not found",
172 + writeJSON(w, sdk.APIResponse{
173 + Success: false,
174 + Message: "lease not found",
175 })
176 return
177 }
178 if subtle.ConstantTimeCompare([]byte(strings.TrimSpace(entry.Lease.ReverseToken)), []byte(strings.TrimSpace(renewReq.ReverseToken))) != 1 {
243 - writeJSON(w, map[string]any{
244 - "success": false,
245 - "message": "unauthorized lease renewal",
179 + writeJSON(w, sdk.APIResponse{
180 + Success: false,
181 + Message: "unauthorized lease renewal",
182 })
183 return
184 }
@@ -250,17 +186,18 @@ func (r *SDKRegistry) handleRenew(w http.ResponseWriter, req *http.Request, serv
186 // Update expiration
187 entry.Lease.Expires = time.Now().Add(30 * time.Second)
188 if !serv.GetLeaseManager().UpdateLease(entry.Lease) {
253 - writeJSON(w, map[string]any{
254 - "success": false,
255 - "message": "failed to renew lease",
189 + writeJSON(w, sdk.APIResponse{
190 + Success: false,
191 + Message: "failed to renew lease",
192 })
193 return
194 }
195
196 // Re-register route if needed (e.g., router restarted while lease remained active).
197 // Only TLS leases need SNI routes.
262 - if normalizeTLSMode(sdk.TLSMode(entry.Lease.TLSMode)) != sdk.TLSModeNoTLS {
263 - if err := registerSNIRoute(serv, entry.Lease.ID, entry.Lease.Name); err != nil {
198 + if sdk.TLSMode(entry.Lease.TLSMode) != sdk.TLSModeNoTLS {
199 + sniName := strings.ToLower(strings.TrimSpace(entry.Lease.Name)) + "." + serv.BaseHost
200 + if err := serv.GetSNIRouter().RegisterRoute(sniName, entry.Lease.ID, entry.Lease.Name); err != nil {
201 log.Warn().
202 Err(err).
203 Str("lease_id", entry.Lease.ID).
@@ -269,39 +206,11 @@ func (r *SDKRegistry) handleRenew(w http.ResponseWriter, req *http.Request, serv
206 }
207 }
208
272 - writeJSON(w, map[string]any{
273 - "success": true,
209 + writeJSON(w, sdk.APIResponse{
210 + Success: true,
211 })
212 }
213
277 -func registerSNIRoute(serv *portal.RelayServer, leaseID, name string) error {
278 - sniRouter := serv.GetSNIRouter()
279 - if sniRouter == nil {
280 - return nil
281 - }
282 - if serv.BaseHost == "" {
283 - return fmt.Errorf("base domain not configured (set PORTAL_URL)")
284 - }
285 - sniName := strings.ToLower(strings.TrimSpace(name)) + "." + serv.BaseHost
286 - return sniRouter.RegisterRoute(sniName, leaseID, name)
287 -}
288 -
289 -func unregisterSNIRoute(serv *portal.RelayServer, leaseID string) {
290 - sniRouter := serv.GetSNIRouter()
291 - if sniRouter == nil {
292 - return
293 - }
294 - sniRouter.UnregisterRouteByLeaseID(leaseID)
295 -}
296 -
297 -func normalizeTLSMode(mode sdk.TLSMode) sdk.TLSMode {
298 - normalized := sdk.TLSMode(strings.ToLower(strings.TrimSpace(string(mode))))
299 - if normalized == "" {
300 - return sdk.TLSModeNoTLS
301 - }
302 - return normalized
303 -}
304 -
214 // handleDomain returns the relay's base domain for TLS certificate construction.
215 func (r *SDKRegistry) handleDomain(w http.ResponseWriter, req *http.Request, serv *portal.RelayServer) {
216 if serv.BaseHost == "" {
cmd/relay-server/serve.go
+4 -9
@@ -139,14 +139,10 @@ func serveAPI(addr string, serv *portal.RelayServer, admin *Admin, frontend *Fro
139 go func() {
140 var err error
141 if tlsCertFile != "" && tlsKeyFile != "" {
142 - log.Info().
143 - Str("addr", addr).
144 - Str("cert_file", tlsCertFile).
145 - Str("key_file", tlsKeyFile).
146 - Msg("[server] admin https enabled")
142 + log.Info().Str("addr", addr).Str("cert_file", tlsCertFile).Str("key_file", tlsKeyFile).Msg("[server] https api enabled")
143 err = srv.ListenAndServeTLS(tlsCertFile, tlsKeyFile)
144 } else {
149 - log.Info().Msgf("[server] http: %s", addr)
145 + log.Info().Str("addr", addr).Msgf("[server] http api enabled")
146 err = srv.ListenAndServe()
147 }
148 if err != nil && err != http.ErrServerClosed {
@@ -170,12 +166,11 @@ func shouldProxyHTTP(host string, serv *portal.RelayServer) (string, *portal.Lea
166 entry, ok := serv.GetLeaseManager().GetLeaseByName(leaseName)
167 if !ok {
168 log.Debug().Str("lease_name", leaseName).Msg("[proxy] shouldProxyHTTP: lease not found")
173 - // Keep existing behavior: unknown subdomain goes through proxy path and returns 404.
169 return leaseName, nil, true
170 }
171
172 // If TLS mode is no-tls, we can proxy via HTTP.
178 - shouldProxy := normalizeTLSMode(sdk.TLSMode(entry.Lease.TLSMode)) == sdk.TLSModeNoTLS
173 + shouldProxy := sdk.TLSMode(entry.Lease.TLSMode) == sdk.TLSModeNoTLS
174 log.Debug().
175 Str("lease_name", leaseName).
176 Str("tls_mode", entry.Lease.TLSMode).
@@ -194,7 +189,7 @@ func proxyToHTTP(w http.ResponseWriter, r *http.Request, serv *portal.RelayServe
189 return
190 }
191
197 - if normalizeTLSMode(sdk.TLSMode(entry.Lease.TLSMode)) != sdk.TLSModeNoTLS {
192 + if sdk.TLSMode(entry.Lease.TLSMode) != sdk.TLSModeNoTLS {
193 http.Error(w, "TLS enabled requires HTTPS access", http.StatusBadRequest)
194 return
195 }
cmd/relay-server/utils.go
+1 -1
@@ -386,7 +386,7 @@ func (r *leaseRow) fromLeaseEntry(entry *portal.LeaseEntry, admin *Admin, portal
386 }
387
388 kind := "http"
389 - if normalizeTLSMode(sdk.TLSMode(lease.TLSMode)) != sdk.TLSModeNoTLS {
389 + if sdk.TLSMode(lease.TLSMode) != sdk.TLSModeNoTLS {
390 kind = "https"
391 }
392
portal/acme/acme.go
+2 -2
@@ -295,8 +295,8 @@ func certCoversDomains(certFile string, domains []string) (bool, error) {
295 }
296
297 for _, domain := range domains {
298 - if strings.HasPrefix(domain, "*.") {
299 - probeHost := "acme-probe." + strings.TrimPrefix(domain, "*.")
298 + if after, ok := strings.CutPrefix(domain, "*."); ok {
299 + probeHost := "acme-probe." + after
300 if err := cert.VerifyHostname(probeHost); err != nil {
301 return false, nil
302 }
sdk/client.go
+19 -34
@@ -31,7 +31,6 @@ type Client struct {
31 func NewClient(opt ...ClientOption) (*Client, error) {
32 config := &ClientConfig{
33 BootstrapServers: []string{},
34 - ReverseWorkers: 0, // uses defaultReverseWorkers from listener
34 ReverseDialTimeout: 5 * time.Second,
35 TLSMode: TLSModeNoTLS,
36 }
@@ -88,7 +87,7 @@ func (c *Client) Listen(name string, options ...MetadataOption) (net.Listener, e
87 }
88
89 leaseCopy := *lease
91 - listener, listenerErr := NewListener(relayAddr, &leaseCopy, tlsConfig, c.config.ReverseWorkers, c.config.ReverseDialTimeout, listenerCloseFns...)
90 + listener, listenerErr := NewListener(relayAddr, &leaseCopy, tlsConfig, 0, c.config.ReverseDialTimeout, listenerCloseFns...)
91 if listenerErr != nil {
92 for _, closeFn := range listenerCloseFns {
93 if closeFn != nil {
@@ -220,18 +219,12 @@ func (c *Client) buildTLSConfig(relayAddr, leaseName string) (*tls.Config, []fun
219
220 switch tlsMode {
221 case TLSModeSelf:
223 - var cert tls.Certificate
224 - var err error
225 - if c.config.TLSCertificate != nil {
226 - cert = *c.config.TLSCertificate
227 - } else {
228 - if c.config.TLSSelfCertFile == "" || c.config.TLSSelfKeyFile == "" {
229 - return nil, nil, fmt.Errorf("self TLS mode requires certificate/key (WithTLSSelfCertificate or WithTLSSelfCertificateFiles)")
230 - }
231 - cert, err = tls.LoadX509KeyPair(c.config.TLSSelfCertFile, c.config.TLSSelfKeyFile)
232 - if err != nil {
233 - return nil, nil, fmt.Errorf("load self TLS certificate files: %w", err)
234 - }
222 + if c.config.TLSSelfCertFile == "" || c.config.TLSSelfKeyFile == "" {
223 + return nil, nil, fmt.Errorf("self TLS mode requires certificate/key files (WithTLSSelfCertificateFiles)")
224 + }
225 + cert, err := tls.LoadX509KeyPair(c.config.TLSSelfCertFile, c.config.TLSSelfKeyFile)
226 + if err != nil {
227 + return nil, nil, fmt.Errorf("load self TLS certificate files: %w", err)
228 }
229
230 tlsConfig := &tls.Config{MinVersion: tls.VersionTLS12}
@@ -242,35 +235,30 @@ func (c *Client) buildTLSConfig(relayAddr, leaseName string) (*tls.Config, []fun
235 return tlsConfig, nil, nil
236
237 case TLSModeKeyless:
245 - keylessEndpoint := c.config.TLSKeyless.Endpoint
238 + keylessEndpoint := c.config.TLSKeylessEndpoint
239 if keylessEndpoint == "" {
240 keylessEndpoint = relayAddr
241 }
242
250 - keylessKeyID := c.config.TLSKeyless.KeyID
251 - if keylessKeyID == "" {
252 - keylessKeyID = "relay-cert"
253 - }
243 + keylessKeyID := "relay-cert"
244
255 - keylessServerName := c.config.TLSKeyless.ServerName
256 - if keylessServerName == "" {
257 - if parsed, err := url.Parse(keylessEndpoint); err == nil {
258 - keylessServerName = parsed.Hostname()
259 - }
245 + keylessServerName := ""
246 + if parsed, err := url.Parse(keylessEndpoint); err == nil {
247 + keylessServerName = parsed.Hostname()
248 }
249
250 certPEM, rootCAPEM, err := keyless.ResolveMaterials(
251 context.Background(),
252 keylessEndpoint,
253 keylessServerName,
266 - c.config.TLSKeylessCertificatePEM,
267 - c.config.TLSKeyless.RootCAPEM,
254 + nil,
255 + nil,
256 )
257 if err != nil {
258 return nil, nil, fmt.Errorf("prepare keyless materials: %w", err)
259 }
260
273 - baseDomain := c.config.TLSKeyless.BaseDomain
261 + baseDomain := c.config.TLSKeylessBaseDomain
262 if baseDomain == "" {
263 baseDomain = ExtractBaseDomain(relayAddr)
264 }
@@ -286,13 +274,10 @@ func (c *Client) buildTLSConfig(relayAddr, leaseName string) (*tls.Config, []fun
274 }
275
276 remoteSigner, err := keylesstls.NewRemoteSigner(keylesstls.RemoteSignerConfig{
289 - Endpoint: keylessEndpoint,
290 - ServerName: keylessServerName,
291 - KeyID: keylessKeyID,
292 - EnableMTLS: c.config.TLSKeyless.EnableMTLS,
293 - ClientCertPEM: c.config.TLSKeyless.ClientCertPEM,
294 - ClientKeyPEM: c.config.TLSKeyless.ClientKeyPEM,
295 - RootCAPEM: rootCAPEM,
277 + Endpoint: keylessEndpoint,
278 + ServerName: keylessServerName,
279 + KeyID: keylessKeyID,
280 + RootCAPEM: rootCAPEM,
281 }, certPEM)
282 if err != nil {
283 return nil, nil, fmt.Errorf("create keyless remote signer: %w", err)
sdk/listener.go
+3 -6
@@ -64,7 +64,7 @@ func NewListener(relayAddr string, lease *portal.Lease, tlsConfig *tls.Config, r
64 if lease.Name == "" {
65 return nil, fmt.Errorf("lease name is required")
66 }
67 - if strings.TrimSpace(lease.ReverseToken) == "" {
67 + if lease.ReverseToken == "" {
68 return nil, fmt.Errorf("lease reverse token is required")
69 }
70
@@ -389,15 +389,12 @@ func (l *Listener) postJSON(path string, body any) error {
389 return nil
390 }
391
392 - var apiResp struct {
393 - Success *bool `json:"success"`
394 - Message string `json:"message"`
395 - }
392 + var apiResp APIResponse
393 if err := json.Unmarshal(data, &apiResp); err != nil {
394 // Non-JSON success payloads are treated as successful.
395 return nil
396 }
400 - if apiResp.Success != nil && !*apiResp.Success {
397 + if !apiResp.Success {
398 msg := strings.TrimSpace(apiResp.Message)
399 if msg == "" {
400 msg = strings.TrimSpace(string(data))
sdk/types.go
+13 -89
@@ -1,10 +1,7 @@
1 package sdk
2
3 import (
4 - "context"
5 - "crypto/tls"
4 "errors"
7 - "io"
5 "time"
6
7 "gosuda.org/portal/portal"
@@ -29,37 +26,21 @@ const (
26 TLSModeKeyless TLSMode = "keyless"
27 )
28
32 -type TLSKeylessConfig struct {
33 - Endpoint string
34 - ServerName string
35 - BaseDomain string
36 - KeyID string
37 - RootCAPEM []byte
38 - EnableMTLS bool
39 - ClientCertPEM []byte
40 - ClientKeyPEM []byte
41 -}
42 -
29 type ClientConfig struct {
44 - BootstrapServers []string
45 - Dialer func(context.Context, string) (io.ReadWriteCloser, error)
46 - HealthCheckInterval time.Duration // Interval for health checks (default: 10 seconds)
47 - ReconnectMaxRetries int // Maximum reconnection attempts (default: 0 = infinite)
48 - ReconnectInterval time.Duration // Interval between reconnection attempts (default: 5 seconds)
49 - ReverseWorkers int // Number of reverse websocket workers per listener (default: 16)
50 - ReverseDialTimeout time.Duration // Reverse websocket dial timeout (default: 5 seconds)
51 -
52 - // TLS configuration for tunnel server mode
30 + BootstrapServers []string
31 + ReverseDialTimeout time.Duration // Reverse websocket dial timeout (default: 5 seconds)
32 +
33 TLSMode TLSMode
34
55 - // Optional local certificate used in self TLS mode.
56 - TLSCertificate *tls.Certificate
35 + // Self TLS mode certificate/key file paths.
36 TLSSelfCertFile string
37 TLSSelfKeyFile string
38
60 - // Optional certificate chain and remote signer config used by keyless mode.
61 - TLSKeylessCertificatePEM []byte
62 - TLSKeyless TLSKeylessConfig
39 + // Optional keyless overrides.
40 + // If endpoint is empty, SDK uses the relay URL.
41 + TLSKeylessEndpoint string
42 + // If base domain is empty, SDK derives it from relay or signer endpoint.
43 + TLSKeylessBaseDomain string
44 }
45
46 type ClientOption func(*ClientConfig)
@@ -70,51 +51,12 @@ func WithBootstrapServers(servers []string) ClientOption {
51 }
52 }
53
73 -func WithDialer(dialer func(context.Context, string) (io.ReadWriteCloser, error)) ClientOption {
74 - return func(c *ClientConfig) {
75 - c.Dialer = dialer
76 - }
77 -}
78 -
79 -func WithHealthCheckInterval(interval time.Duration) ClientOption {
80 - return func(c *ClientConfig) {
81 - c.HealthCheckInterval = interval
82 - }
83 -}
84 -
85 -func WithReconnectMaxRetries(retries int) ClientOption {
86 - return func(c *ClientConfig) {
87 - c.ReconnectMaxRetries = retries
88 - }
89 -}
90 -
91 -func WithReconnectInterval(interval time.Duration) ClientOption {
92 - return func(c *ClientConfig) {
93 - c.ReconnectInterval = interval
94 - }
95 -}
96 -
97 -func WithReverseWorkers(workers int) ClientOption {
98 - return func(c *ClientConfig) {
99 - c.ReverseWorkers = workers
100 - }
101 -}
102 -
54 func WithReverseDialTimeout(timeout time.Duration) ClientOption {
55 return func(c *ClientConfig) {
56 c.ReverseDialTimeout = timeout
57 }
58 }
59
109 -// WithTLSSelfCertificate enables TLS with a locally managed certificate/key pair.
110 -func WithTLSSelfCertificate(cert tls.Certificate) ClientOption {
111 - return func(c *ClientConfig) {
112 - c.TLSMode = TLSModeSelf
113 - copy := cert
114 - c.TLSCertificate = &copy
115 - }
116 -}
117 -
60 // WithTLSSelfCertificateFiles enables self TLS mode using certificate/key file paths.
61 func WithTLSSelfCertificateFiles(certFile, keyFile string) ClientOption {
62 return func(c *ClientConfig) {
@@ -124,30 +66,12 @@ func WithTLSSelfCertificateFiles(certFile, keyFile string) ClientOption {
66 }
67 }
68
127 -// WithTLSKeyless enables TLS with a local certificate chain and remote keyless signer.
128 -func WithTLSKeyless(certPEM []byte, cfg TLSKeylessConfig) ClientOption {
129 - return func(c *ClientConfig) {
130 - c.TLSMode = TLSModeKeyless
131 - c.TLSKeylessCertificatePEM = append([]byte(nil), certPEM...)
132 - c.TLSKeyless = TLSKeylessConfig{
133 - Endpoint: cfg.Endpoint,
134 - ServerName: cfg.ServerName,
135 - BaseDomain: cfg.BaseDomain,
136 - KeyID: cfg.KeyID,
137 - RootCAPEM: append([]byte(nil), cfg.RootCAPEM...),
138 - EnableMTLS: cfg.EnableMTLS,
139 - ClientCertPEM: append([]byte(nil), cfg.ClientCertPEM...),
140 - ClientKeyPEM: append([]byte(nil), cfg.ClientKeyPEM...),
141 - }
142 - }
143 -}
144 -
145 -// WithTLSKeylessBaseDomain sets a global base domain override for keyless certificate hostname validation.
146 -// If unset, base domain is derived per relay URL.
147 -func WithTLSKeylessBaseDomain(baseDomain string) ClientOption {
69 +// WithTLSKeyless enables keyless TLS mode with optional signer overrides.
70 +func WithTLSKeyless(endpoint, baseDomain string) ClientOption {
71 return func(c *ClientConfig) {
72 c.TLSMode = TLSModeKeyless
150 - c.TLSKeyless.BaseDomain = baseDomain
73 + c.TLSKeylessEndpoint = endpoint
74 + c.TLSKeylessBaseDomain = baseDomain
75 }
76 }
77