File
Blob: tun/client/tunnel.go
| 1 | package client |
| 2 | |
| 3 | import ( |
| 4 | "context" |
| 5 | "errors" |
| 6 | "fmt" |
| 7 | "strings" |
| 8 | "time" |
| 9 | |
| 10 | "go.miragespace.co/specter/spec/protocol" |
| 11 | "go.miragespace.co/specter/spec/rpc" |
| 12 | |
| 13 | "go.uber.org/zap" |
| 14 | ) |
| 15 | |
| 16 | // TunnelSyncResult reports the last publication acknowledgement for a configured |
| 17 | // tunnel. Publication is not an end-to-end health check of its target. |
| 18 | type TunnelSyncResult struct { |
| 19 | Hostname string `json:"hostname"` |
| 20 | Target string `json:"target"` |
| 21 | Published bool `json:"published"` |
| 22 | PublishedEndpoints int `json:"publishedEndpoints"` |
| 23 | Error string `json:"error,omitempty"` |
| 24 | } |
| 25 | |
| 26 | type SyncResult struct { |
| 27 | Applied bool `json:"applied"` |
| 28 | Saved bool `json:"saved"` |
| 29 | Error string `json:"error,omitempty"` |
| 30 | AttemptedAt *time.Time `json:"attemptedAt,omitempty"` |
| 31 | Tunnels []TunnelSyncResult `json:"tunnels"` |
| 32 | } |
| 33 | |
| 34 | func (r SyncResult) pendingPublication() bool { |
| 35 | for _, tunnel := range r.Tunnels { |
| 36 | if tunnel.Error != "" { |
| 37 | return true |
| 38 | } |
| 39 | } |
| 40 | return false |
| 41 | } |
| 42 | |
| 43 | type publicationState struct { |
| 44 | endpoints string |
| 45 | published int |
| 46 | } |
| 47 | |
| 48 | // ConfigSaveError means that the operation changed live state, but the updated |
| 49 | // configuration could not be saved. Reloading the old file may undo that change. |
| 50 | type ConfigSaveError struct{ Err error } |
| 51 | |
| 52 | func (e *ConfigSaveError) Error() string { |
| 53 | return fmt.Sprintf("change applied, but configuration was not saved: %v", e.Err) |
| 54 | } |
| 55 | |
| 56 | func (e *ConfigSaveError) Unwrap() error { return e.Err } |
| 57 | |
| 58 | func (c *Client) requestHostname(ctx context.Context) (string, error) { |
| 59 | connected := c.getConnectedNodes() |
| 60 | if len(connected) == 0 { |
| 61 | return "", fmt.Errorf("no rpc candidates available") |
| 62 | } |
| 63 | // Generation is not idempotent. After an ambiguous failure, the next sync |
| 64 | // queries registered hostnames before attempting to generate another one. |
| 65 | callCtx, cancel := context.WithTimeout(ctx, rpcTimeout) |
| 66 | defer cancel() |
| 67 | candidate := connected[c.hostnameNextCandidate%uint64(len(connected))] |
| 68 | resp, err := c.tunnelClient.GenerateHostname(rpc.WithNode(callCtx, candidate), &protocol.GenerateHostnameRequest{}) |
| 69 | if err != nil { |
| 70 | c.hostnameNextCandidate++ |
| 71 | return "", err |
| 72 | } |
| 73 | if resp.GetHostname() == "" { |
| 74 | c.hostnameNextCandidate++ |
| 75 | return "", fmt.Errorf("server returned an empty generated hostname") |
| 76 | } |
| 77 | return resp.GetHostname(), nil |
| 78 | } |
| 79 | |
| 80 | func endpointSignature(nodes []*protocol.Node) string { |
| 81 | var signature strings.Builder |
| 82 | for _, node := range nodes { |
| 83 | fmt.Fprintf(&signature, "%d:%s;", node.GetId(), node.GetAddress()) |
| 84 | } |
| 85 | return signature.String() |
| 86 | } |
| 87 | |
| 88 | // SyncConfigTunnels explicitly republishes the current configuration. Automatic |
| 89 | // retries use the successful acknowledgements from this attempt to retry only |
| 90 | // unfinished work. |
| 91 | func (c *Client) SyncConfigTunnels(ctx context.Context) SyncResult { |
| 92 | c.syncMu.Lock() |
| 93 | defer c.syncMu.Unlock() |
| 94 | return c.syncConfigTunnels(ctx, true) |
| 95 | } |
| 96 | |
| 97 | // syncConfigTunnels requires syncMu to serialize reload, removal, and retries. |
| 98 | func (c *Client) syncConfigTunnels(ctx context.Context, force bool) SyncResult { |
| 99 | ctx, cancel := context.WithTimeout(ctx, 30*time.Second) |
| 100 | defer cancel() |
| 101 | if force || c.publication == nil { |
| 102 | c.publication = make(map[string]publicationState) |
| 103 | } |
| 104 | c.configMu.RLock() |
| 105 | tunnels := append([]Tunnel{}, c.Configuration.Tunnels...) |
| 106 | c.configMu.RUnlock() |
| 107 | |
| 108 | now := time.Now() |
| 109 | result := SyncResult{Applied: true, AttemptedAt: &now, Tunnels: make([]TunnelSyncResult, len(tunnels))} |
| 110 | var syncErrors []error |
| 111 | for i, tunnel := range tunnels { |
| 112 | result.Tunnels[i] = TunnelSyncResult{Hostname: tunnel.Hostname, Target: tunnel.Target} |
| 113 | } |
| 114 | |
| 115 | c.Logger.Info("Synchronizing tunnels in config file with specter", zap.Int("tunnels", len(tunnels))) |
| 116 | |
| 117 | c.syncStateMu.RLock() |
| 118 | result.Saved = !c.lastSync.Applied || c.lastSync.Saved |
| 119 | c.syncStateMu.RUnlock() |
| 120 | hostnameAssigned := false |
| 121 | needsHostname := false |
| 122 | for _, tunnel := range tunnels { |
| 123 | needsHostname = needsHostname || tunnel.Hostname == "" |
| 124 | } |
| 125 | available := make([]string, 0) |
| 126 | var lookupErr error |
| 127 | if force || needsHostname { |
| 128 | registered, err := c.GetRegisteredHostnames(ctx) |
| 129 | if err != nil { |
| 130 | lookupErr = fmt.Errorf("querying registered hostnames: %w", err) |
| 131 | syncErrors = append(syncErrors, lookupErr) |
| 132 | } else { |
| 133 | inUse := make(map[string]bool) |
| 134 | for _, tunnel := range tunnels { |
| 135 | inUse[tunnel.Hostname] = true |
| 136 | } |
| 137 | for _, hostname := range registered { |
| 138 | if !strings.Contains(hostname, ".") && !inUse[hostname] { |
| 139 | available = append(available, hostname) |
| 140 | inUse[hostname] = true |
| 141 | } |
| 142 | } |
| 143 | } |
| 144 | } |
| 145 | |
| 146 | connected := c.getConnectedNodes() |
| 147 | endpoints := endpointSignature(connected) |
| 148 | var generationErr error |
| 149 | for i := range tunnels { |
| 150 | tunnel := &tunnels[i] |
| 151 | outcome := &result.Tunnels[i] |
| 152 | if tunnel.Hostname == "" { |
| 153 | if lookupErr != nil { |
| 154 | outcome.Error = lookupErr.Error() |
| 155 | continue |
| 156 | } |
| 157 | var err error |
| 158 | if len(available) > 0 { |
| 159 | tunnel.Hostname, available = available[0], available[1:] |
| 160 | } else if generationErr != nil { |
| 161 | err = fmt.Errorf("hostname generation deferred until registered names are reconciled: %w", generationErr) |
| 162 | } else { |
| 163 | tunnel.Hostname, err = c.requestHostname(ctx) |
| 164 | generationErr = err |
| 165 | } |
| 166 | if err != nil { |
| 167 | outcome.Error = fmt.Sprintf("requesting hostname: %v", err) |
| 168 | syncErrors = append(syncErrors, fmt.Errorf("%s: %s", tunnel.Target, outcome.Error)) |
| 169 | continue |
| 170 | } |
| 171 | outcome.Hostname = tunnel.Hostname |
| 172 | hostnameAssigned = true |
| 173 | } |
| 174 | |
| 175 | if previous, ok := c.publication[tunnel.Hostname]; !force && ok && previous.endpoints == endpoints { |
| 176 | outcome.Published = true |
| 177 | outcome.PublishedEndpoints = previous.published |
| 178 | continue |
| 179 | } |
| 180 | // A fresh attempt can partially replace the numbered routing slots. |
| 181 | // Its failure invalidates any earlier acknowledgement, even if a later |
| 182 | // retry returns to the same gateway ordering as that acknowledgement. |
| 183 | delete(c.publication, tunnel.Hostname) |
| 184 | published, err := c.publishTunnel(ctx, tunnel.Hostname, connected) |
| 185 | outcome.PublishedEndpoints = len(published) |
| 186 | outcome.Published = len(published) > 0 |
| 187 | if err == nil && len(published) < len(connected) { |
| 188 | err = fmt.Errorf("published %d of %d connected gateways", len(published), len(connected)) |
| 189 | } |
| 190 | if err == nil && len(published) == 0 { |
| 191 | err = fmt.Errorf("no gateways acknowledged publication") |
| 192 | } |
| 193 | if err != nil { |
| 194 | outcome.Error = err.Error() |
| 195 | syncErrors = append(syncErrors, fmt.Errorf("%s: %w", tunnel.Hostname, err)) |
| 196 | c.Logger.Error("Failed to fully publish tunnel", zap.String("hostname", tunnel.Hostname), zap.Error(err)) |
| 197 | continue |
| 198 | } |
| 199 | c.publication[tunnel.Hostname] = publicationState{endpoints: endpoints, published: len(published)} |
| 200 | fqdn := tunnel.Hostname |
| 201 | if !strings.Contains(fqdn, ".") { |
| 202 | fqdn = fmt.Sprintf("%s.%s", fqdn, c.rootDomain.Load()) |
| 203 | } |
| 204 | c.Logger.Info("Tunnel published", zap.String("hostname", fqdn), zap.String("target", tunnel.Target), zap.Int("published", len(published))) |
| 205 | } |
| 206 | |
| 207 | // Retain generated hostnames in live state even when publication or saving |
| 208 | // fails, so retrying cannot generate a second name for the same tunnel. |
| 209 | // Publication-only retries must not overwrite edits waiting in the YAML |
| 210 | // file for an explicit reload. |
| 211 | if force || hostnameAssigned { |
| 212 | err := c.RebuildTunnels(tunnels) |
| 213 | result.Saved = err == nil |
| 214 | syncErrors = append(syncErrors, err) |
| 215 | } else if !result.Saved { |
| 216 | syncErrors = append(syncErrors, errors.New("configuration changes are not saved")) |
| 217 | } |
| 218 | if err := errors.Join(syncErrors...); err != nil { |
| 219 | result.Error = err.Error() |
| 220 | } |
| 221 | c.recordSyncResult(result) |
| 222 | return result |
| 223 | } |
| 224 | |
| 225 | func (c *Client) recordSyncResult(result SyncResult) { |
| 226 | c.syncStateMu.Lock() |
| 227 | defer c.syncStateMu.Unlock() |
| 228 | c.lastSync = result |
| 229 | c.lastSync.Tunnels = append([]TunnelSyncResult{}, result.Tunnels...) |
| 230 | if !result.pendingPublication() { |
| 231 | c.syncBackoff = 0 |
| 232 | c.nextSync = time.Time{} |
| 233 | return |
| 234 | } |
| 235 | if c.syncBackoff == 0 { |
| 236 | c.syncBackoff = checkInterval |
| 237 | } else { |
| 238 | c.syncBackoff *= 2 |
| 239 | } |
| 240 | if c.syncBackoff > 5*time.Minute { |
| 241 | c.syncBackoff = 5 * time.Minute |
| 242 | } |
| 243 | c.nextSync = time.Now().Add(c.syncBackoff) |
| 244 | } |
| 245 | |
| 246 | func (c *Client) publishTunnel(ctx context.Context, hostname string, connected []*protocol.Node) ([]*protocol.Node, error) { |
| 247 | ctx, cancel := context.WithTimeout(ctx, rpcTimeout) |
| 248 | defer cancel() |
| 249 | resp, err := retryRPC(c, ctx, func(node *protocol.Node) (*protocol.PublishTunnelResponse, error) { |
| 250 | ctx = rpc.WithNode(ctx, node) |
| 251 | return c.tunnelClient.PublishTunnel(ctx, &protocol.PublishTunnelRequest{ |
| 252 | Hostname: hostname, |
| 253 | Servers: connected, |
| 254 | }) |
| 255 | }) |
| 256 | if err != nil { |
| 257 | return nil, err |
| 258 | } |
| 259 | return resp.GetPublished(), nil |
| 260 | } |
| 261 | |
| 262 | func (c *Client) GetRegisteredHostnames(ctx context.Context) ([]string, error) { |
| 263 | ctx, cancel := context.WithTimeout(ctx, rpcTimeout) |
| 264 | defer cancel() |
| 265 | resp, err := retryRPC(c, ctx, func(node *protocol.Node) (*protocol.RegisteredHostnamesResponse, error) { |
| 266 | ctx = rpc.WithNode(ctx, node) |
| 267 | return c.tunnelClient.RegisteredHostnames(ctx, &protocol.RegisteredHostnamesRequest{}) |
| 268 | }) |
| 269 | if err != nil { |
| 270 | return nil, err |
| 271 | } |
| 272 | return resp.GetHostnames(), nil |
| 273 | } |
| 274 | |
| 275 | func (c *Client) RebuildTunnels(tunnels []Tunnel) error { |
| 276 | c.configMu.Lock() |
| 277 | defer c.configMu.Unlock() |
| 278 | |
| 279 | next := *c.Configuration |
| 280 | next.Tunnels = append([]Tunnel{}, tunnels...) |
| 281 | if err := next.validate(); err != nil { |
| 282 | return err |
| 283 | } |
| 284 | diff := diffTunnels(c.Configuration.Tunnels, next.Tunnels) |
| 285 | c.closeOutdatedProxies(diff...) |
| 286 | c.Configuration.Tunnels = next.Tunnels |
| 287 | c.Configuration.buildRouter(diff...) |
| 288 | if err := c.Configuration.writeFile(); err != nil { |
| 289 | return &ConfigSaveError{Err: err} |
| 290 | } |
| 291 | return nil |
| 292 | } |
| 293 | |
| 294 | func (c *Client) tunnelRemovalWrapper(tunnel Tunnel, fn func() error) error { |
| 295 | c.syncMu.Lock() |
| 296 | defer c.syncMu.Unlock() |
| 297 | if err := fn(); err != nil { |
| 298 | return err |
| 299 | } |
| 300 | |
| 301 | c.configMu.Lock() |
| 302 | defer c.configMu.Unlock() |
| 303 | |
| 304 | var index int = -1 |
| 305 | for i, t := range c.Configuration.Tunnels { |
| 306 | if t.Hostname == tunnel.Hostname { |
| 307 | index = i |
| 308 | break |
| 309 | } |
| 310 | } |
| 311 | if index == -1 { |
| 312 | return nil |
| 313 | } |
| 314 | |
| 315 | c.closeOutdatedProxies(tunnel) |
| 316 | |
| 317 | c.Configuration.Tunnels = append(c.Configuration.Tunnels[:index], c.Configuration.Tunnels[index+1:]...) |
| 318 | c.Configuration.validate() |
| 319 | c.Configuration.buildRouter(tunnel) |
| 320 | delete(c.publication, tunnel.Hostname) |
| 321 | |
| 322 | var saveErr error |
| 323 | if err := c.Configuration.writeFile(); err != nil { |
| 324 | saveErr = &ConfigSaveError{Err: err} |
| 325 | } |
| 326 | c.syncStateMu.RLock() |
| 327 | result := c.lastSync |
| 328 | c.syncStateMu.RUnlock() |
| 329 | remaining := make([]TunnelSyncResult, 0, len(c.Configuration.Tunnels)) |
| 330 | var pendingErrors []error |
| 331 | for _, tunnel := range c.Configuration.Tunnels { |
| 332 | outcome := TunnelSyncResult{Hostname: tunnel.Hostname, Target: tunnel.Target, Error: "publication pending"} |
| 333 | for _, previous := range result.Tunnels { |
| 334 | if previous.Hostname == tunnel.Hostname && previous.Target == tunnel.Target { |
| 335 | outcome = previous |
| 336 | break |
| 337 | } |
| 338 | } |
| 339 | remaining = append(remaining, outcome) |
| 340 | if outcome.Error != "" { |
| 341 | pendingErrors = append(pendingErrors, fmt.Errorf("%s: %s", tunnel.Hostname, outcome.Error)) |
| 342 | } |
| 343 | } |
| 344 | pendingErrors = append(pendingErrors, saveErr) |
| 345 | result.Tunnels, result.Applied, result.Saved, result.Error = remaining, true, saveErr == nil, "" |
| 346 | if err := errors.Join(pendingErrors...); err != nil { |
| 347 | result.Error = err.Error() |
| 348 | } |
| 349 | c.recordSyncResult(result) |
| 350 | return saveErr |
| 351 | } |
| 352 | |
| 353 | func (c *Client) UnpublishTunnel(ctx context.Context, tunnel Tunnel) error { |
| 354 | ctx, cancel := context.WithTimeout(ctx, rpcTimeout) |
| 355 | defer cancel() |
| 356 | err := c.tunnelRemovalWrapper(tunnel, func() error { |
| 357 | _, err := retryRPC(c, ctx, func(node *protocol.Node) (*protocol.UnpublishTunnelResponse, error) { |
| 358 | ctx = rpc.WithNode(ctx, node) |
| 359 | return c.tunnelClient.UnpublishTunnel(ctx, &protocol.UnpublishTunnelRequest{ |
| 360 | Hostname: tunnel.Hostname, |
| 361 | }) |
| 362 | }) |
| 363 | return err |
| 364 | }) |
| 365 | return err |
| 366 | } |
| 367 | |
| 368 | func (c *Client) ReleaseTunnel(ctx context.Context, tunnel Tunnel) error { |
| 369 | ctx, cancel := context.WithTimeout(ctx, rpcTimeout) |
| 370 | defer cancel() |
| 371 | return c.tunnelRemovalWrapper(tunnel, func() error { |
| 372 | _, err := retryRPC(c, ctx, func(node *protocol.Node) (*protocol.ReleaseTunnelResponse, error) { |
| 373 | ctx = rpc.WithNode(ctx, node) |
| 374 | return c.tunnelClient.ReleaseTunnel(ctx, &protocol.ReleaseTunnelRequest{ |
| 375 | Hostname: tunnel.Hostname, |
| 376 | }) |
| 377 | }) |
| 378 | return err |
| 379 | }) |
| 380 | } |
| 381 | |
| 382 | func (c *Client) closeOutdatedProxies(tunnels ...Tunnel) { |
| 383 | for _, t := range tunnels { |
| 384 | proxy, loaded := c.proxies.LoadAndDelete(t.Hostname) |
| 385 | if loaded { |
| 386 | c.Logger.Info("Shutting down proxy", zap.String("hostname", t.Hostname), zap.String("target", t.Target)) |
| 387 | proxy.acceptor.Close() |
| 388 | proxy.forwarder.Close() |
| 389 | } |
| 390 | } |
| 391 | } |
| 392 | |
| 393 | func diffTunnels(old, new []Tunnel) []Tunnel { |
| 394 | diff := make([]Tunnel, 0) |
| 395 | oldMap := map[string]Tunnel{} |
| 396 | newMap := map[string]Tunnel{} |
| 397 | for _, o := range old { |
| 398 | if o.Hostname == "" { |
| 399 | continue |
| 400 | } |
| 401 | oldMap[o.Hostname] = o |
| 402 | } |
| 403 | for _, n := range new { |
| 404 | if n.Hostname == "" { |
| 405 | continue |
| 406 | } |
| 407 | newMap[n.Hostname] = n |
| 408 | } |
| 409 | // if new != old |
| 410 | for hostname, tunnel := range newMap { |
| 411 | oldTunnel, ok := oldMap[hostname] |
| 412 | if ok && (oldTunnel.Target != tunnel.Target || |
| 413 | oldTunnel.Insecure != tunnel.Insecure || |
| 414 | oldTunnel.ProxyHeaderTimeout != tunnel.ProxyHeaderTimeout || |
| 415 | oldTunnel.ProxyHeaderHost != tunnel.ProxyHeaderHost || |
| 416 | oldTunnel.ProxyHeaderMode != tunnel.ProxyHeaderMode) { |
| 417 | diff = append(diff, oldTunnel) |
| 418 | } |
| 419 | } |
| 420 | // if old is gone |
| 421 | for hostname, tunnel := range oldMap { |
| 422 | if _, ok := newMap[hostname]; !ok { |
| 423 | diff = append(diff, tunnel) |
| 424 | } |
| 425 | } |
| 426 | return diff |
| 427 | } |