From c4474e6e41ebc8c5cfabe48484cdb408e896c3c4 Mon Sep 17 00:00:00 2001 From: Owen Date: Fri, 25 Sep 2026 11:51:10 -0400 Subject: [PATCH] Support updating sites --- api/api.go | 69 +++++++++++++++------- olm/connect.go | 4 +- olm/gateway.go | 133 +++++++++++++++++++++++++++++++++++++---- olm/olm.go | 42 +++++++------ olm/types.go | 7 +++ peers/gateway_test.go | 81 +++++++++++++++++++++++++ peers/manager.go | 134 +++++++++++++++++++++++++++++++++++++++--- 7 files changed, 407 insertions(+), 63 deletions(-) create mode 100644 peers/gateway_test.go diff --git a/api/api.go b/api/api.go index 3170797..5ed7bbb 100644 --- a/api/api.go +++ b/api/api.go @@ -13,25 +13,30 @@ import ( "github.com/fosrl/newt/network" ) -// ConnectionRequest defines the structure for an incoming connection request +// ConnectionRequest defines the structure for an incoming connection request. +// GatewaySiteResourceId is the numeric ID of the gateway-mode site resource +// GatewaySiteIds were selected from; it is required when GatewaySiteIds is +// non-empty, and is how olm later matches server-pushed gateway updates to the +// resource the user actually connected through. type ConnectionRequest struct { - ID string `json:"id"` - Secret string `json:"secret"` - Endpoint string `json:"endpoint"` - UserToken string `json:"userToken,omitempty"` - MTU int `json:"mtu,omitempty"` - DNS string `json:"dns,omitempty"` - DNSProxyIP string `json:"dnsProxyIP,omitempty"` - UpstreamDNS []string `json:"upstreamDNS,omitempty"` - InterfaceName string `json:"interfaceName,omitempty"` - Holepunch bool `json:"holepunch,omitempty"` - TlsClientCert string `json:"tlsClientCert,omitempty"` - PingInterval string `json:"pingInterval,omitempty"` - PingTimeout string `json:"pingTimeout,omitempty"` - OrgID string `json:"orgId,omitempty"` - MatchDomains []string `json:"matchDomains,omitempty"` - SubnetRouter bool `json:"subnetRouter,omitempty"` - GatewaySiteIds []int `json:"gatewaySiteIds,omitempty"` + ID string `json:"id"` + Secret string `json:"secret"` + Endpoint string `json:"endpoint"` + UserToken string `json:"userToken,omitempty"` + MTU int `json:"mtu,omitempty"` + DNS string `json:"dns,omitempty"` + DNSProxyIP string `json:"dnsProxyIP,omitempty"` + UpstreamDNS []string `json:"upstreamDNS,omitempty"` + InterfaceName string `json:"interfaceName,omitempty"` + Holepunch bool `json:"holepunch,omitempty"` + TlsClientCert string `json:"tlsClientCert,omitempty"` + PingInterval string `json:"pingInterval,omitempty"` + PingTimeout string `json:"pingTimeout,omitempty"` + OrgID string `json:"orgId,omitempty"` + MatchDomains []string `json:"matchDomains,omitempty"` + SubnetRouter bool `json:"subnetRouter,omitempty"` + GatewaySiteResourceId int `json:"gatewaySiteResourceId,omitempty"` + GatewaySiteIds []int `json:"gatewaySiteIds,omitempty"` } // SwitchOrgRequest defines the structure for switching organizations @@ -41,9 +46,12 @@ type SwitchOrgRequest struct { // GatewayRequest defines the structure for a "select gateway" request: the // set of site IDs that should act as the gateway (full-tunnel/default-route) -// for all tunnel traffic. Every ID must already be a tracked/connected peer. +// for all tunnel traffic, plus the numeric ID of the gateway site resource +// they were selected from (used to match later server-pushed updates). Every +// site ID must already be a tracked/connected peer. type GatewayRequest struct { - SiteIds []int `json:"siteIds"` + SiteResourceId int `json:"siteResourceId"` + SiteIds []int `json:"siteIds"` } // PowerModeRequest represents a request to change power mode @@ -94,6 +102,8 @@ type StatusResponse struct { ExitNodeStatus *ExitNodeStatus `json:"exitNode,omitempty"` GatewayActive bool `json:"gatewayActive,omitempty"` GatewaySiteIds []int `json:"gatewaySiteIds,omitempty"` + + GatewaySiteResourceId int `json:"gatewaySiteResourceId,omitempty"` // the gateway site resource the selection belongs to; 0 when inactive } type MetadataChangeRequest struct { @@ -137,6 +147,8 @@ type API struct { gatewayActive bool gatewaySiteIds []int + gatewaySiteResourceId int + version string agent string orgID string @@ -462,10 +474,11 @@ func (s *API) ClearExitNodeStatus() { // SetGatewayStatus records the current gateway (full-tunnel/default-route) // state for exposure via the status endpoint. -func (s *API) SetGatewayStatus(active bool, siteIds []int) { +func (s *API) SetGatewayStatus(active bool, siteResourceId int, siteIds []int) { s.statusMu.Lock() defer s.statusMu.Unlock() s.gatewayActive = active + s.gatewaySiteResourceId = siteResourceId s.gatewaySiteIds = siteIds } @@ -497,6 +510,10 @@ func (s *API) handleConnect(w http.ResponseWriter, r *http.Request) { http.Error(w, "Missing required fields: id, secret, and endpoint must be provided", http.StatusBadRequest) return } + if len(req.GatewaySiteIds) > 0 && req.GatewaySiteResourceId <= 0 { + http.Error(w, "Missing required field: gatewaySiteResourceId must be provided with gatewaySiteIds", http.StatusBadRequest) + return + } // Call the connect handler if set if s.onConnect != nil { @@ -536,6 +553,8 @@ func (s *API) handleStatus(w http.ResponseWriter, r *http.Request) { ExitNodeStatus: s.exitNodeStatus, GatewayActive: s.gatewayActive, GatewaySiteIds: s.gatewaySiteIds, + + GatewaySiteResourceId: s.gatewaySiteResourceId, } s.statusMu.RUnlock() @@ -706,6 +725,8 @@ func (s *API) GetStatus() StatusResponse { ExitNodeStatus: s.exitNodeStatus, GatewayActive: s.gatewayActive, GatewaySiteIds: s.gatewaySiteIds, + + GatewaySiteResourceId: s.gatewaySiteResourceId, } } @@ -803,12 +824,16 @@ func (s *API) handleSelectGateway(w http.ResponseWriter, r *http.Request) { return } + if req.SiteResourceId <= 0 { + http.Error(w, "Missing required field: siteResourceId must be provided", http.StatusBadRequest) + return + } if len(req.SiteIds) == 0 { http.Error(w, "Missing required field: siteIds must be provided", http.StatusBadRequest) return } - logger.Info("Received select-gateway request via API: siteIds=%v", req.SiteIds) + logger.Info("Received select-gateway request via API: siteResourceId=%d siteIds=%v", req.SiteResourceId, req.SiteIds) if s.onSelectGateway != nil { if err := s.onSelectGateway(req); err != nil { diff --git a/olm/connect.go b/olm/connect.go index aa70b35..9a2c4d7 100644 --- a/olm/connect.go +++ b/olm/connect.go @@ -312,7 +312,7 @@ func (o *Olm) handleConnect(msg websocket.WSMessage) { o.registered = true if len(o.tunnelConfig.GatewaySiteIds) > 0 { - o.applyPendingGatewayConfig(o.tunnelConfig.GatewaySiteIds) + o.applyPendingGatewayConfig(o.tunnelConfig.GatewaySiteResourceId, o.tunnelConfig.GatewaySiteIds) } // Start ping monitor now that we are registered and connected @@ -393,7 +393,7 @@ func (o *Olm) handleTerminate(msg websocket.WSMessage) { o.apiServer.SetConnectionStatus(false) o.apiServer.SetRegistered(false) o.apiServer.ClearPeerStatuses() - o.apiServer.SetGatewayStatus(false, nil) + o.apiServer.SetGatewayStatus(false, 0, nil) network.ClearNetworkSettings() diff --git a/olm/gateway.go b/olm/gateway.go index 3c89ced..56b18a8 100644 --- a/olm/gateway.go +++ b/olm/gateway.go @@ -1,21 +1,41 @@ package olm import ( + "encoding/json" "fmt" "net/url" "github.com/fosrl/newt/logger" + "github.com/fosrl/olm/peers" + "github.com/fosrl/olm/websocket" ) +// GatewaySitesUpdateData is the payload of the server's +// "olm/wg/gateway/sites/update" message: sites that were added to / removed +// from the gateway site resource SiteResourceId. +type GatewaySitesUpdateData struct { + SiteResourceId int `json:"siteResourceId"` + AddedSiteIds []int `json:"addedSiteIds"` + RemovedSiteIds []int `json:"removedSiteIds"` +} + +// GatewayDisableData is the payload of the server's "olm/wg/gateway/disable" +// message: the gateway site resource SiteResourceId can no longer be used as +// the gateway (deleted, disabled, or this client lost access to it). +type GatewayDisableData struct { + SiteResourceId int `json:"siteResourceId"` +} + // SelectGateway designates siteIds as the gateway (full-tunnel/default-route) -// candidate set. Requires the tunnel to already be registered/connected; -// rejects otherwise. Every site ID must already be a tracked peer, or the -// call is rejected outright by the peer manager. -func (o *Olm) SelectGateway(siteIds []int) error { +// candidate set, selected from the gateway site resource siteResourceId. +// Requires the tunnel to already be registered/connected; rejects otherwise. +// Every site ID must already be a tracked peer, or the call is rejected +// outright by the peer manager. +func (o *Olm) SelectGateway(siteResourceId int, siteIds []int) error { if !o.registered { return fmt.Errorf("cannot select gateway: not registered/connected") } - return o.applySelectGateway(siteIds) + return o.applySelectGateway(siteResourceId, siteIds) } // DisableGateway fully clears gateway state. Requires the tunnel to already @@ -31,7 +51,7 @@ func (o *Olm) DisableGateway() error { if err := pm.ClearGateway(); err != nil { return err } - o.apiServer.SetGatewayStatus(false, nil) + o.apiServer.SetGatewayStatus(false, 0, nil) return nil } @@ -39,24 +59,113 @@ func (o *Olm) DisableGateway() error { // to the peer manager. Shared by SelectGateway (API-invoked, already // registered-checked by the caller) and applyPendingGatewayConfig // (StartTunnel-time, called after registration completes). -func (o *Olm) applySelectGateway(siteIds []int) error { +func (o *Olm) applySelectGateway(siteResourceId int, siteIds []int) error { pm := o.getPeerManager() if pm == nil { return fmt.Errorf("cannot select gateway: tunnel not running") } - if err := pm.SetGateway(siteIds, extractControlEndpointHost(o.tunnelConfig.Endpoint)); err != nil { + if err := pm.SetGateway(siteResourceId, siteIds, extractControlEndpointHost(o.tunnelConfig.Endpoint)); err != nil { return err } - o.apiServer.SetGatewayStatus(true, siteIds) + o.apiServer.SetGatewayStatus(true, siteResourceId, siteIds) return nil } -// applyPendingGatewayConfig applies TunnelConfig.GatewaySiteIds once sites +// refreshGatewayStatus re-publishes the peer manager's current gateway state +// to the status endpoint after a server-pushed change. +func (o *Olm) refreshGatewayStatus(pm *peers.PeerManager) { + active, siteResourceId, siteIds := pm.GetGatewayState() + if !active { + o.apiServer.SetGatewayStatus(false, 0, nil) + return + } + o.apiServer.SetGatewayStatus(true, siteResourceId, siteIds) +} + +// handleGatewaySitesUpdate handles the server's "olm/wg/gateway/sites/update" +// message, sent when sites are added to or removed from a gateway site +// resource (via the API or a blueprint). The message is ignored unless it is +// for the same site resource the current gateway was selected from - a site +// added to some other gateway resource must not join our candidate set. +func (o *Olm) handleGatewaySitesUpdate(msg websocket.WSMessage) { + logger.Debug("Received gateway sites update message: %v", msg.Data) + + if !o.tunnelRunning || !o.registered { + logger.Debug("Tunnel not running/registered, ignoring gateway sites update message") + return + } + + jsonData, err := json.Marshal(msg.Data) + if err != nil { + logger.Error("Error marshaling data: %v", err) + return + } + + var data GatewaySitesUpdateData + if err := json.Unmarshal(jsonData, &data); err != nil { + logger.Error("Error unmarshaling gateway sites update data: %v", err) + return + } + + pm := o.getPeerManager() + if pm == nil { + logger.Debug("Ignoring gateway sites update message: peerManager is nil (shutdown in progress)") + return + } + + matched, _, _ := pm.UpdateGatewaySites(data.SiteResourceId, data.AddedSiteIds, data.RemovedSiteIds) + if !matched { + logger.Debug("Ignoring gateway sites update for site resource %d: not the active gateway resource", data.SiteResourceId) + return + } + o.refreshGatewayStatus(pm) +} + +// handleGatewayDisable handles the server's "olm/wg/gateway/disable" message, +// sent when the gateway site resource is deleted, disabled, changed to a +// different mode, or this client loses access to it. Ignored unless it is for +// the site resource the current gateway was selected from. +func (o *Olm) handleGatewayDisable(msg websocket.WSMessage) { + logger.Debug("Received gateway disable message: %v", msg.Data) + + if !o.tunnelRunning || !o.registered { + logger.Debug("Tunnel not running/registered, ignoring gateway disable message") + return + } + + jsonData, err := json.Marshal(msg.Data) + if err != nil { + logger.Error("Error marshaling data: %v", err) + return + } + + var data GatewayDisableData + if err := json.Unmarshal(jsonData, &data); err != nil { + logger.Error("Error unmarshaling gateway disable data: %v", err) + return + } + + pm := o.getPeerManager() + if pm == nil { + logger.Debug("Ignoring gateway disable message: peerManager is nil (shutdown in progress)") + return + } + + if !pm.ClearGatewayForResource(data.SiteResourceId) { + logger.Debug("Ignoring gateway disable for site resource %d: not the active gateway resource", data.SiteResourceId) + return + } + logger.Info("Gateway disabled: site resource %d is no longer available", data.SiteResourceId) + o.refreshGatewayStatus(pm) +} + +// applyPendingGatewayConfig applies TunnelConfig.GatewaySiteIds (selected from +// the gateway site resource siteResourceId) once sites // are tracked peers (called from handleConnect after o.registered is set // true). IDs that never showed up as tracked peers are logged and dropped; // if none show up at all, gateway is not established and this is logged // clearly, without failing tunnel startup. -func (o *Olm) applyPendingGatewayConfig(requestedSiteIds []int) { +func (o *Olm) applyPendingGatewayConfig(siteResourceId int, requestedSiteIds []int) { pm := o.getPeerManager() if pm == nil { return @@ -76,7 +185,7 @@ func (o *Olm) applyPendingGatewayConfig(requestedSiteIds []int) { return } - if err := o.applySelectGateway(tracked); err != nil { + if err := o.applySelectGateway(siteResourceId, tracked); err != nil { logger.Error("Failed to establish gateway from StartTunnel config: %v", err) } } diff --git a/olm/olm.go b/olm/olm.go index 6cd6b6f..60fd1f1 100644 --- a/olm/olm.go +++ b/olm/olm.go @@ -289,20 +289,21 @@ func (o *Olm) registerAPICallbacks() { logger.Info("Received connection request via HTTP: id=%s, endpoint=%s", req.ID, req.Endpoint) tunnelConfig := TunnelConfig{ - Endpoint: req.Endpoint, - ID: req.ID, - Secret: req.Secret, - UserToken: req.UserToken, - MTU: req.MTU, - DNS: req.DNS, - UpstreamDNS: req.UpstreamDNS, - InterfaceName: req.InterfaceName, - Holepunch: req.Holepunch, - TlsClientCert: req.TlsClientCert, - OrgID: req.OrgID, - MatchDomains: req.MatchDomains, - SubnetRouter: req.SubnetRouter, - GatewaySiteIds: req.GatewaySiteIds, + Endpoint: req.Endpoint, + ID: req.ID, + Secret: req.Secret, + UserToken: req.UserToken, + MTU: req.MTU, + DNS: req.DNS, + UpstreamDNS: req.UpstreamDNS, + InterfaceName: req.InterfaceName, + Holepunch: req.Holepunch, + TlsClientCert: req.TlsClientCert, + OrgID: req.OrgID, + MatchDomains: req.MatchDomains, + SubnetRouter: req.SubnetRouter, + GatewaySiteIds: req.GatewaySiteIds, + GatewaySiteResourceId: req.GatewaySiteResourceId, } var err error @@ -409,8 +410,8 @@ func (o *Olm) registerAPICallbacks() { }, // onSelectGateway func(req api.GatewayRequest) error { - logger.Info("Received select-gateway request via API: siteIds=%v", req.SiteIds) - return o.SelectGateway(req.SiteIds) + logger.Info("Received select-gateway request via API: siteResourceId=%d siteIds=%v", req.SiteResourceId, req.SiteIds) + return o.SelectGateway(req.SiteResourceId, req.SiteIds) }, // onDisableGateway func() error { @@ -636,6 +637,11 @@ func (o *Olm) StartTunnel(config TunnelConfig) { o.websocket.RegisterHandler("olm/wg/exitnode/disconnect", o.handleExitNodeDisconnect) o.websocket.RegisterHandler("olm/wg/exitnode/data/update", o.handleExitNodeUpdateData) + // Handlers for the server to push changes to the gateway site resource the + // client selected (sites added/removed, or the resource going away) + o.websocket.RegisterHandler("olm/wg/gateway/sites/update", o.handleGatewaySitesUpdate) + o.websocket.RegisterHandler("olm/wg/gateway/disable", o.handleGatewayDisable) + // Handler for the server to push a live DNS config override (upstream DNS, // tunnel DNS, override DNS, match domains) after registration, mirroring the // DNSConfig field sent on the initial "olm/wg/connect" message. @@ -842,7 +848,7 @@ func (o *Olm) StartTunnel(config TunnelConfig) { o.apiServer.SetRegistered(false) o.apiServer.ClearOlmError() o.apiServer.ClearPeerStatuses() - o.apiServer.SetGatewayStatus(false, nil) + o.apiServer.SetGatewayStatus(false, 0, nil) network.ClearNetworkSettings() o.Close() @@ -1061,7 +1067,7 @@ func (o *Olm) StopTunnel() error { o.apiServer.SetConnectionStatus(false) o.apiServer.SetRegistered(false) o.apiServer.ClearOlmError() - o.apiServer.SetGatewayStatus(false, nil) + o.apiServer.SetGatewayStatus(false, 0, nil) network.ClearNetworkSettings() o.apiServer.ClearPeerStatuses() diff --git a/olm/types.go b/olm/types.go index 8b19cd9..06907b8 100644 --- a/olm/types.go +++ b/olm/types.go @@ -187,4 +187,11 @@ type TunnelConfig struct { // starts, for callers that want a gateway already established rather // than issuing a separate SelectGateway API call after connecting. GatewaySiteIds []int + + // GatewaySiteResourceId is the numeric ID (not the niceId, which can be + // renamed) of the gateway-mode site resource GatewaySiteIds were selected + // from. Required when GatewaySiteIds is non-empty; it is what lets olm + // apply server-pushed gateway updates only for the resource the user + // actually connected through. + GatewaySiteResourceId int } diff --git a/peers/gateway_test.go b/peers/gateway_test.go new file mode 100644 index 0000000..1a7a4ea --- /dev/null +++ b/peers/gateway_test.go @@ -0,0 +1,81 @@ +package peers + +import ( + "reflect" + "testing" + "time" +) + +// newGatewayTestManager returns a PeerManager with gateway mode marked active +// for siteResourceId with the given candidate sites, without touching the OS +// routing table or a WireGuard device. Only usable for paths that don't need +// either (sites that aren't tracked peers, no owner of the gateway CIDR). +func newGatewayTestManager(siteResourceId int, siteIds ...int) *PeerManager { + pm := &PeerManager{ + peers: make(map[int]SiteConfig), + allowedIPOwners: make(map[string]int), + allowedIPClaims: make(map[string]map[int]bool), + lastOwnerChange: make(map[string]time.Time), + gatewaySiteIds: make(map[int]bool), + gatewayExcludedIPs: make(map[string]int), + gatewayExtraEndpoints: make(map[string]bool), + gatewayActive: true, + gatewaySiteResourceId: siteResourceId, + } + for _, id := range siteIds { + pm.gatewaySiteIds[id] = true + } + return pm +} + +func TestUpdateGatewaySitesIgnoresOtherResource(t *testing.T) { + pm := newGatewayTestManager(5, 1) + + matched, active, siteIds := pm.UpdateGatewaySites(6, []int{2}, nil) + if matched { + t.Fatalf("update for a different resource must not match") + } + if !active || !reflect.DeepEqual(siteIds, []int{1}) { + t.Fatalf("state must be unchanged, got active=%v siteIds=%v", active, siteIds) + } +} + +func TestUpdateGatewaySitesInactive(t *testing.T) { + pm := newGatewayTestManager(5, 1) + pm.gatewayActive = false + + if matched, _, _ := pm.UpdateGatewaySites(5, []int{2}, nil); matched { + t.Fatalf("update must not match when gateway mode is inactive") + } +} + +func TestUpdateGatewaySitesAddRemove(t *testing.T) { + pm := newGatewayTestManager(5, 1) + + matched, active, siteIds := pm.UpdateGatewaySites(5, []int{2, 3}, nil) + if !matched || !active || !reflect.DeepEqual(siteIds, []int{1, 2, 3}) { + t.Fatalf("add: matched=%v active=%v siteIds=%v", matched, active, siteIds) + } + + matched, active, siteIds = pm.UpdateGatewaySites(5, nil, []int{2}) + if !matched || !active || !reflect.DeepEqual(siteIds, []int{1, 3}) { + t.Fatalf("remove: matched=%v active=%v siteIds=%v", matched, active, siteIds) + } + + // removed wins over added if a message lists an ID in both + _, _, siteIds = pm.UpdateGatewaySites(5, []int{4}, []int{4}) + if !reflect.DeepEqual(siteIds, []int{1, 3}) { + t.Fatalf("add+remove of the same ID must leave it out, got %v", siteIds) + } +} + +func TestClearGatewayForResourceIgnoresOtherResource(t *testing.T) { + pm := newGatewayTestManager(5, 1) + + if pm.ClearGatewayForResource(6) { + t.Fatalf("clearing for a different resource must not match") + } + if active, id, _ := pm.GetGatewayState(); !active || id != 5 { + t.Fatalf("state must be unchanged, got active=%v resource=%d", active, id) + } +} diff --git a/peers/manager.go b/peers/manager.go index 6e45f5d..c44cbde 100644 --- a/peers/manager.go +++ b/peers/manager.go @@ -3,6 +3,7 @@ package peers import ( "fmt" "net" + "sort" "strconv" "strings" "sync" @@ -76,11 +77,17 @@ type PeerManager struct { // it any more. gatewayControlIP is the control-plane (Pangolin server) // endpoint, resolved once at activation and excluded for the lifetime of // the gateway so the management connection doesn't depend on the current - // gateway-owner site's own uplink. - gatewayActive bool - gatewaySiteIds map[int]bool - gatewayExcludedIPs map[string]int - gatewayControlIP string + // gateway-owner site's own uplink. gatewaySiteResourceId is the numeric + // ID (not the niceId, which can be renamed) of the gateway-mode site + // resource the candidate set was selected from, so server-pushed + // add/remove/disable updates (see UpdateGatewaySites) are only applied when + // they concern that resource and not some other gateway resource that + // happens to share a site. + gatewayActive bool + gatewaySiteResourceId int + gatewaySiteIds map[int]bool + gatewayExcludedIPs map[string]int + gatewayControlIP string // gatewayExtraEndpoints tracks "host:port" (or already-resolved "ip:port") // endpoints registered by callers outside the normal site-peer lifecycle - @@ -392,8 +399,10 @@ func (pm *PeerManager) releaseGatewayClaimLocked(siteId int) { } // SetGateway designates siteIds as the gateway (full-tunnel/default-route) -// candidate set. Every siteId must already be a tracked peer, or the call is -// rejected outright (no partial application). On first activation this +// candidate set, selected from the gateway site resource siteResourceId (the +// server only tells us about changes to that one resource - see +// UpdateGatewaySites). Every siteId must already be a tracked peer, or the +// call is rejected outright (no partial application). On first activation this // installs the OS-level gateway route plus every bypass route needed so the // tunnel's own traffic (control-plane endpoint, every tracked peer's active // endpoint) isn't captured by it; subsequent calls only change which sites @@ -401,10 +410,13 @@ func (pm *PeerManager) releaseGatewayClaimLocked(siteId int) { // claim/optimizer machinery - exactly like remote subnets. controlEndpointHost // is the Pangolin server host olm is registered against (bare host, port // optional); always excluded regardless of which sites are selected. -func (pm *PeerManager) SetGateway(siteIds []int, controlEndpointHost string) error { +func (pm *PeerManager) SetGateway(siteResourceId int, siteIds []int, controlEndpointHost string) error { pm.mu.Lock() defer pm.mu.Unlock() + if siteResourceId <= 0 { + return fmt.Errorf("a valid gateway site resource ID must be provided") + } if len(siteIds) == 0 { return fmt.Errorf("at least one site ID must be provided") } @@ -441,11 +453,114 @@ func (pm *PeerManager) SetGateway(siteIds []int, controlEndpointHost string) err } } pm.gatewaySiteIds = newSet + pm.gatewaySiteResourceId = siteResourceId - logger.Info("Gateway set to sites %v", siteIds) + logger.Info("Gateway set to sites %v (site resource %d)", siteIds, siteResourceId) return nil } +// GetGatewayState returns whether gateway mode is active, the site resource ID +// it was selected from, and the current candidate site IDs (sorted). +func (pm *PeerManager) GetGatewayState() (active bool, siteResourceId int, siteIds []int) { + pm.mu.RLock() + defer pm.mu.RUnlock() + return pm.gatewayActive, pm.gatewaySiteResourceId, pm.gatewaySiteIdsSortedLocked() +} + +// gatewaySiteIdsSortedLocked returns the gateway candidate set as a sorted +// slice, for stable status output. Must be called with pm.mu held. +func (pm *PeerManager) gatewaySiteIdsSortedLocked() []int { + ids := make([]int, 0, len(pm.gatewaySiteIds)) + for id := range pm.gatewaySiteIds { + ids = append(ids, id) + } + sort.Ints(ids) + return ids +} + +// UpdateGatewaySites applies a server-pushed change to the gateway candidate +// set: addedSiteIds/removedSiteIds are the sites that were added to / removed +// from the gateway site resource siteResourceId. It is a no-op (matched=false) +// unless gateway mode is active AND was selected from that exact resource, so +// a site added to some other gateway resource is never pulled into the +// candidate set. If the update leaves the candidate set empty, gateway mode is +// cleared entirely (an installed default route with no owning peer would just +// blackhole traffic). Sites that aren't tracked peers yet are recorded as +// intent only - AddPeer claims the gateway CIDR for them once their peer +// arrives (the server sends the peer add and this update independently, so +// either order is possible). Returns the resulting gateway state. +func (pm *PeerManager) UpdateGatewaySites(siteResourceId int, addedSiteIds, removedSiteIds []int) (matched bool, active bool, siteIds []int) { + pm.mu.Lock() + defer pm.mu.Unlock() + + if !pm.gatewayActive || pm.gatewaySiteResourceId != siteResourceId { + return false, pm.gatewayActive, pm.gatewaySiteIdsSortedLocked() + } + + removed := make(map[int]bool, len(removedSiteIds)) + for _, id := range removedSiteIds { + removed[id] = true + } + + // Work out the resulting set first so we can tell up front if it would be + // empty, and so removed wins over added if a message lists an ID in both. + newSet := make(map[int]bool, len(pm.gatewaySiteIds)+len(addedSiteIds)) + for id := range pm.gatewaySiteIds { + if !removed[id] { + newSet[id] = true + } + } + for _, id := range addedSiteIds { + if !removed[id] { + newSet[id] = true + } + } + + if len(newSet) == 0 { + logger.Info("Gateway: all sites removed from site resource %d, clearing gateway", siteResourceId) + pm.clearGatewayLocked() + return true, false, nil + } + + // Claim added sites BEFORE releasing removed ones, so a swap doesn't leave + // a window with no owner of the gateway CIDR (same ordering rationale as + // handleWgPeerUpdateData). + for id := range newSet { + if pm.gatewaySiteIds[id] { + continue + } + pm.gatewaySiteIds[id] = true + if _, tracked := pm.peers[id]; tracked { + pm.claimGatewayClaimLocked(id) + } + } + for id := range removed { + if !pm.gatewaySiteIds[id] { + continue + } + delete(pm.gatewaySiteIds, id) + pm.releaseGatewayClaimLocked(id) + } + + logger.Info("Gateway sites for site resource %d are now %v", siteResourceId, pm.gatewaySiteIdsSortedLocked()) + return true, true, pm.gatewaySiteIdsSortedLocked() +} + +// ClearGatewayForResource fully clears gateway state, but only if gateway mode +// was selected from the site resource siteResourceId (e.g. that resource was +// deleted, disabled, or this client lost access to it). Returns whether it +// matched and cleared. +func (pm *PeerManager) ClearGatewayForResource(siteResourceId int) bool { + pm.mu.Lock() + defer pm.mu.Unlock() + + if !pm.gatewayActive || pm.gatewaySiteResourceId != siteResourceId { + return false + } + pm.clearGatewayLocked() + return true +} + // clearGatewayLocked is ClearGateway's body, split out so Close() (which // already holds pm.mu) can reuse it without re-locking. Must be called with // pm.mu held. @@ -457,6 +572,7 @@ func (pm *PeerManager) clearGatewayLocked() { pm.releaseGatewayClaimLocked(id) } pm.gatewaySiteIds = make(map[int]bool) + pm.gatewaySiteResourceId = 0 pm.deactivateGatewayLocked() pm.gatewayActive = false logger.Info("Gateway cleared")