feat: migrate hop route functionality to auth package and update related references
Kim committed
Apr 17, 2026 at 11:37 UTC
eb0d575c501eb91c2e97eda3093f3d9e5bcb53b3
5 files changed
+83
-190
portal/api_server.go
+2
-4
@@ -19,7 +19,6 @@ import (
19
"github.com/gosuda/portal-tunnel/v2/portal/auth"
20
"github.com/gosuda/portal-tunnel/v2/portal/discovery"
21
"github.com/gosuda/portal-tunnel/v2/portal/keyless"
22
- "github.com/gosuda/portal-tunnel/v2/portal/overlay"
22
"github.com/gosuda/portal-tunnel/v2/portal/transport"
23
"github.com/gosuda/portal-tunnel/v2/types"
24
"github.com/gosuda/portal-tunnel/v2/utils"
@@ -490,8 +489,8 @@ func (s *Server) handleHop(w http.ResponseWriter, r *http.Request) {
489
if !ok {
490
return
491
}
493
- route, err := overlay.VerifyHopRoute(r.Method, route)
494
- if errors.Is(err, overlay.ErrHopRouteSignatureInvalid) {
492
+ route, err := auth.VerifyHopRoute(r.Method, route)
493
+ if errors.Is(err, auth.ErrHopRouteSignatureInvalid) {
494
utils.WriteAPIError(w, http.StatusForbidden, types.APIErrorCodeUnauthorized, "hop route signature is invalid")
495
return
496
}
@@ -503,7 +502,6 @@ func (s *Server) handleHop(w http.ResponseWriter, r *http.Request) {
502
utils.WriteAPIError(w, http.StatusForbidden, types.APIErrorCodeUnauthorized, "hop route relay url does not match receiving relay")
503
return
504
}
506
-
505
if r.Method == http.MethodDelete {
506
s.registry.DeleteHopRoute(&route)
507
utils.WriteAPIData(w, http.StatusOK, map[string]any{})
portal/auth/hop_route.go
renamed
+18
-41
@@ -1,4 +1,4 @@
1
-package overlay
1
+package auth
2
3
import (
4
"crypto/sha256"
@@ -24,16 +24,29 @@ func SignHopRoute(method string, route types.HopRoute, identity types.Identity,
24
if err != nil {
25
return types.HopRoute{}, err
26
}
27
- owner, err := deriveHopRouteOwner(identity, route)
27
+ ownerScope := struct {
28
+ RelayURL string `json:"relay_url"`
29
+ MatchHostname string `json:"match_hostname"`
30
+ MatchToken string `json:"match_token"`
31
+ }{
32
+ RelayURL: route.RelayURL,
33
+ MatchHostname: route.MatchHostname,
34
+ MatchToken: route.MatchToken,
35
+ }
36
+ encodedOwnerScope, err := json.Marshal(ownerScope)
37
if err != nil {
38
return types.HopRoute{}, err
39
}
31
- route.OwnerPublicKey = owner.PublicKey
32
-
33
- route, err = normalizeHopRoute(route, true)
40
+ ownerToken, err := identity.DeriveToken("hop-owner:" + string(encodedOwnerScope) + ":0")
41
+ if err != nil {
42
+ return types.HopRoute{}, err
43
+ }
44
+ ownerSeed := sha256.Sum256([]byte(ownerToken))
45
+ owner, err := utils.ResolveSecp256k1Identity(hex.EncodeToString(ownerSeed[:]))
46
if err != nil {
47
return types.HopRoute{}, err
48
}
49
+ route.OwnerPublicKey = owner.PublicKey
50
payload, err := types.HopRouteBytes(method, route)
51
if err != nil {
52
return types.HopRoute{}, err
@@ -88,39 +101,3 @@ func normalizeHopRoute(route types.HopRoute, requireOwner bool) (types.HopRoute,
101
route.Signature = strings.TrimSpace(route.Signature)
102
return route, nil
103
}
91
-
92
-func deriveHopRouteOwner(identity types.Identity, route types.HopRoute) (types.Identity, error) {
93
- nonce, err := hopRouteOwnerNonce(route)
94
- if err != nil {
95
- return types.Identity{}, err
96
- }
97
- for counter := range 4 {
98
- token, err := identity.DeriveToken(fmt.Sprintf("hop-owner:%s:%d", nonce, counter))
99
- if err != nil {
100
- return types.Identity{}, err
101
- }
102
- seed := sha256.Sum256([]byte(token))
103
- owner, err := utils.ResolveSecp256k1Identity(hex.EncodeToString(seed[:]))
104
- if err == nil {
105
- return owner, nil
106
- }
107
- }
108
- return types.Identity{}, errors.New("derive hop route owner key")
109
-}
110
-
111
-func hopRouteOwnerNonce(route types.HopRoute) (string, error) {
112
- payload := struct {
113
- RelayURL string `json:"relay_url"`
114
- MatchHostname string `json:"match_hostname"`
115
- MatchToken string `json:"match_token"`
116
- }{
117
- RelayURL: route.RelayURL,
118
- MatchHostname: route.MatchHostname,
119
- MatchToken: route.MatchToken,
120
- }
121
- encoded, err := json.Marshal(payload)
122
- if err != nil {
123
- return "", err
124
- }
125
- return string(encoded), nil
126
-}
portal/overlay/hop_mux.go
+53
-142
@@ -3,7 +3,6 @@ package overlay
3
import (
4
"context"
5
"encoding/binary"
6
- "encoding/json"
6
"errors"
7
"fmt"
8
"io"
@@ -13,18 +12,11 @@ import (
12
"time"
13
14
"github.com/hashicorp/yamux"
16
- "github.com/rs/zerolog/log"
15
)
16
17
const (
20
- hopProtocolVersion = 1
21
- hopPrefaceLimit = 4 << 10
22
- hopIncomingBuffer = 128
23
- defaultPrefaceTimeout = 2 * time.Second
24
-)
25
-
26
-const (
27
- hopModeTLSStream = "tls-stream"
18
+ maxHopTokenBytes = 256
19
+ defaultTokenTimeout = 2 * time.Second
20
)
21
22
type HopMux struct {
@@ -37,15 +29,10 @@ type HopMux struct {
29
done chan struct{}
30
}
31
40
-type hopPreface struct {
41
- Version int `json:"version"`
42
- Mode string `json:"mode"`
43
- Token string `json:"token"`
44
-}
45
-
32
type HopStream struct {
47
- Conn net.Conn
48
- Token string
33
+ Conn net.Conn
34
+ Token string
35
+ RemoteAddr string
36
}
37
38
func NewHopMux(overlay *Overlay) (*HopMux, error) {
@@ -59,17 +46,13 @@ func NewHopMux(overlay *Overlay) (*HopMux, error) {
46
return &HopMux{
47
listener: listener,
48
overlay: overlay,
62
- incoming: make(chan HopStream, hopIncomingBuffer),
49
+ incoming: make(chan HopStream),
50
outbound: make(map[string]*yamux.Session),
51
done: make(chan struct{}),
52
}, nil
53
}
54
55
func (m *HopMux) Serve(ctx context.Context) error {
69
- if m == nil || m.listener == nil {
70
- <-ctx.Done()
71
- return nil
72
- }
56
go func() {
57
<-ctx.Done()
58
_ = m.Close()
@@ -92,10 +75,6 @@ func (m *HopMux) Serve(ctx context.Context) error {
75
}
76
77
func (m *HopMux) Accept(ctx context.Context) (HopStream, error) {
95
- if m == nil {
96
- <-ctx.Done()
97
- return HopStream{}, ctx.Err()
98
- }
78
select {
79
case stream := <-m.incoming:
80
return stream, nil
@@ -105,14 +84,7 @@ func (m *HopMux) Accept(ctx context.Context) (HopStream, error) {
84
}
85
86
func (m *HopMux) Close() error {
108
- if m == nil {
109
- return nil
110
- }
111
-
87
m.mu.Lock()
113
- if m.done == nil {
114
- m.done = make(chan struct{})
115
- }
88
select {
89
case <-m.done:
90
m.mu.Unlock()
@@ -122,19 +94,12 @@ func (m *HopMux) Close() error {
94
close(m.done)
95
sessions := make([]*yamux.Session, 0, len(m.outbound))
96
for _, session := range m.outbound {
125
- if session == nil {
126
- continue
127
- }
97
sessions = append(sessions, session)
98
}
99
m.outbound = make(map[string]*yamux.Session)
131
- listener := m.listener
100
m.mu.Unlock()
101
134
- var closeErr error
135
- if listener != nil {
136
- closeErr = errors.Join(closeErr, listener.Close())
137
- }
102
+ closeErr := m.listener.Close()
103
for _, session := range sessions {
104
closeErr = errors.Join(closeErr, session.Close())
105
}
@@ -147,23 +112,6 @@ func (m *HopMux) OpenStream(ctx context.Context, overlayIPv4, token string) (net
112
return nil, errors.New("next hop token is required")
113
}
114
150
- stream, err := m.openYamuxStream(ctx, overlayIPv4)
151
- if err != nil {
152
- return nil, err
153
- }
154
- preface := hopPreface{
155
- Version: hopProtocolVersion,
156
- Mode: hopModeTLSStream,
157
- Token: token,
158
- }
159
- if err := writeFramedJSON(stream, preface, hopPrefaceLimit); err != nil {
160
- _ = stream.Close()
161
- return nil, err
162
- }
163
- return stream, nil
164
-}
165
-
166
-func (m *HopMux) openYamuxStream(ctx context.Context, overlayIPv4 string) (*yamux.Stream, error) {
115
overlayIPv4 = strings.TrimSpace(overlayIPv4)
116
if overlayIPv4 == "" {
117
return nil, errors.New("next hop overlay ipv4 is required")
@@ -175,10 +123,30 @@ func (m *HopMux) openYamuxStream(ctx context.Context, overlayIPv4 string) (*yamu
123
}
124
stream, err := session.OpenStream()
125
if err != nil {
178
- m.forgetSession(overlayIPv4, session)
126
+ m.mu.Lock()
127
+ if m.outbound[overlayIPv4] == session {
128
+ delete(m.outbound, overlayIPv4)
129
+ }
130
+ m.mu.Unlock()
131
_ = session.Close()
132
return nil, err
133
}
134
+
135
+ payload := []byte(token)
136
+ if len(payload) > maxHopTokenBytes {
137
+ _ = stream.Close()
138
+ return nil, errors.New("next hop token is too large")
139
+ }
140
+ frame := make([]byte, 4+len(payload))
141
+ binary.BigEndian.PutUint32(frame[:4], uint32(len(payload)))
142
+ copy(frame[4:], payload)
143
+ if n, err := stream.Write(frame); err != nil {
144
+ _ = stream.Close()
145
+ return nil, err
146
+ } else if n != len(frame) {
147
+ _ = stream.Close()
148
+ return nil, io.ErrShortWrite
149
+ }
150
return stream, nil
151
}
152
@@ -199,9 +167,6 @@ func (m *HopMux) session(ctx context.Context, overlayIPv4 string) (*yamux.Sessio
167
delete(m.outbound, overlayIPv4)
168
m.mu.Unlock()
169
202
- if m.overlay == nil || m.overlay.stack == nil {
203
- return nil, errors.New("overlay is not initialized")
204
- }
170
addr := net.JoinHostPort(overlayIPv4, fmt.Sprintf("%d", DefaultPeerYamuxPort))
171
conn, err := m.overlay.stack.DialContext(ctx, "tcp", addr)
172
if err != nil {
@@ -232,19 +197,10 @@ func (m *HopMux) session(ctx context.Context, overlayIPv4 string) (*yamux.Sessio
197
return session, nil
198
}
199
235
-func (m *HopMux) forgetSession(overlayIPv4 string, session *yamux.Session) {
236
- m.mu.Lock()
237
- defer m.mu.Unlock()
238
- if m.outbound[overlayIPv4] == session {
239
- delete(m.outbound, overlayIPv4)
240
- }
241
-}
242
-
200
func (m *HopMux) serveSession(ctx context.Context, conn net.Conn) {
201
session, err := yamux.Server(conn, hopYamuxConfig())
202
if err != nil {
203
_ = conn.Close()
247
- log.Warn().Err(err).Msg("create hop yamux session")
204
return
205
}
206
m.mu.Lock()
@@ -282,38 +238,41 @@ func (m *HopMux) serveSession(ctx context.Context, conn net.Conn) {
238
}
239
240
func (m *HopMux) handleStream(ctx context.Context, stream *yamux.Stream) {
285
- _ = stream.SetReadDeadline(time.Now().Add(defaultPrefaceTimeout))
286
- var preface hopPreface
287
- err := readFramedJSON(stream, &preface, hopPrefaceLimit)
288
- _ = stream.SetReadDeadline(time.Time{})
289
- if err != nil {
241
+ _ = stream.SetReadDeadline(time.Now().Add(defaultTokenTimeout))
242
+ defer stream.SetReadDeadline(time.Time{})
243
+
244
+ var size [4]byte
245
+ if _, err := io.ReadFull(stream, size[:]); err != nil {
246
_ = stream.Close()
247
return
248
}
293
- if preface.Version != hopProtocolVersion {
249
+ n := binary.BigEndian.Uint32(size[:])
250
+ if n == 0 || n > uint32(maxHopTokenBytes) {
251
_ = stream.Close()
252
return
253
}
297
- switch preface.Mode {
298
- case hopModeTLSStream:
299
- if strings.TrimSpace(preface.Token) == "" {
300
- _ = stream.Close()
301
- return
302
- }
303
- m.deliver(ctx, HopStream{
304
- Conn: stream,
305
- Token: preface.Token,
306
- })
307
- default:
254
+ payload := make([]byte, n)
255
+ if _, err := io.ReadFull(stream, payload); err != nil {
256
+ _ = stream.Close()
257
+ return
258
+ }
259
+ token := strings.TrimSpace(string(payload))
260
+ if token == "" {
261
_ = stream.Close()
262
+ return
263
+ }
264
+ remoteAddr := ""
265
+ if stream.RemoteAddr() != nil {
266
+ remoteAddr = stream.RemoteAddr().String()
267
}
310
-}
311
-
312
-func (m *HopMux) deliver(ctx context.Context, stream HopStream) {
268
select {
314
- case m.incoming <- stream:
269
+ case m.incoming <- HopStream{
270
+ Conn: stream,
271
+ Token: token,
272
+ RemoteAddr: remoteAddr,
273
+ }:
274
case <-ctx.Done():
316
- _ = stream.Conn.Close()
275
+ _ = stream.Close()
276
}
277
}
278
@@ -325,51 +284,3 @@ func hopYamuxConfig() *yamux.Config {
284
cfg.StreamCloseTimeout = 5 * time.Minute
285
return cfg
286
}
328
-
329
-func writeFramedJSON(w io.Writer, value any, limit int) error {
330
- payload, err := json.Marshal(value)
331
- if err != nil {
332
- return err
333
- }
334
- if len(payload) == 0 || len(payload) > limit {
335
- return errors.New("frame size is invalid")
336
- }
337
- var size [4]byte
338
- binary.BigEndian.PutUint32(size[:], uint32(len(payload)))
339
- if err := writeAll(w, size[:]); err != nil {
340
- return err
341
- }
342
- return writeAll(w, payload)
343
-}
344
-
345
-func readFramedJSON(r io.Reader, dst any, limit int) error {
346
- var size [4]byte
347
- if _, err := io.ReadFull(r, size[:]); err != nil {
348
- return err
349
- }
350
- n := binary.BigEndian.Uint32(size[:])
351
- if n == 0 || n > uint32(limit) {
352
- return errors.New("frame size is invalid")
353
- }
354
- payload := make([]byte, n)
355
- if _, err := io.ReadFull(r, payload); err != nil {
356
- return err
357
- }
358
- return json.Unmarshal(payload, dst)
359
-}
360
-
361
-func writeAll(w io.Writer, p []byte) error {
362
- for len(p) > 0 {
363
- n, err := w.Write(p)
364
- if n > 0 {
365
- p = p[n:]
366
- }
367
- if err != nil {
368
- return err
369
- }
370
- if n == 0 {
371
- return io.ErrShortWrite
372
- }
373
- }
374
- return nil
375
-}
portal/server.go
+8
-1
@@ -519,7 +519,14 @@ func (s *Server) runHopMux(ctx context.Context) error {
519
}
520
go func(stream overlay.HopStream) {
521
record, ok := s.registry.RecordByHopToken(stream.Token, time.Now())
522
- if !ok || !s.bridgeLeaseConn(groupCtx, stream.Conn, record) {
522
+ if !ok {
523
+ log.Warn().Str("remote_addr", stream.RemoteAddr).Msg("hop stream rejected")
524
+ _ = stream.Conn.Close()
525
+ return
526
+ }
527
+ log.Info().Str("remote_addr", stream.RemoteAddr).Bool("forward", record.isHopForward()).Msg("hop stream received")
528
+ if !s.bridgeLeaseConn(groupCtx, stream.Conn, record) {
529
+ log.Warn().Str("remote_addr", stream.RemoteAddr).Msg("hop stream bridge failed")
530
_ = stream.Conn.Close()
531
}
532
}(stream)
sdk/api_client.go
+2
-2
@@ -11,7 +11,7 @@ import (
11
"strings"
12
"time"
13
14
- "github.com/gosuda/portal-tunnel/v2/portal/overlay"
14
+ "github.com/gosuda/portal-tunnel/v2/portal/auth"
15
"github.com/gosuda/portal-tunnel/v2/types"
16
"github.com/gosuda/portal-tunnel/v2/utils"
17
)
@@ -227,7 +227,7 @@ func (l *listener) syncHopRoutes(ctx context.Context, method string, expiresAt t
227
228
var syncErr error
229
for _, unsignedRoute := range orderedRoutes {
230
- route, err := overlay.SignHopRoute(method, unsignedRoute, l.identity, expiresAt)
230
+ route, err := auth.SignHopRoute(method, unsignedRoute, l.identity, expiresAt)
231
if err != nil {
232
if method == http.MethodDelete {
233
syncErr = errors.Join(syncErr, err)