Skip to content
File

Blob: tun/client/dialer/dialer.go

go204 lines
1package dialer
2 
3import (
4 "context"
5 "crypto/tls"
6 "io"
7 "net"
8 "sync/atomic"
9 "time"
10 
11 "go.miragespace.co/specter/overlay"
12 "go.miragespace.co/specter/spec/protocol"
13 "go.miragespace.co/specter/spec/transport"
14 "go.miragespace.co/specter/spec/tun"
15 
16 "github.com/libp2p/go-yamux/v4"
17 "github.com/quic-go/quic-go"
18 "go.uber.org/zap"
19)
20 
21type DialerConfig struct {
22 Logger *zap.Logger
23 Parsed *ParsedApex
24 InsecureSkipVerify bool
25 NoReconnection bool
26}
27 
28type TransportDialer interface {
29 Dial() (net.Conn, error)
30 Remote() net.Addr
31}
32 
33// TODO: can we unit test this somehow
34var rebootstrapRetry = time.Second * 5
35 
36type bootstrapFn func() (net.Addr, error)
37 
38func TLSDialer(ctx context.Context, dCfg DialerConfig) (net.Addr, TransportDialer, error) {
39 var (
40 bootstrap bootstrapFn
41 aSession atomic.Value
42 )
43 
44 clientTLSConf := &tls.Config{
45 ServerName: dCfg.Parsed.Host,
46 InsecureSkipVerify: dCfg.InsecureSkipVerify,
47 NextProtos: []string{
48 tun.ALPN(protocol.Link_TCP),
49 },
50 }
51 
52 override := GetServerNameOverride(ctx)
53 if override != "" {
54 clientTLSConf.ServerName = override
55 }
56 
57 dialer := &tls.Dialer{
58 Config: clientTLSConf,
59 }
60 
61 bootstrap = func() (net.Addr, error) {
62 dCfg.Logger.Debug("bootstrapping yamux")
63 openCtx, cancel := context.WithTimeout(ctx, transport.ConnectTimeout)
64 defer cancel()
65 
66 conn, err := dialer.DialContext(openCtx, "tcp", dCfg.Parsed.String())
67 if err != nil {
68 return nil, err
69 }
70 
71 cfg := yamux.DefaultConfig()
72 cfg.LogOutput = io.Discard
73 session, err := yamux.Client(conn, cfg, nil)
74 if err != nil {
75 return nil, err
76 }
77 
78 aSession.Store(session)
79 if !dCfg.NoReconnection {
80 go rebootstrap(ctx, dCfg.Logger, bootstrap, session.CloseChan())
81 }
82 go func() {
83 select {
84 case <-ctx.Done():
85 session.Close()
86 case <-session.CloseChan():
87 return
88 }
89 }()
90 
91 return conn.RemoteAddr(), nil
92 }
93 
94 remote, err := bootstrap()
95 if err != nil {
96 return nil, nil, err
97 }
98 
99 return remote, &tlsDialer{aSession, ctx}, nil
100}
101 
102type tlsDialer struct {
103 aSession atomic.Value
104 ctx context.Context
105}
106 
107var _ TransportDialer = (*tlsDialer)(nil)
108 
109func (t *tlsDialer) Dial() (net.Conn, error) {
110 session := t.aSession.Load().(*yamux.Session)
111 return session.OpenStream(t.ctx)
112}
113 
114func (t *tlsDialer) Remote() net.Addr {
115 session := t.aSession.Load().(*yamux.Session)
116 return session.RemoteAddr()
117}
118 
119func QuicDialer(ctx context.Context, dCfg DialerConfig) (net.Addr, TransportDialer, error) {
120 var (
121 bootstrap bootstrapFn
122 aQuic atomic.Value
123 )
124 
125 clientTLSConf := &tls.Config{
126 ServerName: dCfg.Parsed.Host,
127 InsecureSkipVerify: dCfg.InsecureSkipVerify,
128 NextProtos: []string{
129 tun.ALPN(protocol.Link_TCP),
130 },
131 }
132 
133 override := GetServerNameOverride(ctx)
134 if override != "" {
135 clientTLSConf.ServerName = override
136 }
137 
138 bootstrap = func() (net.Addr, error) {
139 dCfg.Logger.Debug("bootstrapping quic")
140 q, err := quic.DialAddrEarly(ctx, dCfg.Parsed.String(), clientTLSConf, &quic.Config{
141 KeepAlivePeriod: time.Second * 5,
142 HandshakeIdleTimeout: transport.ConnectTimeout,
143 MaxIdleTimeout: time.Second * 30,
144 EnableDatagrams: true,
145 })
146 if err != nil {
147 return nil, err
148 }
149 
150 aQuic.Store(q)
151 if !dCfg.NoReconnection {
152 go rebootstrap(ctx, dCfg.Logger, bootstrap, q.Context().Done())
153 }
154 
155 return q.RemoteAddr(), nil
156 }
157 
158 remote, err := bootstrap()
159 if err != nil {
160 return nil, nil, err
161 }
162 
163 return remote, &quicDialer{aQuic}, nil
164}
165 
166type quicDialer struct {
167 aQuic atomic.Value // quic.Connection
168}
169 
170var _ TransportDialer = (*quicDialer)(nil)
171 
172func (q *quicDialer) Dial() (net.Conn, error) {
173 qConn := q.aQuic.Load().(*quic.Conn)
174 r, err := qConn.OpenStream()
175 if err != nil {
176 return nil, err
177 }
178 return overlay.WrapQuicConnection(r, qConn), nil
179}
180 
181func (q *quicDialer) Remote() net.Addr {
182 qConn := q.aQuic.Load().(*quic.Conn)
183 return qConn.RemoteAddr()
184}
185 
186func rebootstrap(ctx context.Context, logger *zap.Logger, fn bootstrapFn, exit <-chan struct{}) {
187 logger.Debug("re-bootstrap started")
188 select {
189 case <-ctx.Done():
190 return
191 case <-exit:
192 logger.Info("Disconnected from gateway, re-bootstrapping")
193 AGAIN:
194 logger.Debug("calling bootstrapFn")
195 remote, err := fn()
196 if err != nil {
197 logger.Error("Failed to re-bootstrap, retrying", zap.Error(err))
198 time.Sleep(rebootstrapRetry)
199 goto AGAIN
200 }
201 logger.Info("Connection to gateway re-established", zap.String("via", remote.String()))
202 }
203}