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)