main
go 513 lines 11.4 KB
Raw
1 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/identity"
22 "github.com/gosuda/portal-tunnel/v2/types"
23 "github.com/gosuda/portal-tunnel/v2/utils"
24 )
25
26 const (
27 DefaultMTU = 1420
28 DefaultListenPort = 51820
29 DefaultPeerAPIHTTPPort = 7777
30 DefaultPeerYamuxPort = 7778
31 DefaultPersistentKeepalive = 25
32
33 maxHopTokenBytes = 256
34 defaultTokenTimeout = 2 * time.Second
35 )
36
37 type Config struct {
38 PrivateKey string
39 PublicKey string
40 ListenPort int
41 }
42
43 func (c Config) Copy() Config {
44 return Config{
45 PrivateKey: c.PrivateKey,
46 PublicKey: c.PublicKey,
47 ListenPort: c.ListenPort,
48 }
49 }
50
51 func NormalizeConfig(cfg Config) (Config, error) {
52 configured := strings.TrimSpace(cfg.PrivateKey) != "" ||
53 strings.TrimSpace(cfg.PublicKey) != "" ||
54 cfg.ListenPort != 0
55 if !configured {
56 return cfg, nil
57 }
58
59 if strings.TrimSpace(cfg.PrivateKey) == "" {
60 return Config{}, errors.New("wireguard private key is required when relay overlay is enabled")
61 }
62
63 privateKey, err := identity.NormalizeWireGuardPrivateKey(cfg.PrivateKey)
64 if err != nil {
65 return Config{}, fmt.Errorf("normalize wireguard private key: %w", err)
66 }
67 publicKey, err := identity.WireGuardPublicKeyFromPrivate(privateKey)
68 if err != nil {
69 return Config{}, fmt.Errorf("derive wireguard public key: %w", err)
70 }
71 if configuredPublicKey := strings.TrimSpace(cfg.PublicKey); configuredPublicKey != "" && configuredPublicKey != publicKey {
72 return Config{}, errors.New("wireguard public key does not match private key")
73 }
74
75 cfg.PrivateKey = privateKey
76 cfg.PublicKey = publicKey
77 if cfg.ListenPort == 0 {
78 cfg.ListenPort = DefaultListenPort
79 }
80 if cfg.ListenPort < 0 || cfg.ListenPort > 65535 {
81 return Config{}, errors.New("wireguard listen port is invalid")
82 }
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
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
110 }
111 publicKey := strings.TrimSpace(cfg.PublicKey)
112 if publicKey == "" {
113 return nil, errors.New("wireguard public key is required")
114 }
115
116 stack, err := newStack(cfg)
117 if err != nil {
118 return nil, err
119 }
120
121 listener, err := stack.ListenTCP(DefaultPeerAPIHTTPPort)
122 if err != nil {
123 _ = stack.Close()
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,
137 }
138
139 client := utils.NewHTTPClient(
140 utils.WithHTTPDialContext(stack.DialContext),
141 utils.WithHTTPTLSHandshakeTimeout(10*time.Second),
142 utils.WithHTTPMaxIdleConns(100),
143 utils.WithHTTPIdleConnTimeout(90*time.Second),
144 utils.WithHTTPResponseHeaderTimeout(30*time.Second),
145 utils.WithHTTPExpectContinueTimeout(1*time.Second),
146 utils.WithoutHTTP2(),
147 )
148
149 publicCfg := cfg.Copy()
150 publicCfg.PrivateKey = ""
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 {
167 if o == nil {
168 return Config{}
169 }
170 return o.cfg.Copy()
171 }
172
173 func (o *Overlay) Serve(ctx context.Context) error {
174 if o == nil {
175 return nil
176 }
177
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 err
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 }
408 }
409
410 func (o *Overlay) Shutdown(ctx context.Context) error {
411 if o == nil {
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)
424 if err != nil && !errors.Is(err, http.ErrServerClosed) {
425 shutdownErr = errors.Join(shutdownErr, err)
426 }
427 }
428 if o.client != nil {
429 o.client.CloseIdleConnections()
430 }
431 if o.listener != nil {
432 err := o.listener.Close()
433 if err != nil && !errors.Is(err, net.ErrClosed) {
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 }
454 return shutdownErr
455 }
456
457 func (o *Overlay) Client() *http.Client {
458 if o == nil || o.stack == nil {
459 return nil
460 }
461 return o.client
462 }
463
464 func (o *Overlay) DiscoverRelay(ctx context.Context, relay types.RelayDescriptor) (types.DiscoveryResponse, error) {
465 if o == nil || o.stack == nil {
466 return types.DiscoveryResponse{}, errors.New("overlay is not initialized")
467 }
468 if !relay.HasOverlayPeer() {
469 return types.DiscoveryResponse{}, errors.New("relay wireguard overlay metadata is required")
470 }
471 overlayIPv4, err := identity.DeriveWireGuardOverlayIPv4(relay.WireGuardPublicKey)
472 if err != nil {
473 return types.DiscoveryResponse{}, err
474 }
475
476 var resp types.DiscoveryResponse
477 baseURL := &url.URL{
478 Scheme: "http",
479 Host: net.JoinHostPort(overlayIPv4, fmt.Sprintf("%d", DefaultPeerAPIHTTPPort)),
480 }
481 if err := utils.HTTPDoAPIPath(ctx, o.Client(), baseURL, http.MethodGet, types.PathDiscovery, nil, nil, &resp); err != nil {
482 return types.DiscoveryResponse{}, err
483 }
484 return resp, nil
485 }
486
487 func (o *Overlay) Sync(relays []types.RelayDescriptor) error {
488 if o == nil || o.stack == nil {
489 return nil
490 }
491
492 peers := make([]types.RelayDescriptor, 0, len(relays))
493 for _, desc := range relays {
494 if !desc.HasOverlayPeer() {
495 continue
496 }
497 if desc.WireGuardPublicKey == o.cfg.PublicKey {
498 continue
499 }
500 peers = append(peers, desc)
501 }
502 sort.Slice(peers, func(i, j int) bool {
503 return peers[i].WireGuardPublicKey < peers[j].WireGuardPublicKey
504 })
505 return o.stack.ApplyPeers(peers)
506 }
507
508 func hopYamuxConfig() *yamux.Config {
509 cfg := yamux.DefaultConfig()
510 cfg.Logger = nil
511 cfg.MaxStreamWindowSize = 16 * 1024 * 1024
512 return cfg
513 }