Skip to content

Commit b811d09

Browse files
committed
tidy codes
1 parent 1432292 commit b811d09

3 files changed

Lines changed: 261 additions & 282 deletions

File tree

portal/api_server.go

Lines changed: 47 additions & 66 deletions
Original file line numberDiff line numberDiff line change
@@ -240,7 +240,11 @@ func (s *Server) handleRenew(w http.ResponseWriter, r *http.Request) {
240240
return
241241
}
242242

243-
resp, err := s.renewLease(req, clientIP)
243+
ttl := s.cfg.LeaseTTL
244+
if req.TTL > 0 {
245+
ttl = time.Duration(req.TTL) * time.Second
246+
}
247+
record, err := s.registry.Renew(strings.TrimSpace(req.LeaseID), req.ReverseToken, ttl, clientIP, utils.SanitizeReportedIP(req.ReportedIP))
244248
if err != nil {
245249
status, code := http.StatusBadRequest, types.APIErrorCodeInvalidRequest
246250
if errors.Is(err, errLeaseNotFound) {
@@ -256,7 +260,7 @@ func (s *Server) handleRenew(w http.ResponseWriter, r *http.Request) {
256260
return
257261
}
258262

259-
utils.WriteAPIData(w, http.StatusOK, resp)
263+
utils.WriteAPIData(w, http.StatusOK, types.RenewResponse{LeaseID: record.ID, ExpiresAt: record.ExpiresAt})
260264
}
261265

262266
func (s *Server) handleUnregister(w http.ResponseWriter, r *http.Request) {
@@ -271,7 +275,8 @@ func (s *Server) handleUnregister(w http.ResponseWriter, r *http.Request) {
271275
return
272276
}
273277

274-
if err := s.unregisterLease(req); err != nil {
278+
record, err := s.registry.Unregister(strings.TrimSpace(req.LeaseID), req.ReverseToken)
279+
if err != nil {
275280
status, code := http.StatusBadRequest, types.APIErrorCodeInvalidRequest
276281
if errors.Is(err, errLeaseNotFound) {
277282
status, code = http.StatusNotFound, types.APIErrorCodeLeaseNotFound
@@ -282,6 +287,9 @@ func (s *Server) handleUnregister(w http.ResponseWriter, r *http.Request) {
282287
utils.WriteAPIError(w, status, code, err.Error())
283288
return
284289
}
290+
if record != nil {
291+
record.Close()
292+
}
285293

286294
utils.WriteAPIOK(w, http.StatusOK)
287295
}
@@ -304,30 +312,22 @@ func (s *Server) handleConnect(w http.ResponseWriter, r *http.Request) {
304312
return
305313
}
306314

307-
lease, err := s.registry.FindByID(leaseID)
308-
if err == nil && !s.registry.policy.IsLeaseRoutable(lease.ID) {
309-
err = errLeaseRejected
310-
}
311-
if err == nil && !utils.TokenMatches(lease.ReverseToken, token) {
312-
err = errUnauthorized
313-
}
314-
if err == nil && lease.stream == nil {
315-
err = errTransportMismatch
316-
}
317-
switch {
318-
case errors.Is(err, errLeaseNotFound):
315+
lease, err := s.admitLeaseByID(leaseID, token, false)
316+
switch err {
317+
case nil:
318+
case errLeaseNotFound:
319319
utils.WriteAPIError(w, http.StatusNotFound, types.APIErrorCodeLeaseNotFound, err.Error())
320320
return
321-
case errors.Is(err, errLeaseRejected):
321+
case errLeaseRejected:
322322
utils.WriteAPIError(w, http.StatusForbidden, types.APIErrorCodeLeaseRejected, "lease is not approved for routing")
323323
return
324-
case errors.Is(err, errUnauthorized):
324+
case errUnauthorized:
325325
utils.WriteAPIError(w, http.StatusForbidden, types.APIErrorCodeUnauthorized, err.Error())
326326
return
327-
case errors.Is(err, errTransportMismatch):
327+
case errTransportMismatch:
328328
utils.WriteAPIError(w, http.StatusConflict, types.APIErrorCodeTransportMismatch, "lease does not support stream transport")
329329
return
330-
case err != nil:
330+
default:
331331
utils.WriteAPIError(w, http.StatusBadRequest, types.APIErrorCodeInvalidRequest, err.Error())
332332
return
333333
}
@@ -396,34 +396,26 @@ func (s *Server) handleQUICTunnelConn(conn *quic.Conn) {
396396
return
397397
}
398398

399-
lease, err := s.registry.FindByID(msg.LeaseID)
400-
if err == nil && !s.registry.policy.IsLeaseRoutable(lease.ID) {
401-
err = errLeaseRejected
402-
}
403-
if err == nil && !utils.TokenMatches(lease.ReverseToken, msg.ReverseToken) {
404-
err = errUnauthorized
405-
}
406-
if err == nil && (lease.stream == nil || lease.datagram == nil) {
407-
err = errTransportMismatch
408-
}
409-
switch {
410-
case errors.Is(err, errLeaseNotFound):
399+
lease, err := s.admitLeaseByID(msg.LeaseID, msg.ReverseToken, true)
400+
switch err {
401+
case nil:
402+
case errLeaseNotFound:
411403
_ = json.NewEncoder(stream).Encode(types.QUICControlResponse{OK: false, Error: types.APIErrorCodeLeaseNotFound})
412404
_ = conn.CloseWithError(1, "lease not found")
413405
return
414-
case errors.Is(err, errUnauthorized):
415-
_ = json.NewEncoder(stream).Encode(types.QUICControlResponse{OK: false, Error: types.APIErrorCodeUnauthorized})
416-
_ = conn.CloseWithError(1, "unauthorized")
417-
return
418-
case errors.Is(err, errLeaseRejected):
406+
case errLeaseRejected:
419407
_ = json.NewEncoder(stream).Encode(types.QUICControlResponse{OK: false, Error: types.APIErrorCodeLeaseRejected})
420408
_ = conn.CloseWithError(1, "lease rejected")
421409
return
422-
case errors.Is(err, errTransportMismatch):
410+
case errUnauthorized:
411+
_ = json.NewEncoder(stream).Encode(types.QUICControlResponse{OK: false, Error: types.APIErrorCodeUnauthorized})
412+
_ = conn.CloseWithError(1, "unauthorized")
413+
return
414+
case errTransportMismatch:
423415
_ = json.NewEncoder(stream).Encode(types.QUICControlResponse{OK: false, Error: types.APIErrorCodeTransportMismatch})
424416
_ = conn.CloseWithError(1, "transport mismatch")
425417
return
426-
case err != nil:
418+
default:
427419
_ = json.NewEncoder(stream).Encode(types.QUICControlResponse{OK: false, Error: types.APIErrorCodeInvalidRequest})
428420
_ = conn.CloseWithError(1, "invalid control message")
429421
return
@@ -445,6 +437,23 @@ func (s *Server) handleQUICTunnelConn(conn *quic.Conn) {
445437
Msg("quic tunnel connected")
446438
}
447439

440+
func (s *Server) admitLeaseByID(leaseID, token string, requireDatagram bool) (*leaseRecord, error) {
441+
lease, err := s.registry.FindByID(leaseID)
442+
if err != nil {
443+
return nil, err
444+
}
445+
if !s.registry.policy.IsLeaseRoutable(lease.ID) {
446+
return nil, errLeaseRejected
447+
}
448+
if !utils.TokenMatches(lease.ReverseToken, token) {
449+
return nil, errUnauthorized
450+
}
451+
if lease.stream == nil || (requireDatagram && lease.datagram == nil) {
452+
return nil, errTransportMismatch
453+
}
454+
return lease, nil
455+
}
456+
448457
func (s *Server) registerLease(req types.RegisterRequest, clientIP string) (types.RegisterResponse, error) {
449458
name, err := utils.NormalizeDNSLabel(req.Name)
450459
if err != nil {
@@ -541,34 +550,6 @@ func (s *Server) registerLease(req types.RegisterRequest, clientIP string) (type
541550
return resp, nil
542551
}
543552

544-
func (s *Server) renewLease(req types.RenewRequest, clientIP string) (types.RenewResponse, error) {
545-
if s.registry.policy.IPFilter().IsIPBanned(clientIP) {
546-
return types.RenewResponse{}, errIPBanned
547-
}
548-
549-
ttl := s.cfg.LeaseTTL
550-
if req.TTL > 0 {
551-
ttl = time.Duration(req.TTL) * time.Second
552-
}
553-
record, err := s.registry.Renew(strings.TrimSpace(req.LeaseID), req.ReverseToken, ttl, clientIP, utils.SanitizeReportedIP(req.ReportedIP))
554-
if err != nil {
555-
return types.RenewResponse{}, err
556-
}
557-
558-
return types.RenewResponse{LeaseID: record.ID, ExpiresAt: record.ExpiresAt}, nil
559-
}
560-
561-
func (s *Server) unregisterLease(req types.UnregisterRequest) error {
562-
record, err := s.registry.Unregister(strings.TrimSpace(req.LeaseID), req.ReverseToken)
563-
if err != nil {
564-
return err
565-
}
566-
if record != nil {
567-
record.Close()
568-
}
569-
return nil
570-
}
571-
572553
func (s *Server) runAPIServer() error {
573554
err := s.apiServer.Serve(s.apiListener)
574555
if err == nil || errors.Is(err, http.ErrServerClosed) || errors.Is(err, net.ErrClosed) {

portal/server.go

Lines changed: 25 additions & 27 deletions
Original file line numberDiff line numberDiff line change
@@ -158,7 +158,6 @@ func NewServer(cfg ServerConfig) (*Server, error) {
158158
Str("owner_private_key", ownerIdentity.PrivateKey).
159159
Msg("generated relay owner private key; set OWNER_PRIVATE_KEY unique identity")
160160
}
161-
cfg.OwnerPrivateKey = ""
162161

163162
runtime := policy.NewRuntime()
164163
runtime.SetUDPPolicy(cfg.UDPPortCount > 0, 0)
@@ -491,32 +490,6 @@ func (s *Server) runSNIListener(ctx context.Context) error {
491490
}
492491
}
493492

494-
func (s *Server) startOverlay() error {
495-
peerMux := http.NewServeMux()
496-
peerMux.HandleFunc(types.PathRoot, s.handleRoot)
497-
peerMux.HandleFunc(types.PathHealthz, s.handleHealthz)
498-
peerMux.HandleFunc(types.PathDiscovery, func(w http.ResponseWriter, r *http.Request) {
499-
if !s.DiscoveryEnabled() {
500-
http.NotFound(w, r)
501-
return
502-
}
503-
s.handleRelayDiscovery(w, r)
504-
})
505-
506-
overlay, err := wireguard.NewOverlay(s.wgConfig, peerMux)
507-
if err != nil {
508-
return fmt.Errorf("start wireguard overlay: %w", err)
509-
}
510-
511-
if err := overlay.Sync(s.cfg.PortalURL, s.relaySet.Snapshot()); err != nil {
512-
_ = overlay.Shutdown(context.Background())
513-
return fmt.Errorf("sync wireguard peers: %w", err)
514-
}
515-
516-
s.overlay = overlay
517-
return nil
518-
}
519-
520493
func (s *Server) startQUICTunnelListener(apiTLS keyless.TLSMaterialConfig) error {
521494
if len(apiTLS.KeyPEM) == 0 {
522495
return fmt.Errorf("quic tunnel requires api tls key")
@@ -565,6 +538,31 @@ func (s *Server) runQUICTunnelListener(listener *quic.Listener) error {
565538
}
566539
}
567540

541+
func (s *Server) startOverlay() error {
542+
peerMux := http.NewServeMux()
543+
peerMux.HandleFunc(types.PathRoot, s.handleRoot)
544+
peerMux.HandleFunc(types.PathHealthz, s.handleHealthz)
545+
peerMux.HandleFunc(types.PathDiscovery, func(w http.ResponseWriter, r *http.Request) {
546+
if !s.DiscoveryEnabled() {
547+
http.NotFound(w, r)
548+
return
549+
}
550+
s.handleRelayDiscovery(w, r)
551+
})
552+
553+
overlay, err := wireguard.NewOverlay(s.wgConfig, peerMux)
554+
if err != nil {
555+
return fmt.Errorf("start wireguard overlay: %w", err)
556+
}
557+
558+
if err := overlay.Sync(s.cfg.PortalURL, s.relaySet.Snapshot()); err != nil {
559+
_ = overlay.Shutdown(context.Background())
560+
return fmt.Errorf("sync wireguard peers: %w", err)
561+
}
562+
563+
s.overlay = overlay
564+
return nil
565+
}
568566
func (s *Server) runRelayDiscoveryLoop(ctx context.Context) error {
569567
ticker := time.NewTicker(defaultDiscoveryInterval)
570568
defer ticker.Stop()

0 commit comments

Comments
 (0)