portal: split api server

Kim committed Mar 11, 2026 at 10:17 UTC 6fbcda4c3cfe6e90b549dc02348faa8815960ff1
5 files changed +527 -491
portal/api_server.go new
+491
@@ -0,0 +1,491 @@
1 +package portal
2 +
3 +import (
4 + "crypto/subtle"
5 + "crypto/tls"
6 + "encoding/json"
7 + "errors"
8 + "fmt"
9 + "io"
10 + "net"
11 + "net/http"
12 + "strings"
13 + "time"
14 +
15 + "github.com/rs/zerolog/log"
16 +
17 + "github.com/gosuda/portal/v2/portal/keyless"
18 + "github.com/gosuda/portal/v2/portal/policy"
19 + "github.com/gosuda/portal/v2/types"
20 +)
21 +
22 +var (
23 + errLeaseNotFound = errors.New("lease not found")
24 + errIPBanned = errors.New("request denied because source IP is banned")
25 + errUnauthorized = errors.New(types.APIErrorCodeUnauthorized)
26 + errHostnameConflict = errors.New("hostname already registered")
27 +)
28 +
29 +func (s *Server) newAPIServer(listener net.Listener) (net.Listener, *http.Server, io.Closer, error) {
30 + apiServer := &http.Server{
31 + Handler: s.wrapAPIHandler(s.apiHandler()),
32 + ReadHeaderTimeout: 10 * time.Second,
33 + TLSNextProto: make(map[string]func(*http.Server, *tls.Conn, http.Handler)),
34 + }
35 +
36 + apiCloser, err := keyless.AttachToHTTPServer(apiServer, s.cfg.APITLS)
37 + if err != nil {
38 + return nil, nil, nil, fmt.Errorf("configure api tls: %w", err)
39 + }
40 +
41 + return tls.NewListener(listener, apiServer.TLSConfig), apiServer, apiCloser, nil
42 +}
43 +
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
57 +}
58 +
59 +func (s *Server) handleRoot(w http.ResponseWriter, _ *http.Request) {
60 + writeAPIData(w, http.StatusOK, map[string]any{
61 + "service": "portal-relay",
62 + "root": s.cfg.RootHost,
63 + })
64 +}
65 +
66 +func (s *Server) handleHealthz(w http.ResponseWriter, _ *http.Request) {
67 + writeAPIData(w, http.StatusOK, map[string]any{"status": "ok"})
68 +}
69 +
70 +func (s *Server) handleDomain(w http.ResponseWriter, r *http.Request) {
71 + if r.Method != http.MethodGet {
72 + writeAPIError(w, http.StatusMethodNotAllowed, types.APIErrorCodeMethodNotAllowed, "method not allowed")
73 + return
74 + }
75 +
76 + name := r.URL.Query().Get("name")
77 + writeAPIData(w, http.StatusOK, types.DomainResponse{
78 + RootHost: s.cfg.RootHost,
79 + SuggestedHostname: suggestHostname(name, s.cfg.RootHost),
80 + Version: types.SDKProtocolVersion,
81 + })
82 +}
83 +
84 +func (s *Server) handleRegister(w http.ResponseWriter, r *http.Request) {
85 + if r.Method != http.MethodPost {
86 + writeAPIError(w, http.StatusMethodNotAllowed, types.APIErrorCodeMethodNotAllowed, "method not allowed")
87 + return
88 + }
89 +
90 + clientIP := s.clientIPFromRequest(r)
91 + if s.isClientIPBanned(clientIP) {
92 + writeAPIError(w, http.StatusForbidden, types.APIErrorCodeIPBanned, "request denied because source IP is banned")
93 + return
94 + }
95 +
96 + var req types.RegisterRequest
97 + if err := decodeJSONBody(w, r, &req); err != nil {
98 + writeAPIError(w, http.StatusBadRequest, types.APIErrorCodeInvalidJSON, err.Error())
99 + return
100 + }
101 +
102 + resp, err := s.registerLease(req, clientIP)
103 + if err != nil {
104 + status, code := http.StatusBadRequest, types.APIErrorCodeInvalidRequest
105 + if errors.Is(err, errHostnameConflict) {
106 + status, code = http.StatusConflict, types.APIErrorCodeHostnameConflict
107 + }
108 + if errors.Is(err, errIPBanned) {
109 + status, code = http.StatusForbidden, types.APIErrorCodeIPBanned
110 + }
111 + writeAPIError(w, status, code, err.Error())
112 + return
113 + }
114 +
115 + writeAPIData(w, http.StatusCreated, resp)
116 +}
117 +
118 +func (s *Server) handleRenew(w http.ResponseWriter, r *http.Request) {
119 + if r.Method != http.MethodPost {
120 + writeAPIError(w, http.StatusMethodNotAllowed, types.APIErrorCodeMethodNotAllowed, "method not allowed")
121 + return
122 + }
123 +
124 + clientIP := s.clientIPFromRequest(r)
125 + if s.isClientIPBanned(clientIP) {
126 + writeAPIError(w, http.StatusForbidden, types.APIErrorCodeIPBanned, "request denied because source IP is banned")
127 + return
128 + }
129 +
130 + var req types.RenewRequest
131 + if err := decodeJSONBody(w, r, &req); err != nil {
132 + writeAPIError(w, http.StatusBadRequest, types.APIErrorCodeInvalidJSON, err.Error())
133 + return
134 + }
135 +
136 + resp, err := s.renewLease(req, clientIP)
137 + if err != nil {
138 + status, code := http.StatusBadRequest, types.APIErrorCodeInvalidRequest
139 + if errors.Is(err, errLeaseNotFound) {
140 + status, code = http.StatusNotFound, types.APIErrorCodeLeaseNotFound
141 + }
142 + if errors.Is(err, errUnauthorized) {
143 + status, code = http.StatusForbidden, types.APIErrorCodeUnauthorized
144 + }
145 + if errors.Is(err, errIPBanned) {
146 + status, code = http.StatusForbidden, types.APIErrorCodeIPBanned
147 + }
148 + writeAPIError(w, status, code, err.Error())
149 + return
150 + }
151 +
152 + writeAPIData(w, http.StatusOK, resp)
153 +}
154 +
155 +func (s *Server) handleUnregister(w http.ResponseWriter, r *http.Request) {
156 + if r.Method != http.MethodPost {
157 + writeAPIError(w, http.StatusMethodNotAllowed, types.APIErrorCodeMethodNotAllowed, "method not allowed")
158 + return
159 + }
160 +
161 + var req types.UnregisterRequest
162 + if err := decodeJSONBody(w, r, &req); err != nil {
163 + writeAPIError(w, http.StatusBadRequest, types.APIErrorCodeInvalidJSON, err.Error())
164 + return
165 + }
166 +
167 + if err := s.unregisterLease(req); err != nil {
168 + status, code := http.StatusBadRequest, types.APIErrorCodeInvalidRequest
169 + if errors.Is(err, errLeaseNotFound) {
170 + status, code = http.StatusNotFound, types.APIErrorCodeLeaseNotFound
171 + }
172 + if errors.Is(err, errUnauthorized) {
173 + status, code = http.StatusForbidden, types.APIErrorCodeUnauthorized
174 + }
175 + writeAPIError(w, status, code, err.Error())
176 + return
177 + }
178 +
179 + writeAPIOK(w, http.StatusOK)
180 +}
181 +
182 +func (s *Server) handleConnect(w http.ResponseWriter, r *http.Request) {
183 + if r.Method != http.MethodGet {
184 + writeAPIError(w, http.StatusMethodNotAllowed, types.APIErrorCodeMethodNotAllowed, "method not allowed")
185 + return
186 + }
187 + if r.ProtoMajor != 1 {
188 + writeAPIError(w, http.StatusHTTPVersionNotSupported, types.APIErrorCodeHTTP11Only, "reverse connect requires HTTP/1.1")
189 + return
190 + }
191 +
192 + leaseID := strings.TrimSpace(r.URL.Query().Get("lease_id"))
193 + token := strings.TrimSpace(r.Header.Get(types.HeaderReverseToken))
194 + clientIP := s.clientIPFromRequest(r)
195 + if s.isClientIPBanned(clientIP) {
196 + writeAPIError(w, http.StatusForbidden, types.APIErrorCodeIPBanned, "request denied because source IP is banned")
197 + return
198 + }
199 +
200 + lease, err := s.findLeaseByID(leaseID)
201 + if err != nil {
202 + writeAPIError(w, http.StatusNotFound, types.APIErrorCodeLeaseNotFound, err.Error())
203 + return
204 + }
205 + if !s.isLeaseRoutable(lease) {
206 + writeAPIError(w, http.StatusForbidden, types.APIErrorCodeLeaseRejected, "lease is not approved for routing")
207 + return
208 + }
209 + if authErr := s.authorizeLeaseToken(lease, token); authErr != nil {
210 + writeAPIError(w, http.StatusForbidden, types.APIErrorCodeUnauthorized, authErr.Error())
211 + return
212 + }
213 +
214 + hijacker, ok := w.(http.Hijacker)
215 + if !ok {
216 + writeAPIError(w, http.StatusInternalServerError, types.APIErrorCodeHijackUnsupported, "hijacking is not supported")
217 + return
218 + }
219 +
220 + conn, rw, err := hijacker.Hijack()
221 + if err != nil {
222 + writeAPIError(w, http.StatusInternalServerError, types.APIErrorCodeHijackFailed, err.Error())
223 + return
224 + }
225 +
226 + if _, err := rw.WriteString("HTTP/1.1 200 OK\r\nContent-Length: 0\r\nConnection: keep-alive\r\n\r\n"); err != nil {
227 + _ = conn.Close()
228 + return
229 + }
230 + if err := rw.Flush(); err != nil {
231 + _ = conn.Close()
232 + return
233 + }
234 +
235 + session := newReverseSession(conn, s.cfg.IdleKeepaliveInterval)
236 + if err := lease.Broker.Offer(session); err != nil {
237 + log.Warn().
238 + Err(err).
239 + Str("component", "relay-server").
240 + Str("lease_id", lease.ID).
241 + Str("lease_name", lease.Name).
242 + Str("remote_addr", session.RemoteAddr()).
243 + Msg("sdk reverse rejected")
244 + _ = session.Close()
245 + return
246 + }
247 +
248 + s.touchLease(lease.ID, clientIP)
249 + log.Info().
250 + Str("component", "relay-server").
251 + Str("lease_id", lease.ID).
252 + Str("lease_name", lease.Name).
253 + Str("remote_addr", session.RemoteAddr()).
254 + Int("ready", lease.Broker.ReadyCount()).
255 + Msg("sdk reverse connected")
256 +}
257 +
258 +func (s *Server) registerLease(req types.RegisterRequest, clientIP string) (types.RegisterResponse, error) {
259 + if strings.TrimSpace(req.Name) == "" {
260 + return types.RegisterResponse{}, errors.New("name is required")
261 + }
262 + if strings.TrimSpace(req.ReverseToken) == "" {
263 + return types.RegisterResponse{}, errors.New("reverse token is required")
264 + }
265 + if s.isClientIPBanned(clientIP) {
266 + return types.RegisterResponse{}, errIPBanned
267 + }
268 +
269 + hostnames := normalizeHostnames(req.Hostnames)
270 + 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 + }
281 + }
282 +
283 + ttl := s.cfg.LeaseTTL
284 + if req.TTL > 0 {
285 + ttl = time.Duration(req.TTL) * time.Second
286 + }
287 +
288 + leaseID := randomID("lease_")
289 + now := time.Now()
290 + expiresAt := now.Add(ttl)
291 + record := &leaseRecord{
292 + ID: leaseID,
293 + Name: strings.TrimSpace(req.Name),
294 + Hostnames: hostnames,
295 + Metadata: req.Metadata,
296 + ReverseToken: req.ReverseToken,
297 + ExpiresAt: expiresAt,
298 + FirstSeenAt: now,
299 + LastSeenAt: now,
300 + ClientIP: clientIP,
301 + Broker: newLeaseBroker(leaseID, s.cfg.IdleKeepaliveInterval, s.cfg.ReadyQueueLimit),
302 + }
303 +
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)
310 + }
311 +
312 + return types.RegisterResponse{
313 + LeaseID: leaseID,
314 + Hostnames: append([]string(nil), hostnames...),
315 + Metadata: record.Metadata,
316 + ExpiresAt: expiresAt,
317 + ConnectURL: s.connectURL(),
318 + }, nil
319 +}
320 +
321 +func (s *Server) renewLease(req types.RenewRequest, clientIP string) (types.RenewResponse, error) {
322 + if s.isClientIPBanned(clientIP) {
323 + return types.RenewResponse{}, errIPBanned
324 + }
325 +
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 +
337 + ttl := s.cfg.LeaseTTL
338 + if req.TTL > 0 {
339 + ttl = time.Duration(req.TTL) * time.Second
340 + }
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)
346 + }
347 +
348 + return types.RenewResponse{LeaseID: record.ID, ExpiresAt: record.ExpiresAt}, nil
349 +}
350 +
351 +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
361 + }
362 + delete(s.leases, record.ID)
363 + s.mu.Unlock()
364 +
365 + s.routes.DeleteLease(record.Hostnames)
366 + s.cfg.Policy.ForgetLease(record.ID)
367 + record.Broker.Close()
368 + return nil
369 +}
370 +
371 +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
382 +}
383 +
384 +func (s *Server) authorizeLeaseToken(record *leaseRecord, token string) error {
385 + if record == nil {
386 + return errLeaseNotFound
387 + }
388 + if !tokenMatches(record.ReverseToken, token) {
389 + return errUnauthorized
390 + }
391 + return nil
392 +}
393 +
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 +
410 +func (s *Server) runAPIServer() error {
411 + err := s.apiServer.Serve(s.apiListener)
412 + if err == nil || errors.Is(err, http.ErrServerClosed) || errors.Is(err, net.ErrClosed) {
413 + return nil
414 + }
415 + return err
416 +}
417 +
418 +func (s *Server) connectURL() string {
419 + base := strings.TrimRight(s.cfg.PortalURL, "/")
420 + if base == "" && s.apiListener != nil {
421 + return "https://" + HostPortOrLoopback(s.apiListener.Addr().String()) + types.PathSDKConnect
422 + }
423 + return base + types.PathSDKConnect
424 +}
425 +
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 +
433 +func decodeJSONBody(w http.ResponseWriter, r *http.Request, dst any) error {
434 + r.Body = http.MaxBytesReader(w, r.Body, defaultControlBodyLimit)
435 + defer r.Body.Close()
436 + return json.NewDecoder(r.Body).Decode(dst)
437 +}
438 +
439 +func normalizeHostnames(hosts []string) []string {
440 + seen := make(map[string]struct{}, len(hosts))
441 + out := make([]string, 0, len(hosts))
442 + for _, host := range hosts {
443 + host = normalizeHostname(host)
444 + if host == "" {
445 + continue
446 + }
447 + if _, ok := seen[host]; ok {
448 + continue
449 + }
450 + seen[host] = struct{}{}
451 + out = append(out, host)
452 + }
453 + return out
454 +}
455 +
456 +func tokenMatches(expected, actual string) bool {
457 + if len(expected) == 0 || len(actual) == 0 {
458 + return false
459 + }
460 + return subtle.ConstantTimeCompare([]byte(expected), []byte(actual)) == 1
461 +}
462 +
463 +func writeAPIData(w http.ResponseWriter, status int, data any) {
464 + w.Header().Set("Content-Type", "application/json")
465 + w.WriteHeader(status)
466 + _ = json.NewEncoder(w).Encode(types.APIEnvelope[any]{OK: true, Data: data})
467 +}
468 +
469 +func writeAPIOK(w http.ResponseWriter, status int) {
470 + writeAPIData(w, status, map[string]any{})
471 +}
472 +
473 +func writeAPIError(w http.ResponseWriter, status int, code, message string) {
474 + w.Header().Set("Content-Type", "application/json")
475 + w.WriteHeader(status)
476 + _ = json.NewEncoder(w).Encode(types.APIEnvelope[any]{
477 + OK: false,
478 + Error: &types.APIError{Code: code, Message: message},
479 + })
480 +}
481 +
482 +func (s *Server) clientIPFromRequest(r *http.Request) string {
483 + if r == nil {
484 + return ""
485 + }
486 + return policy.ExtractClientIP(r, s.cfg.TrustProxyHeaders)
487 +}
488 +
489 +func (s *Server) isClientIPBanned(clientIP string) bool {
490 + return s.cfg.Policy.IPFilter().IsIPBanned(clientIP)
491 +}
portal/server.go
+24 -469
@@ -2,9 +2,6 @@ package portal
2
3 import (
4 "context"
5 - "crypto/subtle"
6 - "crypto/tls"
7 - "encoding/json"
5 "errors"
6 "fmt"
7 "io"
@@ -14,16 +11,24 @@ import (
11 "sync"
12 "time"
13
17 - "github.com/rs/zerolog/log"
18 - "golang.org/x/sync/errgroup"
19 -
14 "github.com/gosuda/keyless_tls/relay/l4"
15 + "golang.org/x/sync/errgroup"
16
17 "github.com/gosuda/portal/v2/portal/keyless"
18 "github.com/gosuda/portal/v2/portal/policy"
19 "github.com/gosuda/portal/v2/types"
20 )
21
22 +const (
23 + defaultLeaseTTL = 30 * time.Second
24 + defaultClaimTimeout = 10 * time.Second
25 + defaultIdleKeepalive = 15 * time.Second
26 + defaultReadyQueueLimit = 8
27 + defaultClientHelloWait = 2 * time.Second
28 + defaultControlBodyLimit = 4 << 20
29 + defaultSessionWriteLimit = 5 * time.Second
30 +)
31 +
32 type ServerConfig struct {
33 APIHandlerWrapper func(http.Handler) http.Handler
34 KeylessSignerHandler http.Handler
@@ -47,8 +52,7 @@ type Server struct {
52 apiTLSClose io.Closer
53 apiListener net.Listener
54 apiServer *http.Server
50 - ctxDone <-chan struct{}
51 - baseContext func() context.Context
55 + ctx context.Context
56 cancel context.CancelFunc
57 group *errgroup.Group
58 routes *routeTable
@@ -144,25 +148,19 @@ func (s *Server) Start(ctx context.Context) error {
148
149 group, groupCtx := errgroup.WithContext(serverCtx)
150
147 - apiServer := &http.Server{
148 - Handler: s.wrapAPIHandler(s.apiHandler()),
149 - ReadHeaderTimeout: 10 * time.Second,
150 - TLSNextProto: make(map[string]func(*http.Server, *tls.Conn, http.Handler)),
151 - }
152 - apiCloser, err := keyless.AttachToHTTPServer(apiServer, s.cfg.APITLS)
151 + wrappedAPIListener, apiServer, apiCloser, err := s.newAPIServer(apiListener)
152 if err != nil {
153 _ = apiListener.Close()
154 _ = sniListener.Close()
155 cancel()
157 - return fmt.Errorf("configure api tls: %w", err)
156 + return err
157 }
158
160 - s.apiListener = tls.NewListener(apiListener, apiServer.TLSConfig)
159 + s.apiListener = wrappedAPIListener
160 s.sniListener = sniListener
161 s.apiServer = apiServer
162 s.apiTLSClose = apiCloser
164 - s.baseContext = func() context.Context { return groupCtx }
165 - s.ctxDone = groupCtx.Done()
163 + s.ctx = groupCtx
164 s.cancel = cancel
165 s.group = group
166
@@ -245,366 +243,6 @@ func (s *Server) ListLeases() []LeaseSnapshot {
243 return out
244 }
245
248 -func (s *Server) apiHandler() http.Handler {
249 - mux := http.NewServeMux()
250 - if s.cfg.KeylessSignerHandler != nil {
251 - mux.Handle(types.PathV1Sign, s.cfg.KeylessSignerHandler)
252 - }
253 - mux.HandleFunc(types.PathHealthz, s.handleHealthz)
254 - mux.HandleFunc(types.PathSDKDomain, s.handleDomain)
255 - mux.HandleFunc(types.PathSDKRegister, s.handleRegister)
256 - mux.HandleFunc(types.PathSDKRenew, s.handleRenew)
257 - mux.HandleFunc(types.PathSDKUnregister, s.handleUnregister)
258 - mux.HandleFunc(types.PathSDKConnect, s.handleConnect)
259 - mux.HandleFunc("/", s.handleRoot)
260 - return mux
261 -}
262 -
263 -func (s *Server) handleRoot(w http.ResponseWriter, _ *http.Request) {
264 - writeAPIData(w, http.StatusOK, map[string]any{
265 - "service": "portal-relay",
266 - "root": s.cfg.RootHost,
267 - })
268 -}
269 -
270 -func (s *Server) handleHealthz(w http.ResponseWriter, _ *http.Request) {
271 - writeAPIData(w, http.StatusOK, map[string]any{"status": "ok"})
272 -}
273 -
274 -func (s *Server) handleDomain(w http.ResponseWriter, r *http.Request) {
275 - if r.Method != http.MethodGet {
276 - writeAPIError(w, http.StatusMethodNotAllowed, types.APIErrorCodeMethodNotAllowed, "method not allowed")
277 - return
278 - }
279 - name := r.URL.Query().Get("name")
280 - writeAPIData(w, http.StatusOK, types.DomainResponse{
281 - RootHost: s.cfg.RootHost,
282 - SuggestedHostname: suggestHostname(name, s.cfg.RootHost),
283 - Version: types.SDKProtocolVersion,
284 - })
285 -}
286 -
287 -func (s *Server) handleRegister(w http.ResponseWriter, r *http.Request) {
288 - if r.Method != http.MethodPost {
289 - writeAPIError(w, http.StatusMethodNotAllowed, types.APIErrorCodeMethodNotAllowed, "method not allowed")
290 - return
291 - }
292 - clientIP := s.clientIPFromRequest(r)
293 - if s.isClientIPBanned(clientIP) {
294 - writeAPIError(w, http.StatusForbidden, types.APIErrorCodeIPBanned, "request denied because source IP is banned")
295 - return
296 - }
297 - var req types.RegisterRequest
298 - if err := decodeJSONBody(w, r, &req); err != nil {
299 - writeAPIError(w, http.StatusBadRequest, types.APIErrorCodeInvalidJSON, err.Error())
300 - return
301 - }
302 - resp, err := s.registerLease(req, clientIP)
303 - if err != nil {
304 - status, code := http.StatusBadRequest, types.APIErrorCodeInvalidRequest
305 - if errors.Is(err, errHostnameConflict) {
306 - status, code = http.StatusConflict, types.APIErrorCodeHostnameConflict
307 - }
308 - if errors.Is(err, errIPBanned) {
309 - status, code = http.StatusForbidden, types.APIErrorCodeIPBanned
310 - }
311 - writeAPIError(w, status, code, err.Error())
312 - return
313 - }
314 - writeAPIData(w, http.StatusCreated, resp)
315 -}
316 -
317 -func (s *Server) handleRenew(w http.ResponseWriter, r *http.Request) {
318 - if r.Method != http.MethodPost {
319 - writeAPIError(w, http.StatusMethodNotAllowed, types.APIErrorCodeMethodNotAllowed, "method not allowed")
320 - return
321 - }
322 - clientIP := s.clientIPFromRequest(r)
323 - if s.isClientIPBanned(clientIP) {
324 - writeAPIError(w, http.StatusForbidden, types.APIErrorCodeIPBanned, "request denied because source IP is banned")
325 - return
326 - }
327 - var req types.RenewRequest
328 - if err := decodeJSONBody(w, r, &req); err != nil {
329 - writeAPIError(w, http.StatusBadRequest, types.APIErrorCodeInvalidJSON, err.Error())
330 - return
331 - }
332 - resp, err := s.renewLease(req, clientIP)
333 - if err != nil {
334 - status, code := http.StatusBadRequest, types.APIErrorCodeInvalidRequest
335 - if errors.Is(err, errLeaseNotFound) {
336 - status, code = http.StatusNotFound, types.APIErrorCodeLeaseNotFound
337 - }
338 - if errors.Is(err, errUnauthorized) {
339 - status, code = http.StatusForbidden, types.APIErrorCodeUnauthorized
340 - }
341 - if errors.Is(err, errIPBanned) {
342 - status, code = http.StatusForbidden, types.APIErrorCodeIPBanned
343 - }
344 - writeAPIError(w, status, code, err.Error())
345 - return
346 - }
347 - writeAPIData(w, http.StatusOK, resp)
348 -}
349 -
350 -func (s *Server) handleUnregister(w http.ResponseWriter, r *http.Request) {
351 - if r.Method != http.MethodPost {
352 - writeAPIError(w, http.StatusMethodNotAllowed, types.APIErrorCodeMethodNotAllowed, "method not allowed")
353 - return
354 - }
355 - var req types.UnregisterRequest
356 - if err := decodeJSONBody(w, r, &req); err != nil {
357 - writeAPIError(w, http.StatusBadRequest, types.APIErrorCodeInvalidJSON, err.Error())
358 - return
359 - }
360 - if err := s.unregisterLease(req); err != nil {
361 - status, code := http.StatusBadRequest, types.APIErrorCodeInvalidRequest
362 - if errors.Is(err, errLeaseNotFound) {
363 - status, code = http.StatusNotFound, types.APIErrorCodeLeaseNotFound
364 - }
365 - if errors.Is(err, errUnauthorized) {
366 - status, code = http.StatusForbidden, types.APIErrorCodeUnauthorized
367 - }
368 - writeAPIError(w, status, code, err.Error())
369 - return
370 - }
371 - writeAPIOK(w, http.StatusOK)
372 -}
373 -
374 -func (s *Server) handleConnect(w http.ResponseWriter, r *http.Request) {
375 - if r.Method != http.MethodGet {
376 - writeAPIError(w, http.StatusMethodNotAllowed, types.APIErrorCodeMethodNotAllowed, "method not allowed")
377 - return
378 - }
379 - if r.ProtoMajor != 1 {
380 - writeAPIError(w, http.StatusHTTPVersionNotSupported, types.APIErrorCodeHTTP11Only, "reverse connect requires HTTP/1.1")
381 - return
382 - }
383 -
384 - leaseID := strings.TrimSpace(r.URL.Query().Get("lease_id"))
385 - token := strings.TrimSpace(r.Header.Get(types.HeaderReverseToken))
386 - clientIP := s.clientIPFromRequest(r)
387 - if s.isClientIPBanned(clientIP) {
388 - writeAPIError(w, http.StatusForbidden, types.APIErrorCodeIPBanned, "request denied because source IP is banned")
389 - return
390 - }
391 -
392 - lease, err := s.findLeaseByID(leaseID)
393 - if err != nil {
394 - writeAPIError(w, http.StatusNotFound, types.APIErrorCodeLeaseNotFound, err.Error())
395 - return
396 - }
397 - if !s.isLeaseRoutable(lease) {
398 - writeAPIError(w, http.StatusForbidden, types.APIErrorCodeLeaseRejected, "lease is not approved for routing")
399 - return
400 - }
401 - if authErr := s.authorizeLeaseToken(lease, token); authErr != nil {
402 - writeAPIError(w, http.StatusForbidden, types.APIErrorCodeUnauthorized, authErr.Error())
403 - return
404 - }
405 -
406 - hijacker, ok := w.(http.Hijacker)
407 - if !ok {
408 - writeAPIError(w, http.StatusInternalServerError, types.APIErrorCodeHijackUnsupported, "hijacking is not supported")
409 - return
410 - }
411 -
412 - conn, rw, err := hijacker.Hijack()
413 - if err != nil {
414 - writeAPIError(w, http.StatusInternalServerError, types.APIErrorCodeHijackFailed, err.Error())
415 - return
416 - }
417 -
418 - if _, err := rw.WriteString("HTTP/1.1 200 OK\r\nContent-Length: 0\r\nConnection: keep-alive\r\n\r\n"); err != nil {
419 - _ = conn.Close()
420 - return
421 - }
422 - if err := rw.Flush(); err != nil {
423 - _ = conn.Close()
424 - return
425 - }
426 -
427 - session := newReverseSession(conn, s.cfg.IdleKeepaliveInterval)
428 - if err := lease.Broker.Offer(session); err != nil {
429 - log.Warn().
430 - Err(err).
431 - Str("component", "relay-server").
432 - Str("lease_id", lease.ID).
433 - Str("lease_name", lease.Name).
434 - Str("remote_addr", session.RemoteAddr()).
435 - Msg("sdk reverse rejected")
436 - _ = session.Close()
437 - return
438 - }
439 - s.touchLease(lease.ID, clientIP)
440 - log.Info().
441 - Str("component", "relay-server").
442 - Str("lease_id", lease.ID).
443 - Str("lease_name", lease.Name).
444 - Str("remote_addr", session.RemoteAddr()).
445 - Int("ready", lease.Broker.ReadyCount()).
446 - Msg("sdk reverse connected")
447 -}
448 -
449 -func (s *Server) registerLease(req types.RegisterRequest, clientIP string) (types.RegisterResponse, error) {
450 - if strings.TrimSpace(req.Name) == "" {
451 - return types.RegisterResponse{}, errors.New("name is required")
452 - }
453 - if strings.TrimSpace(req.ReverseToken) == "" {
454 - return types.RegisterResponse{}, errors.New("reverse token is required")
455 - }
456 - if s.isClientIPBanned(clientIP) {
457 - return types.RegisterResponse{}, errIPBanned
458 - }
459 -
460 - hostnames := normalizeHostnames(req.Hostnames)
461 - if len(hostnames) == 0 {
462 - hostnames = []string{suggestHostname(req.Name, s.cfg.RootHost)}
463 - }
464 -
465 - s.mu.Lock()
466 - defer s.mu.Unlock()
467 -
468 - for _, host := range hostnames {
469 - if owner := s.findLeaseByHostnameLocked(host); owner != nil {
470 - return types.RegisterResponse{}, fmt.Errorf("%w: %s", errHostnameConflict, host)
471 - }
472 - }
473 -
474 - ttl := s.cfg.LeaseTTL
475 - if req.TTL > 0 {
476 - ttl = time.Duration(req.TTL) * time.Second
477 - }
478 -
479 - leaseID := randomID("lease_")
480 - now := time.Now()
481 - expiresAt := now.Add(ttl)
482 - record := &leaseRecord{
483 - ID: leaseID,
484 - Name: strings.TrimSpace(req.Name),
485 - Hostnames: hostnames,
486 - Metadata: req.Metadata,
487 - ReverseToken: req.ReverseToken,
488 - ExpiresAt: expiresAt,
489 - FirstSeenAt: now,
490 - LastSeenAt: now,
491 - ClientIP: clientIP,
492 - Broker: newLeaseBroker(leaseID, s.cfg.IdleKeepaliveInterval, s.cfg.ReadyQueueLimit),
493 - }
494 -
495 - s.leases[leaseID] = record
496 - for _, host := range hostnames {
497 - s.routes.Set(host, leaseID)
498 - }
499 - if strings.TrimSpace(clientIP) != "" {
500 - s.cfg.Policy.IPFilter().RegisterLeaseIP(leaseID, clientIP)
501 - }
502 -
503 - return types.RegisterResponse{
504 - LeaseID: leaseID,
505 - Hostnames: append([]string(nil), hostnames...),
506 - Metadata: record.Metadata,
507 - ExpiresAt: expiresAt,
508 - ConnectURL: s.connectURL(),
509 - }, nil
510 -}
511 -
512 -func (s *Server) renewLease(req types.RenewRequest, clientIP string) (types.RenewResponse, error) {
513 - if s.isClientIPBanned(clientIP) {
514 - return types.RenewResponse{}, errIPBanned
515 - }
516 -
517 - s.mu.Lock()
518 - defer s.mu.Unlock()
519 -
520 - record, ok := s.leases[strings.TrimSpace(req.LeaseID)]
521 - if !ok {
522 - return types.RenewResponse{}, errLeaseNotFound
523 - }
524 - if !tokenMatches(record.ReverseToken, req.ReverseToken) {
525 - return types.RenewResponse{}, errUnauthorized
526 - }
527 -
528 - ttl := s.cfg.LeaseTTL
529 - if req.TTL > 0 {
530 - ttl = time.Duration(req.TTL) * time.Second
531 - }
532 - record.ExpiresAt = time.Now().Add(ttl)
533 - record.LastSeenAt = time.Now()
534 - if strings.TrimSpace(clientIP) != "" {
535 - record.ClientIP = clientIP
536 - s.cfg.Policy.IPFilter().RegisterLeaseIP(record.ID, clientIP)
537 - }
538 - return types.RenewResponse{LeaseID: record.ID, ExpiresAt: record.ExpiresAt}, nil
539 -}
540 -
541 -func (s *Server) unregisterLease(req types.UnregisterRequest) error {
542 - s.mu.Lock()
543 - record, ok := s.leases[strings.TrimSpace(req.LeaseID)]
544 - if !ok {
545 - s.mu.Unlock()
546 - return errLeaseNotFound
547 - }
548 - if !tokenMatches(record.ReverseToken, req.ReverseToken) {
549 - s.mu.Unlock()
550 - return errUnauthorized
551 - }
552 - delete(s.leases, record.ID)
553 - s.mu.Unlock()
554 -
555 - s.routes.DeleteLease(record.Hostnames)
556 - s.cfg.Policy.ForgetLease(record.ID)
557 - record.Broker.Close()
558 - return nil
559 -}
560 -
561 -func (s *Server) findLeaseByID(leaseID string) (*leaseRecord, error) {
562 - s.mu.RLock()
563 - record, ok := s.leases[strings.TrimSpace(leaseID)]
564 - s.mu.RUnlock()
565 - if !ok {
566 - return nil, errLeaseNotFound
567 - }
568 - if time.Now().After(record.ExpiresAt) {
569 - return nil, errLeaseNotFound
570 - }
571 - return record, nil
572 -}
573 -
574 -func (s *Server) authorizeLeaseToken(record *leaseRecord, token string) error {
575 - if record == nil {
576 - return errLeaseNotFound
577 - }
578 - if !tokenMatches(record.ReverseToken, token) {
579 - return errUnauthorized
580 - }
581 - return nil
582 -}
583 -
584 -func (s *Server) findLeaseByHostnameLocked(host string) *leaseRecord {
585 - host = normalizeHostname(host)
586 - now := time.Now()
587 - for _, lease := range s.leases {
588 - if now.After(lease.ExpiresAt) {
589 - continue
590 - }
591 - for _, candidate := range lease.Hostnames {
592 - if normalizeHostname(candidate) == host {
593 - return lease
594 - }
595 - }
596 - }
597 - return nil
598 -}
599 -
600 -func (s *Server) runAPIServer() error {
601 - err := s.apiServer.Serve(s.apiListener)
602 - if err == nil || errors.Is(err, http.ErrServerClosed) || errors.Is(err, net.ErrClosed) {
603 - return nil
604 - }
605 - return err
606 -}
607 -
246 func (s *Server) runSNIListener() error {
247 for {
248 conn, err := s.sniListener.Accept()
@@ -681,9 +319,10 @@ func (s *Server) runLeaseJanitor() error {
319 ticker := time.NewTicker(5 * time.Second)
320 defer ticker.Stop()
321
322 + ctx := s.context()
323 for {
324 select {
686 - case <-s.ctxDone:
325 + case <-ctx.Done():
326 return nil
327 case <-ticker.C:
328 s.cleanupExpiredLeases()
@@ -712,86 +351,15 @@ func (s *Server) cleanupExpiredLeases() {
351 }
352
353 func (s *Server) watchContext() error {
715 - if s.ctxDone == nil {
354 + if s.ctx == nil {
355 return nil
356 }
718 - <-s.ctxDone
357 + <-s.ctx.Done()
358 shutdownCtx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
359 defer cancel()
360 return s.Shutdown(shutdownCtx)
361 }
362
724 -func (s *Server) connectURL() string {
725 - base := strings.TrimRight(s.cfg.PortalURL, "/")
726 - if base == "" && s.apiListener != nil {
727 - return "https://" + HostPortOrLoopback(s.apiListener.Addr().String()) + types.PathSDKConnect
728 - }
729 - return base + types.PathSDKConnect
730 -}
731 -
732 -func (s *Server) wrapAPIHandler(base http.Handler) http.Handler {
733 - if s.cfg.APIHandlerWrapper == nil {
734 - return base
735 - }
736 - return s.cfg.APIHandlerWrapper(base)
737 -}
738 -
739 -var (
740 - errLeaseNotFound = errors.New("lease not found")
741 - errIPBanned = errors.New("request denied because source IP is banned")
742 - errUnauthorized = errors.New(types.APIErrorCodeUnauthorized)
743 - errHostnameConflict = errors.New("hostname already registered")
744 -)
745 -
746 -func decodeJSONBody(w http.ResponseWriter, r *http.Request, dst any) error {
747 - r.Body = http.MaxBytesReader(w, r.Body, defaultControlBodyLimit)
748 - defer r.Body.Close()
749 - return json.NewDecoder(r.Body).Decode(dst)
750 -}
751 -
752 -func normalizeHostnames(hosts []string) []string {
753 - seen := make(map[string]struct{}, len(hosts))
754 - out := make([]string, 0, len(hosts))
755 - for _, host := range hosts {
756 - host = normalizeHostname(host)
757 - if host == "" {
758 - continue
759 - }
760 - if _, ok := seen[host]; ok {
761 - continue
762 - }
763 - seen[host] = struct{}{}
764 - out = append(out, host)
765 - }
766 - return out
767 -}
768 -
769 -func tokenMatches(expected, actual string) bool {
770 - if len(expected) == 0 || len(actual) == 0 {
771 - return false
772 - }
773 - return subtle.ConstantTimeCompare([]byte(expected), []byte(actual)) == 1
774 -}
775 -
776 -func writeAPIData(w http.ResponseWriter, status int, data any) {
777 - w.Header().Set("Content-Type", "application/json")
778 - w.WriteHeader(status)
779 - _ = json.NewEncoder(w).Encode(types.APIEnvelope[any]{OK: true, Data: data})
780 -}
781 -
782 -func writeAPIOK(w http.ResponseWriter, status int) {
783 - writeAPIData(w, status, map[string]any{})
784 -}
785 -
786 -func writeAPIError(w http.ResponseWriter, status int, code, message string) {
787 - w.Header().Set("Content-Type", "application/json")
788 - w.WriteHeader(status)
789 - _ = json.NewEncoder(w).Encode(types.APIEnvelope[any]{
790 - OK: false,
791 - Error: &types.APIError{Code: code, Message: message},
792 - })
793 -}
794 -
363 func bridgeConns(left, right net.Conn) {
364 defer left.Close()
365 defer right.Close()
@@ -820,20 +388,18 @@ func closeWrite(conn net.Conn) {
388 }
389
390 func (s *Server) context() context.Context {
823 - if s.baseContext != nil {
824 - if ctx := s.baseContext(); ctx != nil {
825 - return ctx
826 - }
391 + if s.ctx != nil {
392 + return s.ctx
393 }
394 return context.Background()
395 }
396
397 func (s *Server) isClosed() bool {
832 - if s.ctxDone == nil {
398 + if s.ctx == nil {
399 return false
400 }
401 select {
836 - case <-s.ctxDone:
402 + case <-s.ctx.Done():
403 return true
404 default:
405 return false
@@ -863,17 +429,6 @@ func (s *Server) snapshotForLease(record *leaseRecord) LeaseSnapshot {
429 }
430 }
431
866 -func (s *Server) clientIPFromRequest(r *http.Request) string {
867 - if r == nil {
868 - return ""
869 - }
870 - return policy.ExtractClientIP(r, s.cfg.TrustProxyHeaders)
871 -}
872 -
873 -func (s *Server) isClientIPBanned(clientIP string) bool {
874 - return s.cfg.Policy.IPFilter().IsIPBanned(clientIP)
875 -}
876 -
432 func (s *Server) isLeaseRoutable(record *leaseRecord) bool {
433 if record == nil {
434 return false
portal/utils.go
-10
@@ -9,16 +9,6 @@ import (
9 "time"
10 )
11
12 -const (
13 - defaultLeaseTTL = 30 * time.Second
14 - defaultClaimTimeout = 10 * time.Second
15 - defaultIdleKeepalive = 15 * time.Second
16 - defaultReadyQueueLimit = 8
17 - defaultClientHelloWait = 2 * time.Second
18 - defaultControlBodyLimit = 4 << 20
19 - defaultSessionWriteLimit = 5 * time.Second
20 -)
21 -
12 func PortalRootHost(portalURL string) string {
13 u, err := url.Parse(strings.TrimSpace(portalURL))
14 if err != nil || u.Host == "" {
sdk/api_client.go renamed
+10 -10
@@ -33,7 +33,7 @@ const (
33 defaultHTTPShutdownTimeout = 5 * time.Second
34 )
35
36 -type relayClient struct {
36 +type apiClient struct {
37 baseURL *url.URL
38 httpClient *http.Client
39 rawTLSConfig *tls.Config
@@ -44,7 +44,7 @@ type relayClient struct {
44 metadata types.LeaseMetadata
45 }
46
47 -func newRelayClient(ctx context.Context, relayURL string, cfg ListenerConfig) (*relayClient, error) {
47 +func newApiClient(ctx context.Context, relayURL string, cfg ListenerConfig) (*apiClient, error) {
48 name := strings.TrimSpace(cfg.Name)
49 if name == "" {
50 return nil, errors.New("listener name is required")
@@ -111,7 +111,7 @@ func newRelayClient(ctx context.Context, relayURL string, cfg ListenerConfig) (*
111 ForceAttemptHTTP2: false,
112 }
113
114 - api := &relayClient{
114 + api := &apiClient{
115 baseURL: baseURL,
116 httpClient: &http.Client{Transport: transport, Timeout: requestTimeout},
117 rawTLSConfig: baseTLS,
@@ -129,7 +129,7 @@ func newRelayClient(ctx context.Context, relayURL string, cfg ListenerConfig) (*
129 return api, nil
130 }
131
132 -func (a *relayClient) close() {
132 +func (a *apiClient) close() {
133 if a == nil || a.httpClient == nil {
134 return
135 }
@@ -138,7 +138,7 @@ func (a *relayClient) close() {
138 }
139 }
140
141 -func (a *relayClient) registerLease(ctx context.Context, hostnames []string, ttl time.Duration) (types.RegisterResponse, error) {
141 +func (a *apiClient) registerLease(ctx context.Context, hostnames []string, ttl time.Duration) (types.RegisterResponse, error) {
142 if len(hostnames) == 0 {
143 hostnames = a.hostnames
144 }
@@ -156,7 +156,7 @@ func (a *relayClient) registerLease(ctx context.Context, hostnames []string, ttl
156 return resp, nil
157 }
158
159 -func (a *relayClient) ensureCompatible(ctx context.Context) error {
159 +func (a *apiClient) ensureCompatible(ctx context.Context) error {
160 var resp types.DomainResponse
161 if err := a.doJSON(ctx, http.MethodGet, types.PathSDKDomain, nil, &resp); err != nil {
162 return fmt.Errorf("check relay compatibility: %w", err)
@@ -167,7 +167,7 @@ func (a *relayClient) ensureCompatible(ctx context.Context) error {
167 return nil
168 }
169
170 -func (a *relayClient) renewLease(ctx context.Context, leaseID string, ttl time.Duration) error {
170 +func (a *apiClient) renewLease(ctx context.Context, leaseID string, ttl time.Duration) error {
171 return a.doJSON(ctx, http.MethodPost, types.PathSDKRenew, types.RenewRequest{
172 LeaseID: leaseID,
173 ReverseToken: a.reverseToken,
@@ -175,14 +175,14 @@ func (a *relayClient) renewLease(ctx context.Context, leaseID string, ttl time.D
175 }, &types.RenewResponse{})
176 }
177
178 -func (a *relayClient) unregisterLease(ctx context.Context, leaseID string) error {
178 +func (a *apiClient) unregisterLease(ctx context.Context, leaseID string) error {
179 return a.doJSON(ctx, http.MethodPost, types.PathSDKUnregister, types.UnregisterRequest{
180 LeaseID: leaseID,
181 ReverseToken: a.reverseToken,
182 }, nil)
183 }
184
185 -func (a *relayClient) openReverseSession(ctx context.Context, leaseID string) (net.Conn, error) {
185 +func (a *apiClient) openReverseSession(ctx context.Context, leaseID string) (net.Conn, error) {
186 dialer := &tls.Dialer{
187 NetDialer: &net.Dialer{Timeout: a.dialTimeout},
188 Config: a.rawTLSConfig.Clone(),
@@ -230,7 +230,7 @@ func (a *relayClient) openReverseSession(ctx context.Context, leaseID string) (n
230 return wrapBufferedConn(conn, reader), nil
231 }
232
233 -func (a *relayClient) doJSON(ctx context.Context, method, path string, payload any, out any) error {
233 +func (a *apiClient) doJSON(ctx context.Context, method, path string, payload any, out any) error {
234 var body io.Reader
235 if payload != nil {
236 buf, err := json.Marshal(payload)
sdk/listener.go
+2 -2
@@ -43,7 +43,7 @@ type Listener struct {
43 handshakeTimeout time.Duration
44 ctx context.Context
45 cancel context.CancelFunc
46 - api *relayClient
46 + api *apiClient
47 accepted chan net.Conn
48 leaseID string
49 hostnames []string
@@ -81,7 +81,7 @@ func NewListener(ctx context.Context, relayURL string, cfg ListenerConfig) (*Lis
81 retryWait = defaultRetryWait
82 }
83
84 - api, err := newRelayClient(listenerCtx, relayURL, cfg)
84 + api, err := newApiClient(listenerCtx, relayURL, cfg)
85 if err != nil {
86 cancel()
87 return nil, err