refact: rollback multi relay

rabbitprincess committed Mar 7, 2026 at 21:41 UTC c3c29bc7fcd2e572e80ffc1c85fa6b4ee56b795c
9 files changed +418 -889
cmd/demo-app/handler.go
+7 -5
@@ -45,12 +45,14 @@ func handleWebSocket(conn *websocket.Conn) {
45 }
46 }
47
48 -func handleCookies(w http.ResponseWriter, _ *http.Request) {
48 +func handleCookies(w http.ResponseWriter, r *http.Request) {
49 + secure := r != nil && r.TLS != nil
50 +
51 for _, cookie := range []*http.Cookie{
50 - {Name: "session_id", Value: "abc123", Path: "/", MaxAge: 3600},
51 - {Name: "auth_token", Value: "secret456", Path: "/", MaxAge: 3600},
52 - {Name: "csrf_token", Value: "xyz789", Path: "/", MaxAge: 3600},
53 - {Name: "user_pref", Value: "dark_mode", Path: "/", MaxAge: 86400},
52 + {Name: "session_id", Value: "abc123", Path: "/", MaxAge: 3600, Secure: secure},
53 + {Name: "auth_token", Value: "secret456", Path: "/", MaxAge: 3600, Secure: secure},
54 + {Name: "csrf_token", Value: "xyz789", Path: "/", MaxAge: 3600, Secure: secure},
55 + {Name: "user_pref", Value: "dark_mode", Path: "/", MaxAge: 86400, Secure: secure},
56 } {
57 http.SetCookie(w, cookie)
58 }
cmd/demo-app/main.go
+8 -3
@@ -31,7 +31,7 @@ func main() {
31 log.Logger = log.Output(zerolog.ConsoleWriter{Out: os.Stdout, TimeFormat: time.RFC3339})
32 logger := log.With().Str("component", "demo-app").Logger()
33
34 - flag.StringVar(&flagServerURL, "server-url", "https://localhost:4017", "relay API URLs (comma-separated, https only)")
34 + flag.StringVar(&flagServerURL, "server-url", "https://localhost:4017", "relay API URL (https only)")
35 flag.StringVar(&flagAddr, "addr", "127.0.0.1:8092", "local demo HTTP listen address (disable if empty)")
36 flag.StringVar(&flagName, "name", "demo-app", "backend display name")
37 flag.StringVar(&flagDesc, "description", "Portal demo connectivity app", "lease description")
@@ -53,7 +53,7 @@ func runDemo() error {
53 ctx, stop := signal.NotifyContext(context.Background(), os.Interrupt, syscall.SIGTERM)
54 defer stop()
55
56 - sdkClient, err := sdk.NewClient(sdk.ClientConfig{RelayURLs: sdk.SplitCSV(flagServerURL)})
56 + sdkClient, err := sdk.NewClient(sdk.ClientConfig{RelayURL: flagServerURL})
57 if err != nil {
58 return fmt.Errorf("new client: %w", err)
59 }
@@ -74,7 +74,12 @@ func runDemo() error {
74 }
75 defer listener.Close()
76
77 - logger.Info().Strs("public_urls", listener.PublicURLs()).Str("local_addr", flagAddr).Msg("demo app registered with relay")
77 + logger.Info().
78 + Str("relay", flagServerURL).
79 + Str("lease_id", listener.LeaseID()).
80 + Strs("public_urls", listener.PublicURLs()).
81 + Str("local_addr", flagAddr).
82 + Msg("demo app registered with relay")
83 if err := sdk.RunHTTP(ctx, listener, newHandler(), flagAddr); err != nil {
84 return err
85 }
cmd/portal-tunnel/main.go
+29 -46
@@ -7,13 +7,13 @@ import (
7 "fmt"
8 "os"
9 "os/signal"
10 + "sync"
11 "sync/atomic"
12 "syscall"
13 "time"
14
15 "github.com/rs/zerolog"
16 "github.com/rs/zerolog/log"
16 - "golang.org/x/sync/errgroup"
17
18 "github.com/gosuda/portal/v2/sdk"
19 "github.com/gosuda/portal/v2/types"
@@ -62,18 +62,16 @@ func runTunnel() error {
62 ctx, stop := signal.NotifyContext(context.Background(), os.Interrupt, syscall.SIGTERM)
63 defer stop()
64
65 - relayURLs := sdk.SplitCSV(flagRelayURLs)
65 + relayURLs, err := normalizeRelayURLs(flagRelayURLs)
66 + if err != nil {
67 + return err
68 + }
69 if len(relayURLs) == 0 {
70 return errors.New("no relay URLs provided")
71 }
72
70 - targetAddr, err := normalizeTargetAddr(flagHost)
71 - if err != nil {
72 - return fmt.Errorf("invalid --host value %q: %w", flagHost, err)
73 - }
74 -
73 logger.Info().
76 - Str("local", targetAddr).
74 + Str("local", flagHost).
75 Int("relay_count", len(relayURLs)).
76 Strs("relays", relayURLs).
77 Msg("starting portal tunnel")
@@ -89,67 +87,52 @@ func runTunnel() error {
87 },
88 }
89
92 - client, err := sdk.NewClient(sdk.ClientConfig{RelayURLs: relayURLs})
90 + runtimes, err := startRelayRuntimes(ctx, relayURLs, listenReq)
91 if err != nil {
94 - return fmt.Errorf("service %s: failed to create client: %w", flagName, err)
92 + return fmt.Errorf("service %s: failed to start relays: %w", flagName, err)
93 }
96 - defer client.Close()
94
98 - listener, err := client.Listen(ctx, listenReq)
99 - if err != nil {
100 - return fmt.Errorf("service %s: failed to start listener: %w", flagName, err)
101 - }
102 - defer listener.Close()
95 + var connWG sync.WaitGroup
96 + var connCount atomic.Int64
97 + relayDone := make(chan relayLoopResult, len(runtimes))
98
104 - for _, entry := range listener.Entries() {
99 + for _, runtime := range runtimes {
100 logger.Info().
106 - Str("relay", entry.RelayURL).
107 - Str("lease_id", entry.LeaseID).
108 - Strs("public_urls", entry.PublicURLs()).
101 + Str("relay", runtime.relayURL).
102 + Str("lease_id", runtime.listener.LeaseID()).
103 + Strs("public_urls", runtime.listener.PublicURLs()).
104 Msg("relay tunnel ready")
105 + go runtime.run(ctx, flagHost, &connWG, &connCount, relayDone)
106 }
107
112 - var connGroup errgroup.Group
113 - var connCount atomic.Int64
114 - group, groupCtx := errgroup.WithContext(ctx)
115 - group.Go(func() error {
116 - if err := runProxyLoop(groupCtx, listener, targetAddr, &connGroup, &connCount); err != nil {
117 - return fmt.Errorf("relay accept loop: %w", err)
118 - }
119 - return nil
120 - })
121 - group.Go(func() error {
122 - <-groupCtx.Done()
123 - if err := listener.Close(); err != nil {
124 - return fmt.Errorf("listener close: %w", err)
125 - }
126 - return nil
127 - })
128 -
129 - waitErr := group.Wait()
108 + waitErr := waitForRelayLoops(ctx, relayDone, len(runtimes))
109 + if waitErr != nil {
110 + stop()
111 + }
112 + closeErr := closeRelayRuntimes(runtimes)
113 if waitErr != nil {
114 logger.Error().Err(waitErr).Msg("relay supervisor exited with error")
115 }
116 + if closeErr != nil {
117 + logger.Error().Err(closeErr).Msg("relay shutdown failed")
118 + }
119
120 if ctx.Err() != nil {
121 logger.Info().Msg("tunnel shutting down")
122 }
123
138 - done := make(chan error, 1)
124 + done := make(chan struct{})
125 go func() {
140 - done <- connGroup.Wait()
126 + connWG.Wait()
127 + close(done)
128 }()
129
130 select {
144 - case err := <-done:
145 - if err != nil {
146 - logger.Error().Err(err).Msg("proxy connection group failed")
147 - waitErr = errors.Join(waitErr, err)
148 - }
131 + case <-done:
132 case <-time.After(5 * time.Second):
133 logger.Warn().Msg("tunnel shutdown timeout; connections still active")
134 }
135
136 logger.Info().Msg("tunnel shutdown complete")
154 - return waitErr
137 + return errors.Join(waitErr, closeErr)
138 }
cmd/portal-tunnel/relays.go
+160 -17
@@ -13,51 +13,194 @@ import (
13 "time"
14
15 "github.com/rs/zerolog/log"
16 - "golang.org/x/sync/errgroup"
16
17 "github.com/gosuda/portal/v2/sdk"
18 )
19
21 -var bufferPool = sync.Pool{
22 - New: func() any {
23 - b := make([]byte, 64*1024)
24 - return &b
25 - },
20 +type relayRuntime struct {
21 + client *sdk.Client
22 + listener *sdk.Listener
23 + relayURL string
24 }
25
28 -func runProxyLoop(ctx context.Context, listener *sdk.Listener, targetAddr string, connGroup *errgroup.Group, connCount *atomic.Int64) error {
29 - logger := log.With().Str("component", "portal-tunnel").Logger()
26 +type relayLoopResult struct {
27 + err error
28 + leaseID string
29 + relayURL string
30 +}
31 +
32 +func (r *relayRuntime) run(ctx context.Context, localAddr string, connWG *sync.WaitGroup, connCount *atomic.Int64, done chan<- relayLoopResult) {
33 + logger := log.With().
34 + Str("component", "portal-tunnel").
35 + Str("relay", r.relayURL).
36 + Str("lease_id", r.listener.LeaseID()).
37 + Logger()
38 +
39 + var runErr error
40 + defer func() {
41 + done <- relayLoopResult{
42 + leaseID: r.listener.LeaseID(),
43 + relayURL: r.relayURL,
44 + err: runErr,
45 + }
46 + }()
47
48 for {
32 - relayConn, entry, err := listener.AcceptEntry()
49 + relayConn, err := r.listener.Accept()
50 if err != nil {
51 if errors.Is(err, net.ErrClosed) || errors.Is(err, context.Canceled) || ctx.Err() != nil {
35 - err = nil
52 + return
53 }
37 - return err
54 + runErr = err
55 + return
56 }
57
58 connID := connCount.Add(1)
59 logger.Info().
60 Int64("conn_id", connID).
61 Str("remote_addr", relayConn.RemoteAddr().String()).
44 - Str("relay", entry.RelayURL).
45 - Str("lease_id", entry.LeaseID).
62 Msg("accepted relay connection")
63
48 - connGroup.Go(func() error {
49 - if err := proxyConnection(ctx, targetAddr, relayConn); err != nil {
64 + connWG.Add(1)
65 + go func(connID int64, relayConn net.Conn) {
66 + defer connWG.Done()
67 + if err := proxyConnection(ctx, localAddr, relayConn); err != nil {
68 logger.Error().Err(err).Int64("conn_id", connID).Msg("proxy connection failed")
69 }
70 logger.Info().Int64("conn_id", connID).Msg("proxy connection closed")
53 - return nil
71 + }(connID, relayConn)
72 + }
73 +}
74 +
75 +func startRelayRuntimes(ctx context.Context, relayURLs []string, req sdk.ListenRequest) ([]*relayRuntime, error) {
76 + runtimes := make([]*relayRuntime, 0, len(relayURLs))
77 + for _, relayURL := range relayURLs {
78 + client, err := sdk.NewClient(sdk.ClientConfig{RelayURL: relayURL})
79 + if err != nil {
80 + _ = closeRelayRuntimes(runtimes)
81 + return nil, fmt.Errorf("create relay client %s: %w", relayURL, err)
82 + }
83 +
84 + listener, err := client.Listen(ctx, req)
85 + if err != nil {
86 + client.Close()
87 + _ = closeRelayRuntimes(runtimes)
88 + return nil, fmt.Errorf("register relay lease %s: %w", relayURL, err)
89 + }
90 +
91 + runtimes = append(runtimes, &relayRuntime{
92 + relayURL: relayURL,
93 + client: client,
94 + listener: listener,
95 })
96 }
97 + return runtimes, nil
98 +}
99 +
100 +func closeRelayRuntimes(runtimes []*relayRuntime) error {
101 + var closeErr error
102 + for _, runtime := range runtimes {
103 + if runtime == nil {
104 + continue
105 + }
106 + if runtime.listener != nil {
107 + closeErr = errors.Join(closeErr, runtime.listener.Close())
108 + }
109 + if runtime.client != nil {
110 + runtime.client.Close()
111 + }
112 + }
113 + return closeErr
114 +}
115 +
116 +func waitForRelayLoops(ctx context.Context, done <-chan relayLoopResult, relayCount int) error {
117 + logger := log.With().Str("component", "portal-tunnel").Logger()
118 + active := relayCount
119 +
120 + for active > 0 {
121 + result := <-done
122 + active--
123 +
124 + switch {
125 + case result.err != nil:
126 + logger.Error().
127 + Err(result.err).
128 + Str("relay", result.relayURL).
129 + Str("lease_id", result.leaseID).
130 + Int("remaining_relays", active).
131 + Msg("relay accept loop stopped")
132 + case ctx.Err() != nil:
133 + logger.Info().
134 + Str("relay", result.relayURL).
135 + Str("lease_id", result.leaseID).
136 + Int("remaining_relays", active).
137 + Msg("relay accept loop stopped during shutdown")
138 + default:
139 + logger.Warn().
140 + Str("relay", result.relayURL).
141 + Str("lease_id", result.leaseID).
142 + Int("remaining_relays", active).
143 + Msg("relay accept loop stopped")
144 + }
145 + }
146 +
147 + select {
148 + case <-ctx.Done():
149 + return nil
150 + default:
151 + }
152 + return errors.New("all relay listeners stopped")
153 }
154
58 -func proxyConnection(ctx context.Context, targetAddr string, relayConn net.Conn) error {
155 +func normalizeRelayURLs(raw string) ([]string, error) {
156 + seen := make(map[string]struct{})
157 + var relayURLs []string
158 + for _, relayURL := range sdk.SplitCSV(raw) {
159 + normalized, err := normalizeRelayURL(relayURL)
160 + if err != nil {
161 + return nil, err
162 + }
163 + if _, ok := seen[normalized]; ok {
164 + continue
165 + }
166 + seen[normalized] = struct{}{}
167 + relayURLs = append(relayURLs, normalized)
168 + }
169 + return relayURLs, nil
170 +}
171 +
172 +func normalizeRelayURL(raw string) (string, error) {
173 + u, err := url.Parse(strings.TrimSpace(raw))
174 + if err != nil {
175 + return "", fmt.Errorf("parse relay url: %w", err)
176 + }
177 + if !strings.EqualFold(u.Scheme, "https") {
178 + return "", fmt.Errorf("relay url must use https: %q", raw)
179 + }
180 + if u.Host == "" {
181 + return "", fmt.Errorf("relay url host is empty: %q", raw)
182 + }
183 + u.Path = strings.TrimRight(u.Path, "/")
184 + u.RawQuery = ""
185 + u.Fragment = ""
186 + return u.String(), nil
187 +}
188 +
189 +var bufferPool = sync.Pool{
190 + New: func() any {
191 + b := make([]byte, 64*1024)
192 + return &b
193 + },
194 +}
195 +
196 +func proxyConnection(ctx context.Context, localAddr string, relayConn net.Conn) error {
197 defer relayConn.Close()
198
199 + targetAddr, err := normalizeTargetAddr(localAddr)
200 + if err != nil {
201 + return fmt.Errorf("invalid --host value %q: %w", localAddr, err)
202 + }
203 +
204 dialer := &net.Dialer{Timeout: 5 * time.Second}
205 localConn, err := dialer.DialContext(ctx, "tcp", targetAddr)
206 if err != nil {
cmd/portal-tunnel/relays_test.go deleted
-43
@@ -1,43 +0,0 @@
1 -package main
2 -
3 -import "testing"
4 -
5 -func TestNormalizeTargetAddr(t *testing.T) {
6 - t.Parallel()
7 -
8 - tests := []struct {
9 - name string
10 - raw string
11 - want string
12 - wantErr bool
13 - }{
14 - {name: "host and port", raw: "localhost:8080", want: "localhost:8080"},
15 - {name: "host only", raw: "localhost", want: "localhost:80"},
16 - {name: "http url", raw: "http://localhost:8080", want: "localhost:8080"},
17 - {name: "https url", raw: "https://127.0.0.1", want: "127.0.0.1:80"},
18 - {name: "ipv6 host", raw: "::1", want: "[::1]:80"},
19 - {name: "url with path", raw: "http://localhost:8080/app", wantErr: true},
20 - {name: "url with query", raw: "http://localhost:8080/?x=1", wantErr: true},
21 - {name: "empty", raw: " ", wantErr: true},
22 - }
23 -
24 - for _, tt := range tests {
25 - t.Run(tt.name, func(t *testing.T) {
26 - t.Parallel()
27 -
28 - got, err := normalizeTargetAddr(tt.raw)
29 - if tt.wantErr {
30 - if err == nil {
31 - t.Fatalf("normalizeTargetAddr(%q) error = nil, want error", tt.raw)
32 - }
33 - return
34 - }
35 - if err != nil {
36 - t.Fatalf("normalizeTargetAddr(%q) error = %v", tt.raw, err)
37 - }
38 - if got != tt.want {
39 - t.Fatalf("normalizeTargetAddr(%q) = %q, want %q", tt.raw, got, tt.want)
40 - }
41 - })
42 - }
43 -}
sdk/client.go
+60 -159
@@ -29,98 +29,40 @@ const (
29 defaultLeaseTTL = 2 * time.Minute
30 defaultRenewBefore = 30 * time.Second
31 defaultReadyTarget = 1
32 - defaultRetryDelay = 5 * time.Second
32 defaultHTTPShutdownTimeout = 5 * time.Second
33 )
34
35 // ClientConfig configures the SDK client.
36 type ClientConfig struct {
38 - RelayURLs []string
37 + RelayURL string
38 RootCAPEM []byte
39 }
40
41 type Client struct {
43 - clients []*relayClient
44 -}
45 -
46 -type relayClient struct {
47 - baseURL *url.URL
48 - httpClient *http.Client
49 - rawTLSConfig *tls.Config
42 + baseURL *url.URL
43 + httpClient *http.Client
44 + rawTLSConfig *tls.Config
45 + dialTimeout time.Duration
46 + handshakeTimeout time.Duration
47 + leaseTTL time.Duration
48 + renewBefore time.Duration
49 + readyTarget int
50 }
51
52 func NewClient(cfg ClientConfig) (*Client, error) {
53 - relayURLs, err := normalizeRelayURLs(cfg.RelayURLs)
53 + baseURL, err := url.Parse(strings.TrimSpace(cfg.RelayURL))
54 if err != nil {
55 - return nil, err
56 - }
57 -
58 - clients := make([]*relayClient, 0, len(relayURLs))
59 - for _, relayURL := range relayURLs {
60 - client, err := newRelayClient(cfg, relayURL)
61 - if err != nil {
62 - for _, existing := range clients {
63 - existing.Close()
64 - }
65 - return nil, err
66 - }
67 - clients = append(clients, client)
68 - }
69 -
70 - return &Client{clients: clients}, nil
71 -}
72 -
73 -func (c *Client) Close() {
74 - if c == nil {
75 - return
76 - }
77 -
78 - for _, client := range c.clients {
79 - client.Close()
80 - }
81 -}
82 -
83 -func (c *Client) Listen(ctx context.Context, req ListenRequest) (*Listener, error) {
84 - if strings.TrimSpace(req.Name) == "" {
85 - return nil, errors.New("listener name is required")
86 - }
87 - if len(c.clients) == 0 {
88 - return nil, errors.New("no relay urls configured")
89 - }
90 -
91 - listener := newListener(ctx)
92 -
93 - entries := make([]*listenerLease, 0, len(c.clients))
94 - acceptedCap := 0
95 - for _, client := range c.clients {
96 - entry, err := client.listenEntry(listener, req)
97 - if err != nil {
98 - listener.cancel()
99 - return nil, errors.Join(err, closeListenerEntries(entries))
100 - }
101 - entries = append(entries, entry)
102 - acceptedCap += entry.readyTarget
103 - }
104 -
105 - if acceptedCap <= 0 {
106 - acceptedCap = len(entries)
55 + return nil, fmt.Errorf("parse relay url: %w", err)
56 }
108 -
109 - listener.accepted = make(chan acceptedConn, acceptedCap)
110 - listener.entries = entries
111 - listener.activeCount = len(entries)
112 - for _, entry := range entries {
113 - entry.start()
57 + if !strings.EqualFold(baseURL.Scheme, "https") {
58 + return nil, fmt.Errorf("relay url must use https: %q", cfg.RelayURL)
59 }
115 -
116 - return listener, nil
117 -}
118 -
119 -func newRelayClient(cfg ClientConfig, relayURL string) (*relayClient, error) {
120 - baseURL, err := url.Parse(relayURL)
121 - if err != nil {
122 - return nil, fmt.Errorf("parse relay url: %w", err)
60 + if baseURL.Host == "" {
61 + return nil, fmt.Errorf("relay url host is empty: %q", cfg.RelayURL)
62 }
63 + baseURL.Path = strings.TrimRight(baseURL.Path, "/")
64 + baseURL.RawQuery = ""
65 + baseURL.Fragment = ""
66
67 if len(cfg.RootCAPEM) == 0 && isLocalRelayHost(baseURL.Hostname()) {
68 bootstrapCtx, cancel := context.WithTimeout(context.Background(), defaultDialTimeout+defaultHandshakeTimeout)
@@ -150,55 +92,22 @@ func newRelayClient(cfg ClientConfig, relayURL string) (*relayClient, error) {
92 ForceAttemptHTTP2: false,
93 }
94
153 - return &relayClient{
95 + return &Client{
96 baseURL: baseURL,
97 httpClient: &http.Client{
98 Transport: transport,
99 Timeout: defaultRequestTimeout,
100 },
159 - rawTLSConfig: baseTLS,
101 + rawTLSConfig: baseTLS,
102 + dialTimeout: defaultDialTimeout,
103 + handshakeTimeout: defaultHandshakeTimeout,
104 + leaseTTL: defaultLeaseTTL,
105 + renewBefore: defaultRenewBefore,
106 + readyTarget: defaultReadyTarget,
107 }, nil
108 }
109
163 -func normalizeRelayURLs(rawURLs []string) ([]string, error) {
164 - if len(rawURLs) == 0 {
165 - return nil, errors.New("relay url is required")
166 - }
167 -
168 - seen := make(map[string]struct{}, len(rawURLs))
169 - relayURLs := make([]string, 0, len(rawURLs))
170 - for _, raw := range rawURLs {
171 - normalized, err := normalizeRelayURL(raw)
172 - if err != nil {
173 - return nil, err
174 - }
175 - if _, ok := seen[normalized]; ok {
176 - continue
177 - }
178 - seen[normalized] = struct{}{}
179 - relayURLs = append(relayURLs, normalized)
180 - }
181 - return relayURLs, nil
182 -}
183 -
184 -func normalizeRelayURL(raw string) (string, error) {
185 - baseURL, err := url.Parse(strings.TrimSpace(raw))
186 - if err != nil {
187 - return "", fmt.Errorf("parse relay url: %w", err)
188 - }
189 - if !strings.EqualFold(baseURL.Scheme, "https") {
190 - return "", fmt.Errorf("relay url must use https: %q", raw)
191 - }
192 - if baseURL.Host == "" {
193 - return "", fmt.Errorf("relay url host is empty: %q", raw)
194 - }
195 - baseURL.Path = strings.TrimRight(baseURL.Path, "/")
196 - baseURL.RawQuery = ""
197 - baseURL.Fragment = ""
198 - return baseURL.String(), nil
199 -}
200 -
201 -func (c *relayClient) Close() {
110 +func (c *Client) Close() {
111 if c == nil || c.httpClient == nil {
112 return
113 }
@@ -207,7 +116,14 @@ func (c *relayClient) Close() {
116 }
117 }
118
210 -func (c *relayClient) listenEntry(listener *Listener, req ListenRequest) (*listenerLease, error) {
119 +func (c *Client) Listen(ctx context.Context, req ListenRequest) (*Listener, error) {
120 + if strings.TrimSpace(req.Name) == "" {
121 + return nil, errors.New("listener name is required")
122 + }
123 + if ctx == nil {
124 + ctx = context.Background()
125 + }
126 +
127 reverseToken := strings.TrimSpace(req.ReverseToken)
128 if reverseToken == "" {
129 reverseToken = randomToken()
@@ -215,12 +131,13 @@ func (c *relayClient) listenEntry(listener *Listener, req ListenRequest) (*liste
131
132 readyTarget := req.ReadyTarget
133 if readyTarget <= 0 {
218 - readyTarget = defaultReadyTarget
134 + readyTarget = c.readyTarget
135 }
136 leaseTTL := req.LeaseTTL
137 if leaseTTL <= 0 {
222 - leaseTTL = defaultLeaseTTL
138 + leaseTTL = c.leaseTTL
139 }
140 + acceptedCap := max(readyTarget*2, 1)
141
142 registerReq := types.RegisterRequest{
143 Name: req.Name,
@@ -232,7 +149,7 @@ func (c *relayClient) listenEntry(listener *Listener, req ListenRequest) (*liste
149 }
150
151 var registerResp types.RegisterResponse
235 - if err := c.doJSON(listener.ctx, http.MethodPost, types.PathSDKRegister, registerReq, &registerResp); err != nil {
152 + if err := c.doJSON(ctx, http.MethodPost, types.PathSDKRegister, registerReq, &registerResp); err != nil {
153 return nil, err
154 }
155
@@ -242,25 +159,31 @@ func (c *relayClient) listenEntry(listener *Listener, req ListenRequest) (*liste
159 return nil, err
160 }
161
245 - return &listenerLease{
246 - parent: listener,
247 - client: c,
248 - info: ListenerEntry{
249 - RelayURL: c.baseURL.String(),
250 - LeaseID: registerResp.LeaseID,
251 - Hostnames: append([]string(nil), registerResp.Hostnames...),
252 - Metadata: cloneLeaseMetadata(registerResp.Metadata),
253 - },
162 + listenerCtx, cancel := context.WithCancel(ctx)
163 + listener := &Listener{
164 + client: c,
165 + baseContext: func() context.Context { return listenerCtx },
166 + ctxDone: listenerCtx.Done(),
167 + cancel: cancel,
168 + leaseID: registerResp.LeaseID,
169 + hostnames: append([]string(nil), registerResp.Hostnames...),
170 + metadata: cloneLeaseMetadata(registerResp.Metadata),
171 reverseToken: reverseToken,
172 leaseTTL: leaseTTL,
173 readyTarget: readyTarget,
174 tlsConfig: tlsConf,
175 tlsCloser: tlsCloser,
259 - active: true,
260 - }, nil
176 + accepted: make(chan net.Conn, acceptedCap),
177 + signal: make(chan struct{}, 1),
178 + }
179 +
180 + go listener.runSupervisor()
181 + go listener.runRenewLoop()
182 + listener.notify()
183 + return listener, nil
184 }
185
263 -func (c *relayClient) doJSON(ctx context.Context, method, path string, payload any, out any) error {
186 +func (c *Client) doJSON(ctx context.Context, method, path string, payload any, out any) error {
187 var body io.Reader
188 if payload != nil {
189 buf, err := json.Marshal(payload)
@@ -306,7 +229,7 @@ func (c *relayClient) doJSON(ctx context.Context, method, path string, payload a
229 return json.Unmarshal(envelope.Data, out)
230 }
231
309 -func (c *relayClient) renewLease(ctx context.Context, leaseID, reverseToken string, ttl time.Duration) error {
232 +func (c *Client) renewLease(ctx context.Context, leaseID, reverseToken string, ttl time.Duration) error {
233 return c.doJSON(ctx, http.MethodPost, types.PathSDKRenew, types.RenewRequest{
234 LeaseID: leaseID,
235 ReverseToken: reverseToken,
@@ -314,16 +237,16 @@ func (c *relayClient) renewLease(ctx context.Context, leaseID, reverseToken stri
237 }, &types.RenewResponse{})
238 }
239
317 -func (c *relayClient) unregisterLease(ctx context.Context, leaseID, reverseToken string) error {
240 +func (c *Client) unregisterLease(ctx context.Context, leaseID, reverseToken string) error {
241 return c.doJSON(ctx, http.MethodPost, types.PathSDKUnregister, types.UnregisterRequest{
242 LeaseID: leaseID,
243 ReverseToken: reverseToken,
244 }, nil)
245 }
246
324 -func (c *relayClient) openReverseSession(ctx context.Context, leaseID, reverseToken string) (net.Conn, error) {
247 +func (c *Client) openReverseSession(ctx context.Context, leaseID, reverseToken string) (net.Conn, error) {
248 dialer := &tls.Dialer{
326 - NetDialer: &net.Dialer{Timeout: defaultDialTimeout},
249 + NetDialer: &net.Dialer{Timeout: c.dialTimeout},
250 Config: c.rawTLSConfig.Clone(),
251 }
252
@@ -390,28 +313,6 @@ func decodeAPIResponseError(resp *http.Response) error {
313 }
314 }
315
393 -func isLeaseResetError(err error) bool {
394 - var apiErr *types.APIRequestError
395 - if !errors.As(err, &apiErr) {
396 - return false
397 - }
398 - return apiErr.Code == types.APIErrorCodeLeaseNotFound
399 -}
400 -
401 -func isTerminalEntryError(err error) bool {
402 - var apiErr *types.APIRequestError
403 - if !errors.As(err, &apiErr) {
404 - return false
405 - }
406 -
407 - switch apiErr.Code {
408 - case types.APIErrorCodeIPBanned, types.APIErrorCodeUnauthorized:
409 - return true
410 - default:
411 - return false
412 - }
413 -}
414 -
316 func buildRootCAs(rootCAPEM []byte) (*x509.CertPool, error) {
317 if len(rootCAPEM) == 0 {
318 return nil, nil
sdk/client_test.go
+5 -45
@@ -22,17 +22,13 @@ func TestNewClientAutoTrustsLocalhostRelayCertificate(t *testing.T) {
22 }))
23 defer server.Close()
24
25 - client, err := NewClient(ClientConfig{RelayURLs: []string{server.URL}})
25 + client, err := NewClient(ClientConfig{RelayURL: server.URL})
26 if err != nil {
27 t.Fatalf("NewClient() error = %v", err)
28 }
29 defer client.Close()
30
31 - if len(client.clients) != 1 {
32 - t.Fatalf("client count = %d, want 1", len(client.clients))
33 - }
34 -
35 - resp, err := client.clients[0].httpClient.Get(server.URL)
31 + resp, err := client.httpClient.Get(server.URL)
32 if err != nil {
33 t.Fatalf("httpClient.Get() error = %v", err)
34 }
@@ -43,43 +39,6 @@ func TestNewClientAutoTrustsLocalhostRelayCertificate(t *testing.T) {
39 }
40 }
41
46 -func TestNewClientSupportsDedupedRelayURLs(t *testing.T) {
47 - t.Parallel()
48 -
49 - serverA := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
50 - w.WriteHeader(http.StatusOK)
51 - }))
52 - defer serverA.Close()
53 -
54 - serverB := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
55 - w.WriteHeader(http.StatusOK)
56 - }))
57 - defer serverB.Close()
58 -
59 - client, err := NewClient(ClientConfig{
60 - RelayURLs: []string{serverA.URL, serverB.URL},
61 - })
62 - if err != nil {
63 - t.Fatalf("NewClient() error = %v", err)
64 - }
65 - defer client.Close()
66 -
67 - if len(client.clients) != 2 {
68 - t.Fatalf("client count = %d, want 2", len(client.clients))
69 - }
70 -
71 - for i, relayClient := range client.clients {
72 - resp, err := relayClient.httpClient.Get(relayClient.baseURL.String())
73 - if err != nil {
74 - t.Fatalf("client[%d].httpClient.Get() error = %v", i, err)
75 - }
76 - _ = resp.Body.Close()
77 - if resp.StatusCode != http.StatusOK {
78 - t.Fatalf("client[%d] status = %d, want %d", i, resp.StatusCode, http.StatusOK)
79 - }
80 - }
81 -}
82 -
42 func TestOpenReverseSessionPreservesAPIErrorCode(t *testing.T) {
43 t.Parallel()
44
@@ -116,7 +75,7 @@ func TestOpenReverseSessionPreservesAPIErrorCode(t *testing.T) {
75 t.Fatalf("server client transport type = %T, want *http.Transport", server.Client().Transport)
76 }
77
119 - relayClient := &relayClient{
78 + client := &Client{
79 baseURL: baseURL,
80 httpClient: &http.Client{
81 Transport: transport.Clone(),
@@ -128,9 +87,10 @@ func TestOpenReverseSessionPreservesAPIErrorCode(t *testing.T) {
87 RootCAs: transport.TLSClientConfig.RootCAs,
88 NextProtos: []string{"http/1.1"},
89 },
90 + dialTimeout: defaultDialTimeout,
91 }
92
133 - _, err = relayClient.openReverseSession(context.Background(), "lease-123", "tok_123")
93 + _, err = client.openReverseSession(context.Background(), "lease-123", "tok_123")
94 if err == nil {
95 t.Fatal("openReverseSession() error = nil, want APIRequestError")
96 }
sdk/listener.go
+119 -355
@@ -10,8 +10,6 @@ import (
10 "sync"
11 "time"
12
13 - "golang.org/x/sync/errgroup"
14 -
13 "github.com/gosuda/portal/v2/types"
14 )
15
@@ -24,98 +22,36 @@ type ListenRequest struct {
22 LeaseTTL time.Duration
23 }
24
27 -type ListenerEntry struct {
28 - RelayURL string
29 - LeaseID string
30 - Hostnames []string
31 - Metadata types.LeaseMetadata
32 -}
33 -
34 -// PublicURLs returns the HTTPS URLs exposed by one relay-specific lease.
35 -func (e ListenerEntry) PublicURLs() []string {
36 - urls := make([]string, 0, len(e.Hostnames))
37 - for _, host := range e.Hostnames {
38 - urls = append(urls, "https://"+host)
39 - }
40 - return urls
41 -}
42 -
43 -func (e ListenerEntry) clone() ListenerEntry {
44 - e.Hostnames = append([]string(nil), e.Hostnames...)
45 - e.Metadata = cloneLeaseMetadata(e.Metadata)
46 - return e
47 -}
48 -
25 type Listener struct {
50 - ctx context.Context
51 - cancel context.CancelFunc
52 - accepted chan acceptedConn
53 - entries []*listenerLease
54 - workers errgroup.Group
55 - closeOnce sync.Once
56 -
57 - mu sync.RWMutex
58 - activeCount int
59 - terminalErr error
60 -}
61 -
62 -type listenerLease struct {
63 - parent *Listener
64 - client *relayClient
26 + tlsCloser io.Closer
27 + tlsConfig *tls.Config
28 + baseContext func() context.Context
29 + ctxDone <-chan struct{}
30 + cancel context.CancelFunc
31 + client *Client
32 + signal chan struct{}
33 + accepted chan net.Conn
34 + leaseID string
35 reverseToken string
36 + hostnames []string
37 + metadata types.LeaseMetadata
38 readyTarget int
39 leaseTTL time.Duration
40
69 - mu sync.RWMutex
70 - info ListenerEntry
71 - tlsConfig *tls.Config
72 - tlsCloser io.Closer
73 - active bool
74 - terminalErr error
75 -}
76 -
77 -type acceptedConn struct {
78 - conn net.Conn
79 - entry ListenerEntry
80 -}
81 -
82 -type sessionSnapshot struct {
83 - info ListenerEntry
84 - tlsConfig *tls.Config
85 - active bool
86 - terminalErr error
87 -}
88 -
89 -func newListener(ctx context.Context) *Listener {
90 - if ctx == nil {
91 - ctx = context.Background()
92 - }
93 -
94 - listenerCtx, cancel := context.WithCancel(ctx)
95 - return &Listener{
96 - ctx: listenerCtx,
97 - cancel: cancel,
98 - }
41 + activeSessions int
42 + closeOnce sync.Once
43 + mu sync.Mutex
44 }
45
46 func (l *Listener) Accept() (net.Conn, error) {
102 - conn, _, err := l.AcceptEntry()
103 - return conn, err
104 -}
105 -
106 -// AcceptEntry returns the next accepted connection plus relay-specific lease
107 -// metadata for callers that need to distinguish which relay claimed it.
108 -func (l *Listener) AcceptEntry() (net.Conn, ListenerEntry, error) {
109 - for {
110 - select {
111 - case <-l.ctx.Done():
112 - return nil, ListenerEntry{}, l.closeError()
113 - case accepted := <-l.accepted:
114 - if accepted.conn == nil {
115 - return nil, ListenerEntry{}, l.closeError()
116 - }
117 - return accepted.conn, accepted.entry.clone(), nil
47 + select {
48 + case <-l.ctxDone:
49 + return nil, net.ErrClosed
50 + case conn := <-l.accepted:
51 + if conn == nil {
52 + return nil, net.ErrClosed
53 }
54 + return conn, nil
55 }
56 }
57
@@ -123,164 +59,106 @@ func (l *Listener) Close() error {
59 var closeErr error
60 l.closeOnce.Do(func() {
61 l.cancel()
126 - closeErr = errors.Join(l.workers.Wait(), closeListenerEntries(l.entries))
127 - l.drainAccepted()
62 +
63 + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
64 + defer cancel()
65 + if err := l.client.unregisterLease(ctx, l.leaseID, l.reverseToken); err != nil {
66 + closeErr = err
67 + }
68 + if l.tlsCloser != nil {
69 + closeErr = errors.Join(closeErr, l.tlsCloser.Close())
70 + }
71 })
72 return closeErr
73 }
74
75 func (l *Listener) Addr() net.Addr {
133 - if entry, ok := l.singleEntry(); ok {
134 - return listenerAddr("portal:" + entry.LeaseID)
135 - }
136 - return listenerAddr("portal:multi")
76 + return listenerAddr("portal:" + l.leaseID)
77 }
78
139 -// Entries returns relay-specific lease details for advanced multi-relay callers.
140 -func (l *Listener) Entries() []ListenerEntry {
141 - entries := make([]ListenerEntry, 0, len(l.entries))
142 - for _, entry := range l.entries {
143 - info, ok := entry.snapshotInfo()
144 - if !ok {
145 - continue
146 - }
147 - entries = append(entries, info)
148 - }
149 - return entries
79 +func (l *Listener) LeaseID() string {
80 + return l.leaseID
81 }
82
152 -func (l *Listener) singleEntry() (ListenerEntry, bool) {
153 - entries := l.Entries()
154 - if len(entries) != 1 {
155 - return ListenerEntry{}, false
156 - }
157 - return entries[0], true
83 +func (l *Listener) Hostnames() []string {
84 + return append([]string(nil), l.hostnames...)
85 +}
86 +
87 +func (l *Listener) Metadata() types.LeaseMetadata {
88 + return cloneLeaseMetadata(l.metadata)
89 }
90
160 -// PublicURLs returns all public HTTPS URLs exposed by the listener.
91 func (l *Listener) PublicURLs() []string {
162 - var urls []string
163 - for _, entry := range l.Entries() {
164 - urls = append(urls, entry.PublicURLs()...)
92 + urls := make([]string, 0, len(l.hostnames))
93 + for _, host := range l.hostnames {
94 + urls = append(urls, "https://"+host)
95 }
96 return urls
97 }
98
169 -func closeListenerEntries(entries []*listenerLease) error {
170 - if len(entries) == 0 {
171 - return nil
172 - }
173 -
174 - var closeErr error
175 - var mu sync.Mutex
176 - var group errgroup.Group
177 - group.SetLimit(min(len(entries), 4))
99 +func (l *Listener) runSupervisor() {
100 + for {
101 + select {
102 + case <-l.ctxDone:
103 + return
104 + case <-l.signal:
105 + }
106
179 - for _, entry := range entries {
180 - if entry == nil {
181 - continue
107 + for l.reserveSessionSlot() {
108 + go l.runSession()
109 }
183 - group.Go(func() error {
184 - ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
185 - defer cancel()
186 -
187 - if err := entry.close(ctx); err != nil {
188 - mu.Lock()
189 - closeErr = errors.Join(closeErr, err)
190 - mu.Unlock()
191 - }
192 - return nil
193 - })
110 }
195 -
196 - _ = group.Wait()
197 -
198 - return closeErr
111 }
112
201 -func (l *listenerLease) start() {
202 - if l == nil || l.parent == nil {
203 - return
113 +func (l *Listener) runRenewLoop() {
114 + interval := l.leaseTTL / 2
115 + if interval <= 0 {
116 + interval = 30 * time.Second
117 }
205 -
206 - l.parent.workers.Go(func() error {
207 - l.runRenewLoop()
208 - return nil
209 - })
210 - for i := 0; i < l.readyTarget; i++ {
211 - l.parent.workers.Go(func() error {
212 - l.runSessionWorker()
213 - return nil
214 - })
118 + if l.client.renewBefore > 0 && l.leaseTTL > l.client.renewBefore {
119 + interval = l.leaseTTL - l.client.renewBefore
120 + }
121 + if interval <= 0 {
122 + interval = 30 * time.Second
123 }
216 -}
124
218 -func (l *listenerLease) runRenewLoop() {
219 - ticker := time.NewTicker(l.renewInterval())
125 + ticker := time.NewTicker(interval)
126 defer ticker.Stop()
127
128 for {
129 select {
224 - case <-l.parent.ctx.Done():
130 + case <-l.ctxDone:
131 return
132 case <-ticker.C:
227 - if l.shouldStop() {
228 - return
229 - }
230 -
231 - ctx, cancel := context.WithTimeout(l.parent.ctx, 10*time.Second)
232 - err := l.client.renewLease(ctx, l.leaseID(), l.reverseToken, l.leaseTTL)
133 + ctx, cancel := context.WithTimeout(l.context(), 10*time.Second)
134 + _ = l.client.renewLease(ctx, l.leaseID, l.reverseToken, l.leaseTTL)
135 cancel()
234 - if err == nil {
235 - continue
236 - }
237 - if isLeaseResetError(err) || isTerminalEntryError(err) {
238 - l.stop(err)
239 - return
240 - }
136 }
137 }
138 }
139
245 -func (l *listenerLease) runSessionWorker() {
246 - for {
247 - if l.shouldStop() {
248 - return
249 - }
250 -
251 - snapshot := l.sessionSnapshot()
252 - if !snapshot.active {
253 - return
254 - }
140 +func (l *Listener) runSession() {
141 + defer l.releaseSessionSlot()
142
256 - conn, err := l.client.openReverseSession(l.parent.ctx, snapshot.info.LeaseID, l.reverseToken)
257 - if err != nil {
258 - if isLeaseResetError(err) || isTerminalEntryError(err) {
259 - l.stop(err)
260 - return
261 - }
262 - sleepOrDone(l.parent.ctx, defaultRetryDelay)
263 - continue
264 - }
143 + sessionCtx := l.context()
144 + conn, err := l.client.openReverseSession(sessionCtx, l.leaseID, l.reverseToken)
145 + if err != nil {
146 + sleepOrDone(sessionCtx, time.Second)
147 + return
148 + }
149
266 - if err := l.parent.awaitActivation(conn, snapshot.tlsConfig, snapshot.info); err != nil {
267 - _ = conn.Close()
268 - if errors.Is(err, context.Canceled) || errors.Is(err, net.ErrClosed) {
269 - return
270 - }
271 - if isLeaseResetError(err) || isTerminalEntryError(err) {
272 - l.stop(err)
273 - return
274 - }
275 - sleepOrDone(l.parent.ctx, defaultRetryDelay)
150 + if err := l.awaitActivation(conn); err != nil {
151 + _ = conn.Close()
152 + if !errors.Is(err, context.Canceled) && !errors.Is(err, net.ErrClosed) {
153 + sleepOrDone(sessionCtx, time.Second)
154 }
155 }
156 }
157
280 -func (l *Listener) awaitActivation(conn net.Conn, tlsConfig *tls.Config, info ListenerEntry) error {
158 +func (l *Listener) awaitActivation(conn net.Conn) error {
159 var marker [1]byte
160 for {
283 - _ = conn.SetReadDeadline(time.Now().Add(2 * defaultHandshakeTimeout))
161 + _ = conn.SetReadDeadline(time.Now().Add(2 * l.client.handshakeTimeout))
162 if _, err := io.ReadFull(conn, marker[:]); err != nil {
163 return err
164 }
@@ -290,189 +168,54 @@ func (l *Listener) awaitActivation(conn net.Conn, tlsConfig *tls.Config, info Li
168 case types.MarkerKeepalive:
169 continue
170 case types.MarkerTLSStart:
293 - return l.activate(conn, tlsConfig, info)
171 + return l.activate(conn)
172 default:
173 return fmt.Errorf("unexpected reverse marker: 0x%02x", marker[0])
174 }
175 }
176 }
177
300 -func (l *Listener) activate(conn net.Conn, tlsConfig *tls.Config, info ListenerEntry) error {
301 - if tlsConfig == nil {
302 - return errors.New("missing tls config")
303 - }
304 -
305 - tlsConn := tls.Server(conn, tlsConfig.Clone())
306 - handshakeCtx, cancel := context.WithTimeout(l.ctx, defaultHandshakeTimeout)
178 +func (l *Listener) activate(conn net.Conn) error {
179 + tlsConn := tls.Server(conn, l.tlsConfig.Clone())
180 + handshakeCtx, cancel := context.WithTimeout(l.context(), l.client.handshakeTimeout)
181 defer cancel()
182 if err := tlsConn.HandshakeContext(handshakeCtx); err != nil {
183 return err
184 }
185
186 select {
313 - case <-l.ctx.Done():
187 + case <-l.ctxDone:
188 _ = tlsConn.Close()
315 - return l.closeError()
316 - case l.accepted <- acceptedConn{conn: tlsConn, entry: info.clone()}:
317 - return nil
318 - }
319 -}
320 -
321 -func (l *listenerLease) close(ctx context.Context) error {
322 - if l == nil {
189 + return l.context().Err()
190 + case l.accepted <- tlsConn:
191 return nil
192 }
325 -
326 - var closeErr error
327 - if l.client != nil {
328 - if err := l.client.unregisterLease(ctx, l.info.LeaseID, l.reverseToken); err != nil {
329 - closeErr = errors.Join(closeErr, err)
330 - }
331 - }
332 - if l.tlsCloser != nil {
333 - closeErr = errors.Join(closeErr, l.tlsCloser.Close())
334 - }
335 - return closeErr
336 -}
337 -
338 -func (l *listenerLease) renewInterval() time.Duration {
339 - interval := l.leaseTTL / 2
340 - if interval <= 0 {
341 - interval = 30 * time.Second
342 - }
343 - if defaultRenewBefore > 0 && l.leaseTTL > defaultRenewBefore {
344 - interval = l.leaseTTL - defaultRenewBefore
345 - }
346 - if interval <= 0 {
347 - interval = 30 * time.Second
348 - }
349 - return interval
193 }
194
352 -func (l *Listener) fail(err error) {
195 +func (l *Listener) reserveSessionSlot() bool {
196 l.mu.Lock()
354 - if err != nil && l.terminalErr == nil {
355 - l.terminalErr = err
356 - }
357 - l.mu.Unlock()
358 - l.cancel()
359 -}
360 -
361 -func (l *Listener) closeError() error {
362 - l.mu.RLock()
363 - terminalErr := l.terminalErr
364 - l.mu.RUnlock()
365 - if terminalErr != nil {
366 - return terminalErr
367 - }
368 - if err := l.ctx.Err(); err != nil {
369 - if errors.Is(err, context.Canceled) {
370 - return net.ErrClosed
371 - }
372 - return err
373 - }
374 - return net.ErrClosed
375 -}
376 -
377 -func (l *Listener) isClosed() bool {
378 - if l == nil || l.ctx == nil {
197 + defer l.mu.Unlock()
198 + if l.isClosed() {
199 return false
200 }
381 - select {
382 - case <-l.ctx.Done():
383 - return true
384 - default:
201 + if l.activeSessions >= l.readyTarget {
202 return false
203 }
204 + l.activeSessions++
205 + return true
206 }
207
389 -func (l *Listener) drainAccepted() {
390 - for {
391 - select {
392 - case accepted := <-l.accepted:
393 - if accepted.conn != nil {
394 - _ = accepted.conn.Close()
395 - }
396 - default:
397 - return
398 - }
399 - }
400 -}
401 -
402 -func (l *Listener) entryStopped(_ *listenerLease, err error) {
208 +func (l *Listener) releaseSessionSlot() {
209 l.mu.Lock()
404 - if l.activeCount > 0 {
405 - l.activeCount--
406 - }
407 - shouldFail := l.activeCount == 0
210 + l.activeSessions--
211 l.mu.Unlock()
409 -
410 - if shouldFail {
411 - l.fail(err)
412 - }
413 -}
414 -
415 -func (l *listenerLease) snapshotInfo() (ListenerEntry, bool) {
416 - snapshot := l.sessionSnapshot()
417 - if !snapshot.active {
418 - return ListenerEntry{}, false
419 - }
420 - return snapshot.info, true
421 -}
422 -
423 -func (l *listenerLease) leaseID() string {
424 - l.mu.RLock()
425 - defer l.mu.RUnlock()
426 - return l.info.LeaseID
427 -}
428 -
429 -func (l *listenerLease) sessionSnapshot() sessionSnapshot {
430 - l.mu.RLock()
431 - defer l.mu.RUnlock()
432 -
433 - return sessionSnapshot{
434 - info: l.info.clone(),
435 - tlsConfig: l.tlsConfig,
436 - active: l.active,
437 - terminalErr: l.terminalErr,
438 - }
439 -}
440 -
441 -func (l *listenerLease) shutdownState() (bool, error) {
442 - l.mu.RLock()
443 - defer l.mu.RUnlock()
444 - return !l.active, l.terminalErr
445 -}
446 -
447 -func (l *listenerLease) shouldStop() bool {
448 - if l == nil {
449 - return true
450 - }
451 - if l.parent != nil && l.parent.isClosed() {
452 - return true
453 - }
454 - stopped, _ := l.shutdownState()
455 - return stopped
212 + l.notify()
213 }
214
458 -func (l *listenerLease) stop(err error) {
459 - if l == nil {
460 - return
461 - }
462 -
463 - l.mu.Lock()
464 - if !l.active {
465 - l.mu.Unlock()
466 - return
467 - }
468 - l.active = false
469 - if err != nil && l.terminalErr == nil {
470 - l.terminalErr = err
471 - }
472 - l.mu.Unlock()
473 -
474 - if l.parent != nil {
475 - l.parent.entryStopped(l, err)
215 +func (l *Listener) notify() {
216 + select {
217 + case l.signal <- struct{}{}:
218 + default:
219 }
220 }
221
@@ -490,6 +233,27 @@ type listenerAddr string
233 func (a listenerAddr) Network() string { return "portal" }
234 func (a listenerAddr) String() string { return string(a) }
235
236 +func (l *Listener) context() context.Context {
237 + if l.baseContext != nil {
238 + if ctx := l.baseContext(); ctx != nil {
239 + return ctx
240 + }
241 + }
242 + return context.Background()
243 +}
244 +
245 +func (l *Listener) isClosed() bool {
246 + if l.ctxDone == nil {
247 + return false
248 + }
249 + select {
250 + case <-l.ctxDone:
251 + return true
252 + default:
253 + return false
254 + }
255 +}
256 +
257 func cloneLeaseMetadata(metadata types.LeaseMetadata) types.LeaseMetadata {
258 metadata.Tags = append([]string(nil), metadata.Tags...)
259 return metadata
sdk/listener_test.go
+30 -216
@@ -1,145 +1,75 @@
1 package sdk
2
3 import (
4 - "context"
5 - "errors"
4 "net"
5 "testing"
6
7 "github.com/gosuda/portal/v2/types"
8 )
9
12 -func TestListenerSingleEntryAccessors(t *testing.T) {
10 +func TestListenerAccessors(t *testing.T) {
11 t.Parallel()
12
15 - listener := newListener(context.Background())
16 - listener.entries = []*listenerLease{
17 - {
18 - info: ListenerEntry{
19 - RelayURL: "https://relay.example.com",
20 - LeaseID: "lease-1",
21 - Hostnames: []string{"app.relay.example.com"},
22 - },
23 - active: true,
13 + listener := &Listener{
14 + leaseID: "lease-1",
15 + hostnames: []string{"app.relay.example.com"},
16 + metadata: types.LeaseMetadata{
17 + Owner: "alice",
18 + Tags: []string{"one", "two"},
19 },
20 }
21
27 - entries := listener.Entries()
28 - if len(entries) != 1 {
29 - t.Fatalf("Entries() len = %d, want 1", len(entries))
30 - }
22 if listener.Addr().String() != "portal:lease-1" {
23 t.Fatalf("Addr().String() = %q, want %q", listener.Addr().String(), "portal:lease-1")
24 }
34 -
35 - publicURLs := listener.PublicURLs()
36 - if len(publicURLs) != 1 || publicURLs[0] != "https://app.relay.example.com" {
37 - t.Fatalf("PublicURLs() = %#v, want [https://app.relay.example.com]", publicURLs)
25 + if listener.LeaseID() != "lease-1" {
26 + t.Fatalf("LeaseID() = %q, want %q", listener.LeaseID(), "lease-1")
27 }
39 -}
40 -
41 -func TestListenerMultiEntryAccessors(t *testing.T) {
42 - t.Parallel()
28
44 - listener := newListener(context.Background())
45 - listener.entries = []*listenerLease{
46 - {
47 - info: ListenerEntry{
48 - RelayURL: "https://relay-a.example.com",
49 - LeaseID: "lease-a",
50 - Hostnames: []string{"a.example.com"},
51 - },
52 - active: true,
53 - },
54 - {
55 - info: ListenerEntry{
56 - RelayURL: "https://relay-b.example.com",
57 - LeaseID: "lease-b",
58 - Hostnames: []string{"b.example.com"},
59 - },
60 - active: true,
61 - },
29 + hostnames := listener.Hostnames()
30 + if len(hostnames) != 1 || hostnames[0] != "app.relay.example.com" {
31 + t.Fatalf("Hostnames() = %#v, want [app.relay.example.com]", hostnames)
32 }
33
64 - if listener.Addr().String() != "portal:multi" {
65 - t.Fatalf("Addr().String() = %q, want %q", listener.Addr().String(), "portal:multi")
34 + metadata := listener.Metadata()
35 + if metadata.Owner != "alice" {
36 + t.Fatalf("Metadata().Owner = %q, want %q", metadata.Owner, "alice")
37 }
67 -
68 - entries := listener.Entries()
69 - if len(entries) != 2 {
70 - t.Fatalf("Entries() len = %d, want 2", len(entries))
38 + if len(metadata.Tags) != 2 {
39 + t.Fatalf("Metadata().Tags len = %d, want 2", len(metadata.Tags))
40 }
41
42 publicURLs := listener.PublicURLs()
74 - if len(publicURLs) != 2 {
75 - t.Fatalf("PublicURLs() len = %d, want 2", len(publicURLs))
76 - }
77 -}
78 -
79 -func TestListenerEntriesSkipInactiveLeases(t *testing.T) {
80 - t.Parallel()
81 -
82 - listener := newListener(context.Background())
83 - listener.entries = []*listenerLease{
84 - {
85 - info: ListenerEntry{
86 - RelayURL: "https://relay-a.example.com",
87 - LeaseID: "lease-a",
88 - Hostnames: []string{"a.example.com"},
89 - },
90 - active: true,
91 - },
92 - {
93 - info: ListenerEntry{
94 - RelayURL: "https://relay-b.example.com",
95 - LeaseID: "lease-b",
96 - Hostnames: []string{"b.example.com"},
97 - },
98 - active: false,
99 - },
100 - }
101 -
102 - entries := listener.Entries()
103 - if len(entries) != 1 {
104 - t.Fatalf("Entries() len = %d, want 1", len(entries))
105 - }
106 - if entries[0].LeaseID != "lease-a" {
107 - t.Fatalf("Entries()[0].LeaseID = %q, want %q", entries[0].LeaseID, "lease-a")
43 + if len(publicURLs) != 1 || publicURLs[0] != "https://app.relay.example.com" {
44 + t.Fatalf("PublicURLs() = %#v, want [https://app.relay.example.com]", publicURLs)
45 }
46 }
47
111 -func TestListenerAcceptEntry(t *testing.T) {
48 +func TestListenerAccept(t *testing.T) {
49 t.Parallel()
50
114 - listener := newListener(context.Background())
115 - listener.accepted = make(chan acceptedConn, 2)
116 -
51 + done := make(chan struct{})
52 serverConn1, clientConn1 := net.Pipe()
53 defer clientConn1.Close()
54 serverConn2, clientConn2 := net.Pipe()
55 defer clientConn2.Close()
56
122 - listener.accepted <- acceptedConn{
123 - conn: serverConn1,
124 - entry: ListenerEntry{
125 - RelayURL: "https://relay.example.com",
126 - LeaseID: "lease-1",
127 - Hostnames: []string{"app.relay.example.com"},
128 - },
57 + listener := &Listener{
58 + ctxDone: done,
59 + accepted: make(chan net.Conn, 2),
60 + hostnames: []string{"app.relay.example.com"},
61 }
130 - listener.accepted <- acceptedConn{conn: serverConn2}
62 + listener.accepted <- serverConn1
63 + listener.accepted <- serverConn2
64
132 - conn, entry, err := listener.AcceptEntry()
65 + conn, err := listener.Accept()
66 if err != nil {
134 - t.Fatalf("AcceptEntry() error = %v", err)
67 + t.Fatalf("Accept() error = %v", err)
68 }
69 defer conn.Close()
70
71 if conn != serverConn1 {
139 - t.Fatal("AcceptEntry() did not return the original connection")
140 - }
141 - if entry.LeaseID != "lease-1" {
142 - t.Fatalf("AcceptEntry().LeaseID = %q, want %q", entry.LeaseID, "lease-1")
72 + t.Fatal("Accept() did not return the original connection")
73 }
74
75 plainConn, err := listener.Accept()
@@ -151,119 +81,3 @@ func TestListenerAcceptEntry(t *testing.T) {
81 t.Fatal("Accept() did not return the original connection")
82 }
83 }
154 -
155 -func TestListenerAcceptEntryReturnsTerminalError(t *testing.T) {
156 - t.Parallel()
157 -
158 - listener := newListener(context.Background())
159 - listener.accepted = make(chan acceptedConn, 1)
160 -
161 - wantErr := &types.APIRequestError{Code: types.APIErrorCodeUnauthorized, Message: "bad reverse token"}
162 - listener.fail(wantErr)
163 -
164 - conn, entry, err := listener.AcceptEntry()
165 - if conn != nil {
166 - t.Fatal("AcceptEntry() conn != nil, want nil")
167 - }
168 - if entry.RelayURL != "" || entry.LeaseID != "" || len(entry.Hostnames) != 0 || len(entry.Metadata.Tags) != 0 {
169 - t.Fatalf("AcceptEntry() entry = %#v, want zero value", entry)
170 - }
171 - if !errors.Is(err, wantErr) {
172 - t.Fatalf("AcceptEntry() error = %v, want %v", err, wantErr)
173 - }
174 -}
175 -
176 -func TestListenerStopEntryKeepsOtherLeasesActive(t *testing.T) {
177 - t.Parallel()
178 -
179 - listener := newListener(context.Background())
180 - entryA := &listenerLease{
181 - parent: listener,
182 - info: ListenerEntry{
183 - RelayURL: "https://relay-a.example.com",
184 - LeaseID: "lease-a",
185 - Hostnames: []string{"a.example.com"},
186 - },
187 - active: true,
188 - }
189 - entryB := &listenerLease{
190 - parent: listener,
191 - info: ListenerEntry{
192 - RelayURL: "https://relay-b.example.com",
193 - LeaseID: "lease-b",
194 - Hostnames: []string{"b.example.com"},
195 - },
196 - active: true,
197 - }
198 - listener.entries = []*listenerLease{entryA, entryB}
199 - listener.activeCount = 2
200 -
201 - entryA.stop(&types.APIRequestError{Code: types.APIErrorCodeUnauthorized, Message: "stopped"})
202 -
203 - if listener.isClosed() {
204 - t.Fatal("listener closed after stopping one entry, want active")
205 - }
206 -
207 - entries := listener.Entries()
208 - if len(entries) != 1 {
209 - t.Fatalf("Entries() len = %d, want 1", len(entries))
210 - }
211 - if entries[0].LeaseID != "lease-b" {
212 - t.Fatalf("Entries()[0].LeaseID = %q, want %q", entries[0].LeaseID, "lease-b")
213 - }
214 -}
215 -
216 -func TestListenerStopLastEntryCancelsListener(t *testing.T) {
217 - t.Parallel()
218 -
219 - listener := newListener(context.Background())
220 - entry := &listenerLease{
221 - parent: listener,
222 - info: ListenerEntry{
223 - RelayURL: "https://relay.example.com",
224 - LeaseID: "lease-1",
225 - Hostnames: []string{"app.relay.example.com"},
226 - },
227 - active: true,
228 - }
229 - listener.entries = []*listenerLease{entry}
230 - listener.activeCount = 1
231 - listener.accepted = make(chan acceptedConn, 1)
232 -
233 - wantErr := &types.APIRequestError{Code: types.APIErrorCodeLeaseNotFound, Message: "lease disappeared"}
234 - entry.stop(wantErr)
235 -
236 - if !listener.isClosed() {
237 - t.Fatal("listener is still active after stopping last entry")
238 - }
239 -
240 - _, _, err := listener.AcceptEntry()
241 - if !errors.Is(err, wantErr) {
242 - t.Fatalf("AcceptEntry() error = %v, want %v", err, wantErr)
243 - }
244 -}
245 -
246 -func TestListenerEntryCloneDeepCopiesMetadataTags(t *testing.T) {
247 - t.Parallel()
248 -
249 - entry := ListenerEntry{
250 - RelayURL: "https://relay.example.com",
251 - LeaseID: "lease-1",
252 - Hostnames: []string{"app.relay.example.com"},
253 - Metadata: types.LeaseMetadata{
254 - Owner: "alice",
255 - Tags: []string{"one", "two"},
256 - },
257 - }
258 -
259 - clone := entry.clone()
260 - clone.Hostnames[0] = "changed.example.com"
261 - clone.Metadata.Tags[0] = "changed"
262 -
263 - if entry.Hostnames[0] != "app.relay.example.com" {
264 - t.Fatalf("entry.Hostnames[0] = %q, want %q", entry.Hostnames[0], "app.relay.example.com")
265 - }
266 - if entry.Metadata.Tags[0] != "one" {
267 - t.Fatalf("entry.Metadata.Tags[0] = %q, want %q", entry.Metadata.Tags[0], "one")
268 - }
269 -}