File
Blob: tun/client/lightweight_token.go
| 1 | package client |
| 2 | |
| 3 | import ( |
| 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 | |
| 19 | type 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. |
| 26 | type tokenConnections struct { |
| 27 | client *LightweightClient |
| 28 | slots [tun.NumRedundantLinks]*tokenConnection |
| 29 | candidates []*protocol.Node |
| 30 | wake chan struct{} |
| 31 | backoff time.Duration |
| 32 | } |
| 33 | |
| 34 | func (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 | |
| 52 | func (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 | |
| 162 | func (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 | |
| 187 | func (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 | |
| 236 | func (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 | |
| 248 | func (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 | |
| 254 | func (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 | |
| 263 | func (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 | |
| 273 | func (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 | |
| 282 | func 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 | } |