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, ®isterResp); err != nil {
152
+ if err := c.doJSON(ctx, http.MethodPost, types.PathSDKRegister, registerReq, ®isterResp); 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
-}