Refactor listener and exposure handling for improved performance and clarity
rabbitprincess committed
Apr 14, 2026 at 22:53 UTC
deb614b1b89bdecde3b36c38e09dba02ad99b694
6 files changed
+353
-520
portal/transport/stream_client.go
+11
-16
@@ -44,13 +44,13 @@ func (s *ClientStream) Accept(done <-chan struct{}) (net.Conn, error) {
44
45
func (s *ClientStream) RunSession(
46
ctx context.Context,
47
- open func(context.Context) (net.Conn, error),
48
- currentTLSConfig func() *tls.Config,
47
+ conn net.Conn,
48
+ tlsConfig *tls.Config,
49
) (bool, error) {
50
if s == nil {
51
return false, net.ErrClosed
52
}
53
- return s.runSession(ctx, open, currentTLSConfig)
53
+ return s.runSession(ctx, conn, tlsConfig)
54
}
55
56
func (s *ClientStream) ActiveSessions() int {
@@ -81,12 +81,11 @@ func (s *ClientStream) Drain() {
81
82
func (s *ClientStream) runSession(
83
ctx context.Context,
84
- open func(context.Context) (net.Conn, error),
85
- currentTLSConfig func() *tls.Config,
84
+ conn net.Conn,
85
+ tlsConfig *tls.Config,
86
) (bool, error) {
87
- conn, err := open(ctx)
88
- if err != nil {
89
- return false, err
87
+ if conn == nil {
88
+ return false, net.ErrClosed
89
}
90
s.sessionOpened()
91
defer s.sessionClosed()
@@ -104,7 +103,7 @@ func (s *ClientStream) runSession(
103
case types.MarkerKeepalive:
104
continue
105
case types.MarkerTLSStart:
107
- if err := s.activate(ctx, conn, currentTLSConfig); err != nil {
106
+ if err := s.activate(ctx, conn, tlsConfig); err != nil {
107
_ = conn.Close()
108
return true, err
109
}
@@ -122,16 +121,12 @@ func (s *ClientStream) runSession(
121
}
122
}
123
125
-func (s *ClientStream) activate(ctx context.Context, conn net.Conn, currentTLSConfig func() *tls.Config) error {
126
- var tlsCfg *tls.Config
127
- if currentTLSConfig != nil {
128
- tlsCfg = currentTLSConfig()
129
- }
130
- if tlsCfg == nil {
124
+func (s *ClientStream) activate(ctx context.Context, conn net.Conn, tlsConfig *tls.Config) error {
125
+ if tlsConfig == nil {
126
return errors.New("tls config is unavailable")
127
}
128
134
- tlsConn := tls.Server(conn, tlsCfg)
129
+ tlsConn := tls.Server(conn, tlsConfig)
130
handshakeCtx, cancel := context.WithTimeout(ctx, s.handshakeTimeout)
131
defer cancel()
132
if err := tlsConn.HandshakeContext(handshakeCtx); err != nil {
sdk/api_client.go
+50
-290
@@ -1,21 +1,14 @@
1
package sdk
2
3
import (
4
- "bufio"
5
- "bytes"
4
"context"
7
- "crypto/tls"
8
- "encoding/json"
5
"errors"
6
"fmt"
11
- "io"
7
"net"
8
"net/http"
9
"strings"
10
"time"
11
17
- "github.com/quic-go/quic-go"
18
-
12
"github.com/gosuda/portal-tunnel/v2/types"
13
"github.com/gosuda/portal-tunnel/v2/utils"
14
)
@@ -33,145 +26,37 @@ const (
26
27
var errRelayIncompatible = errors.New("relay is incompatible")
28
36
-func closeIdleHTTPClient(controlHTTPClient *http.Client) {
37
- if controlHTTPClient == nil {
38
- return
39
- }
40
- if transport, ok := controlHTTPClient.Transport.(*http.Transport); ok {
41
- transport.CloseIdleConnections()
42
- }
43
-}
44
-
29
// resetTransport tears down the cached HTTP client and TLS config so the next
30
// API call creates fresh TCP connections. Call this after detecting a system
31
// sleep/wake cycle where pooled connections are almost certainly dead.
32
func (l *listener) resetTransport() {
49
- if l == nil {
50
- return
51
- }
52
- l.mu.Lock()
53
- controlHTTPClient := l.controlHTTPClient
54
- l.controlHTTPClient = nil
55
- l.controlTLSConfig = nil
56
- l.mu.Unlock()
57
- closeIdleHTTPClient(controlHTTPClient)
58
-}
59
-
60
-func (l *listener) registerLease(ctx context.Context, ttl time.Duration, udpEnabled, tcpEnabled bool) (types.RegisterResponse, error) {
61
- if err := l.ensureControlHTTPClient(ctx); err != nil {
62
- return types.RegisterResponse{}, err
63
- }
64
- l.mu.Lock()
65
- controlHTTPClient := l.controlHTTPClient
66
- relayURL := l.relayURL
67
- identity := l.identity.Copy()
68
- metadata := l.metadata.Copy()
69
- l.mu.Unlock()
70
- if controlHTTPClient == nil {
71
- return types.RegisterResponse{}, errors.New("relay http client is unavailable")
72
- }
73
-
74
- var challenge types.RegisterChallengeResponse
75
- challengeReq := types.RegisterChallengeRequest{
76
- Identity: identity,
77
- Metadata: metadata,
78
- TTL: int(ttl / time.Second),
79
- UDPEnabled: udpEnabled,
80
- TCPEnabled: tcpEnabled,
81
- }
82
- if err := utils.HTTPDoAPIPath(ctx, controlHTTPClient, relayURL, http.MethodPost, types.PathSDKRegisterChallenge, challengeReq, nil, &challenge); err != nil {
83
- return types.RegisterResponse{}, err
84
- }
85
-
86
- signature, err := utils.SignEthereumPersonalMessage(challenge.SIWEMessage, identity.PrivateKey)
87
- if err != nil {
88
- return types.RegisterResponse{}, err
89
- }
90
-
91
- var resp types.RegisterResponse
92
- if err := utils.HTTPDoAPIPath(ctx, controlHTTPClient, relayURL, http.MethodPost, types.PathSDKRegister, types.RegisterRequest{
93
- ChallengeID: challenge.ChallengeID,
94
- SIWEMessage: challenge.SIWEMessage,
95
- SIWESignature: signature,
96
- ReportedIP: l.reportedIP(ctx),
97
- }, nil, &resp); err != nil {
98
- return types.RegisterResponse{}, err
33
+ if l.httpClient != nil {
34
+ if transport, ok := l.httpClient.Transport.(*http.Transport); ok {
35
+ transport.CloseIdleConnections()
36
+ }
37
}
100
- return resp, nil
38
+ l.httpClient = nil
39
+ l.tlsConfig = nil
40
}
41
103
-func (l *listener) ensureControlHTTPClient(ctx context.Context) error {
104
- if l == nil {
105
- return errors.New("listener is unavailable")
106
- }
107
- l.mu.Lock()
108
- if l.controlHTTPClient != nil && l.controlTLSConfig != nil {
109
- l.mu.Unlock()
42
+func (l *listener) initHTTPTransport(ctx context.Context) error {
43
+ if l.httpClient != nil {
44
return nil
45
}
112
- relayURL := l.relayURL
113
- requestTimeout := l.requestTimeout
114
- l.mu.Unlock()
115
- if relayURL == nil {
116
- return errors.New("relay url is unavailable")
117
- }
46
47
bootstrapCtx, cancel := context.WithTimeout(ctx, defaultDialTimeout+defaultHandshakeTimeout)
48
defer cancel()
49
122
- controlTLSConfig, controlHTTPClient, err := utils.NewHTTPTLSClient(bootstrapCtx, relayURL, requestTimeout)
50
+ tlsConfig, httpClient, err := utils.NewHTTPTLSClient(bootstrapCtx, l.relayURL, l.requestTimeout)
51
if err != nil {
52
return err
53
}
126
- if err := l.ensureCompatible(ctx, controlHTTPClient); err != nil {
127
- closeIdleHTTPClient(controlHTTPClient)
128
- return err
129
- }
130
-
131
- l.mu.Lock()
132
- if l.controlHTTPClient != nil && l.controlTLSConfig != nil {
133
- l.mu.Unlock()
134
- closeIdleHTTPClient(controlHTTPClient)
135
- return nil
136
- }
137
- oldControlHTTPClient := l.controlHTTPClient
138
- l.controlHTTPClient = controlHTTPClient
139
- l.controlTLSConfig = controlTLSConfig
140
- l.mu.Unlock()
141
- closeIdleHTTPClient(oldControlHTTPClient)
142
-
143
- return nil
144
-}
145
-
146
-func (l *listener) reportedIP(ctx context.Context) string {
147
- l.mu.Lock()
148
- if l.resolvedPublicIP != "" {
149
- ip := l.resolvedPublicIP
150
- l.mu.Unlock()
151
- return ip
152
- }
153
- l.mu.Unlock()
154
-
155
- ip := utils.ResolvePublicIP(ctx)
156
- l.mu.Lock()
157
- if l.resolvedPublicIP == "" {
158
- l.resolvedPublicIP = ip
159
- }
160
- ip = l.resolvedPublicIP
161
- l.mu.Unlock()
162
- return ip
163
-}
164
-
165
-func (l *listener) ensureCompatible(ctx context.Context, controlHTTPClient *http.Client) error {
166
- l.mu.Lock()
167
- relayURL := l.relayURL
168
- l.mu.Unlock()
169
- if relayURL == nil {
170
- return errors.New("relay url is unavailable")
171
- }
54
173
- var resp types.DomainResponse
174
- if err := utils.HTTPDoAPIPath(ctx, controlHTTPClient, relayURL, http.MethodGet, types.PathSDKDomain, nil, nil, &resp); err != nil {
55
+ var domainResp types.DomainResponse
56
+ if err := utils.HTTPDoAPIPath(ctx, httpClient, l.relayURL, http.MethodGet, types.PathSDKDomain, nil, nil, &domainResp); err != nil {
57
+ if transport, ok := httpClient.Transport.(*http.Transport); ok {
58
+ transport.CloseIdleConnections()
59
+ }
60
err = fmt.Errorf("check relay compatibility: %w", err)
61
var netErr net.Error
62
var apiErr *types.APIRequestError
@@ -183,31 +68,54 @@ func (l *listener) ensureCompatible(ctx context.Context, controlHTTPClient *http
68
}
69
return fmt.Errorf("%w: %w", errRelayIncompatible, err)
70
}
186
- protocolVersion := strings.TrimSpace(resp.ProtocolVersion)
71
+ protocolVersion := strings.TrimSpace(domainResp.ProtocolVersion)
72
if protocolVersion != types.SDKVersion {
73
+ if transport, ok := httpClient.Transport.(*http.Transport); ok {
74
+ transport.CloseIdleConnections()
75
+ }
76
return fmt.Errorf("%w: relay sdk protocol version mismatch: relay=%q client=%q", errRelayIncompatible, protocolVersion, types.SDKVersion)
77
}
78
+
79
+ l.httpClient = httpClient
80
+ l.tlsConfig = tlsConfig
81
return nil
82
}
83
193
-func (l *listener) renewRegisteredLease(ctx context.Context, ttl time.Duration, accessToken string) (types.RenewResponse, error) {
194
- if err := l.ensureControlHTTPClient(ctx); err != nil {
195
- return types.RenewResponse{}, err
84
+func (l *listener) registerLease(ctx context.Context, ttl time.Duration, udpEnabled, tcpEnabled bool) (types.RegisterResponse, error) {
85
+ var challenge types.RegisterChallengeResponse
86
+ if err := utils.HTTPDoAPIPath(ctx, l.httpClient, l.relayURL, http.MethodPost, types.PathSDKRegisterChallenge, types.RegisterChallengeRequest{
87
+ Identity: l.identity,
88
+ Metadata: l.metadata,
89
+ TTL: int(ttl / time.Second),
90
+ UDPEnabled: udpEnabled,
91
+ TCPEnabled: tcpEnabled,
92
+ }, nil, &challenge); err != nil {
93
+ return types.RegisterResponse{}, err
94
+ }
95
+
96
+ signature, err := utils.SignEthereumPersonalMessage(challenge.SIWEMessage, l.identity.PrivateKey)
97
+ if err != nil {
98
+ return types.RegisterResponse{}, err
99
}
100
198
- l.mu.Lock()
199
- controlHTTPClient := l.controlHTTPClient
200
- relayURL := l.relayURL
201
- l.mu.Unlock()
202
- if controlHTTPClient == nil {
203
- return types.RenewResponse{}, errors.New("relay http client is unavailable")
101
+ var resp types.RegisterResponse
102
+ if err := utils.HTTPDoAPIPath(ctx, l.httpClient, l.relayURL, http.MethodPost, types.PathSDKRegister, types.RegisterRequest{
103
+ ChallengeID: challenge.ChallengeID,
104
+ SIWEMessage: challenge.SIWEMessage,
105
+ SIWESignature: signature,
106
+ ReportedIP: utils.ResolvePublicIP(ctx),
107
+ }, nil, &resp); err != nil {
108
+ return types.RegisterResponse{}, err
109
}
110
+ return resp, nil
111
+}
112
113
+func (l *listener) renewRegisteredLease(ctx context.Context, ttl time.Duration, accessToken string) (types.RenewResponse, error) {
114
var resp types.RenewResponse
207
- if err := utils.HTTPDoAPIPath(ctx, controlHTTPClient, relayURL, http.MethodPost, types.PathSDKRenew, types.RenewRequest{
115
+ if err := utils.HTTPDoAPIPath(ctx, l.httpClient, l.relayURL, http.MethodPost, types.PathSDKRenew, types.RenewRequest{
116
AccessToken: accessToken,
117
TTL: int(ttl / time.Second),
210
- ReportedIP: l.reportedIP(ctx),
118
+ ReportedIP: utils.ResolvePublicIP(ctx),
119
}, nil, &resp); err != nil {
120
return types.RenewResponse{}, err
121
}
@@ -215,155 +123,7 @@ func (l *listener) renewRegisteredLease(ctx context.Context, ttl time.Duration,
123
}
124
125
func (l *listener) unregisterLease(ctx context.Context, accessToken string) error {
218
- if err := l.ensureControlHTTPClient(ctx); err != nil {
219
- return err
220
- }
221
- l.mu.Lock()
222
- controlHTTPClient := l.controlHTTPClient
223
- relayURL := l.relayURL
224
- l.mu.Unlock()
225
- if controlHTTPClient == nil {
226
- return errors.New("relay http client is unavailable")
227
- }
228
- return utils.HTTPDoAPIPath(ctx, controlHTTPClient, relayURL, http.MethodPost, types.PathSDKUnregister, types.UnregisterRequest{
126
+ return utils.HTTPDoAPIPath(ctx, l.httpClient, l.relayURL, http.MethodPost, types.PathSDKUnregister, types.UnregisterRequest{
127
AccessToken: accessToken,
128
}, nil, nil)
129
}
232
-
233
-func (l *listener) openReverseSession(ctx context.Context, accessToken string) (net.Conn, error) {
234
- if err := l.ensureControlHTTPClient(ctx); err != nil {
235
- return nil, err
236
- }
237
- l.mu.Lock()
238
- controlTLSConfig := l.controlTLSConfig
239
- relayURL := l.relayURL
240
- dialTimeout := l.dialTimeout
241
- l.mu.Unlock()
242
- if controlTLSConfig == nil {
243
- return nil, errors.New("relay tls config is unavailable")
244
- }
245
-
246
- dialer := &tls.Dialer{
247
- NetDialer: &net.Dialer{Timeout: dialTimeout},
248
- Config: controlTLSConfig.Clone(),
249
- }
250
-
251
- conn, err := dialer.DialContext(ctx, "tcp", utils.EnsurePort(relayURL.Host))
252
- if err != nil {
253
- return nil, err
254
- }
255
-
256
- req := &http.Request{
257
- Method: http.MethodGet,
258
- URL: utils.ResolveAPIURL(relayURL, types.PathSDKConnect),
259
- Host: relayURL.Host,
260
- Header: make(http.Header),
261
- }
262
- req.Header.Set(types.HeaderAccessToken, accessToken)
263
- req.Header.Set("Connection", "keep-alive")
264
-
265
- if writeErr := req.Write(conn); writeErr != nil {
266
- _ = conn.Close()
267
- return nil, writeErr
268
- }
269
-
270
- reader := bufio.NewReader(conn)
271
- resp, err := http.ReadResponse(reader, req)
272
- if err != nil {
273
- _ = conn.Close()
274
- return nil, err
275
- }
276
- defer resp.Body.Close()
277
-
278
- if resp.StatusCode != http.StatusOK {
279
- apiErr := utils.DecodeAPIRequestError(resp)
280
- _ = conn.Close()
281
- return nil, apiErr
282
- }
283
-
284
- return wrapBufferedConn(conn, reader), nil
285
-}
286
-
287
-type bufferedConn struct {
288
- net.Conn
289
- reader *bytes.Reader
290
-}
291
-
292
-func wrapBufferedConn(conn net.Conn, reader *bufio.Reader) net.Conn {
293
- if reader == nil || reader.Buffered() == 0 {
294
- return conn
295
- }
296
- buf := make([]byte, reader.Buffered())
297
- if _, err := io.ReadFull(reader, buf); err != nil {
298
- return conn
299
- }
300
- return &bufferedConn{Conn: conn, reader: bytes.NewReader(buf)}
301
-}
302
-
303
-func (c *bufferedConn) Read(p []byte) (int, error) {
304
- if c.reader != nil && c.reader.Len() > 0 {
305
- return c.reader.Read(p)
306
- }
307
- return c.Conn.Read(p)
308
-}
309
-
310
-// openQUICSession opens a QUIC connection to the relay for datagram transport.
311
-func (l *listener) openQUICSession(ctx context.Context, accessToken string, sniPort int) (*quic.Conn, error) {
312
- if err := l.ensureControlHTTPClient(ctx); err != nil {
313
- return nil, err
314
- }
315
-
316
- l.mu.Lock()
317
- controlTLSConfig := l.controlTLSConfig
318
- relayURL := l.relayURL
319
- l.mu.Unlock()
320
- if controlTLSConfig == nil {
321
- return nil, errors.New("relay tls config is unavailable")
322
- }
323
-
324
- tlsConf := controlTLSConfig.Clone()
325
- tlsConf.NextProtos = []string{"portal-tunnel"}
326
-
327
- quicConf := &quic.Config{
328
- EnableDatagrams: true,
329
- KeepAlivePeriod: 15 * time.Second,
330
- MaxIdleTimeout: 60 * time.Second,
331
- }
332
-
333
- host := strings.TrimSpace(relayURL.Hostname())
334
- if host == "" {
335
- host = strings.TrimSpace(relayURL.Host)
336
- }
337
- dialAddr := net.JoinHostPort(host, fmt.Sprintf("%d", sniPort))
338
- conn, err := quic.DialAddr(ctx, dialAddr, tlsConf, quicConf)
339
- if err != nil {
340
- return nil, fmt.Errorf("quic dial: %w", err)
341
- }
342
-
343
- stream, err := conn.OpenStreamSync(ctx)
344
- if err != nil {
345
- _ = conn.CloseWithError(1, "stream open failed")
346
- return nil, fmt.Errorf("open control stream: %w", err)
347
- }
348
-
349
- controlMsg := types.QUICControlMessage{
350
- AccessToken: accessToken,
351
- }
352
- if err := json.NewEncoder(stream).Encode(controlMsg); err != nil {
353
- _ = conn.CloseWithError(1, "control write failed")
354
- return nil, fmt.Errorf("write control: %w", err)
355
- }
356
-
357
- _ = stream.SetReadDeadline(time.Now().Add(10 * time.Second))
358
- var resp types.QUICControlResponse
359
- if err := json.NewDecoder(io.LimitReader(stream, 4096)).Decode(&resp); err != nil {
360
- _ = conn.CloseWithError(1, "control read failed")
361
- return nil, fmt.Errorf("read control response: %w", err)
362
- }
363
- if !resp.OK {
364
- _ = conn.CloseWithError(1, resp.Error)
365
- return nil, fmt.Errorf("quic connect rejected: %s", resp.Error)
366
- }
367
-
368
- return conn, nil
369
-}
sdk/expose.go
+11
-6
@@ -115,7 +115,7 @@ func Expose(ctx context.Context, cfg ExposeConfig) (*Exposure, error) {
115
tcpEnabled: cfg.TCPEnabled,
116
banMITM: cfg.BanMITM,
117
maxActiveRelays: cfg.MaxActiveRelays,
118
- metadata: cfg.Metadata.Copy(),
118
+ metadata: cfg.Metadata,
119
accepted: make(chan net.Conn, max(len(relayURLs)*defaultReadyTarget*2, 1)),
120
datagrams: make(chan types.DatagramFrame, max(len(relayURLs)*32, 1)),
121
relaySet: discovery.NewRelaySet(relayURLs),
@@ -160,7 +160,7 @@ func (e *Exposure) Addr() net.Addr {
160
}
161
162
func (e *Exposure) Identity() types.Identity {
163
- return e.identity.Copy()
163
+ return e.identity
164
}
165
166
func (e *Exposure) AcceptDatagram() (types.DatagramFrame, error) {
@@ -414,12 +414,12 @@ func (e *Exposure) reconcileRelayListeners(failOnError bool) error {
414
retryCount = 0
415
}
416
listener, err := newListener(context.Background(), relayURL, listenerConfig{
417
- Identity: e.identity.Copy(),
417
+ Identity: e.identity,
418
UDPEnabled: e.udpEnabled,
419
TCPEnabled: e.tcpEnabled,
420
BanMITM: e.banMITM,
421
RetryCount: retryCount,
422
- Metadata: e.metadata.Copy(),
422
+ Metadata: e.metadata,
423
relaySet: e.relaySet,
424
})
425
if err != nil {
@@ -485,7 +485,7 @@ func (e *Exposure) runListenerAcceptLoop(listener *listener) {
485
log.Warn().
486
Err(err).
487
Str("relay_url", relayURL).
488
- Str("address", listener.address()).
488
+ Str("address", listener.identity.Address).
489
Msg("datagram accept failed")
490
return
491
}
@@ -509,7 +509,12 @@ func (e *Exposure) runListenerAcceptLoop(listener *listener) {
509
for {
510
conn, err := listener.Accept()
511
if err != nil {
512
- if listener.closed() || errors.Is(err, net.ErrClosed) {
512
+ select {
513
+ case <-listener.doneCh:
514
+ return
515
+ default:
516
+ }
517
+ if errors.Is(err, net.ErrClosed) {
518
return
519
}
520
log.Warn().Err(err).Str("relay_url", relayURL).Msg("exposure listener accept failed")
sdk/listener.go
+236
-175
@@ -1,8 +1,11 @@
1
package sdk
2
3
import (
4
+ "bufio"
5
+ "bytes"
6
"context"
7
"crypto/tls"
8
+ "encoding/json"
9
"errors"
10
"fmt"
11
"io"
@@ -13,6 +16,7 @@ import (
16
"sync"
17
"time"
18
19
+ "github.com/quic-go/quic-go"
20
"github.com/rs/zerolog/log"
21
22
"github.com/gosuda/portal-tunnel/v2/portal/discovery"
@@ -42,40 +46,35 @@ type listenerConfig struct {
46
var errLeaseRefreshRequired = errors.New("lease refresh required")
47
48
type listener struct {
45
- relayURL *url.URL
46
- controlHTTPClient *http.Client
47
- controlTLSConfig *tls.Config
48
- dialTimeout time.Duration
49
- requestTimeout time.Duration
50
- resolvedPublicIP string
51
- accessToken string
52
- expiresAt time.Time
53
- sniPort int
54
-
55
- cancel context.CancelFunc
56
- doneCh <-chan struct{}
57
-
58
- readyTarget int
59
- retryCount int
60
- retryWait time.Duration
61
- leaseTTL time.Duration
62
- renewBefore time.Duration
49
+ cancel context.CancelFunc
50
+ doneCh <-chan struct{}
51
+ closeOnce sync.Once
52
+
53
+ relayURL *url.URL
54
+ identity types.Identity
55
+ metadata types.LeaseMetadata
56
+ relaySet *discovery.RelaySet
57
+ tcpEnabled bool
58
+ dialTimeout time.Duration
59
+ requestTimeout time.Duration
60
+ readyTarget int
61
+ retryCount int
62
+ retryWait time.Duration
63
+ leaseTTL time.Duration
64
+ renewBefore time.Duration
65
66
stream *transport.ClientStream
67
datagram *transport.ClientDatagram
68
mitmManager *mitmManager
69
68
- closeOnce sync.Once
69
-
70
- banMITM bool
71
- tcpEnabled bool
72
- identity types.Identity
73
- relaySet *discovery.RelaySet
74
- mu sync.Mutex
75
- hostname string
76
- udpAddr string
77
- metadata types.LeaseMetadata
70
+ httpClient *http.Client
71
+ tlsConfig *tls.Config
72
73
+ hostname string
74
+ udpAddr string
75
+ accessToken string
76
+ expiresAt time.Time
77
+ sniPort int
78
tenantTLSConfig *tls.Config
79
tenantTLSCloser io.Closer
80
}
@@ -104,9 +103,13 @@ func newListener(ctx context.Context, relayURL string, cfg listenerConfig) (*lis
103
}
104
105
l := &listener{
107
- doneCh: listenerCtx.Done(),
106
cancel: cancel,
107
+ doneCh: listenerCtx.Done(),
108
relayURL: relayurl,
109
+ identity: cfg.Identity,
110
+ metadata: cfg.Metadata,
111
+ relaySet: cfg.relaySet,
112
+ tcpEnabled: cfg.TCPEnabled,
113
dialTimeout: dialTimeout,
114
requestTimeout: requestTimeout,
115
readyTarget: readyTarget,
@@ -114,20 +117,15 @@ func newListener(ctx context.Context, relayURL string, cfg listenerConfig) (*lis
117
retryWait: retryWait,
118
leaseTTL: leaseTTL,
119
renewBefore: renewBefore,
117
- identity: cfg.Identity.Copy(),
118
- metadata: cfg.Metadata.Copy(),
119
- banMITM: cfg.BanMITM,
120
- tcpEnabled: cfg.TCPEnabled,
121
- relaySet: cfg.relaySet,
120
}
123
- l.mitmManager = newMITMManager(listenerCtx, l)
121
+ l.mitmManager = newMITMManager(listenerCtx, l, cfg.BanMITM)
122
l.stream = transport.NewClientStream(readyTarget, handshakeTimeout)
123
if cfg.UDPEnabled {
124
l.datagram = transport.NewClientDatagram(func(err error) {
125
log.Info().
126
Err(err).
127
Str("component", "sdk-datagram-plane").
130
- Str("address", l.address()).
128
+ Str("address", l.identity.Address).
129
Msg("quic datagram plane disconnected; waiting to reconnect")
130
})
131
}
@@ -159,7 +157,7 @@ func (l *listener) run(ctx context.Context) {
157
log.Error().
158
Err(err).
159
Str("relay_url", relayURL).
162
- Str("address", l.address()).
160
+ Str("address", l.identity.Address).
161
Msg("lease registration failed; closing listener")
162
_ = l.Close()
163
return
@@ -174,7 +172,7 @@ func (l *listener) run(ctx context.Context) {
172
173
retries = 0
174
publicURL := l.publicURL()
177
- event := log.Info().Str("address", l.address())
175
+ event := log.Info().Str("address", l.identity.Address)
176
if publicURL != "" {
177
event.Msg("service ready at " + publicURL)
178
} else {
@@ -187,7 +185,7 @@ func (l *listener) run(ctx context.Context) {
185
}
186
187
if errors.Is(err, errLeaseRefreshRequired) {
190
- _, _, _, tenantTLSCloser := l.clearLease("lease refresh required")
188
+ _, _, tenantTLSCloser := l.clearLease("lease refresh required")
189
if tenantTLSCloser != nil {
190
_ = tenantTLSCloser.Close()
191
}
@@ -199,7 +197,7 @@ func (l *listener) run(ctx context.Context) {
197
log.Error().
198
Err(err).
199
Str("relay_url", relayURL).
202
- Str("address", l.address()).
200
+ Str("address", l.identity.Address).
201
Msg("listener connection retry budget exhausted; closing listener")
202
_ = l.Close()
203
return
@@ -213,21 +211,16 @@ func (l *listener) Close() error {
211
l.cancel()
212
}
213
216
- identity, registered, accessToken, tenantTLSCloser := l.clearLease("")
217
-
218
- l.mu.Lock()
219
- stream := l.stream
220
- datagram := l.datagram
221
- l.mu.Unlock()
214
+ registered, accessToken, tenantTLSCloser := l.clearLease("")
215
223
- if stream != nil {
224
- stream.Drain()
216
+ if l.stream != nil {
217
+ l.stream.Drain()
218
}
226
- if datagram != nil {
227
- datagram.Close()
219
+ if l.datagram != nil {
220
+ l.datagram.Close()
221
}
222
230
- if registered && identity.Key() != "" && strings.TrimSpace(accessToken) != "" {
223
+ if registered && l.identity.Key() != "" && accessToken != "" {
224
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
225
closeErr = errors.Join(closeErr, l.unregisterLease(ctx, accessToken))
226
cancel()
@@ -240,29 +233,24 @@ func (l *listener) Close() error {
233
return closeErr
234
}
235
243
-func (l *listener) clearLease(reason string) (types.Identity, bool, string, io.Closer) {
244
- l.mu.Lock()
245
- identity := l.identity.Copy()
236
+func (l *listener) clearLease(reason string) (bool, string, io.Closer) {
237
registered := l.hostname != ""
238
accessToken := l.accessToken
239
tenantTLSCloser := l.tenantTLSCloser
249
- datagram := l.datagram
240
l.hostname = ""
241
l.udpAddr = ""
252
- l.tenantTLSConfig = nil
253
- l.tenantTLSCloser = nil
242
l.accessToken = ""
243
l.expiresAt = time.Time{}
244
l.sniPort = 0
257
- l.mu.Unlock()
258
-
245
+ l.tenantTLSConfig = nil
246
+ l.tenantTLSCloser = nil
247
if l.mitmManager != nil {
248
l.mitmManager.reset()
249
}
262
- if datagram != nil && reason != "" {
263
- datagram.Clear(reason)
250
+ if l.datagram != nil && reason != "" {
251
+ l.datagram.Clear(reason)
252
}
265
- return identity, registered, accessToken, tenantTLSCloser
253
+ return registered, accessToken, tenantTLSCloser
254
}
255
256
func (l *listener) Accept() (net.Conn, error) {
@@ -280,7 +268,7 @@ func (l *listener) Accept() (net.Conn, error) {
268
log.Debug().
269
Err(handleErr).
270
Str("relay_url", l.relayURL.String()).
283
- Str("address", l.address()).
271
+ Str("address", l.identity.Address).
272
Msg("mitm self-probe handling failed")
273
}
274
if handled {
@@ -291,7 +279,7 @@ func (l *listener) Accept() (net.Conn, error) {
279
}
280
281
func (l *listener) acceptDatagram() (types.DatagramFrame, error) {
294
- if l == nil || l.datagram == nil {
282
+ if l.datagram == nil {
283
return types.DatagramFrame{}, net.ErrClosed
284
}
285
@@ -301,65 +289,52 @@ func (l *listener) acceptDatagram() (types.DatagramFrame, error) {
289
}
290
291
frame.Payload = append([]byte(nil), frame.Payload...)
304
- l.mu.Lock()
305
- frame.Address = l.identity.Address
292
frame.UDPAddr = l.udpAddr
293
+ frame.Address = l.identity.Address
294
if l.relayURL != nil {
295
frame.RelayURL = l.relayURL.String()
296
}
310
- l.mu.Unlock()
297
return frame, nil
298
}
299
300
func (l *listener) sendDatagram(frame types.DatagramFrame) error {
315
- if l == nil || l.datagram == nil {
301
+ if l.datagram == nil {
302
return net.ErrClosed
303
}
304
319
- l.mu.Lock()
320
- datagram := l.datagram
321
- address := l.identity.Address
322
- l.mu.Unlock()
323
-
324
- if address == "" || datagram == nil {
305
+ if l.identity.Address == "" {
306
return net.ErrClosed
307
}
327
- if frameAddress := strings.TrimSpace(frame.Address); frameAddress != "" && frameAddress != address {
308
+ if frameAddress := strings.TrimSpace(frame.Address); frameAddress != "" && frameAddress != l.identity.Address {
309
return errors.New("datagram frame targets stale address")
310
}
330
- return datagram.Send(frame.FlowID, frame.Payload)
311
+ return l.datagram.Send(frame.FlowID, frame.Payload)
312
}
313
314
func (l *listener) datagramReady() (string, bool, bool) {
334
- if l == nil || l.datagram == nil {
315
+ if l.datagram == nil {
316
return "", false, false
317
}
318
338
- l.mu.Lock()
319
hostname := l.hostname
320
udpAddr := l.udpAddr
341
- datagram := l.datagram
342
- l.mu.Unlock()
343
-
344
- ready := datagram != nil && datagram.Connected() && udpAddr != ""
345
- pending := !ready && !l.closed() && (hostname == "" || udpAddr != "")
321
+ ready := l.datagram.Connected() && udpAddr != ""
322
+ closed := false
323
+ select {
324
+ case <-l.doneCh:
325
+ closed = true
326
+ default:
327
+ }
328
+ pending := !ready && !closed && (hostname == "" || udpAddr != "")
329
return udpAddr, ready, pending
330
}
331
349
-func (l *listener) address() string {
350
- l.mu.Lock()
351
- defer l.mu.Unlock()
352
- return l.identity.Address
353
-}
354
-
332
func (l *listener) publicURL() string {
356
- if l == nil || l.relayURL == nil {
333
+ if l.relayURL == nil {
334
return ""
335
}
336
360
- l.mu.Lock()
337
hostname := l.hostname
362
- l.mu.Unlock()
338
if hostname == "" {
339
return ""
340
}
@@ -380,13 +355,15 @@ func (l *listener) publicURL() string {
355
}
356
357
func (l *listener) runLease(ctx context.Context) error {
383
- l.mu.Lock()
384
- identity := l.identity.Copy()
385
- accessToken := l.accessToken
386
- sniPort := l.sniPort
358
+ registered := l.hostname != ""
359
tlsConfig := l.tenantTLSConfig
360
+ if !registered {
361
+ if ctx.Err() != nil {
362
+ return ctx.Err()
363
+ }
364
+ return errLeaseRefreshRequired
365
+ }
366
readyTarget := l.readyTarget
389
- l.mu.Unlock()
367
368
leaseCtx, cancel := context.WithCancel(ctx)
369
defer cancel()
@@ -395,7 +372,7 @@ func (l *listener) runLease(ctx context.Context) error {
372
if l.stream != nil && readyTarget > 0 {
373
for range readyTarget {
374
go func() {
398
- if err := l.runReverseSessionLoop(leaseCtx, accessToken, tlsConfig); err != nil {
375
+ if err := l.runReverseSessionLoop(leaseCtx, tlsConfig); err != nil {
376
select {
377
case errCh <- err:
378
case <-leaseCtx.Done():
@@ -405,7 +382,7 @@ func (l *listener) runLease(ctx context.Context) error {
382
}
383
}
384
if l.datagram != nil {
408
- go l.runDatagramLoop(leaseCtx, identity, accessToken, sniPort)
385
+ go l.runDatagramLoop(leaseCtx)
386
}
387
go func() {
388
if err := l.runRenewLoop(leaseCtx); err != nil {
@@ -425,25 +402,26 @@ func (l *listener) runLease(ctx context.Context) error {
402
}
403
}
404
428
-func (l *listener) runReverseSessionLoop(ctx context.Context, accessToken string, tlsConfig *tls.Config) error {
405
+func (l *listener) runReverseSessionLoop(ctx context.Context, tlsConfig *tls.Config) error {
406
if l.stream == nil {
407
return nil
408
}
409
410
var retries int
411
for {
435
- claimed, err := l.stream.RunSession(
436
- ctx,
437
- func(ctx context.Context) (net.Conn, error) {
438
- if strings.TrimSpace(accessToken) == "" {
439
- return nil, errors.New("access token is not available")
440
- }
441
- return l.openReverseSession(ctx, accessToken)
442
- },
443
- func() *tls.Config {
444
- return tlsConfig
445
- },
446
- )
412
+ conn, err := l.openReverseSession(ctx)
413
+ if err != nil {
414
+ if errors.Is(err, context.Canceled) || errors.Is(err, net.ErrClosed) {
415
+ return nil
416
+ }
417
+ retries++
418
+ if !l.waitRetry(ctx, "reverse session connect", err, retries) {
419
+ return err
420
+ }
421
+ continue
422
+ }
423
+
424
+ claimed, err := l.stream.RunSession(ctx, conn, tlsConfig)
425
switch {
426
case err == nil:
427
retries = 0
@@ -460,10 +438,11 @@ func (l *listener) runReverseSessionLoop(ctx context.Context, accessToken string
438
}
439
}
440
463
-func (l *listener) runDatagramLoop(ctx context.Context, identity types.Identity, accessToken string, sniPort int) {
441
+func (l *listener) runDatagramLoop(ctx context.Context) {
442
if l.datagram == nil {
443
return
444
}
445
+ address := l.identity.Address
446
447
for {
448
select {
@@ -473,12 +452,12 @@ func (l *listener) runDatagramLoop(ctx context.Context, identity types.Identity,
452
default:
453
}
454
476
- conn, err := l.openQUICSession(ctx, accessToken, sniPort)
455
+ conn, err := l.openQUICSession(ctx)
456
if err != nil {
457
log.Info().
458
Err(err).
459
Str("component", "sdk-datagram-plane").
481
- Str("address", identity.Address).
460
+ Str("address", address).
461
Msg("quic datagram plane unavailable; retrying")
462
if !utils.SleepOrDone(ctx, 2*time.Second) {
463
l.datagram.Clear("lease stopped")
@@ -489,7 +468,7 @@ func (l *listener) runDatagramLoop(ctx context.Context, identity types.Identity,
468
469
log.Info().
470
Str("component", "sdk-datagram-plane").
492
- Str("address", identity.Address).
471
+ Str("address", address).
472
Str("remote_addr", conn.RemoteAddr().String()).
473
Msg("quic tunnel connected")
474
@@ -501,7 +480,7 @@ func (l *listener) runDatagramLoop(ctx context.Context, identity types.Identity,
480
log.Info().
481
Err(err).
482
Str("component", "sdk-datagram-plane").
504
- Str("address", identity.Address).
483
+ Str("address", address).
484
Msg("quic datagram plane did not bind cleanly; retrying")
485
if !utils.SleepOrDone(ctx, time.Second) {
486
return
@@ -522,16 +501,138 @@ func (l *listener) runDatagramLoop(ctx context.Context, identity types.Identity,
501
}
502
}
503
525
-func (l *listener) runRenewLoop(ctx context.Context) error {
526
- interval := l.leaseTTL / 2
527
- if interval <= 0 {
528
- interval = 30 * time.Second
504
+func (l *listener) openReverseSession(ctx context.Context) (net.Conn, error) {
505
+ if l.accessToken == "" {
506
+ return nil, errors.New("access token is not available")
507
}
530
- if l.renewBefore > 0 && l.leaseTTL > l.renewBefore {
531
- interval = l.leaseTTL - l.renewBefore
508
+
509
+ dialer := &tls.Dialer{
510
+ NetDialer: &net.Dialer{Timeout: l.dialTimeout},
511
+ Config: l.tlsConfig.Clone(),
512
}
533
- if interval <= 0 {
534
- interval = 30 * time.Second
513
+
514
+ conn, err := dialer.DialContext(ctx, "tcp", utils.EnsurePort(l.relayURL.Host))
515
+ if err != nil {
516
+ return nil, err
517
+ }
518
+
519
+ req := &http.Request{
520
+ Method: http.MethodGet,
521
+ URL: utils.ResolveAPIURL(l.relayURL, types.PathSDKConnect),
522
+ Host: l.relayURL.Host,
523
+ Header: make(http.Header),
524
+ }
525
+ req.Header.Set(types.HeaderAccessToken, l.accessToken)
526
+ req.Header.Set("Connection", "keep-alive")
527
+
528
+ if writeErr := req.Write(conn); writeErr != nil {
529
+ _ = conn.Close()
530
+ return nil, writeErr
531
+ }
532
+
533
+ reader := bufio.NewReader(conn)
534
+ resp, err := http.ReadResponse(reader, req)
535
+ if err != nil {
536
+ _ = conn.Close()
537
+ return nil, err
538
+ }
539
+ defer resp.Body.Close()
540
+
541
+ if resp.StatusCode != http.StatusOK {
542
+ apiErr := utils.DecodeAPIRequestError(resp)
543
+ _ = conn.Close()
544
+ return nil, apiErr
545
+ }
546
+
547
+ return wrapBufferedConn(conn, reader), nil
548
+}
549
+
550
+type bufferedConn struct {
551
+ net.Conn
552
+ reader *bytes.Reader
553
+}
554
+
555
+func wrapBufferedConn(conn net.Conn, reader *bufio.Reader) net.Conn {
556
+ if reader == nil || reader.Buffered() == 0 {
557
+ return conn
558
+ }
559
+ buf := make([]byte, reader.Buffered())
560
+ if _, err := io.ReadFull(reader, buf); err != nil {
561
+ return conn
562
+ }
563
+ return &bufferedConn{Conn: conn, reader: bytes.NewReader(buf)}
564
+}
565
+
566
+func (c *bufferedConn) Read(p []byte) (int, error) {
567
+ if c.reader != nil && c.reader.Len() > 0 {
568
+ return c.reader.Read(p)
569
+ }
570
+ return c.Conn.Read(p)
571
+}
572
+
573
+func (l *listener) openQUICSession(ctx context.Context) (*quic.Conn, error) {
574
+ if l.accessToken == "" {
575
+ return nil, errors.New("access token is not available")
576
+ }
577
+ if l.sniPort <= 0 {
578
+ return nil, errors.New("sni port is not available")
579
+ }
580
+
581
+ tlsConf := l.tlsConfig.Clone()
582
+ tlsConf.NextProtos = []string{"portal-tunnel"}
583
+
584
+ quicConf := &quic.Config{
585
+ EnableDatagrams: true,
586
+ KeepAlivePeriod: 15 * time.Second,
587
+ MaxIdleTimeout: 60 * time.Second,
588
+ }
589
+
590
+ host := strings.TrimSpace(l.relayURL.Hostname())
591
+ if host == "" {
592
+ host = strings.TrimSpace(l.relayURL.Host)
593
+ }
594
+ dialAddr := net.JoinHostPort(host, fmt.Sprintf("%d", l.sniPort))
595
+ conn, err := quic.DialAddr(ctx, dialAddr, tlsConf, quicConf)
596
+ if err != nil {
597
+ return nil, fmt.Errorf("quic dial: %w", err)
598
+ }
599
+
600
+ stream, err := conn.OpenStreamSync(ctx)
601
+ if err != nil {
602
+ _ = conn.CloseWithError(1, "stream open failed")
603
+ return nil, fmt.Errorf("open control stream: %w", err)
604
+ }
605
+
606
+ controlMsg := types.QUICControlMessage{
607
+ AccessToken: l.accessToken,
608
+ }
609
+ if err := json.NewEncoder(stream).Encode(controlMsg); err != nil {
610
+ _ = conn.CloseWithError(1, "control write failed")
611
+ return nil, fmt.Errorf("write control: %w", err)
612
+ }
613
+
614
+ _ = stream.SetReadDeadline(time.Now().Add(10 * time.Second))
615
+ var resp types.QUICControlResponse
616
+ if err := json.NewDecoder(io.LimitReader(stream, 4096)).Decode(&resp); err != nil {
617
+ _ = conn.CloseWithError(1, "control read failed")
618
+ return nil, fmt.Errorf("read control response: %w", err)
619
+ }
620
+ if !resp.OK {
621
+ _ = conn.CloseWithError(1, resp.Error)
622
+ return nil, fmt.Errorf("quic connect rejected: %s", resp.Error)
623
+ }
624
+
625
+ return conn, nil
626
+}
627
+
628
+func (l *listener) runRenewLoop(ctx context.Context) error {
629
+ leaseTTL := l.leaseTTL
630
+ if leaseTTL <= 0 {
631
+ leaseTTL = defaultLeaseTTL
632
+ }
633
+ interval := leaseTTL / 2
634
+ if l.renewBefore > 0 && l.renewBefore < leaseTTL {
635
+ interval = leaseTTL - l.renewBefore
636
}
637
638
const wakeThreshold = 10 * time.Second
@@ -555,7 +656,7 @@ func (l *listener) runRenewLoop(ctx context.Context) error {
656
log.Info().
657
Dur("expected", interval).
658
Dur("actual", elapsed).
558
- Str("address", l.address()).
659
+ Str("address", l.identity.Address).
660
Msg("system sleep/wake detected; resetting transport and re-registering")
661
return errLeaseRefreshRequired
662
}
@@ -582,11 +683,8 @@ func (l *listener) runRenewLoop(ctx context.Context) error {
683
}
684
685
func (l *listener) renewLease(ctx context.Context) error {
585
- l.mu.Lock()
686
expiresAt := l.expiresAt
587
- accessToken := strings.TrimSpace(l.accessToken)
588
- l.mu.Unlock()
589
-
687
+ accessToken := l.accessToken
688
if accessToken == "" || !time.Now().Before(expiresAt) {
689
return errLeaseRefreshRequired
690
}
@@ -605,16 +703,18 @@ func (l *listener) renewLease(ctx context.Context) error {
703
if resp.AccessToken == "" {
704
return errors.New("relay did not return renewed access token")
705
}
608
- l.mu.Lock()
706
if l.accessToken == accessToken {
707
l.accessToken = resp.AccessToken
708
l.expiresAt = resp.ExpiresAt
709
}
613
- l.mu.Unlock()
710
return nil
711
}
712
713
func (l *listener) registerAndConfigure(ctx context.Context) error {
714
+ if err := l.initHTTPTransport(ctx); err != nil {
715
+ return err
716
+ }
717
+
718
resp, err := l.registerLease(ctx, l.leaseTTL, l.datagram != nil, l.tcpEnabled)
719
if err != nil {
720
return err
@@ -628,14 +728,10 @@ func (l *listener) registerAndConfigure(ctx context.Context) error {
728
_ = l.unregisterLease(context.Background(), resp.AccessToken)
729
return err
730
}
631
- l.mu.Lock()
632
- localIdentity := l.identity.Copy()
633
- l.mu.Unlock()
634
- if registeredIdentity.Key() != localIdentity.Key() {
731
+ if registeredIdentity.Key() != l.identity.Key() {
732
_ = l.unregisterLease(context.Background(), resp.AccessToken)
733
return errors.New("relay returned mismatched lease identity")
734
}
638
- resp.Identity = registeredIdentity
735
if l.datagram != nil && !resp.UDPEnabled {
736
_ = l.unregisterLease(context.Background(), resp.AccessToken)
737
return &types.APIRequestError{
@@ -660,33 +756,19 @@ func (l *listener) registerAndConfigure(ctx context.Context) error {
756
}
757
return ctx.Err()
758
}
663
-
664
- l.mu.Lock()
665
- if ctx.Err() != nil {
666
- l.mu.Unlock()
667
- _ = l.unregisterLease(context.Background(), resp.AccessToken)
668
- if tenantTLSCloser != nil {
669
- _ = tenantTLSCloser.Close()
670
- }
671
- return ctx.Err()
672
- }
673
- oldCloser := l.tenantTLSCloser
759
datagram := l.datagram
675
- l.identity.Name = resp.Identity.Name
676
- l.identity.Address = resp.Identity.Address
760
+ oldCloser := l.tenantTLSCloser
761
l.hostname = resp.Hostname
762
l.udpAddr = resp.UDPAddr
763
l.accessToken = resp.AccessToken
764
l.expiresAt = resp.ExpiresAt
681
- if l.datagram != nil {
765
+ if datagram != nil {
766
l.sniPort = resp.SNIPort
767
} else {
768
l.sniPort = 0
769
}
770
l.tenantTLSConfig = tlsConf
771
l.tenantTLSCloser = tenantTLSCloser
688
- l.mu.Unlock()
689
-
772
if oldCloser != nil {
773
_ = oldCloser.Close()
774
}
@@ -712,7 +794,7 @@ func (l *listener) waitRetry(ctx context.Context, operation string, err error, r
794
logger := log.With().
795
Str("relay_url", relayURL).
796
Str("operation", operation).
715
- Str("address", l.address()).
797
+ Str("address", l.identity.Address).
798
Logger()
799
800
if l.retryCount > 0 && retries > l.retryCount {
@@ -745,24 +827,3 @@ type listenerAddr string
827
828
func (a listenerAddr) Network() string { return "portal" }
829
func (a listenerAddr) String() string { return string(a) }
748
-
749
-func (l *listener) closed() bool {
750
- select {
751
- case <-l.doneCh:
752
- return true
753
- default:
754
- return false
755
- }
756
-}
757
-
758
-func (l *listener) ban() {
759
- relayURL := ""
760
- if l.relayURL != nil {
761
- relayURL = l.relayURL.String()
762
- }
763
- if l.relaySet != nil && relayURL != "" {
764
- l.relaySet.UnconfirmRelayURL(relayURL)
765
- l.relaySet.BanRelayURL(relayURL)
766
- }
767
- _ = l.Close()
768
-}
sdk/mitm.go
+29
-19
@@ -52,6 +52,7 @@ type mitmProbeResult struct {
52
type mitmManager struct {
53
ctx context.Context
54
listener *listener
55
+ ban bool
56
57
mu sync.Mutex
58
pending map[string]*mitmProbePending
@@ -59,9 +60,10 @@ type mitmManager struct {
60
lastAt time.Time
61
}
62
62
-func newMITMManager(ctx context.Context, listener *listener) *mitmManager {
63
+func newMITMManager(ctx context.Context, listener *listener, ban bool) *mitmManager {
64
return &mitmManager{
65
ctx: ctx,
66
+ ban: ban,
67
listener: listener,
68
pending: make(map[string]*mitmProbePending),
69
}
@@ -86,18 +88,14 @@ func (m *mitmManager) probeTLSPassthrough(ctx context.Context) (MITMProbeReport,
88
return MITMProbeReport{}, errors.New("listener is not registered")
89
}
90
89
- l.mu.Lock()
90
- hostname := l.hostname
91
- leaseTLSConfig := l.tenantTLSConfig
92
- l.mu.Unlock()
93
- if hostname == "" {
91
+ if l.hostname == "" {
92
return MITMProbeReport{}, errors.New("listener hostname is unavailable")
93
}
94
95
report := MITMProbeReport{
96
RelayURL: l.relayURL.String(),
97
PublicURL: publicURL,
100
- Address: l.address(),
98
+ Address: l.identity.Address,
99
}
100
101
probeCtx, cancel := context.WithTimeout(ctx, defaultMITMProbeTimeout)
@@ -115,14 +113,14 @@ func (m *mitmManager) probeTLSPassthrough(ctx context.Context) (MITMProbeReport,
113
}
114
115
probeTLSConf := &tls.Config{
118
- ServerName: hostname,
116
+ ServerName: l.hostname,
117
InsecureSkipVerify: true,
118
}
121
- if leaseTLSConfig != nil {
122
- probeTLSConf.MinVersion = leaseTLSConfig.MinVersion
123
- probeTLSConf.MaxVersion = leaseTLSConfig.MaxVersion
124
- if len(leaseTLSConfig.NextProtos) > 0 {
125
- probeTLSConf.NextProtos = append([]string(nil), leaseTLSConfig.NextProtos...)
119
+ if l.tenantTLSConfig != nil {
120
+ probeTLSConf.MinVersion = l.tenantTLSConfig.MinVersion
121
+ probeTLSConf.MaxVersion = l.tenantTLSConfig.MaxVersion
122
+ if len(l.tenantTLSConfig.NextProtos) > 0 {
123
+ probeTLSConf.NextProtos = append([]string(nil), l.tenantTLSConfig.NextProtos...)
124
}
125
}
126
@@ -199,8 +197,10 @@ func (m *mitmManager) probeDialAddress(publicURL string) (string, error) {
197
198
func (m *mitmManager) maybeStart() {
199
l := m.listener
202
- if l.closed() {
200
+ select {
201
+ case <-l.doneCh:
202
return
203
+ default:
204
}
205
206
m.mu.Lock()
@@ -229,12 +229,18 @@ func (m *mitmManager) logResult(report MITMProbeReport, err error) {
229
if l == nil {
230
return
231
}
232
+ closed := false
233
+ select {
234
+ case <-l.doneCh:
235
+ closed = true
236
+ default:
237
+ }
238
relayURL := ""
239
if l.relayURL != nil {
240
relayURL = l.relayURL.String()
241
}
242
switch {
237
- case l.closed():
243
+ case closed:
244
return
245
case err != nil:
246
if errors.Is(err, context.Canceled) || errors.Is(err, net.ErrClosed) {
@@ -243,7 +249,7 @@ func (m *mitmManager) logResult(report MITMProbeReport, err error) {
249
log.Warn().
250
Err(err).
251
Str("relay_url", relayURL).
246
- Str("address", l.address()).
252
+ Str("address", l.identity.Address).
253
Msg("tls passthrough self-probe failed")
254
case report.Reason == types.MITMProbeReasonProbeTimeout:
255
log.Warn().
@@ -253,14 +259,18 @@ func (m *mitmManager) logResult(report MITMProbeReport, err error) {
259
Msg("tls self-probe timed out before passthrough could be verified")
260
case report.Detected:
261
event := log.Warn().
256
- Bool("ban_mitm", l.banMITM).
262
+ Bool("ban_mitm", m.ban).
263
Str("reason", report.Reason).
264
Str("relay_url", report.RelayURL).
265
Str("public_url", report.PublicURL).
266
Str("address", report.Address)
261
- if l.banMITM {
267
+ if m.ban {
268
event.Msg("tls termination suspected by self-probe; banning relay")
263
- l.ban()
269
+ if l.relaySet != nil && report.RelayURL != "" {
270
+ l.relaySet.UnconfirmRelayURL(report.RelayURL)
271
+ l.relaySet.BanRelayURL(report.RelayURL)
272
+ }
273
+ _ = l.Close()
274
return
275
}
276
event.Msg("tls termination suspected by self-probe")
sdk/mitm_test.go
+16
-14
@@ -27,7 +27,7 @@ func TestMITMProbeConnMatchesExporter(t *testing.T) {
27
defer closeMITMProbeTLSConn(serverConn)
28
29
listener := &listener{}
30
- listener.mitmManager = newMITMManager(context.Background(), listener)
30
+ listener.mitmManager = newMITMManager(context.Background(), listener, false)
31
32
nonce := make([]byte, 16)
33
if _, err := rand.Read(nonce); err != nil {
@@ -87,7 +87,7 @@ func TestMITMProbeConnDetectsExporterMismatch(t *testing.T) {
87
defer closeMITMProbeTLSConn(serverConn)
88
89
listener := &listener{}
90
- listener.mitmManager = newMITMManager(context.Background(), listener)
90
+ listener.mitmManager = newMITMManager(context.Background(), listener, false)
91
92
nonce := make([]byte, 16)
93
if _, err := rand.Read(nonce); err != nil {
@@ -145,7 +145,7 @@ func TestMITMProbeConnPassesThroughNormalTraffic(t *testing.T) {
145
defer closeMITMProbeTLSConn(serverConn)
146
147
listener := &listener{}
148
- listener.mitmManager = newMITMManager(context.Background(), listener)
148
+ listener.mitmManager = newMITMManager(context.Background(), listener, false)
149
150
type handleResult struct {
151
conn net.Conn
@@ -216,10 +216,9 @@ func TestMITMProbeDetectionBansListener(t *testing.T) {
216
close(doneCh)
217
}
218
},
219
- doneCh: doneCh,
220
- banMITM: true,
219
+ doneCh: doneCh,
220
}
222
- listener.mitmManager = newMITMManager(context.Background(), listener)
221
+ listener.mitmManager = newMITMManager(context.Background(), listener, true)
222
223
listener.mitmManager.logResult(MITMProbeReport{
224
RelayURL: relayURL.String(),
@@ -232,8 +231,10 @@ func TestMITMProbeDetectionBansListener(t *testing.T) {
231
t.Fatal("relay still active after mitm detection")
232
}
233
}
235
- if !listener.closed() {
236
- t.Fatal("listener.closed() = false, want true")
234
+ select {
235
+ case <-listener.doneCh:
236
+ default:
237
+ t.Fatal("listener.doneCh is open, want closed")
238
}
239
}
240
@@ -248,9 +249,8 @@ func TestMITMProbeDetectionWarnsWithoutBanningListener(t *testing.T) {
249
relayURL: relayURL,
250
relaySet: mustRelaySet(t, relayURL.String()),
251
doneCh: doneCh,
251
- banMITM: false,
252
}
253
- listener.mitmManager = newMITMManager(context.Background(), listener)
253
+ listener.mitmManager = newMITMManager(context.Background(), listener, false)
254
255
listener.mitmManager.logResult(MITMProbeReport{
256
RelayURL: relayURL.String(),
@@ -265,8 +265,10 @@ func TestMITMProbeDetectionWarnsWithoutBanningListener(t *testing.T) {
265
if len(activeRelayURLs) != 1 || activeRelayURLs[0] != relayURL.String() {
266
t.Fatalf("ActiveRelayURLs() = %v, want [%q]", activeRelayURLs, relayURL.String())
267
}
268
- if listener.closed() {
269
- t.Fatal("listener.closed() = true, want false")
268
+ select {
269
+ case <-listener.doneCh:
270
+ t.Fatal("listener.doneCh is closed, want open")
271
+ default:
272
}
273
}
274
@@ -279,7 +281,7 @@ func TestMITMProbeDialAddressUsesRelayHostForLocalRelay(t *testing.T) {
281
listener := &listener{
282
relayURL: relayURL,
283
}
282
- listener.mitmManager = newMITMManager(context.Background(), listener)
284
+ listener.mitmManager = newMITMManager(context.Background(), listener, false)
285
286
got, err := listener.mitmManager.probeDialAddress("https://bravo-gecko-disco.localhost:4017")
287
if err != nil {
@@ -299,7 +301,7 @@ func TestMITMProbeDialAddressUsesPublicURLForRemoteRelay(t *testing.T) {
301
listener := &listener{
302
relayURL: relayURL,
303
}
302
- listener.mitmManager = newMITMManager(context.Background(), listener)
304
+ listener.mitmManager = newMITMManager(context.Background(), listener, false)
305
306
got, err := listener.mitmManager.probeDialAddress("https://bravo-gecko-disco.example")
307
if err != nil {