add mergelistener
rabbitprincess committed
Mar 7, 2026 at 22:17 UTC
93e4b8f0717274b0eff135ae3ca52a2415616364
7 files changed
+334
-96
Makefile
+1
-1
@@ -31,7 +31,7 @@ lint-auto:
31
golangci-lint run --fix $(GO_PACKAGES)
32
33
test:
34
- go test -v -race -coverprofile=coverage.out $(GO_PACKAGES)
34
+ go test -v -coverprofile=coverage.out $(GO_PACKAGES)
35
36
vuln:
37
govulncheck $(GO_PACKAGES)
cmd/portal-tunnel/main.go
+15
-4
@@ -5,6 +5,7 @@ import (
5
"errors"
6
"flag"
7
"fmt"
8
+ "net"
9
"os"
10
"os/signal"
11
"sync"
@@ -94,7 +95,7 @@ func runTunnel() error {
95
96
var connWG sync.WaitGroup
97
var connCount atomic.Int64
97
- relayDone := make(chan relayLoopResult, len(runtimes))
98
+ listeners := make([]net.Listener, 0, len(runtimes))
99
100
for _, runtime := range runtimes {
101
logger.Info().
@@ -102,14 +103,24 @@ func runTunnel() error {
103
Str("lease_id", runtime.listener.LeaseID()).
104
Strs("public_urls", runtime.listener.PublicURLs()).
105
Msg("relay tunnel ready")
105
- go runtime.run(ctx, flagHost, &connWG, &connCount, relayDone)
106
+ listeners = append(listeners, runtime.listener)
107
}
108
108
- waitErr := waitForRelayLoops(ctx, relayDone, len(runtimes))
109
+ relayListener, err := sdk.MergeListeners(listeners...)
110
+ if err != nil {
111
+ closeErr := closeRelayRuntimes(runtimes)
112
+ return errors.Join(fmt.Errorf("merge relay listeners: %w", err), closeErr)
113
+ }
114
+ go func() {
115
+ <-ctx.Done()
116
+ _ = relayListener.Close()
117
+ }()
118
+
119
+ waitErr := proxyRelayConnections(ctx, relayListener, flagHost, &connWG, &connCount)
120
if waitErr != nil {
121
stop()
122
}
112
- closeErr := closeRelayRuntimes(runtimes)
123
+ closeErr := errors.Join(relayListener.Close(), closeRelayRuntimes(runtimes))
124
if waitErr != nil {
125
logger.Error().Err(waitErr).Msg("relay supervisor exited with error")
126
}
cmd/portal-tunnel/relays.go
+26
-81
@@ -23,55 +23,6 @@ type relayRuntime struct {
23
relayURL string
24
}
25
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 {
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 {
52
- return
53
- }
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()).
62
- Msg("accepted relay connection")
63
-
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")
71
- }(connID, relayConn)
72
- }
73
-}
74
-
26
func startRelayRuntimes(ctx context.Context, relayURLs []string, req sdk.ListenRequest) ([]*relayRuntime, error) {
27
runtimes := make([]*relayRuntime, 0, len(relayURLs))
28
for _, relayURL := range relayURLs {
@@ -113,43 +64,37 @@ func closeRelayRuntimes(runtimes []*relayRuntime) error {
64
return closeErr
65
}
66
116
-func waitForRelayLoops(ctx context.Context, done <-chan relayLoopResult, relayCount int) error {
67
+func proxyRelayConnections(ctx context.Context, relayListener net.Listener, localAddr string, connWG *sync.WaitGroup, connCount *atomic.Int64) error {
68
logger := log.With().Str("component", "portal-tunnel").Logger()
118
- active := relayCount
119
-
120
- for active > 0 {
121
- result := <-done
122
- active--
69
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")
70
+ for {
71
+ relayConn, err := relayListener.Accept()
72
+ if err != nil {
73
+ switch {
74
+ case ctx.Err() != nil || errors.Is(err, context.Canceled):
75
+ return nil
76
+ case errors.Is(err, net.ErrClosed):
77
+ return errors.New("all relay listeners stopped")
78
+ default:
79
+ return err
80
+ }
81
}
145
- }
82
147
- select {
148
- case <-ctx.Done():
149
- return nil
150
- default:
83
+ connID := connCount.Add(1)
84
+ logger.Info().
85
+ Int64("conn_id", connID).
86
+ Str("remote_addr", relayConn.RemoteAddr().String()).
87
+ Msg("accepted relay connection")
88
+
89
+ connWG.Add(1)
90
+ go func(connID int64, relayConn net.Conn) {
91
+ defer connWG.Done()
92
+ if err := proxyConnection(ctx, localAddr, relayConn); err != nil {
93
+ logger.Error().Err(err).Int64("conn_id", connID).Msg("proxy connection failed")
94
+ }
95
+ logger.Info().Int64("conn_id", connID).Msg("proxy connection closed")
96
+ }(connID, relayConn)
97
}
152
- return errors.New("all relay listeners stopped")
98
}
99
100
func normalizeRelayURLs(raw string) ([]string, error) {
sdk/client.go
+3
-3
@@ -210,7 +210,7 @@ func (c *Client) Listen(ctx context.Context, req ListenRequest) (*Listener, erro
210
211
registerReq := types.RegisterRequest{
212
Name: req.Name,
213
- Hostnames: append([]string(nil), req.Hostnames...),
213
+ Hostnames: req.Hostnames,
214
Metadata: req.Metadata,
215
ReverseToken: reverseToken,
216
TLS: true,
@@ -235,8 +235,8 @@ func (c *Client) Listen(ctx context.Context, req ListenRequest) (*Listener, erro
235
ctxDone: listenerCtx.Done(),
236
cancel: cancel,
237
leaseID: registerResp.LeaseID,
238
- hostnames: append([]string(nil), registerResp.Hostnames...),
239
- metadata: cloneLeaseMetadata(registerResp.Metadata),
238
+ hostnames: registerResp.Hostnames,
239
+ metadata: registerResp.Metadata,
240
reverseToken: reverseToken,
241
leaseTTL: leaseTTL,
242
readyTarget: readyTarget,
sdk/helper.go
+133
@@ -7,6 +7,7 @@ import (
7
"net"
8
"net/http"
9
"strings"
10
+ "sync"
11
12
"golang.org/x/sync/errgroup"
13
)
@@ -68,6 +69,138 @@ func RunHTTP(ctx context.Context, relayListener net.Listener, handler http.Handl
69
return group.Wait()
70
}
71
72
+// MergeListeners fans in multiple listeners into one net.Listener. It keeps
73
+// serving accepts from remaining listeners when one listener stops, and returns
74
+// a terminal error only after all source listeners have stopped.
75
+func MergeListeners(listeners ...net.Listener) (net.Listener, error) {
76
+ if len(listeners) == 0 {
77
+ return nil, errors.New("at least one listener is required")
78
+ }
79
+
80
+ merged := &mergedListener{
81
+ listeners: make([]net.Listener, 0, len(listeners)),
82
+ accepted: make(chan net.Conn),
83
+ closed: make(chan struct{}),
84
+ }
85
+ for i, listener := range listeners {
86
+ if listener == nil {
87
+ return nil, fmt.Errorf("listener %d is nil", i)
88
+ }
89
+ merged.listeners = append(merged.listeners, listener)
90
+ }
91
+
92
+ merged.addr = merged.buildAddr()
93
+ merged.active = len(merged.listeners)
94
+ for _, listener := range merged.listeners {
95
+ source := listener
96
+ go merged.runAcceptLoop(source)
97
+ }
98
+ return merged, nil
99
+}
100
+
101
+type mergedListener struct {
102
+ listeners []net.Listener
103
+ accepted chan net.Conn
104
+ closed chan struct{}
105
+ addr net.Addr
106
+
107
+ closeOnce sync.Once
108
+ mu sync.Mutex
109
+ active int
110
+ terminalErr error
111
+}
112
+
113
+func (l *mergedListener) Accept() (net.Conn, error) {
114
+ conn, ok := <-l.accepted
115
+ if ok {
116
+ return conn, nil
117
+ }
118
+
119
+ l.mu.Lock()
120
+ defer l.mu.Unlock()
121
+ if l.terminalErr == nil {
122
+ return nil, net.ErrClosed
123
+ }
124
+ return nil, l.terminalErr
125
+}
126
+
127
+func (l *mergedListener) Close() error {
128
+ var closeErr error
129
+ l.closeOnce.Do(func() {
130
+ close(l.closed)
131
+ for _, listener := range l.listeners {
132
+ err := listener.Close()
133
+ if errors.Is(err, net.ErrClosed) {
134
+ err = nil
135
+ }
136
+ closeErr = errors.Join(closeErr, err)
137
+ }
138
+ l.recordTerminalError(closeErr)
139
+ })
140
+ return closeErr
141
+}
142
+
143
+func (l *mergedListener) Addr() net.Addr {
144
+ return l.addr
145
+}
146
+
147
+func (l *mergedListener) buildAddr() net.Addr {
148
+ if len(l.listeners) == 1 {
149
+ return l.listeners[0].Addr()
150
+ }
151
+
152
+ parts := make([]string, 0, len(l.listeners))
153
+ for _, listener := range l.listeners {
154
+ parts = append(parts, listener.Addr().String())
155
+ }
156
+ return listenerAddr("merged:" + strings.Join(parts, ","))
157
+}
158
+
159
+func (l *mergedListener) runAcceptLoop(listener net.Listener) {
160
+ for {
161
+ conn, err := listener.Accept()
162
+ if err != nil {
163
+ if !errors.Is(err, net.ErrClosed) {
164
+ l.recordTerminalError(fmt.Errorf("accept %s: %w", listener.Addr().String(), err))
165
+ }
166
+ l.finishWorker()
167
+ return
168
+ }
169
+
170
+ select {
171
+ case l.accepted <- conn:
172
+ case <-l.closed:
173
+ _ = conn.Close()
174
+ l.finishWorker()
175
+ return
176
+ }
177
+ }
178
+}
179
+
180
+func (l *mergedListener) finishWorker() {
181
+ l.mu.Lock()
182
+ l.active--
183
+ last := l.active == 0
184
+ if last && l.terminalErr == nil {
185
+ l.terminalErr = net.ErrClosed
186
+ }
187
+ l.mu.Unlock()
188
+
189
+ if last {
190
+ close(l.accepted)
191
+ }
192
+}
193
+
194
+func (l *mergedListener) recordTerminalError(err error) {
195
+ if err == nil {
196
+ return
197
+ }
198
+
199
+ l.mu.Lock()
200
+ l.terminalErr = errors.Join(l.terminalErr, err)
201
+ l.mu.Unlock()
202
+}
203
+
204
// SplitCSV splits a comma-separated string, trimming whitespace and dropping
205
// empty entries.
206
func SplitCSV(raw string) []string {
sdk/helper_test.go
+154
@@ -2,13 +2,167 @@ package sdk
2
3
import (
4
"context"
5
+ "errors"
6
"io"
7
"net"
8
"net/http"
9
+ "sort"
10
"testing"
11
"time"
12
)
13
14
+func TestMergeListenersRequiresInput(t *testing.T) {
15
+ t.Parallel()
16
+
17
+ listener, err := MergeListeners()
18
+ if err == nil {
19
+ t.Fatal("MergeListeners() error = nil, want error")
20
+ }
21
+ if listener != nil {
22
+ t.Fatalf("MergeListeners() listener = %#v, want nil", listener)
23
+ }
24
+}
25
+
26
+func TestMergeListenersAcceptsFromAllSources(t *testing.T) {
27
+ t.Parallel()
28
+
29
+ listener1, err := net.Listen("tcp", "127.0.0.1:0")
30
+ if err != nil {
31
+ t.Fatalf("Listen() error = %v", err)
32
+ }
33
+ defer listener1.Close()
34
+
35
+ listener2, err := net.Listen("tcp", "127.0.0.1:0")
36
+ if err != nil {
37
+ t.Fatalf("Listen() error = %v", err)
38
+ }
39
+ defer listener2.Close()
40
+
41
+ merged, err := MergeListeners(listener1, listener2)
42
+ if err != nil {
43
+ t.Fatalf("MergeListeners() error = %v", err)
44
+ }
45
+ defer merged.Close()
46
+
47
+ client1, err := net.Dial("tcp", listener1.Addr().String())
48
+ if err != nil {
49
+ t.Fatalf("Dial() error = %v", err)
50
+ }
51
+ defer client1.Close()
52
+
53
+ client2, err := net.Dial("tcp", listener2.Addr().String())
54
+ if err != nil {
55
+ t.Fatalf("Dial() error = %v", err)
56
+ }
57
+ defer client2.Close()
58
+
59
+ got := make([]string, 0, 2)
60
+ for range 2 {
61
+ conn, err := merged.Accept()
62
+ if err != nil {
63
+ t.Fatalf("Accept() error = %v", err)
64
+ }
65
+ got = append(got, conn.LocalAddr().String())
66
+ _ = conn.Close()
67
+ }
68
+
69
+ want := []string{listener1.Addr().String(), listener2.Addr().String()}
70
+ sort.Strings(got)
71
+ sort.Strings(want)
72
+ if len(got) != len(want) || got[0] != want[0] || got[1] != want[1] {
73
+ t.Fatalf("accepted local addrs = %v, want %v", got, want)
74
+ }
75
+}
76
+
77
+func TestMergeListenersContinuesAfterOneSourceCloses(t *testing.T) {
78
+ t.Parallel()
79
+
80
+ listener1, err := net.Listen("tcp", "127.0.0.1:0")
81
+ if err != nil {
82
+ t.Fatalf("Listen() error = %v", err)
83
+ }
84
+ defer listener1.Close()
85
+
86
+ listener2, err := net.Listen("tcp", "127.0.0.1:0")
87
+ if err != nil {
88
+ t.Fatalf("Listen() error = %v", err)
89
+ }
90
+ defer listener2.Close()
91
+
92
+ merged, err := MergeListeners(listener1, listener2)
93
+ if err != nil {
94
+ t.Fatalf("MergeListeners() error = %v", err)
95
+ }
96
+ defer merged.Close()
97
+
98
+ if err := listener1.Close(); err != nil && !errors.Is(err, net.ErrClosed) {
99
+ t.Fatalf("listener1.Close() error = %v", err)
100
+ }
101
+
102
+ acceptCh := make(chan net.Conn, 1)
103
+ errCh := make(chan error, 1)
104
+ go func() {
105
+ conn, acceptErr := merged.Accept()
106
+ if acceptErr != nil {
107
+ errCh <- acceptErr
108
+ return
109
+ }
110
+ acceptCh <- conn
111
+ }()
112
+
113
+ client, err := net.Dial("tcp", listener2.Addr().String())
114
+ if err != nil {
115
+ t.Fatalf("Dial() error = %v", err)
116
+ }
117
+ defer client.Close()
118
+
119
+ select {
120
+ case acceptErr := <-errCh:
121
+ t.Fatalf("Accept() error = %v", acceptErr)
122
+ case conn := <-acceptCh:
123
+ if conn.LocalAddr().String() != listener2.Addr().String() {
124
+ t.Fatalf("conn.LocalAddr() = %q, want %q", conn.LocalAddr().String(), listener2.Addr().String())
125
+ }
126
+ _ = conn.Close()
127
+ case <-time.After(3 * time.Second):
128
+ t.Fatal("Accept() did not return connection from surviving listener")
129
+ }
130
+}
131
+
132
+func TestMergeListenersCloseUnblocksAccept(t *testing.T) {
133
+ t.Parallel()
134
+
135
+ listener, err := net.Listen("tcp", "127.0.0.1:0")
136
+ if err != nil {
137
+ t.Fatalf("Listen() error = %v", err)
138
+ }
139
+ defer listener.Close()
140
+
141
+ merged, err := MergeListeners(listener)
142
+ if err != nil {
143
+ t.Fatalf("MergeListeners() error = %v", err)
144
+ }
145
+
146
+ errCh := make(chan error, 1)
147
+ go func() {
148
+ _, acceptErr := merged.Accept()
149
+ errCh <- acceptErr
150
+ }()
151
+
152
+ if err := merged.Close(); err != nil {
153
+ t.Fatalf("Close() error = %v", err)
154
+ }
155
+
156
+ select {
157
+ case acceptErr := <-errCh:
158
+ if !errors.Is(acceptErr, net.ErrClosed) {
159
+ t.Fatalf("Accept() error = %v, want net.ErrClosed", acceptErr)
160
+ }
161
+ case <-time.After(3 * time.Second):
162
+ t.Fatal("Accept() did not unblock after Close()")
163
+ }
164
+}
165
+
166
func TestRunHTTPAppRelayOnly(t *testing.T) {
167
t.Parallel()
168
sdk/listener.go
+2
-7
@@ -81,11 +81,11 @@ func (l *Listener) LeaseID() string {
81
}
82
83
func (l *Listener) Hostnames() []string {
84
- return append([]string(nil), l.hostnames...)
84
+ return l.hostnames
85
}
86
87
func (l *Listener) Metadata() types.LeaseMetadata {
88
- return cloneLeaseMetadata(l.metadata)
88
+ return l.metadata
89
}
90
91
func (l *Listener) PublicURLs() []string {
@@ -253,8 +253,3 @@ func (l *Listener) isClosed() bool {
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
260
-}