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{