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 {