main
go 590 lines 18.3 KB
Raw
1 package portal
2
3 import (
4 "context"
5 "crypto/tls"
6 "errors"
7 "fmt"
8 "io"
9 "net"
10 "net/http"
11 "net/url"
12 "strings"
13 "time"
14
15 "github.com/rs/zerolog/log"
16
17 "github.com/gosuda/portal-tunnel/v2/portal/auth"
18 "github.com/gosuda/portal-tunnel/v2/portal/identity"
19 "github.com/gosuda/portal-tunnel/v2/portal/keyless"
20 "github.com/gosuda/portal-tunnel/v2/portal/x402"
21 "github.com/gosuda/portal-tunnel/v2/types"
22 "github.com/gosuda/portal-tunnel/v2/utils"
23 )
24
25 type apiError struct {
26 code string
27 msg string
28 status int
29 }
30
31 func (e *apiError) Error() string { return e.msg }
32
33 var (
34 errFeatureUnavailable = &apiError{types.APIErrorCodeFeatureUnavailable, "feature unavailable", http.StatusServiceUnavailable}
35 errHostnameConflict = &apiError{types.APIErrorCodeHostnameConflict, "hostname conflict", http.StatusConflict}
36 errIPBanned = &apiError{types.APIErrorCodeIPBanned, "request denied because source IP is banned", http.StatusForbidden}
37 errLeaseNotFound = &apiError{types.APIErrorCodeLeaseNotFound, "lease not found", http.StatusNotFound}
38 errLeaseRejected = &apiError{types.APIErrorCodeLeaseRejected, "lease is not approved for routing", http.StatusForbidden}
39 errTransportMismatch = &apiError{types.APIErrorCodeTransportMismatch, "transport mismatch", http.StatusConflict}
40 errUnauthorized = &apiError{types.APIErrorCodeUnauthorized, "unauthorized", http.StatusForbidden}
41 errUDPDisabled = &apiError{types.APIErrorCodeUDPDisabled, "udp disabled", http.StatusForbidden}
42 errUDPCapacityExceeded = &apiError{types.APIErrorCodeUDPCapacityExceeded, "udp capacity exceeded", http.StatusServiceUnavailable}
43 errUDPPortExhausted = &apiError{types.APIErrorCodeUDPPortExhausted, "no udp ports available", http.StatusServiceUnavailable}
44 errTCPPortDisabled = &apiError{types.APIErrorCodeTCPPortDisabled, "tcp port disabled", http.StatusForbidden}
45 errTCPPortCapacityExceeded = &apiError{types.APIErrorCodeTCPPortCapacityExceeded, "tcp port capacity exceeded", http.StatusServiceUnavailable}
46 errTCPPortExhausted = &apiError{types.APIErrorCodeTCPPortExhausted, "no tcp ports available", http.StatusServiceUnavailable}
47 errRegisterChallengePending = &apiError{types.APIErrorCodeRateLimited, "too many pending register challenges", http.StatusTooManyRequests}
48 )
49
50 func writeAPIErrorResponse(w http.ResponseWriter, err error) {
51 var ae *apiError
52 if errors.As(err, &ae) {
53 utils.WriteAPIError(w, ae.status, ae.code, ae.msg)
54 return
55 }
56 utils.InvalidRequestError(err).Write(w)
57 }
58
59 func (s *Server) newAPIServer(listener net.Listener, apiMux *http.ServeMux, apiTLS keyless.TLSMaterialConfig) (net.Listener, *http.Server, io.Closer, error) {
60 var keylessSignerHandler http.Handler
61 if len(apiTLS.KeyPEM) > 0 {
62 signer, err := keyless.NewSigner(apiTLS.KeyPEM)
63 if err != nil {
64 return nil, nil, nil, fmt.Errorf("configure api signer: %w", err)
65 }
66 keylessSignerHandler = signer.Handler()
67 }
68
69 apiServer := &http.Server{
70 Handler: s.apiHandler(apiMux, keylessSignerHandler),
71 ReadHeaderTimeout: 10 * time.Second,
72 TLSNextProto: make(map[string]func(*http.Server, *tls.Conn, http.Handler)),
73 }
74
75 apiCloser, err := keyless.AttachToHTTPServer(apiServer, apiTLS)
76 if err != nil {
77 return nil, nil, nil, fmt.Errorf("configure api tls: %w", err)
78 }
79
80 return tls.NewListener(listener, apiServer.TLSConfig), apiServer, apiCloser, nil
81 }
82
83 func (s *Server) apiHandler(base *http.ServeMux, keylessSignerHandler http.Handler) http.Handler {
84 if base == nil {
85 base = http.NewServeMux()
86 base.HandleFunc("/{$}", s.handleRoot)
87 }
88
89 return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
90 if utils.HandleAPICORS(w, r) {
91 return
92 }
93 switch strings.TrimSpace(r.URL.Path) {
94 case types.PathHealthz:
95 s.handleHealthz(w, r)
96 case types.PathSDKDomain:
97 s.handleDomain(w, r)
98 case types.PathSDKRegisterChallenge:
99 s.handleRegisterChallenge(w, r)
100 case types.PathSDKRegister:
101 s.handleRegister(w, r)
102 case types.PathSDKRenew:
103 s.handleRenew(w, r)
104 case types.PathSDKUnregister:
105 s.handleUnregister(w, r)
106 case types.PathSDKHop:
107 s.handleHop(w, r)
108 case types.PathSDKConnect:
109 s.handleConnect(w, r)
110 case types.PathDiscovery:
111 if !s.config().DiscoveryEnabled {
112 base.ServeHTTP(w, r)
113 return
114 }
115 s.handleRelayDiscovery(w, r)
116 case types.PathDiscoveryAnnounce:
117 if !s.config().DiscoveryEnabled {
118 base.ServeHTTP(w, r)
119 return
120 }
121 s.handleRelayDiscoveryAnnounce(w, r)
122 case types.PathV1Sign:
123 if keylessSignerHandler == nil {
124 http.NotFound(w, r)
125 return
126 }
127 if err := s.registry.verifySigningAccessToken(r.Header.Get(types.HeaderAccessToken)); err != nil {
128 writeAPIErrorResponse(w, err)
129 return
130 }
131 keylessSignerHandler.ServeHTTP(w, r)
132 default:
133 base.ServeHTTP(w, r)
134 }
135 })
136 }
137
138 func (s *Server) handleRoot(w http.ResponseWriter, _ *http.Request) {
139 utils.WriteAPIData(w, http.StatusOK, map[string]any{
140 "service": "portal-relay",
141 "root": s.identity.Name,
142 })
143 }
144
145 func (s *Server) handleHealthz(w http.ResponseWriter, _ *http.Request) {
146 utils.WriteAPIData(w, http.StatusOK, map[string]any{"status": "ok"})
147 }
148
149 func (s *Server) handleRelayDiscovery(w http.ResponseWriter, r *http.Request) {
150 if !utils.RequireMethod(w, r, http.MethodGet) {
151 return
152 }
153 if s.relaySet == nil {
154 utils.WriteAPIError(w, http.StatusServiceUnavailable, types.APIErrorCodeFeatureUnavailable, "relay discovery disabled")
155 return
156 }
157
158 now := time.Now().UTC()
159 self, err := s.newSelfDescriptor(now)
160 if err != nil {
161 utils.WriteAPIError(w, http.StatusInternalServerError, types.APIErrorCodeInternal, err.Error())
162 return
163 }
164
165 utils.WriteAPIData(w, http.StatusOK, types.DiscoveryResponse{
166 ProtocolVersion: types.DiscoveryVersion,
167 GeneratedAt: now,
168 Relays: s.relaySet.Descriptors(self),
169 })
170 }
171
172 func (s *Server) handleRelayDiscoveryAnnounce(w http.ResponseWriter, r *http.Request) {
173 if !utils.RequireMethod(w, r, http.MethodPost) {
174 return
175 }
176 if s.relaySet == nil {
177 utils.WriteAPIError(w, http.StatusServiceUnavailable, types.APIErrorCodeFeatureUnavailable, "relay discovery disabled")
178 return
179 }
180 clientIP, ok := s.extractAllowedClientIP(w, r)
181 if !ok {
182 return
183 }
184 if !s.announceLimiter.Allow(clientIP) {
185 utils.WriteAPIError(w, http.StatusTooManyRequests, types.APIErrorCodeRateLimited, "announce rate limit exceeded")
186 return
187 }
188
189 req, ok := utils.DecodeJSONRequest[types.DiscoveryAnnounceRequest](w, r, defaultControlBodyLimit)
190 if !ok {
191 return
192 }
193 if req.ProtocolVersion != "" && req.ProtocolVersion != types.DiscoveryVersion {
194 utils.WriteAPIError(w, http.StatusBadRequest, types.APIErrorCodeInvalidRequest,
195 fmt.Sprintf("announce protocol mismatch: relay=%q client=%q", types.DiscoveryVersion, req.ProtocolVersion))
196 return
197 }
198
199 desc, err := identity.NormalizeRelayDescriptor(req.Descriptor)
200 if err != nil {
201 utils.WriteAPIError(w, http.StatusBadRequest, types.APIErrorCodeInvalidRequest, err.Error())
202 return
203 }
204 // Self-announce guard: the relay's own URL is established locally, not
205 // gossiped through the announce endpoint. Validate the normalized URL so
206 // scheme-less inputs are checked the same way signature verification will
207 // check them later.
208 announceURL, err := url.Parse(desc.APIHTTPSAddr)
209 if err != nil {
210 utils.WriteAPIError(w, http.StatusBadRequest, types.APIErrorCodeInvalidRequest, err.Error())
211 return
212 }
213 host := utils.NormalizeHostname(announceURL.Hostname())
214 if utils.IsLocalRelayHost(host) {
215 utils.WriteAPIError(w, http.StatusBadRequest, types.APIErrorCodeInvalidRequest,
216 fmt.Sprintf("self-announce rejected: host %q is local-only", host))
217 return
218 }
219 cfg := s.config()
220 if selfURL, err := utils.NormalizeRelayURL(cfg.PortalURL); err == nil && desc.APIHTTPSAddr == selfURL {
221 utils.WriteAPIError(w, http.StatusBadRequest, types.APIErrorCodeInvalidRequest,
222 fmt.Sprintf("self-announce rejected: %q matches receiving relay url", desc.APIHTTPSAddr))
223 return
224 }
225 if host != "" && host == utils.NormalizeHostname(s.identity.Name) {
226 utils.WriteAPIError(w, http.StatusBadRequest, types.APIErrorCodeInvalidRequest,
227 fmt.Sprintf("self-announce rejected: host %q matches receiving relay host", host))
228 return
229 }
230
231 now := time.Now().UTC()
232 if err := s.relaySet.InsertAnnounced(desc, now); err != nil {
233 utils.WriteAPIError(w, http.StatusBadRequest, types.APIErrorCodeInvalidRequest, err.Error())
234 return
235 }
236
237 log.Info().
238 Str("relay", desc.APIHTTPSAddr).
239 Str("source_ip", clientIP).
240 Msg("relay discovery announce accepted")
241
242 utils.WriteAPIData(w, http.StatusAccepted, types.DiscoveryAnnounceResponse{
243 ProtocolVersion: types.DiscoveryVersion,
244 Accepted: true,
245 })
246 }
247
248 func (s *Server) handleDomain(w http.ResponseWriter, r *http.Request) {
249 if !utils.RequireMethod(w, r, http.MethodGet) {
250 return
251 }
252 cfg := s.config()
253 x402Info := types.X402FacilitatorInfo{Enabled: cfg.X402Enabled}
254 if cfg.X402Enabled {
255 baseURL := strings.TrimRight(cfg.PortalURL, "/")
256 network := x402.Network(cfg.X402Testnet)
257 x402Info.URL = baseURL + types.PathX402Facilitator
258 x402Info.Network = network
259 x402Info.NetworkName = x402.NetworkDisplayName(network)
260 x402Info.SupportedURL = baseURL + types.X402SupportedPath
261 x402Info.PayTo = cfg.X402PayTo
262 }
263
264 utils.WriteAPIData(w, http.StatusOK, types.DomainResponse{
265 ProtocolVersion: types.SDKVersion,
266 ReleaseVersion: types.ReleaseVersion,
267 ENS: s.acmeManager.ENSStatus(),
268 X402: x402Info,
269 })
270 }
271
272 func (s *Server) handleRegister(w http.ResponseWriter, r *http.Request) {
273 if !utils.RequireMethod(w, r, http.MethodPost) {
274 return
275 }
276
277 clientIP, ok := s.extractAllowedClientIP(w, r)
278 if !ok {
279 return
280 }
281
282 req, ok := utils.DecodeJSONRequest[types.RegisterRequest](w, r, defaultControlBodyLimit)
283 if !ok {
284 return
285 }
286
287 challenge, err := s.registry.consumeVerifiedRegisterChallenge(req)
288 if err != nil {
289 switch {
290 case errors.Is(err, auth.ErrRegisterChallengeInvalidSignature):
291 utils.WriteAPIError(w, http.StatusForbidden, types.APIErrorCodeUnauthorized, err.Error())
292 default:
293 utils.InvalidRequestError(err).Write(w)
294 }
295 return
296 }
297
298 record, resp, err := s.registry.Register(challenge.Request, clientIP, req.ReportedIP)
299 if err != nil {
300 writeAPIErrorResponse(w, err)
301 return
302 }
303 dnsCtx, cancel := context.WithTimeout(context.WithoutCancel(r.Context()), defaultClaimTimeout)
304 err = record.syncENSGaslessDNS(dnsCtx, s.acmeManager)
305 cancel()
306 if err != nil {
307 removed, _ := s.registry.Unregister(types.UnregisterRequest{AccessToken: resp.AccessToken})
308 if removed == nil {
309 record.Close()
310 removed = record
311 }
312 cleanupCtx, cleanupCancel := context.WithTimeout(context.WithoutCancel(r.Context()), defaultClaimTimeout)
313 removed.deleteDNS(cleanupCtx, s.acmeManager, false)
314 cleanupCancel()
315 writeAPIErrorResponse(w, err)
316 return
317 }
318 s.registry.promoteECHDNS(record, s.acmeManager, s.config().SNIPort)
319
320 utils.WriteAPIData(w, http.StatusCreated, resp)
321 }
322
323 func (s *Server) handleRegisterChallenge(w http.ResponseWriter, r *http.Request) {
324 if !utils.RequireMethod(w, r, http.MethodPost) {
325 return
326 }
327
328 clientIP, ok := s.extractAllowedClientIP(w, r)
329 if !ok {
330 return
331 }
332
333 req, ok := utils.DecodeJSONRequest[types.RegisterChallengeRequest](w, r, defaultControlBodyLimit)
334 if !ok {
335 return
336 }
337
338 scheme := "https"
339 if r.TLS == nil {
340 scheme = "http"
341 }
342 domain := strings.TrimSpace(r.Host)
343 if domain == "" {
344 domain = s.identity.Name
345 }
346 registerURI := (&url.URL{
347 Scheme: scheme,
348 Host: domain,
349 Path: types.PathSDKRegister,
350 }).String()
351
352 if strings.TrimSpace(req.HopToken) != "" && s.overlay == nil {
353 utils.WriteAPIError(w, http.StatusServiceUnavailable, types.APIErrorCodeFeatureUnavailable, errFeatureUnavailable.Error())
354 return
355 }
356 if req.UDPEnabled && !s.supportsUDP() {
357 utils.WriteAPIError(w, http.StatusServiceUnavailable, types.APIErrorCodeFeatureUnavailable, errFeatureUnavailable.Error())
358 return
359 }
360 if req.TCPEnabled && !s.supportsTCP() {
361 utils.WriteAPIError(w, http.StatusServiceUnavailable, types.APIErrorCodeFeatureUnavailable, errFeatureUnavailable.Error())
362 return
363 }
364
365 resp, err := s.registry.issueRegisterChallenge(req, domain, registerURI, clientIP)
366 if err != nil {
367 writeAPIErrorResponse(w, err)
368 return
369 }
370
371 utils.WriteAPIData(w, http.StatusCreated, resp)
372 }
373
374 func (s *Server) handleRenew(w http.ResponseWriter, r *http.Request) {
375 if !utils.RequireMethod(w, r, http.MethodPost) {
376 return
377 }
378
379 clientIP, ok := s.extractAllowedClientIP(w, r)
380 if !ok {
381 return
382 }
383
384 req, ok := utils.DecodeJSONRequest[types.RenewRequest](w, r, defaultControlBodyLimit)
385 if !ok {
386 return
387 }
388
389 resp, err := s.registry.Renew(req, clientIP)
390 if err != nil {
391 writeAPIErrorResponse(w, err)
392 return
393 }
394
395 utils.WriteAPIData(w, http.StatusOK, resp)
396 }
397
398 func (s *Server) handleUnregister(w http.ResponseWriter, r *http.Request) {
399 if !utils.RequireMethod(w, r, http.MethodPost) {
400 return
401 }
402
403 req, ok := utils.DecodeJSONRequest[types.UnregisterRequest](w, r, defaultControlBodyLimit)
404 if !ok {
405 return
406 }
407 record, err := s.registry.Unregister(req)
408 if err != nil {
409 writeAPIErrorResponse(w, err)
410 return
411 }
412 dnsCtx, cancel := context.WithTimeout(context.WithoutCancel(r.Context()), defaultClaimTimeout)
413 record.deleteDNS(dnsCtx, s.acmeManager, true)
414 cancel()
415
416 utils.WriteAPIData(w, http.StatusOK, map[string]any{})
417 }
418
419 func (s *Server) handleHop(w http.ResponseWriter, r *http.Request) {
420 switch r.Method {
421 case http.MethodPost, http.MethodDelete:
422 default:
423 utils.MethodNotAllowedError().Write(w)
424 return
425 }
426 if s.overlay == nil || s.relaySet == nil {
427 utils.WriteAPIError(w, http.StatusServiceUnavailable, types.APIErrorCodeFeatureUnavailable, errFeatureUnavailable.Error())
428 return
429 }
430 if _, ok := s.extractAllowedClientIP(w, r); !ok {
431 return
432 }
433
434 route, ok := utils.DecodeJSONRequest[types.HopRoute](w, r, defaultControlBodyLimit)
435 if !ok {
436 return
437 }
438 route, err := auth.VerifyHopRoute(r.Method, route)
439 if errors.Is(err, auth.ErrHopRouteSignatureInvalid) {
440 utils.WriteAPIError(w, http.StatusForbidden, types.APIErrorCodeUnauthorized, "hop route signature is invalid")
441 return
442 }
443 if err != nil {
444 utils.InvalidRequestError(err).Write(w)
445 return
446 }
447 cfg := s.config()
448 if route.RelayURL != cfg.PortalURL {
449 utils.WriteAPIError(w, http.StatusForbidden, types.APIErrorCodeUnauthorized, "hop route relay url does not match receiving relay")
450 return
451 }
452 if r.Method == http.MethodDelete {
453 record := s.registry.DeleteHopRoute(&route)
454 dnsCtx, cancel := context.WithTimeout(context.WithoutCancel(r.Context()), defaultClaimTimeout)
455 record.deleteDNS(dnsCtx, s.acmeManager, true)
456 cancel()
457 utils.WriteAPIData(w, http.StatusOK, map[string]any{})
458 return
459 }
460
461 now := time.Now().UTC()
462 if !route.ExpiresAt.UTC().After(now) {
463 utils.InvalidRequestError(errors.New("route expiry must be in the future")).Write(w)
464 return
465 }
466 forwardRelay, err := auth.VerifyRelayDescriptor(route.ForwardRelay)
467 if err != nil {
468 utils.InvalidRequestError(fmt.Errorf("forward relay: %w", err)).Write(w)
469 return
470 }
471 if !forwardRelay.HasOverlayPeer() {
472 utils.InvalidRequestError(errors.New("forward relay wireguard overlay metadata is required")).Write(w)
473 return
474 }
475 route.ForwardRelay = forwardRelay
476 if err := s.relaySet.InsertAnnounced(forwardRelay, now); err != nil {
477 utils.InvalidRequestError(fmt.Errorf("forward relay: %w", err)).Write(w)
478 return
479 }
480 if err := s.overlay.Sync(s.relaySet.OverlayPeerDescriptor()); err != nil {
481 utils.WriteAPIError(w, http.StatusInternalServerError, types.APIErrorCodeInternal, err.Error())
482 return
483 }
484 record, err := s.registry.RegisterHopRoute(&route, now)
485 if err != nil {
486 writeAPIErrorResponse(w, err)
487 return
488 }
489 dnsCtx, cancel := context.WithTimeout(context.WithoutCancel(r.Context()), defaultClaimTimeout)
490 err = record.syncENSGaslessDNS(dnsCtx, s.acmeManager)
491 cancel()
492 if err != nil {
493 removed := s.registry.DeleteHopRoute(&route)
494 if removed == nil {
495 removed = record
496 }
497 cleanupCtx, cleanupCancel := context.WithTimeout(context.WithoutCancel(r.Context()), defaultClaimTimeout)
498 removed.deleteDNS(cleanupCtx, s.acmeManager, false)
499 cleanupCancel()
500 writeAPIErrorResponse(w, err)
501 return
502 }
503 s.registry.promoteECHDNS(record, s.acmeManager, cfg.SNIPort)
504 var accessToken string
505 if record.isPublicEntry() {
506 accessToken, err = s.registry.issueLeaseAccessToken(record, now)
507 if err != nil {
508 writeAPIErrorResponse(w, &apiError{types.APIErrorCodeInternal, err.Error(), http.StatusInternalServerError})
509 return
510 }
511 }
512 utils.WriteAPIData(w, http.StatusOK, types.HopRouteResponse{
513 AccessToken: accessToken,
514 SNIPort: cfg.SNIPort,
515 })
516 }
517
518 func (s *Server) handleConnect(w http.ResponseWriter, r *http.Request) {
519 if !utils.RequireMethod(w, r, http.MethodGet) {
520 return
521 }
522 if r.ProtoMajor != 1 {
523 utils.WriteAPIError(w, http.StatusHTTPVersionNotSupported, types.APIErrorCodeHTTP11Only, "reverse connect requires HTTP/1.1")
524 return
525 }
526
527 token := strings.TrimSpace(r.Header.Get(types.HeaderAccessToken))
528 clientIP, ok := s.extractAllowedClientIP(w, r)
529 if !ok {
530 return
531 }
532
533 lease, err := s.registry.admitLeaseByToken(token, false)
534 if err != nil {
535 writeAPIErrorResponse(w, err)
536 return
537 }
538
539 hijacker, ok := w.(http.Hijacker)
540 if !ok {
541 utils.WriteAPIError(w, http.StatusInternalServerError, types.APIErrorCodeHijackUnsupported, "hijacking is not supported")
542 return
543 }
544
545 conn, rw, err := hijacker.Hijack()
546 if err != nil {
547 utils.WriteAPIError(w, http.StatusInternalServerError, types.APIErrorCodeHijackFailed, err.Error())
548 return
549 }
550
551 if _, err := rw.WriteString("HTTP/1.1 101 Switching Protocols\r\nUpgrade: raw\r\nConnection: Upgrade\r\n\r\n"); err != nil {
552 _ = conn.Close()
553 return
554 }
555 if err := rw.Flush(); err != nil {
556 _ = conn.Close()
557 return
558 }
559
560 remoteAddr := ""
561 if conn.RemoteAddr() != nil {
562 remoteAddr = conn.RemoteAddr().String()
563 }
564 if err := lease.stream.OfferConn(conn); err != nil {
565 log.Warn().
566 Err(err).
567 Str("address", lease.Address).
568 Str("lease_name", lease.Name).
569 Str("remote_addr", remoteAddr).
570 Msg("sdk reverse rejected")
571 return
572 }
573
574 s.registry.Touch(lease.Key(), clientIP, time.Now())
575 log.Info().
576 Str("address", lease.Address).
577 Str("lease_name", lease.Name).
578 Str("remote_addr", remoteAddr).
579 Int("ready", lease.stream.ReadyCount()).
580 Msg("sdk reverse connected")
581 }
582
583 func (s *Server) extractAllowedClientIP(w http.ResponseWriter, r *http.Request) (string, bool) {
584 clientIP := s.registry.policy.ExtractClientIP(r)
585 if !s.registry.policy.IPFilter().IsIPBanned(clientIP) {
586 return clientIP, true
587 }
588 utils.WriteAPIError(w, http.StatusForbidden, types.APIErrorCodeIPBanned, "request denied because source IP is banned")
589 return "", false
590 }