refact: tidy overlay code
Lee Yunjin committed
May 12, 2026 at 12:21 UTC
e1a0dd1b4af385bffcc480ae7db7ec8bac975bb5
6 files changed
+416
-390
portal/api_server.go
+2
-2
@@ -332,7 +332,7 @@ func (s *Server) handleRegisterChallenge(w http.ResponseWriter, r *http.Request)
332
Path: types.PathSDKRegister,
333
}).String()
334
335
- if strings.TrimSpace(req.HopToken) != "" && s.hopMux == nil {
335
+ if strings.TrimSpace(req.HopToken) != "" && s.overlay == nil {
336
utils.WriteAPIError(w, http.StatusServiceUnavailable, types.APIErrorCodeFeatureUnavailable, errFeatureUnavailable.Error())
337
return
338
}
@@ -406,7 +406,7 @@ func (s *Server) handleHop(w http.ResponseWriter, r *http.Request) {
406
utils.MethodNotAllowedError().Write(w)
407
return
408
}
409
- if s.hopMux == nil || s.overlay == nil || s.relaySet == nil {
409
+ if s.overlay == nil || s.relaySet == nil {
410
utils.WriteAPIError(w, http.StatusServiceUnavailable, types.APIErrorCodeFeatureUnavailable, errFeatureUnavailable.Error())
411
return
412
}
portal/auth/hop_route.go
+38
-27
@@ -3,6 +3,7 @@ package auth
3
import (
4
"errors"
5
"fmt"
6
+ "strings"
7
"time"
8
9
"github.com/gosuda/portal-tunnel/v2/types"
@@ -11,19 +12,51 @@ import (
12
13
var ErrHopRouteSignatureInvalid = errors.New("hop route signature is invalid")
14
15
+func normalizeHopRoute(route *types.HopRoute, requireOwner bool) error {
16
+ ownerPublicKey := strings.ToLower(utils.TrimHexPrefix(strings.TrimSpace(route.OwnerPublicKey)))
17
+ if ownerPublicKey != "" {
18
+ if _, err := utils.ParseSecp256k1PublicKeyHex(ownerPublicKey); err != nil {
19
+ return fmt.Errorf("hop route owner public key: %w", err)
20
+ }
21
+ } else if requireOwner {
22
+ return errors.New("hop route owner public key is required")
23
+ }
24
+
25
+ relayURL, err := utils.NormalizeRelayURL(route.RelayURL)
26
+ if err != nil {
27
+ return fmt.Errorf("hop relay url: %w", err)
28
+ }
29
+
30
+ route.OwnerPublicKey = ownerPublicKey
31
+ route.RelayURL = relayURL
32
+ route.PublicHostname = utils.NormalizeHostname(route.PublicHostname)
33
+ route.RouteHostname = utils.NormalizeHostname(route.RouteHostname)
34
+ route.HostnameHash = strings.TrimSpace(route.HostnameHash)
35
+ route.MatchToken = strings.TrimSpace(route.MatchToken)
36
+ route.Metadata = route.Metadata.Copy()
37
+ route.ForwardToken = strings.TrimSpace(route.ForwardToken)
38
+ route.ExpiresAt = route.ExpiresAt.UTC()
39
+ route.Signature = strings.TrimSpace(route.Signature)
40
+ return nil
41
+}
42
+
43
func SignHopRoute(method string, route types.HopRoute, identity types.Identity, expiresAt time.Time) (types.HopRoute, error) {
44
route.ExpiresAt = expiresAt.UTC()
45
route.Signature = ""
17
- route.OwnerPublicKey = identity.PublicKey
46
+ route.OwnerPublicKey = ""
47
19
- route, err := normalizeHopRoute(route, true)
48
+ if err := normalizeHopRoute(&route, false); err != nil {
49
+ return types.HopRoute{}, err
50
+ }
51
+ identity, err := utils.NormalizeStoredIdentity(identity)
52
if err != nil {
53
return types.HopRoute{}, err
54
}
23
- if identity.PrivateKey == "" || identity.PublicKey == "" {
55
+ if strings.TrimSpace(identity.PrivateKey) == "" || strings.TrimSpace(identity.PublicKey) == "" {
56
return types.HopRoute{}, errors.New("hop route owner identity is required")
57
}
58
59
+ route.OwnerPublicKey = identity.PublicKey
60
payload, err := types.HopRouteBytes(method, route)
61
if err != nil {
62
return types.HopRoute{}, err
@@ -36,11 +69,10 @@ func SignHopRoute(method string, route types.HopRoute, identity types.Identity,
69
}
70
71
func VerifyHopRoute(method string, route types.HopRoute) (types.HopRoute, error) {
39
- signature := route.Signature
72
+ signature := strings.TrimSpace(route.Signature)
73
route.Signature = ""
74
42
- route, err := normalizeHopRoute(route, true)
43
- if err != nil {
75
+ if err := normalizeHopRoute(&route, true); err != nil {
76
return types.HopRoute{}, err
77
}
78
payload, err := types.HopRouteBytes(method, route)
@@ -53,24 +85,3 @@ func VerifyHopRoute(method string, route types.HopRoute) (types.HopRoute, error)
85
route.Signature = signature
86
return route, nil
87
}
56
-
57
-func normalizeHopRoute(route types.HopRoute, requireOwner bool) (types.HopRoute, error) {
58
- if route.OwnerPublicKey != "" {
59
- if _, err := utils.ParseSecp256k1PublicKeyHex(route.OwnerPublicKey); err != nil {
60
- return types.HopRoute{}, fmt.Errorf("hop route owner public key: %w", err)
61
- }
62
- } else if requireOwner {
63
- return types.HopRoute{}, errors.New("hop route owner public key is required")
64
- }
65
-
66
- relayURL, err := utils.NormalizeRelayURL(route.RelayURL)
67
- if err != nil {
68
- return types.HopRoute{}, fmt.Errorf("hop relay url: %w", err)
69
- }
70
-
71
- route.RelayURL = relayURL
72
- route.PublicHostname = utils.NormalizeHostname(route.PublicHostname)
73
- route.RouteHostname = utils.NormalizeHostname(route.RouteHostname)
74
- route.ExpiresAt = route.ExpiresAt.UTC()
75
- return route, nil
76
-}
portal/overlay/hop_mux.go
deleted
-237
@@ -1,237 +0,0 @@
1
-package overlay
2
-
3
-import (
4
- "context"
5
- "encoding/binary"
6
- "errors"
7
- "fmt"
8
- "io"
9
- "net"
10
- "sync"
11
- "time"
12
-
13
- "github.com/hashicorp/yamux"
14
-)
15
-
16
-const (
17
- maxHopTokenBytes = 256
18
- defaultTokenTimeout = 2 * time.Second
19
-)
20
-
21
-type HopMux struct {
22
- listener net.Listener
23
- overlay *Overlay
24
- incoming chan HopStream
25
-
26
- mu sync.Mutex
27
- outbound map[string]*yamux.Session
28
- done chan struct{}
29
-}
30
-
31
-type HopStream struct {
32
- Conn net.Conn
33
- Token string
34
- RemoteAddr string
35
-}
36
-
37
-func NewHopMux(overlay *Overlay) (*HopMux, error) {
38
- if overlay == nil || overlay.stack == nil {
39
- return nil, errors.New("overlay is required for multi-hop mux")
40
- }
41
- listener, err := overlay.stack.ListenTCP(DefaultPeerYamuxPort)
42
- if err != nil {
43
- return nil, fmt.Errorf("listen hop yamux: %w", err)
44
- }
45
- return &HopMux{
46
- listener: listener,
47
- overlay: overlay,
48
- incoming: make(chan HopStream, 1024),
49
- outbound: make(map[string]*yamux.Session),
50
- done: make(chan struct{}),
51
- }, nil
52
-}
53
-
54
-func (m *HopMux) Serve(ctx context.Context) error {
55
- go func() {
56
- <-ctx.Done()
57
- _ = m.Close()
58
- }()
59
-
60
- for {
61
- conn, err := m.listener.Accept()
62
- if err != nil {
63
- if errors.Is(err, net.ErrClosed) || ctx.Err() != nil {
64
- return ctx.Err()
65
- }
66
- return fmt.Errorf("accept hop mux connection: %w", err)
67
- }
68
- go m.serveSession(ctx, conn)
69
- }
70
-}
71
-
72
-func (m *HopMux) Accept(ctx context.Context) (HopStream, error) {
73
- select {
74
- case stream := <-m.incoming:
75
- return stream, nil
76
- case <-ctx.Done():
77
- return HopStream{}, ctx.Err()
78
- }
79
-}
80
-
81
-func (m *HopMux) Close() error {
82
- m.mu.Lock()
83
- select {
84
- case <-m.done:
85
- m.mu.Unlock()
86
- return nil
87
- default:
88
- }
89
- close(m.done)
90
- sessions := m.outbound
91
- m.outbound = make(map[string]*yamux.Session)
92
- m.mu.Unlock()
93
-
94
- closeErr := m.listener.Close()
95
- for _, session := range sessions {
96
- closeErr = errors.Join(closeErr, session.Close())
97
- }
98
- return closeErr
99
-}
100
-
101
-func (m *HopMux) OpenStream(ctx context.Context, overlayIPv4, token string) (net.Conn, error) {
102
- if token == "" || overlayIPv4 == "" {
103
- return nil, errors.New("hop token and overlay ipv4 are required")
104
- }
105
-
106
- session, err := m.session(ctx, overlayIPv4)
107
- if err != nil {
108
- return nil, err
109
- }
110
- stream, err := session.OpenStream()
111
- if err != nil {
112
- m.mu.Lock()
113
- if m.outbound[overlayIPv4] == session {
114
- delete(m.outbound, overlayIPv4)
115
- }
116
- m.mu.Unlock()
117
- _ = session.Close()
118
- return nil, err
119
- }
120
-
121
- payload := []byte(token)
122
- if len(payload) > maxHopTokenBytes {
123
- _ = stream.Close()
124
- return nil, errors.New("next hop token is too large")
125
- }
126
- frame := make([]byte, 4+len(payload))
127
- binary.BigEndian.PutUint32(frame[:4], uint32(len(payload)))
128
- copy(frame[4:], payload)
129
-
130
- _ = stream.SetWriteDeadline(time.Now().Add(defaultTokenTimeout))
131
- if _, err := stream.Write(frame); err != nil {
132
- _ = stream.Close()
133
- return nil, err
134
- }
135
- _ = stream.SetWriteDeadline(time.Time{})
136
- return stream, nil
137
-}
138
-
139
-func (m *HopMux) session(ctx context.Context, overlayIPv4 string) (*yamux.Session, error) {
140
- m.mu.Lock()
141
- if session := m.outbound[overlayIPv4]; session != nil && !session.IsClosed() {
142
- m.mu.Unlock()
143
- return session, nil
144
- }
145
- m.mu.Unlock()
146
-
147
- addr := net.JoinHostPort(overlayIPv4, fmt.Sprintf("%d", DefaultPeerYamuxPort))
148
- conn, err := m.overlay.stack.DialContext(ctx, "tcp", addr)
149
- if err != nil {
150
- return nil, err
151
- }
152
-
153
- session, err := yamux.Client(conn, hopYamuxConfig())
154
- if err != nil {
155
- _ = conn.Close()
156
- return nil, err
157
- }
158
-
159
- m.mu.Lock()
160
- defer m.mu.Unlock()
161
- if current := m.outbound[overlayIPv4]; current != nil && !current.IsClosed() {
162
- _ = session.Close()
163
- return current, nil
164
- }
165
- m.outbound[overlayIPv4] = session
166
- return session, nil
167
-}
168
-
169
-func (m *HopMux) serveSession(ctx context.Context, conn net.Conn) {
170
- session, err := yamux.Server(conn, hopYamuxConfig())
171
- if err != nil {
172
- _ = conn.Close()
173
- return
174
- }
175
- defer session.Close()
176
-
177
- go func() {
178
- select {
179
- case <-ctx.Done():
180
- _ = session.Close()
181
- case <-m.done:
182
- _ = session.Close()
183
- }
184
- }()
185
-
186
- for {
187
- stream, err := session.AcceptStream()
188
- if err != nil {
189
- return
190
- }
191
- go m.handleStream(ctx, stream)
192
- }
193
-}
194
-
195
-func (m *HopMux) handleStream(ctx context.Context, stream *yamux.Stream) {
196
- _ = stream.SetReadDeadline(time.Now().Add(defaultTokenTimeout))
197
- defer stream.SetReadDeadline(time.Time{})
198
-
199
- var size [4]byte
200
- if _, err := io.ReadFull(stream, size[:]); err != nil {
201
- _ = stream.Close()
202
- return
203
- }
204
- n := binary.BigEndian.Uint32(size[:])
205
- if n == 0 || n > uint32(maxHopTokenBytes) {
206
- _ = stream.Close()
207
- return
208
- }
209
- payload := make([]byte, n)
210
- if _, err := io.ReadFull(stream, payload); err != nil {
211
- _ = stream.Close()
212
- return
213
- }
214
- token := string(payload)
215
-
216
- remoteAddr := ""
217
- if addr := stream.RemoteAddr(); addr != nil {
218
- remoteAddr = addr.String()
219
- }
220
-
221
- select {
222
- case m.incoming <- HopStream{Conn: stream, Token: token, RemoteAddr: remoteAddr}:
223
- case <-ctx.Done():
224
- _ = stream.Close()
225
- case <-m.done:
226
- _ = stream.Close()
227
- }
228
-}
229
-
230
-func hopYamuxConfig() *yamux.Config {
231
- cfg := yamux.DefaultConfig()
232
- cfg.Logger = nil
233
- cfg.MaxStreamWindowSize = 16 * 1024 * 1024
234
- cfg.StreamOpenTimeout = 75 * time.Second
235
- cfg.StreamCloseTimeout = 5 * time.Minute
236
- return cfg
237
-}
portal/overlay/overlay.go
+302
-14
@@ -2,15 +2,22 @@ package overlay
2
3
import (
4
"context"
5
+ "encoding/binary"
6
"errors"
7
"fmt"
8
+ "io"
9
"net"
10
"net/http"
11
"net/url"
12
"sort"
13
"strings"
14
+ "sync"
15
+ "sync/atomic"
16
"time"
17
18
+ "github.com/hashicorp/yamux"
19
+ "github.com/rs/zerolog/log"
20
+
21
"github.com/gosuda/portal-tunnel/v2/portal/discovery"
22
"github.com/gosuda/portal-tunnel/v2/types"
23
"github.com/gosuda/portal-tunnel/v2/utils"
@@ -22,6 +29,9 @@ const (
29
DefaultPeerAPIHTTPPort = 7777
30
DefaultPeerYamuxPort = 7778
31
DefaultPersistentKeepalive = 25
32
+
33
+ maxHopTokenBytes = 256
34
+ defaultTokenTimeout = 2 * time.Second
35
)
36
37
type Config struct {
@@ -73,15 +83,27 @@ func NormalizeConfig(cfg Config) (Config, error) {
83
return cfg, nil
84
}
85
86
+type HopStream struct {
87
+ Conn net.Conn
88
+ Token string
89
+ RemoteAddr string
90
+}
91
+
92
+type StreamHandler func(ctx context.Context, stream HopStream)
93
type Overlay struct {
94
cfg Config
95
stack *stack
96
listener net.Listener
97
server *http.Server
98
client *http.Client
99
+
100
+ hopListener net.Listener
101
+ streamHandler atomic.Pointer[StreamHandler]
102
+ hopOutbound sync.Map
103
+ hopDone chan struct{}
104
}
105
84
-func NewOverlay(cfg Config, handler http.Handler) (*Overlay, error) {
106
+func NewOverlay(cfg Config, handler http.Handler, streamHandler StreamHandler) (*Overlay, error) {
107
cfg, err := NormalizeConfig(cfg)
108
if err != nil {
109
return nil, err
@@ -102,6 +124,13 @@ func NewOverlay(cfg Config, handler http.Handler) (*Overlay, error) {
124
return nil, err
125
}
126
127
+ hopListener, err := stack.ListenTCP(DefaultPeerYamuxPort)
128
+ if err != nil {
129
+ _ = listener.Close()
130
+ _ = stack.Close()
131
+ return nil, err
132
+ }
133
+
134
server := &http.Server{
135
Handler: handler,
136
ReadHeaderTimeout: 10 * time.Second,
@@ -119,13 +148,19 @@ func NewOverlay(cfg Config, handler http.Handler) (*Overlay, error) {
148
149
publicCfg := cfg.Copy()
150
publicCfg.PrivateKey = ""
122
- return &Overlay{
123
- cfg: publicCfg,
124
- stack: stack,
125
- listener: listener,
126
- server: server,
127
- client: client,
128
- }, nil
151
+ ov := &Overlay{
152
+ cfg: publicCfg,
153
+ stack: stack,
154
+ listener: listener,
155
+ server: server,
156
+ client: client,
157
+ hopListener: hopListener,
158
+ hopDone: make(chan struct{}),
159
+ }
160
+ if streamHandler != nil {
161
+ ov.streamHandler.Store(&streamHandler)
162
+ }
163
+ return ov, nil
164
}
165
166
func (o *Overlay) Config() Config {
@@ -135,16 +170,241 @@ func (o *Overlay) Config() Config {
170
return o.cfg.Copy()
171
}
172
138
-func (o *Overlay) Serve() error {
139
- if o == nil || o.server == nil || o.listener == nil {
173
+func (o *Overlay) Serve(ctx context.Context) error {
174
+ if o == nil {
175
return nil
176
}
177
143
- err := o.server.Serve(o.listener)
144
- if errors.Is(err, http.ErrServerClosed) || errors.Is(err, net.ErrClosed) {
145
- return nil
178
+ if o.server != nil && o.listener != nil {
179
+ go func() {
180
+ err := o.server.Serve(o.listener)
181
+ if err != nil && !errors.Is(err, http.ErrServerClosed) && !errors.Is(err, net.ErrClosed) {
182
+ log.Error().Err(err).Msg("overlay peer api server exited")
183
+ }
184
+ }()
185
+ }
186
+
187
+ if o.hopListener != nil {
188
+ go func() {
189
+ <-ctx.Done()
190
+ _ = o.hopListener.Close()
191
+ }()
192
+
193
+ for {
194
+ conn, err := o.hopListener.Accept()
195
+ if err != nil {
196
+ if errors.Is(err, net.ErrClosed) || ctx.Err() != nil {
197
+ return nil
198
+ }
199
+ return fmt.Errorf("accept hop mux connection: %w", err)
200
+ }
201
+ go o.serveHopSession(ctx, conn)
202
+ }
203
+ }
204
+
205
+ <-ctx.Done()
206
+ return nil
207
+}
208
+
209
+func (o *Overlay) SetStreamHandler(handler StreamHandler) {
210
+ if handler != nil {
211
+ o.streamHandler.Store(&handler)
212
+ } else {
213
+ o.streamHandler.Store(nil)
214
+ }
215
+}
216
+
217
+func (o *Overlay) OpenHopStream(ctx context.Context, overlayIPv4, token string) (net.Conn, error) {
218
+ token = strings.TrimSpace(token)
219
+ if token == "" {
220
+ return nil, errors.New("next hop token is required")
221
+ }
222
+ overlayIPv4 = strings.TrimSpace(overlayIPv4)
223
+ if overlayIPv4 == "" {
224
+ return nil, errors.New("next hop overlay ipv4 is required")
225
+ }
226
+
227
+ var next *yamux.Stream
228
+ var lastErr error
229
+ for {
230
+ session, err := o.getHopSession(ctx, overlayIPv4)
231
+ if err != nil {
232
+ return nil, err
233
+ }
234
+
235
+ type openResult struct {
236
+ stream *yamux.Stream
237
+ err error
238
+ }
239
+ resCh := make(chan openResult, 1)
240
+ go func() {
241
+ s, openErr := session.OpenStream()
242
+ resCh <- openResult{s, openErr}
243
+ }()
244
+
245
+ select {
246
+ case res := <-resCh:
247
+ if res.err != nil {
248
+ _ = session.Close()
249
+ lastErr = res.err
250
+ } else {
251
+ next = res.stream
252
+ }
253
+ case <-ctx.Done():
254
+ return nil, ctx.Err()
255
+ }
256
+
257
+ if next != nil {
258
+ break
259
+ }
260
+ if errors.Is(lastErr, net.ErrClosed) {
261
+ return nil, fmt.Errorf("open next hop stream: %w", lastErr)
262
+ }
263
+ if !utils.SleepOrDone(ctx, 250*time.Millisecond) {
264
+ return nil, fmt.Errorf("open next hop stream within timeout: %w", errors.Join(lastErr, ctx.Err()))
265
+ }
266
+ }
267
+
268
+ payload := []byte(token)
269
+ if len(payload) > maxHopTokenBytes {
270
+ _ = next.Close()
271
+ return nil, errors.New("next hop token is too large")
272
+ }
273
+ frame := make([]byte, 4+len(payload))
274
+ binary.BigEndian.PutUint32(frame[:4], uint32(len(payload)))
275
+ copy(frame[4:], payload)
276
+ if _, err := next.Write(frame); err != nil {
277
+ _ = next.Close()
278
+ return nil, err
279
+ }
280
+ return next, nil
281
+}
282
+
283
+func (o *Overlay) getHopSession(ctx context.Context, overlayIPv4 string) (*yamux.Session, error) {
284
+ select {
285
+ case <-o.hopDone:
286
+ return nil, net.ErrClosed
287
+ default:
288
+ }
289
+
290
+ if val, ok := o.hopOutbound.Load(overlayIPv4); ok {
291
+ session := val.(*yamux.Session)
292
+ if !session.IsClosed() {
293
+ return session, nil
294
+ }
295
+ o.hopOutbound.Delete(overlayIPv4)
296
+ }
297
+
298
+ addr := net.JoinHostPort(overlayIPv4, fmt.Sprintf("%d", DefaultPeerYamuxPort))
299
+ conn, err := o.stack.DialContext(ctx, "tcp", addr)
300
+ if err != nil {
301
+ return nil, err
302
+ }
303
+
304
+ session, err := yamux.Client(conn, hopYamuxConfig())
305
+ if err != nil {
306
+ _ = conn.Close()
307
+ return nil, err
308
+ }
309
+
310
+ select {
311
+ case <-o.hopDone:
312
+ _ = session.Close()
313
+ return nil, net.ErrClosed
314
+ default:
315
+ }
316
+
317
+ if actual, loaded := o.hopOutbound.LoadOrStore(overlayIPv4, session); loaded {
318
+ current := actual.(*yamux.Session)
319
+ if !current.IsClosed() {
320
+ _ = session.Close()
321
+ return current, nil
322
+ }
323
+ // the existing one is closed, replace it
324
+ o.hopOutbound.Store(overlayIPv4, session)
325
+ }
326
+
327
+ // Background cleanup to prevent memory leaks for dead sessions.
328
+ go func(ip string, s *yamux.Session) {
329
+ select {
330
+ case <-s.CloseChan():
331
+ case <-o.hopDone:
332
+ return
333
+ }
334
+ o.hopOutbound.CompareAndDelete(ip, s)
335
+ }(overlayIPv4, session)
336
+
337
+ return session, nil
338
+}
339
+
340
+func (o *Overlay) serveHopSession(ctx context.Context, conn net.Conn) {
341
+ session, err := yamux.Server(conn, hopYamuxConfig())
342
+ if err != nil {
343
+ _ = conn.Close()
344
+ return
345
+ }
346
+
347
+ select {
348
+ case <-o.hopDone:
349
+ _ = session.Close()
350
+ return
351
+ default:
352
+ }
353
+
354
+ go func() {
355
+ select {
356
+ case <-ctx.Done():
357
+ case <-o.hopDone:
358
+ }
359
+ _ = session.Close()
360
+ }()
361
+
362
+ for {
363
+ stream, err := session.AcceptStream()
364
+ if err != nil {
365
+ return
366
+ }
367
+ go func(stream *yamux.Stream) {
368
+ _ = stream.SetReadDeadline(time.Now().Add(defaultTokenTimeout))
369
+ var size [4]byte
370
+ if _, err := io.ReadFull(stream, size[:]); err != nil {
371
+ _ = stream.Close()
372
+ return
373
+ }
374
+ n := binary.BigEndian.Uint32(size[:])
375
+ if n == 0 || n > uint32(maxHopTokenBytes) {
376
+ _ = stream.Close()
377
+ return
378
+ }
379
+ payload := make([]byte, n)
380
+ if _, err := io.ReadFull(stream, payload); err != nil {
381
+ _ = stream.Close()
382
+ return
383
+ }
384
+ _ = stream.SetReadDeadline(time.Time{})
385
+
386
+ token := strings.TrimSpace(string(payload))
387
+ if token == "" {
388
+ _ = stream.Close()
389
+ return
390
+ }
391
+ remoteAddr := ""
392
+ if stream.RemoteAddr() != nil {
393
+ remoteAddr = stream.RemoteAddr().String()
394
+ }
395
+ hopStream := HopStream{
396
+ Conn: stream,
397
+ Token: token,
398
+ RemoteAddr: remoteAddr,
399
+ }
400
+ handlerPtr := o.streamHandler.Load()
401
+ if handlerPtr != nil && *handlerPtr != nil {
402
+ (*handlerPtr)(ctx, hopStream)
403
+ } else {
404
+ _ = stream.Close()
405
+ }
406
+ }(stream)
407
}
147
- return err
408
}
409
410
func (o *Overlay) Shutdown(ctx context.Context) error {
@@ -152,6 +412,12 @@ func (o *Overlay) Shutdown(ctx context.Context) error {
412
return nil
413
}
414
415
+ select {
416
+ case <-o.hopDone:
417
+ default:
418
+ close(o.hopDone)
419
+ }
420
+
421
var shutdownErr error
422
if o.server != nil {
423
err := o.server.Shutdown(ctx)
@@ -168,6 +434,20 @@ func (o *Overlay) Shutdown(ctx context.Context) error {
434
shutdownErr = errors.Join(shutdownErr, err)
435
}
436
}
437
+ if o.hopListener != nil {
438
+ err := o.hopListener.Close()
439
+ if err != nil && !errors.Is(err, net.ErrClosed) {
440
+ shutdownErr = errors.Join(shutdownErr, err)
441
+ }
442
+ }
443
+
444
+ o.hopOutbound.Range(func(key, value any) bool {
445
+ session := value.(*yamux.Session)
446
+ shutdownErr = errors.Join(shutdownErr, session.Close())
447
+ o.hopOutbound.Delete(key)
448
+ return true
449
+ })
450
+
451
if o.stack != nil {
452
shutdownErr = errors.Join(shutdownErr, o.stack.Close())
453
}
@@ -225,3 +505,11 @@ func (o *Overlay) Sync(relays []discovery.RelayState) error {
505
})
506
return o.stack.ApplyPeers(peers)
507
}
508
+
509
+func hopYamuxConfig() *yamux.Config {
510
+ cfg := yamux.DefaultConfig()
511
+ cfg.Logger = nil
512
+ cfg.MaxStreamWindowSize = 16 * 1024 * 1024
513
+ return cfg
514
+}
515
+
portal/server.go
+31
-77
@@ -126,7 +126,6 @@ type Server struct {
126
quicBackhaul *quic.Listener
127
128
overlay *overlay.Overlay
129
- hopMux *overlay.HopMux
129
relaySet *discovery.RelaySet
130
announceLimiter *discovery.AnnounceLimiter
131
registry *leaseRegistry
@@ -184,7 +183,6 @@ func (s *Server) Start(ctx context.Context, apiMux *http.ServeMux) error {
183
var apiCloser io.Closer
184
var pprofListener net.Listener
185
var pprofServer *http.Server
187
- var hopMux *overlay.HopMux
186
var ov *overlay.Overlay
187
var quicBackhaul *quic.Listener
188
defer func() {
@@ -192,9 +190,6 @@ func (s *Server) Start(ctx context.Context, apiMux *http.ServeMux) error {
190
return
191
}
192
acmeManager.Stop()
195
- if hopMux != nil {
196
- _ = hopMux.Close()
197
- }
193
if ov != nil {
194
_ = ov.Shutdown(context.Background())
195
}
@@ -256,10 +251,6 @@ func (s *Server) Start(ctx context.Context, apiMux *http.ServeMux) error {
251
if err != nil {
252
return err
253
}
259
- hopMux, err = overlay.NewHopMux(ov)
260
- if err != nil {
261
- return err
262
- }
254
}
255
if s.cfg.UDPEnabled {
256
quicBackhaul, err = s.newQUICBackhaulListener(apiTLS)
@@ -279,7 +270,6 @@ func (s *Server) Start(ctx context.Context, apiMux *http.ServeMux) error {
270
s.cancel = cancel
271
s.group = group
272
s.overlay = ov
282
- s.hopMux = hopMux
273
s.quicBackhaul = quicBackhaul
274
started = true
275
@@ -289,10 +279,7 @@ func (s *Server) Start(ctx context.Context, apiMux *http.ServeMux) error {
279
}
280
group.Go(func() error { return s.runPublicIngress(groupCtx) })
281
if s.overlay != nil {
292
- group.Go(s.overlay.Serve)
293
- if s.hopMux != nil {
294
- group.Go(func() error { return s.runOverlayIngress(groupCtx) })
295
- }
282
+ group.Go(func() error { return s.overlay.Serve(groupCtx) })
283
}
284
if s.quicBackhaul != nil {
285
group.Go(s.runQUICBackhaulListener)
@@ -318,7 +305,7 @@ func (s *Server) Start(ctx context.Context, apiMux *http.ServeMux) error {
305
Int("max_port", s.cfg.MaxPort).
306
Bool("discovery_enabled", s.cfg.DiscoveryEnabled).
307
Bool("wireguard_enabled", s.overlay != nil).
321
- Bool("multihop_enabled", s.hopMux != nil).
308
+ Bool("multihop_enabled", s.overlay != nil).
309
Bool("udp_enabled", s.quicBackhaul != nil).
310
Bool("tcp_enabled", s.cfg.TCPEnabled).
311
Bool("api_ech_enabled", len(apiTLS.EncryptedClientHelloKeys) > 0).
@@ -410,11 +397,6 @@ func (s *Server) Shutdown(ctx context.Context) error {
397
shutdownErr = err
398
}
399
}
413
- if s.hopMux != nil {
414
- if err := s.hopMux.Close(); err != nil && shutdownErr == nil && !errors.Is(err, net.ErrClosed) {
415
- shutdownErr = err
416
- }
417
- }
400
if s.overlay != nil {
401
if err := s.overlay.Shutdown(ctx); err != nil && shutdownErr == nil {
402
shutdownErr = err
@@ -557,43 +539,6 @@ func (s *Server) runPublicIngress(ctx context.Context) error {
539
}
540
}
541
560
-func (s *Server) runOverlayIngress(ctx context.Context) error {
561
- group, groupCtx := errgroup.WithContext(ctx)
562
- group.Go(func() error { return s.hopMux.Serve(groupCtx) })
563
- group.Go(func() error {
564
- for {
565
- stream, err := s.hopMux.Accept(groupCtx)
566
- if err != nil {
567
- if ctxErr := groupCtx.Err(); ctxErr != nil {
568
- return ctxErr
569
- }
570
- return err
571
- }
572
- go func(stream overlay.HopStream) {
573
- s.registry.mu.RLock()
574
- record := s.registry.recordByHopToken(stream.Token, time.Now())
575
- s.registry.mu.RUnlock()
576
- if record == nil {
577
- log.Warn().Str("remote_addr", stream.RemoteAddr).Msg("hop stream rejected")
578
- _ = stream.Conn.Close()
579
- return
580
- }
581
- hopRole := "exit"
582
- if record.isHopMiddle() {
583
- hopRole = "middle"
584
- }
585
- log.Info().Str("remote_addr", stream.RemoteAddr).Str("hop_role", hopRole).Msg("hop stream received")
586
-
587
- if err := s.bridgeLeaseConn(groupCtx, stream.Conn, record); err != nil {
588
- log.Warn().Err(err).Str("remote_addr", stream.RemoteAddr).Msg("hop stream bridge failed")
589
- _ = stream.Conn.Close()
590
- }
591
- }(stream)
592
- }
593
- })
594
- return group.Wait()
595
-}
596
-
542
func (s *Server) bridgeLeaseConn(ctx context.Context, conn net.Conn, record *leaseRecord) error {
543
if record.isExpired(time.Now()) {
544
return errLeaseNotFound
@@ -608,21 +553,9 @@ func (s *Server) bridgeLeaseConn(ctx context.Context, conn net.Conn, record *lea
553
554
openCtx, cancel := context.WithTimeout(ctx, defaultClaimTimeout)
555
defer cancel()
611
- var next net.Conn
612
- var lastErr error
613
- for {
614
- var err error
615
- next, err = s.hopMux.OpenStream(openCtx, overlayIPv4, forwardToken)
616
- if err == nil {
617
- break
618
- }
619
- lastErr = err
620
- if errors.Is(err, net.ErrClosed) {
621
- return fmt.Errorf("open next hop stream: %w", err)
622
- }
623
- if !utils.SleepOrDone(openCtx, defaultHopOpenRetryWait) {
624
- return fmt.Errorf("open next hop stream within %s: %w", defaultClaimTimeout, errors.Join(lastErr, openCtx.Err()))
625
- }
556
+ next, err := s.overlay.OpenHopStream(openCtx, overlayIPv4, forwardToken)
557
+ if err != nil {
558
+ return fmt.Errorf("open next hop stream: %w", err)
559
}
560
s.proxy.bridge(conn, next, "", nil)
561
return nil
@@ -738,21 +671,42 @@ func (s *Server) startOverlay() (*overlay.Overlay, error) {
671
peerMux.HandleFunc(types.PathDiscovery, s.handleRelayDiscovery)
672
}
673
741
- overlay, err := overlay.NewOverlay(overlay.Config{
674
+ ov, err := overlay.NewOverlay(overlay.Config{
675
PrivateKey: s.identity.WireGuardPrivateKey,
676
PublicKey: s.identity.WireGuardPublicKey,
677
ListenPort: s.cfg.WireGuardPort,
745
- }, peerMux)
678
+ }, peerMux, nil)
679
if err != nil {
680
return nil, fmt.Errorf("start wireguard overlay: %w", err)
681
}
682
750
- if err := overlay.Sync(s.relaySet.OverlayPeerStates()); err != nil {
751
- _ = overlay.Shutdown(context.Background())
683
+ ov.SetStreamHandler(func(ctx context.Context, stream overlay.HopStream) {
684
+ s.registry.mu.RLock()
685
+ record := s.registry.recordByHopToken(stream.Token, time.Now())
686
+ s.registry.mu.RUnlock()
687
+ if record == nil {
688
+ log.Warn().Str("remote_addr", stream.RemoteAddr).Msg("hop stream rejected")
689
+ _ = stream.Conn.Close()
690
+ return
691
+ }
692
+ hopRole := "exit"
693
+ if record.isHopMiddle() {
694
+ hopRole = "middle"
695
+ }
696
+ log.Info().Str("remote_addr", stream.RemoteAddr).Str("hop_role", hopRole).Msg("hop stream received")
697
+
698
+ if err := s.bridgeLeaseConn(ctx, stream.Conn, record); err != nil {
699
+ log.Warn().Err(err).Str("remote_addr", stream.RemoteAddr).Msg("hop stream bridge failed")
700
+ _ = stream.Conn.Close()
701
+ }
702
+ })
703
+
704
+ if err := ov.Sync(s.relaySet.OverlayPeerStates()); err != nil {
705
+ _ = ov.Shutdown(context.Background())
706
return nil, fmt.Errorf("sync wireguard peers: %w", err)
707
}
708
755
- return overlay, nil
709
+ return ov, nil
710
}
711
712
func (s *Server) runRelayDiscoveryLoop(ctx context.Context) error {
sdk/api_client.go
+43
-33
@@ -83,6 +83,45 @@ func (l *listener) initHTTPTransport(ctx context.Context) error {
83
return nil
84
}
85
86
+func (l *listener) buildHopRoutes(hopPath []types.RelayDescriptor, publicHostname, routeHostname string, echConfigList []byte) ([]types.HopRoute, string, error) {
87
+ if len(hopPath) < 2 {
88
+ return nil, "", errors.New("multi-hop requires at least entry and exit relay urls")
89
+ }
90
+ hopRoutes := make([]types.HopRoute, 0, len(hopPath)-1)
91
+ var previousHopToken string
92
+ for i := 0; i < len(hopPath)-1; i++ {
93
+ token, err := l.identity.DeriveToken(
94
+ "hop-token",
95
+ publicHostname,
96
+ strconv.Itoa(i),
97
+ hopPath[i].APIHTTPSAddr,
98
+ hopPath[i+1].APIHTTPSAddr,
99
+ )
100
+ if err != nil {
101
+ return nil, "", err
102
+ }
103
+ forwardToken := "hpt_" + token
104
+ route := types.HopRoute{
105
+ RelayURL: hopPath[i].APIHTTPSAddr,
106
+ ForwardRelay: hopPath[i+1],
107
+ ForwardToken: forwardToken,
108
+ }
109
+ if i == 0 {
110
+ route.PublicHostname = publicHostname
111
+ route.RouteHostname = routeHostname
112
+ route.HostnameHash = utils.HostnameHash(publicHostname)
113
+ route.ECHConfigList = bytes.Clone(echConfigList)
114
+ route.Metadata = l.metadata.Copy()
115
+ route.Metadata.Hide = true
116
+ } else {
117
+ route.MatchToken = previousHopToken
118
+ }
119
+ hopRoutes = append(hopRoutes, route)
120
+ previousHopToken = forwardToken
121
+ }
122
+ return hopRoutes, previousHopToken, nil
123
+}
124
+
125
func (l *listener) registerLease(ctx context.Context, ttl time.Duration, udpEnabled, tcpEnabled bool) (types.RegisterResponse, []types.HopRoute, string, string, error) {
126
var exitHopToken string
127
var publicHostname string
@@ -144,40 +183,11 @@ func (l *listener) registerLease(ctx context.Context, ttl time.Duration, udpEnab
183
}
184
185
if len(l.multiHop) > 0 {
147
- hopRoutes = make([]types.HopRoute, 0, len(hopPath)-1)
148
- var previousHopToken string
149
- for i := 0; i < len(hopPath)-1; i++ {
150
- token, err := l.identity.DeriveToken(
151
- "hop-token",
152
- publicHostname,
153
- strconv.Itoa(i),
154
- hopPath[i].APIHTTPSAddr,
155
- hopPath[i+1].APIHTTPSAddr,
156
- )
157
- if err != nil {
158
- return types.RegisterResponse{}, nil, "", "", err
159
- }
160
- forwardToken := "hpt_" + token
161
- route := types.HopRoute{
162
- RelayURL: hopPath[i].APIHTTPSAddr,
163
- ForwardRelay: hopPath[i+1],
164
- ForwardToken: forwardToken,
165
- }
166
- if i == 0 {
167
- route.PublicHostname = publicHostname
168
- route.RouteHostname = routeHostname
169
- route.HostnameHash = utils.HostnameHash(publicHostname)
170
- route.ECHConfigList = bytes.Clone(echConfigList)
171
- route.Metadata = l.metadata.Copy()
172
- route.Metadata.Hide = true
173
- hopRoutes = append(hopRoutes, route)
174
- } else {
175
- route.MatchToken = previousHopToken
176
- hopRoutes = append(hopRoutes, route)
177
- }
178
- previousHopToken = forwardToken
186
+ var err error
187
+ hopRoutes, exitHopToken, err = l.buildHopRoutes(hopPath, publicHostname, routeHostname, echConfigList)
188
+ if err != nil {
189
+ return types.RegisterResponse{}, nil, "", "", err
190
}
180
- exitHopToken = previousHopToken
191
}
192
193
registerReq := types.RegisterChallengeRequest{