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 -}