mirror of
https://github.com/fosrl/olm.git
synced 2026-10-01 18:29:08 +02:00
Support updating sites
This commit is contained in:
+47
-22
@@ -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 {
|
||||
|
||||
+2
-2
@@ -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()
|
||||
|
||||
|
||||
+121
-12
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
+24
-18
@@ -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()
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
+125
-9
@@ -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")
|
||||
|
||||
Reference in New Issue
Block a user