Skip to content
File

Blob: tun/client/lightweight_token.go

go296 lines
1package client
2 
3import (
4 "context"
5 "errors"
6 "fmt"
7 "net"
8 "time"
9 
10 "go.miragespace.co/specter/spec/protocol"
11 "go.miragespace.co/specter/spec/transport"
12 "go.miragespace.co/specter/spec/tun"
13 "go.miragespace.co/specter/util"
14 
15 "github.com/twitchtv/twirp"
16 "go.uber.org/zap"
17)
18 
19type tokenConnection struct {
20 conn transport.PhysicalConn
21 node *protocol.Node
22}
23 
24// Only the reconnect callback changes slots and candidates. Connection watchers
25// send coalesced notifications; they never mutate connection or publication state.
26type tokenConnections struct {
27 client *LightweightClient
28 slots [tun.NumRedundantLinks]*tokenConnection
29 candidates []*protocol.Node
30 wake chan struct{}
31 backoff time.Duration
32}
33 
34func (l *LightweightClient) runToken(ctx context.Context) error {
35 c := &tokenConnections{
36 client: l,
37 wake: make(chan struct{}, 1),
38 backoff: time.Second,
39 }
40 defer func() {
41 for slot, connection := range c.slots {
42 if connection != nil {
43 connection.conn.Close("token tunnel ended")
44 c.remove(slot)
45 }
46 }
47 }()
48 c.wake <- struct{}{}
49 return runReconnectLoop(ctx, c.wake, c.reconcile)
50}
51 
52func (c *tokenConnections) reconcile(ctx context.Context) (time.Duration, error) {
53 c.prune()
54 if c.vacancy() < 0 {
55 return 0, nil
56 }
57 apex := &protocol.Node{Address: c.client.Apex.String()}
58 var bootstrap transport.PhysicalConn
59 tried := make(map[string]bool)
60 attempts, unsupportedCount := 0, 0
61 if c.count() == 0 {
62 stream, pc, err := c.client.dialAttachment(ctx, apex)
63 if err != nil {
64 if tokenFatal(err) {
65 return 0, err
66 }
67 tried[apex.Address], attempts = true, 1
68 c.client.Logger.Warn("Bootstrap connection failed", zap.Error(err))
69 } else {
70 bootstrap = pc
71 defer func() {
72 if !c.hasConnection(pc, "") {
73 pc.Close("bootstrap attempt ended")
74 }
75 }()
76 // Discovery does not depend on session admission. Reuse this physical
77 // connection for Open even when discovery itself is unavailable.
78 if err := c.discover(ctx, stream); err != nil {
79 c.client.Logger.Warn("Server discovery failed", zap.Error(err))
80 }
81 }
82 } else {
83 for _, connection := range c.slots {
84 if connection == nil {
85 continue
86 }
87 stream, err := connection.conn.OpenStream(protocol.Stream_RPC)
88 if err == nil {
89 err = c.discover(ctx, stream)
90 }
91 if err == nil {
92 break
93 }
94 c.client.Logger.Warn("Server discovery failed", zap.String("server", connection.node.GetAddress()), zap.Error(err))
95 }
96 }
97 candidates := append(append([]*protocol.Node{}, c.candidates...), apex)
98 if bootstrap != nil {
99 candidates = append([]*protocol.Node{apex}, candidates...)
100 }
101 for _, candidate := range candidates {
102 if ctx.Err() != nil {
103 return 0, ctx.Err()
104 }
105 c.prune()
106 slot := c.vacancy()
107 if slot < 0 {
108 break
109 }
110 address := candidate.GetAddress()
111 if tried[address] || c.hasConnection(nil, address) {
112 continue
113 }
114 tried[address] = true
115 attempts++
116 pc := bootstrap
117 var stream net.Conn
118 var err error
119 if pc != nil && address == apex.Address {
120 stream, err = pc.OpenStream(protocol.Stream_RPC)
121 } else {
122 stream, pc, err = c.client.dialAttachment(ctx, candidate)
123 }
124 if err != nil {
125 if tokenFatal(err) {
126 return 0, err
127 }
128 c.client.Logger.Warn("Tunnel connection failed", zap.String("server", address), zap.Error(err))
129 continue
130 }
131 // An apex alias and a discovered address may share the transport's cached
132 // connection. Never issue another Open or close an already owned handle.
133 if c.hasConnection(pc, "") {
134 stream.Close()
135 continue
136 }
137 if err := c.open(ctx, stream, pc, slot); err != nil {
138 pc.Close("token attachment failed")
139 if tokenFatal(err) {
140 return 0, err
141 }
142 if unsupported(err) {
143 unsupportedCount++
144 }
145 c.client.Logger.Warn("Tunnel connection failed", zap.String("server", address), zap.Error(err))
146 }
147 }
148 c.prune()
149 if c.count() > 0 {
150 c.backoff = time.Second
151 return 0, nil
152 }
153 if attempts > 0 && attempts == unsupportedCount {
154 return 0, fmt.Errorf("servers do not support lightweight tunnels")
155 }
156 delay := util.RandomTimeRange(c.backoff)
157 c.backoff = min(2*c.backoff, 30*time.Second)
158 c.client.Logger.Warn("Retrying token tunnel", zap.Duration("backoff", delay))
159 return delay, nil
160}
161 
162func (c *tokenConnections) discover(ctx context.Context, stream net.Conn) error {
163 defer stream.Close()
164 ctx, cancel := context.WithTimeout(ctx, rpcTimeout)
165 defer cancel()
166 resp, err := tunnelClientOnStream(stream).GetNodes(ctx, &protocol.GetNodesRequest{})
167 if err != nil {
168 return err
169 }
170 nodes := make([]*protocol.Node, 0, tun.NumRedundantLinks)
171 seen := make(map[string]bool)
172 for _, node := range resp.GetNodes() {
173 address := node.GetAddress()
174 if address == "" || seen[address] {
175 continue
176 }
177 seen[address] = true
178 nodes = append(nodes, node)
179 if len(nodes) == tun.NumRedundantLinks {
180 break
181 }
182 }
183 c.candidates = nodes
184 return nil
185}
186 
187func (c *tokenConnections) open(ctx context.Context, stream net.Conn, pc transport.PhysicalConn, slot int) error {
188 defer stream.Close()
189 callCtx, cancel := context.WithTimeout(ctx, 15*time.Second)
190 defer cancel()
191 resp, err := tunnelClientOnStream(stream).OpenDelegatedSession(callCtx, &protocol.OpenDelegatedSessionRequest{
192 Token: c.client.Token,
193 RouteSlot: uint32(slot + 1),
194 })
195 if err != nil {
196 return err
197 }
198 c.prune()
199 if c.hasConnection(nil, resp.GetNode().GetAddress()) {
200 return fmt.Errorf("server is already attached")
201 }
202 if ctx.Err() != nil {
203 return ctx.Err()
204 }
205 select {
206 case <-pc.Done():
207 return fmt.Errorf("connection ended during activation: %v", pc.Err())
208 default:
209 }
210 if err := c.client.acceptSession(resp); err != nil {
211 return err
212 }
213 c.slots[slot] = &tokenConnection{
214 conn: pc,
215 node: resp.GetNode(),
216 }
217 expiry := "none"
218 if resp.GetExpiresAt() != 0 {
219 expiry = time.Unix(resp.GetExpiresAt(), 0).UTC().Format(time.RFC3339)
220 }
221 c.client.Logger.Info("Tunnel connection ready", zap.Int("slot", slot+1), zap.String("server", resp.GetNode().GetAddress()), zap.String("grantId", resp.GetGrantId()), zap.String("expiresAt", expiry))
222 go func() {
223 select {
224 case <-ctx.Done():
225 return
226 case <-pc.Done():
227 }
228 select {
229 case c.wake <- struct{}{}:
230 default:
231 }
232 }()
233 return nil
234}
235 
236func (c *tokenConnections) prune() {
237 for slot, connection := range c.slots {
238 if connection != nil {
239 select {
240 case <-connection.conn.Done():
241 c.remove(slot)
242 default:
243 }
244 }
245 }
246}
247 
248func (c *tokenConnections) remove(slot int) {
249 connection := c.slots[slot]
250 c.slots[slot] = nil
251 c.client.Logger.Info("Tunnel disconnected", zap.Int("slot", slot+1), zap.String("server", connection.node.GetAddress()), zap.Error(connection.conn.Err()))
252}
253 
254func (c *tokenConnections) vacancy() int {
255 for slot, connection := range c.slots {
256 if connection == nil {
257 return slot
258 }
259 }
260 return -1
261}
262 
263func (c *tokenConnections) count() int {
264 count := 0
265 for _, connection := range c.slots {
266 if connection != nil {
267 count++
268 }
269 }
270 return count
271}
272 
273func (c *tokenConnections) hasConnection(pc transport.PhysicalConn, address string) bool {
274 for _, connection := range c.slots {
275 if connection != nil && (connection.conn == pc || connection.node.GetAddress() == address) {
276 return true
277 }
278 }
279 return false
280}
281 
282func tokenFatal(err error) bool {
283 var local *attachmentError
284 if errors.As(err, &local) {
285 return true
286 }
287 var rpcError twirp.Error
288 if errors.As(err, &rpcError) {
289 switch rpcError.Code() {
290 case twirp.InvalidArgument, twirp.Unauthenticated, twirp.PermissionDenied, twirp.NotFound:
291 return true
292 }
293 }
294 return false
295}