main
go 127 lines 3.01 KB
Raw
1 package agent
2
3 import (
4 "context"
5 "errors"
6 "fmt"
7 "net"
8 "net/http"
9 "os"
10 "path/filepath"
11 "strings"
12 "time"
13
14 "github.com/rs/zerolog/log"
15
16 "github.com/gosuda/portal-tunnel/v2/portal/auth"
17 "github.com/gosuda/portal-tunnel/v2/utils"
18 )
19
20 func Run(ctx context.Context, cfg Config) error {
21 endpointStateDir := strings.TrimSpace(cfg.Agent.StateDir)
22 if endpointStateDir == "" {
23 return errors.New("agent.state_dir is required")
24 }
25 if err := os.MkdirAll(endpointStateDir, 0o700); err != nil {
26 return err
27 }
28 token := utils.RandomID("agent_")
29
30 runtimeCtx, cancel := context.WithCancel(ctx)
31 defer cancel()
32
33 manager := newManager(cfg, "")
34 walletAuth, err := auth.NewWalletAuthenticator(auth.WalletAuthConfig{
35 AllowedAddresses: cfg.Agent.AllowedWallets,
36 AllowAnyAddress: len(cfg.Agent.AllowedWallets) == 0,
37 Statement: "Sign in to Portal agent",
38 })
39 if err != nil {
40 return err
41 }
42 controlAddr := strings.TrimSpace(cfg.Agent.ControlAddr)
43 if controlAddr == "" {
44 return errors.New("control address is required")
45 }
46 host, _, err := net.SplitHostPort(controlAddr)
47 if err != nil {
48 return fmt.Errorf("control address must be host:port: %w", err)
49 }
50 host = strings.Trim(host, "[]")
51 if host == "" {
52 return errors.New("control address must include a loopback host")
53 }
54 if !strings.EqualFold(host, "localhost") {
55 ip := net.ParseIP(host)
56 if ip == nil || !ip.IsLoopback() {
57 return fmt.Errorf("control address must bind to loopback, got %q", host)
58 }
59 }
60 var listenConfig net.ListenConfig
61 listener, err := listenConfig.Listen(runtimeCtx, "tcp", controlAddr)
62 if err != nil {
63 return err
64 }
65 control := &http.Server{
66 Handler: &controlHandler{
67 manager: manager,
68 token: token,
69 auth: walletAuth,
70 shutdown: cancel,
71 },
72 ReadHeaderTimeout: 5 * time.Second,
73 }
74 listenAddr := listener.Addr().String()
75 manager.controlAddr = listenAddr
76
77 if err := utils.WriteJSONFile(filepath.Join(endpointStateDir, endpointFilename), endpoint{
78 ControlAddr: listenAddr,
79 Token: token,
80 }, 0o600); err != nil {
81 _ = listener.Close()
82 _ = control.Shutdown(context.Background())
83 return err
84 }
85 defer func() {
86 _ = os.Remove(filepath.Join(endpointStateDir, endpointFilename))
87 }()
88
89 manager.Start(runtimeCtx)
90
91 errCh := make(chan error, 1)
92 go func() {
93 err := control.Serve(listener)
94 if errors.Is(err, http.ErrServerClosed) || errors.Is(err, net.ErrClosed) {
95 err = nil
96 }
97 errCh <- err
98 }()
99
100 log.Info().
101 Str("control_addr", listenAddr).
102 Int("tunnel_count", len(cfg.Tunnels)).
103 Msg("portal agent started")
104
105 var serveErr error
106 select {
107 case <-runtimeCtx.Done():
108 case serveErr = <-errCh:
109 cancel()
110 }
111
112 shutdownCtx, shutdownCancel := context.WithTimeout(context.Background(), 15*time.Second)
113 defer shutdownCancel()
114 stopErr := manager.Stop(shutdownCtx)
115 closeErr := control.Shutdown(shutdownCtx)
116 if serveErr == nil {
117 select {
118 case serveErr = <-errCh:
119 default:
120 }
121 }
122 if errors.Is(serveErr, context.Canceled) {
123 serveErr = nil
124 }
125 log.Info().Msg("portal agent stopped")
126 return errors.Join(serveErr, stopErr, closeErr)
127 }