mirror of
https://github.com/fosrl/olm.git
synced 2026-10-01 18:29:08 +02:00
@@ -34,6 +34,14 @@ body:
|
||||
validations:
|
||||
required: true
|
||||
|
||||
- type: textarea
|
||||
attributes:
|
||||
label: AI Disclosure
|
||||
description: |
|
||||
If you used AI to help write this issue, please disclose it here. This is important for transparency and helps maintain the integrity of the issue tracking process.
|
||||
validations:
|
||||
required: true
|
||||
|
||||
- type: textarea
|
||||
attributes:
|
||||
label: Expected Behavior
|
||||
|
||||
@@ -4,6 +4,10 @@ perpetual license to use, modify, and redistribute these contributions under any
|
||||
choose, including both the AGPLv3 and the Fossorial Commercial license terms. I
|
||||
represent that I have the right to grant this license for all contributed content.
|
||||
|
||||
## AI Disclosure
|
||||
|
||||
> Please disclose how AI was used in this pull request. The use of AI does not preclude this from being merged but is an important factor in how we review your request.
|
||||
|
||||
## Description
|
||||
|
||||
|
||||
|
||||
@@ -20,7 +20,7 @@ jobs:
|
||||
- name: Set up Go
|
||||
uses: actions/setup-go@4dc6199c7b1a012772edbd06daecab0f50c9053c # v6.1.0
|
||||
with:
|
||||
go-version: 1.25
|
||||
go-version: 1.26
|
||||
|
||||
- name: Build binaries
|
||||
run: make go-build-release
|
||||
|
||||
+1
-1
@@ -1 +1 @@
|
||||
1.25
|
||||
1.26
|
||||
@@ -1,5 +1,8 @@
|
||||
# Olm
|
||||
|
||||
> [!NOTE]
|
||||
> Olm is being phased out in favor of the [Pangolin CLI](https://github.com/fosrl/cli) and is only meant for advanced use cases.
|
||||
|
||||
Olm is the cross-platform internal networking library built into every Pangolin client. It does the heavy-lifting of connecting the client to Pangolin sites.
|
||||
|
||||
Don't use Olm as a standalone machine client. Instead, use the [Pangolin CLI](https://github.com/fosrl/cli).
|
||||
|
||||
+143
-16
@@ -13,23 +13,31 @@ 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"`
|
||||
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"`
|
||||
DisableRoutesAndAliasesOnExitNode bool `json:"disableRoutesAndAliasesOnExitNode,omitempty"`
|
||||
GatewaySiteResourceId int `json:"gatewaySiteResourceId,omitempty"`
|
||||
GatewaySiteIds []int `json:"gatewaySiteIds,omitempty"`
|
||||
}
|
||||
|
||||
// SwitchOrgRequest defines the structure for switching organizations
|
||||
@@ -37,6 +45,16 @@ type SwitchOrgRequest struct {
|
||||
OrgID string `json:"org_id"`
|
||||
}
|
||||
|
||||
// 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, 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 {
|
||||
SiteResourceId int `json:"siteResourceId"`
|
||||
SiteIds []int `json:"siteIds"`
|
||||
}
|
||||
|
||||
// PowerModeRequest represents a request to change power mode
|
||||
type PowerModeRequest struct {
|
||||
Mode string `json:"mode"` // "normal" or "low"
|
||||
@@ -83,6 +101,10 @@ type StatusResponse struct {
|
||||
PeerStatuses map[int]*PeerStatus `json:"peers,omitempty"`
|
||||
NetworkSettings network.NetworkSettings `json:"networkSettings,omitempty"`
|
||||
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 {
|
||||
@@ -112,6 +134,8 @@ type API struct {
|
||||
onRebind func() error
|
||||
onPowerMode func(PowerModeRequest) error
|
||||
onJITConnect func(JITConnectionRequest) error
|
||||
onSelectGateway func(GatewayRequest) error
|
||||
onDisableGateway func() error
|
||||
|
||||
statusMu sync.RWMutex
|
||||
peerStatuses map[int]*PeerStatus
|
||||
@@ -121,6 +145,10 @@ type API struct {
|
||||
isRegistered bool
|
||||
isTerminated bool
|
||||
olmError *OlmError
|
||||
gatewayActive bool
|
||||
gatewaySiteIds []int
|
||||
|
||||
gatewaySiteResourceId int
|
||||
|
||||
version string
|
||||
agent string
|
||||
@@ -165,6 +193,8 @@ func (s *API) SetHandlers(
|
||||
onRebind func() error,
|
||||
onPowerMode func(PowerModeRequest) error,
|
||||
onJITConnect func(JITConnectionRequest) error,
|
||||
onSelectGateway func(GatewayRequest) error,
|
||||
onDisableGateway func() error,
|
||||
) {
|
||||
s.onConnect = onConnect
|
||||
s.onSwitchOrg = onSwitchOrg
|
||||
@@ -174,6 +204,8 @@ func (s *API) SetHandlers(
|
||||
s.onRebind = onRebind
|
||||
s.onPowerMode = onPowerMode
|
||||
s.onJITConnect = onJITConnect
|
||||
s.onSelectGateway = onSelectGateway
|
||||
s.onDisableGateway = onDisableGateway
|
||||
}
|
||||
|
||||
// Start starts the HTTP server
|
||||
@@ -193,6 +225,8 @@ func (s *API) Start() error {
|
||||
mux.HandleFunc("/rebind", s.handleRebind)
|
||||
mux.HandleFunc("/power-mode", s.handlePowerMode)
|
||||
mux.HandleFunc("/jit-connect", s.handleJITConnect)
|
||||
mux.HandleFunc("/gateway/select", s.handleSelectGateway)
|
||||
mux.HandleFunc("/gateway/disable", s.handleDisableGateway)
|
||||
|
||||
s.server = &http.Server{
|
||||
Handler: mux,
|
||||
@@ -439,6 +473,16 @@ func (s *API) ClearExitNodeStatus() {
|
||||
s.exitNodeStatus = nil
|
||||
}
|
||||
|
||||
// SetGatewayStatus records the current gateway (full-tunnel/default-route)
|
||||
// state for exposure via the status endpoint.
|
||||
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
|
||||
}
|
||||
|
||||
// handleConnect handles the /connect endpoint
|
||||
func (s *API) handleConnect(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodPost {
|
||||
@@ -467,6 +511,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 {
|
||||
@@ -504,6 +552,10 @@ func (s *API) handleStatus(w http.ResponseWriter, r *http.Request) {
|
||||
PeerStatuses: s.peerStatuses,
|
||||
NetworkSettings: network.GetSettings(),
|
||||
ExitNodeStatus: s.exitNodeStatus,
|
||||
GatewayActive: s.gatewayActive,
|
||||
GatewaySiteIds: s.gatewaySiteIds,
|
||||
|
||||
GatewaySiteResourceId: s.gatewaySiteResourceId,
|
||||
}
|
||||
|
||||
s.statusMu.RUnlock()
|
||||
@@ -672,6 +724,10 @@ func (s *API) GetStatus() StatusResponse {
|
||||
PeerStatuses: s.peerStatuses,
|
||||
NetworkSettings: network.GetSettings(),
|
||||
ExitNodeStatus: s.exitNodeStatus,
|
||||
GatewayActive: s.gatewayActive,
|
||||
GatewaySiteIds: s.gatewaySiteIds,
|
||||
|
||||
GatewaySiteResourceId: s.gatewaySiteResourceId,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -753,6 +809,77 @@ func (s *API) handleJITConnect(w http.ResponseWriter, r *http.Request) {
|
||||
})
|
||||
}
|
||||
|
||||
// handleSelectGateway handles the /gateway/select endpoint. It designates
|
||||
// the given site IDs as the gateway (full-tunnel/default-route) candidate
|
||||
// set - the onSelectGateway handler is responsible for rejecting the call if
|
||||
// the tunnel isn't registered/connected or any site ID isn't a tracked peer.
|
||||
func (s *API) handleSelectGateway(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodPost {
|
||||
http.Error(w, "Method not allowed", http.StatusMethodNotAllowed)
|
||||
return
|
||||
}
|
||||
|
||||
var req GatewayRequest
|
||||
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
|
||||
http.Error(w, fmt.Sprintf("Invalid request body: %v", err), http.StatusBadRequest)
|
||||
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: siteResourceId=%d siteIds=%v", req.SiteResourceId, req.SiteIds)
|
||||
|
||||
if s.onSelectGateway != nil {
|
||||
if err := s.onSelectGateway(req); err != nil {
|
||||
http.Error(w, fmt.Sprintf("Select gateway failed: %v", err), http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
} else {
|
||||
http.Error(w, "Select gateway handler not configured", http.StatusNotImplemented)
|
||||
return
|
||||
}
|
||||
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
w.WriteHeader(http.StatusAccepted)
|
||||
_ = json.NewEncoder(w).Encode(map[string]string{
|
||||
"status": "gateway selection accepted",
|
||||
})
|
||||
}
|
||||
|
||||
// handleDisableGateway handles the /gateway/disable endpoint. It fully clears
|
||||
// gateway state; no request body is needed.
|
||||
func (s *API) handleDisableGateway(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodPost {
|
||||
http.Error(w, "Method not allowed", http.StatusMethodNotAllowed)
|
||||
return
|
||||
}
|
||||
|
||||
logger.Info("Received disable-gateway request via API")
|
||||
|
||||
if s.onDisableGateway != nil {
|
||||
if err := s.onDisableGateway(); err != nil {
|
||||
http.Error(w, fmt.Sprintf("Disable gateway failed: %v", err), http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
} else {
|
||||
http.Error(w, "Disable gateway handler not configured", http.StatusNotImplemented)
|
||||
return
|
||||
}
|
||||
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
w.WriteHeader(http.StatusOK)
|
||||
_ = json.NewEncoder(w).Encode(map[string]string{
|
||||
"status": "gateway disabled",
|
||||
})
|
||||
}
|
||||
|
||||
// handlePowerMode handles the /power-mode endpoint
|
||||
// This allows changing the power mode between "normal" and "low"
|
||||
func (s *API) handlePowerMode(w http.ResponseWriter, r *http.Request) {
|
||||
|
||||
@@ -47,12 +47,14 @@ type OlmConfig struct {
|
||||
PingTimeout string `json:"pingTimeout"`
|
||||
|
||||
// Advanced
|
||||
DisableHolepunch bool `json:"disableHolepunch"`
|
||||
TlsClientCert string `json:"tlsClientCert"`
|
||||
OverrideDNS bool `json:"overrideDNS"`
|
||||
TunnelDNS bool `json:"tunnelDNS"`
|
||||
DisableRelay bool `json:"disableRelay"`
|
||||
PreferLocalRoutes bool `json:"preferLocalRoutes"`
|
||||
DisableHolepunch bool `json:"disableHolepunch"`
|
||||
TlsClientCert string `json:"tlsClientCert"`
|
||||
OverrideDNS bool `json:"overrideDNS"`
|
||||
TunnelDNS bool `json:"tunnelDNS"`
|
||||
DisableRelay bool `json:"disableRelay"`
|
||||
PreferLocalRoutes bool `json:"preferLocalRoutes"`
|
||||
SubnetRouter bool `json:"subnetRouter"`
|
||||
DisableRoutesAndAliasesOnExitNode bool `json:"disableRoutesAndAliasesOnExitNode"`
|
||||
// DoNotCreateNewClient bool `json:"doNotCreateNewClient"`
|
||||
|
||||
// Parsed values (not in JSON)
|
||||
@@ -120,6 +122,8 @@ func DefaultConfig() *OlmConfig {
|
||||
config.sources["tunnelDNS"] = string(SourceDefault)
|
||||
config.sources["disableRelay"] = string(SourceDefault)
|
||||
config.sources["preferLocalRoutes"] = string(SourceDefault)
|
||||
config.sources["subnetRouter"] = string(SourceDefault)
|
||||
config.sources["disableRoutesAndAliasesOnExitNode"] = string(SourceDefault)
|
||||
// config.sources["doNotCreateNewClient"] = string(SourceDefault)
|
||||
|
||||
return config
|
||||
@@ -291,6 +295,14 @@ func loadConfigFromEnv(config *OlmConfig) {
|
||||
config.TunnelDNS = true
|
||||
config.sources["tunnelDNS"] = string(SourceEnv)
|
||||
}
|
||||
if val := os.Getenv("SUBNET_ROUTER"); val == "true" {
|
||||
config.SubnetRouter = true
|
||||
config.sources["subnetRouter"] = string(SourceEnv)
|
||||
}
|
||||
if val := os.Getenv("DISABLE_ROUTES_AND_ALIASES"); val == "true" {
|
||||
config.DisableRoutesAndAliasesOnExitNode = true
|
||||
config.sources["disableRoutesAndAliasesOnExitNode"] = string(SourceEnv)
|
||||
}
|
||||
// if val := os.Getenv("DO_NOT_CREATE_NEW_CLIENT"); val == "true" {
|
||||
// config.DoNotCreateNewClient = true
|
||||
// config.sources["doNotCreateNewClient"] = string(SourceEnv)
|
||||
@@ -303,27 +315,29 @@ func loadConfigFromCLI(config *OlmConfig, args []string) (bool, bool, error) {
|
||||
|
||||
// Store original values to detect changes
|
||||
origValues := map[string]interface{}{
|
||||
"endpoint": config.Endpoint,
|
||||
"id": config.ID,
|
||||
"secret": config.Secret,
|
||||
"org": config.OrgID,
|
||||
"userToken": config.UserToken,
|
||||
"mtu": config.MTU,
|
||||
"dns": config.DNS,
|
||||
"upstreamDNS": fmt.Sprintf("%v", config.UpstreamDNS),
|
||||
"matchDomains": fmt.Sprintf("%v", config.MatchDomains),
|
||||
"logLevel": config.LogLevel,
|
||||
"interface": config.InterfaceName,
|
||||
"httpAddr": config.HTTPAddr,
|
||||
"socketPath": config.SocketPath,
|
||||
"pingInterval": config.PingInterval,
|
||||
"pingTimeout": config.PingTimeout,
|
||||
"enableApi": config.EnableAPI,
|
||||
"disableHolepunch": config.DisableHolepunch,
|
||||
"overrideDNS": config.OverrideDNS,
|
||||
"disableRelay": config.DisableRelay,
|
||||
"preferLocalRoutes": config.PreferLocalRoutes,
|
||||
"tunnelDNS": config.TunnelDNS,
|
||||
"endpoint": config.Endpoint,
|
||||
"id": config.ID,
|
||||
"secret": config.Secret,
|
||||
"org": config.OrgID,
|
||||
"userToken": config.UserToken,
|
||||
"mtu": config.MTU,
|
||||
"dns": config.DNS,
|
||||
"upstreamDNS": fmt.Sprintf("%v", config.UpstreamDNS),
|
||||
"matchDomains": fmt.Sprintf("%v", config.MatchDomains),
|
||||
"logLevel": config.LogLevel,
|
||||
"interface": config.InterfaceName,
|
||||
"httpAddr": config.HTTPAddr,
|
||||
"socketPath": config.SocketPath,
|
||||
"pingInterval": config.PingInterval,
|
||||
"pingTimeout": config.PingTimeout,
|
||||
"enableApi": config.EnableAPI,
|
||||
"disableHolepunch": config.DisableHolepunch,
|
||||
"overrideDNS": config.OverrideDNS,
|
||||
"disableRelay": config.DisableRelay,
|
||||
"preferLocalRoutes": config.PreferLocalRoutes,
|
||||
"tunnelDNS": config.TunnelDNS,
|
||||
"subnetRouter": config.SubnetRouter,
|
||||
"disableRoutesAndAliasesOnExitNode": config.DisableRoutesAndAliasesOnExitNode,
|
||||
// "doNotCreateNewClient": config.DoNotCreateNewClient,
|
||||
}
|
||||
|
||||
@@ -351,6 +365,8 @@ func loadConfigFromCLI(config *OlmConfig, args []string) (bool, bool, error) {
|
||||
serviceFlags.BoolVar(&config.DisableRelay, "disable-relay", config.DisableRelay, "Disable relay connections")
|
||||
serviceFlags.BoolVar(&config.PreferLocalRoutes, "prefer-local-routes", config.PreferLocalRoutes, "Add tunnel routes with a high metric so overlapping local/connected routes take precedence (default false)")
|
||||
serviceFlags.BoolVar(&config.TunnelDNS, "tunnel-dns", config.TunnelDNS, "When enabled, DNS queries are routed through the tunnel for remote resolution. To ensure queries are tunneled correctly, you must define the DNS server as a Pangolin resource and enter its address as an Upstream DNS Server. (default false)")
|
||||
serviceFlags.BoolVar(&config.SubnetRouter, "subnet-router", config.SubnetRouter, "Enable this client to act as a subnet router: traffic forwarded from the local network is NATed to this client's own tunnel IP before going out over the tunnel. Linux only, requires CAP_NET_ADMIN. (default false)")
|
||||
serviceFlags.BoolVar(&config.DisableRoutesAndAliasesOnExitNode, "disable-routes-and-aliases", config.DisableRoutesAndAliasesOnExitNode, "Make the exit node take precedence over individual resources: while an exit node is connected, remove routes/aliases for site resources, restoring them once it disconnects. Gateway routes are still added. (default false)")
|
||||
// serviceFlags.BoolVar(&config.DoNotCreateNewClient, "do-not-create-new-client", config.DoNotCreateNewClient, "Do not create new client")
|
||||
|
||||
version := serviceFlags.Bool("version", false, "Print the version")
|
||||
@@ -440,6 +456,12 @@ func loadConfigFromCLI(config *OlmConfig, args []string) (bool, bool, error) {
|
||||
if config.TunnelDNS != origValues["tunnelDNS"].(bool) {
|
||||
config.sources["tunnelDNS"] = string(SourceCLI)
|
||||
}
|
||||
if config.SubnetRouter != origValues["subnetRouter"].(bool) {
|
||||
config.sources["subnetRouter"] = string(SourceCLI)
|
||||
}
|
||||
if config.DisableRoutesAndAliasesOnExitNode != origValues["disableRoutesAndAliasesOnExitNode"].(bool) {
|
||||
config.sources["disableRoutesAndAliasesOnExitNode"] = string(SourceCLI)
|
||||
}
|
||||
// if config.DoNotCreateNewClient != origValues["doNotCreateNewClient"].(bool) {
|
||||
// config.sources["doNotCreateNewClient"] = string(SourceCLI)
|
||||
// }
|
||||
@@ -572,6 +594,14 @@ func mergeConfigs(dest, src *OlmConfig) {
|
||||
dest.PreferLocalRoutes = src.PreferLocalRoutes
|
||||
dest.sources["preferLocalRoutes"] = string(SourceFile)
|
||||
}
|
||||
if src.SubnetRouter {
|
||||
dest.SubnetRouter = src.SubnetRouter
|
||||
dest.sources["subnetRouter"] = string(SourceFile)
|
||||
}
|
||||
if src.DisableRoutesAndAliasesOnExitNode {
|
||||
dest.DisableRoutesAndAliasesOnExitNode = src.DisableRoutesAndAliasesOnExitNode
|
||||
dest.sources["disableRoutesAndAliasesOnExitNode"] = string(SourceFile)
|
||||
}
|
||||
// if src.DoNotCreateNewClient {
|
||||
// dest.DoNotCreateNewClient = src.DoNotCreateNewClient
|
||||
// dest.sources["doNotCreateNewClient"] = string(SourceFile)
|
||||
@@ -665,6 +695,8 @@ func (c *OlmConfig) ShowConfig() {
|
||||
fmt.Printf(" tunnel-dns = %v [%s]\n", c.TunnelDNS, getSource("tunnelDNS"))
|
||||
fmt.Printf(" disable-relay = %v [%s]\n", c.DisableRelay, getSource("disableRelay"))
|
||||
fmt.Printf(" prefer-local-routes = %v [%s]\n", c.PreferLocalRoutes, getSource("preferLocalRoutes"))
|
||||
fmt.Printf(" subnet-router = %v [%s]\n", c.SubnetRouter, getSource("subnetRouter"))
|
||||
fmt.Printf(" disable-routes-and-aliases = %v [%s]\n", c.DisableRoutesAndAliasesOnExitNode, getSource("disableRoutesAndAliasesOnExitNode"))
|
||||
// fmt.Printf(" do-not-create-new-client = %v [%s]\n", c.DoNotCreateNewClient, getSource("doNotCreateNewClient"))
|
||||
if c.TlsClientCert != "" {
|
||||
fmt.Printf(" tls-cert = %s [%s]\n", c.TlsClientCert, getSource("tlsClientCert"))
|
||||
|
||||
@@ -866,6 +866,29 @@ func (p *DNSProxy) SetUpstreamDNS(servers []string) {
|
||||
p.upstreamDNS = servers
|
||||
}
|
||||
|
||||
// SetTunnelDNS changes whether upstream DNS queries are sent over the
|
||||
// WireGuard tunnel (true) or directly via host networking (false). Only
|
||||
// takes effect for queries issued after the call; in-flight queries keep
|
||||
// using whichever path they already started on. Switching to true after the
|
||||
// proxy has already started lazily brings up the tunnel netstack and its
|
||||
// packet-forwarding goroutine if they weren't already running - NewDNSProxy
|
||||
// only does that eagerly when tunnelDns is true from the start.
|
||||
func (p *DNSProxy) SetTunnelDNS(tunnelDNS bool) {
|
||||
if tunnelDNS && p.tunnelStack == nil {
|
||||
if !p.tunnelIP.IsValid() {
|
||||
logger.Warn("Cannot enable tunnel DNS: tunnel IP not set")
|
||||
return
|
||||
}
|
||||
if err := p.initTunnelNetstack(); err != nil {
|
||||
logger.Error("Failed to initialize tunnel netstack for tunnel DNS: %v", err)
|
||||
return
|
||||
}
|
||||
p.wg.Add(1)
|
||||
go p.runTunnelPacketSender()
|
||||
}
|
||||
p.tunnelDNS = tunnelDNS
|
||||
}
|
||||
|
||||
// AddDNSRecord adds a DNS record to the local store
|
||||
// domain should be a domain name (e.g., "example.com" or "example.com.")
|
||||
// ip should be a valid IPv4 or IPv6 address
|
||||
|
||||
@@ -1,15 +1,16 @@
|
||||
module github.com/fosrl/olm
|
||||
|
||||
go 1.25.0
|
||||
go 1.26.0
|
||||
|
||||
require (
|
||||
github.com/Microsoft/go-winio v0.6.2
|
||||
github.com/fosrl/newt v1.16.0
|
||||
github.com/fosrl/newt v1.18.0
|
||||
github.com/godbus/dbus/v5 v5.2.2
|
||||
github.com/google/nftables v0.3.0
|
||||
github.com/gorilla/websocket v1.5.3
|
||||
github.com/miekg/dns v1.1.70
|
||||
golang.org/x/net v0.56.0
|
||||
golang.org/x/sys v0.46.0
|
||||
golang.org/x/net v0.59.0
|
||||
golang.org/x/sys v0.48.0
|
||||
golang.zx2c4.com/wireguard v0.0.0-20250521234502-f333402bd9cb
|
||||
golang.zx2c4.com/wireguard/wgctrl v0.0.0-20241231184526-a9ab2273dd10
|
||||
gvisor.dev/gvisor v0.0.0-20250503011706-39ed1f5ac29c
|
||||
@@ -19,9 +20,11 @@ require (
|
||||
require (
|
||||
github.com/google/btree v1.1.3 // indirect
|
||||
github.com/google/go-cmp v0.7.0 // indirect
|
||||
github.com/mdlayher/netlink v1.7.3-0.20250113171957-fbb4dce95f42 // indirect
|
||||
github.com/mdlayher/socket v0.5.1 // indirect
|
||||
github.com/vishvananda/netlink v1.3.1 // indirect
|
||||
github.com/vishvananda/netns v0.0.5 // indirect
|
||||
golang.org/x/crypto v0.53.0 // indirect
|
||||
golang.org/x/crypto v0.57.0 // indirect
|
||||
golang.org/x/exp v0.0.0-20251113190631-e25ba8c21ef6 // indirect
|
||||
golang.org/x/mod v0.34.0 // indirect
|
||||
golang.org/x/sync v0.20.0 // indirect
|
||||
|
||||
@@ -1,35 +1,41 @@
|
||||
github.com/Microsoft/go-winio v0.6.2 h1:F2VQgta7ecxGYO8k3ZZz3RS8fVIXVxONVUPlNERoyfY=
|
||||
github.com/Microsoft/go-winio v0.6.2/go.mod h1:yd8OoFMLzJbo9gZq8j5qaps8bJ9aShtEA8Ipt1oGCvU=
|
||||
github.com/fosrl/newt v1.16.0 h1:Nf70uNFn/WqHoTvRg5xTPO3xz9vLgh3BsnMVMoniEEw=
|
||||
github.com/fosrl/newt v1.16.0/go.mod h1:l6kWoZPSaXT+ZRUjiyPgwflRqZWYaXpUj9oQ0sOPh4o=
|
||||
github.com/fosrl/newt v1.18.0 h1:fFBksIV6BoI+BGmv6fwo87tXBk1WlxbVoruvupeD6hk=
|
||||
github.com/fosrl/newt v1.18.0/go.mod h1:CwcuQtifgDQeSWSEB3yfqOgheFv9yltVBYQNgDB/DoM=
|
||||
github.com/godbus/dbus/v5 v5.2.2 h1:TUR3TgtSVDmjiXOgAAyaZbYmIeP3DPkld3jgKGV8mXQ=
|
||||
github.com/godbus/dbus/v5 v5.2.2/go.mod h1:3AAv2+hPq5rdnr5txxxRwiGjPXamgoIHgz9FPBfOp3c=
|
||||
github.com/google/btree v1.1.3 h1:CVpQJjYgC4VbzxeGVHfvZrv1ctoYCAI8vbl07Fcxlyg=
|
||||
github.com/google/btree v1.1.3/go.mod h1:qOPhT0dTNdNzV6Z/lhRX0YXUafgPLFUh+gZMl761Gm4=
|
||||
github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8=
|
||||
github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX3N/iU=
|
||||
github.com/google/nftables v0.3.0 h1:bkyZ0cbpVeMHXOrtlFc8ISmfVqq5gPJukoYieyVmITg=
|
||||
github.com/google/nftables v0.3.0/go.mod h1:BCp9FsrbF1Fn/Yu6CLUc9GGZFw/+hsxfluNXXmxBfRM=
|
||||
github.com/gorilla/websocket v1.5.3 h1:saDtZ6Pbx/0u+bgYQ3q96pZgCzfhKXGPqt7kZ72aNNg=
|
||||
github.com/gorilla/websocket v1.5.3/go.mod h1:YR8l580nyteQvAITg2hZ9XVh4b55+EU/adAjf1fMHhE=
|
||||
github.com/mdlayher/netlink v1.7.3-0.20250113171957-fbb4dce95f42 h1:A1Cq6Ysb0GM0tpKMbdCXCIfBclan4oHk1Jb+Hrejirg=
|
||||
github.com/mdlayher/netlink v1.7.3-0.20250113171957-fbb4dce95f42/go.mod h1:BB4YCPDOzfy7FniQ/lxuYQ3dgmM2cZumHbK8RpTjN2o=
|
||||
github.com/mdlayher/socket v0.5.1 h1:VZaqt6RkGkt2OE9l3GcC6nZkqD3xKeQLyfleW/uBcos=
|
||||
github.com/mdlayher/socket v0.5.1/go.mod h1:TjPLHI1UgwEv5J1B5q0zTZq12A/6H7nKmtTanQE37IQ=
|
||||
github.com/miekg/dns v1.1.70 h1:DZ4u2AV35VJxdD9Fo9fIWm119BsQL5cZU1cQ9s0LkqA=
|
||||
github.com/miekg/dns v1.1.70/go.mod h1:+EuEPhdHOsfk6Wk5TT2CzssZdqkmFhf8r+aVyDEToIs=
|
||||
github.com/vishvananda/netlink v1.3.1 h1:3AEMt62VKqz90r0tmNhog0r/PpWKmrEShJU0wJW6bV0=
|
||||
github.com/vishvananda/netlink v1.3.1/go.mod h1:ARtKouGSTGchR8aMwmkzC0qiNPrrWO5JS/XMVl45+b4=
|
||||
github.com/vishvananda/netns v0.0.5 h1:DfiHV+j8bA32MFM7bfEunvT8IAqQ/NzSJHtcmW5zdEY=
|
||||
github.com/vishvananda/netns v0.0.5/go.mod h1:SpkAiCQRtJ6TvvxPnOSyH3BMl6unz3xZlaprSwhNNJM=
|
||||
golang.org/x/crypto v0.53.0 h1:QZ4Muo8THX6CizN2vPPd5fBGHyogrdK9fG4wLPFUsto=
|
||||
golang.org/x/crypto v0.53.0/go.mod h1:DNLU434OwVakk9PzuwV8w62mAJpRJL3vsgcfp4Qnsio=
|
||||
golang.org/x/crypto v0.57.0 h1:3ZVCjf8Ggz7zneR/EHRVx68Ctf+2pmIMP2UFhh9cC6M=
|
||||
golang.org/x/crypto v0.57.0/go.mod h1:Fdz0i5U6CoizGwLda9DttjSk6qlZo25zYNtR+ycvuZA=
|
||||
golang.org/x/exp v0.0.0-20251113190631-e25ba8c21ef6 h1:zfMcR1Cs4KNuomFFgGefv5N0czO2XZpUbxGUy8i8ug0=
|
||||
golang.org/x/exp v0.0.0-20251113190631-e25ba8c21ef6/go.mod h1:46edojNIoXTNOhySWIWdix628clX9ODXwPsQuG6hsK0=
|
||||
golang.org/x/mod v0.34.0 h1:xIHgNUUnW6sYkcM5Jleh05DvLOtwc6RitGHbDk4akRI=
|
||||
golang.org/x/mod v0.34.0/go.mod h1:ykgH52iCZe79kzLLMhyCUzhMci+nQj+0XkbXpNYtVjY=
|
||||
golang.org/x/net v0.56.0 h1:Rw8j/hFzGvJUZwNBXnAtf5sVDVt+65SK2C7IxCxZt5o=
|
||||
golang.org/x/net v0.56.0/go.mod h1:D3Ku6r+V6JROoZK144D2XfMHFcMq/0zSfLelVTCFKec=
|
||||
golang.org/x/net v0.59.0 h1:5zfYln+w5XCxwrnMMJPufRgNoXEaGxl0wo5GqPXyues=
|
||||
golang.org/x/net v0.59.0/go.mod h1:2DA/G1UfVbCpQPeWTmMPGY7Cs2PkBkwu743bVX5PIVg=
|
||||
golang.org/x/sync v0.20.0 h1:e0PTpb7pjO8GAtTs2dQ6jYa5BWYlMuX047Dco/pItO4=
|
||||
golang.org/x/sync v0.20.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0=
|
||||
golang.org/x/sys v0.2.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||
golang.org/x/sys v0.10.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg=
|
||||
golang.org/x/sys v0.46.0 h1:noSf2Fq6F8DBgS+LysIkx7rIExoNHJsxOAtPp4rthXw=
|
||||
golang.org/x/sys v0.46.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw=
|
||||
golang.org/x/sys v0.48.0 h1:bbX/i/6MgT9BVLM9RT1thmxL04yeTAhbEz4SyadbXoo=
|
||||
golang.org/x/sys v0.48.0/go.mod h1:hNLxWAXmnKAxqDtdwIYC4bM9oQPEecfsnNMuSxOs3og=
|
||||
golang.org/x/time v0.12.0 h1:ScB/8o8olJvc+CQPWrK3fPZNfh7qgwCrY0zJmoEQLSE=
|
||||
golang.org/x/time v0.12.0/go.mod h1:CDIdPxbZBQxdj6cxyCIdrNogrJKMJ7pr37NYpMcMDSg=
|
||||
golang.org/x/tools v0.43.0 h1:12BdW9CeB3Z+J/I/wj34VMl8X+fEXBxVR90JeMX5E7s=
|
||||
|
||||
@@ -256,24 +256,26 @@ func runOlmMainWithArgs(ctx context.Context, cancel context.CancelFunc, signalCt
|
||||
|
||||
if config.ID != "" && config.Secret != "" && config.Endpoint != "" {
|
||||
tunnelConfig := olmpkg.TunnelConfig{
|
||||
Endpoint: config.Endpoint,
|
||||
ID: config.ID,
|
||||
Secret: config.Secret,
|
||||
UserToken: config.UserToken,
|
||||
MTU: config.MTU,
|
||||
DNS: config.DNS,
|
||||
UpstreamDNS: config.UpstreamDNS,
|
||||
MatchDomains: config.MatchDomains,
|
||||
InterfaceName: config.InterfaceName,
|
||||
Holepunch: !config.DisableHolepunch,
|
||||
TlsClientCert: config.TlsClientCert,
|
||||
PingIntervalDuration: config.PingIntervalDuration,
|
||||
PingTimeoutDuration: config.PingTimeoutDuration,
|
||||
OrgID: config.OrgID,
|
||||
OverrideDNS: config.OverrideDNS,
|
||||
DisableRelay: config.DisableRelay,
|
||||
PreferLocalRoutes: config.PreferLocalRoutes,
|
||||
EnableUAPI: true,
|
||||
Endpoint: config.Endpoint,
|
||||
ID: config.ID,
|
||||
Secret: config.Secret,
|
||||
UserToken: config.UserToken,
|
||||
MTU: config.MTU,
|
||||
DNS: config.DNS,
|
||||
UpstreamDNS: config.UpstreamDNS,
|
||||
MatchDomains: config.MatchDomains,
|
||||
InterfaceName: config.InterfaceName,
|
||||
Holepunch: !config.DisableHolepunch,
|
||||
TlsClientCert: config.TlsClientCert,
|
||||
PingIntervalDuration: config.PingIntervalDuration,
|
||||
PingTimeoutDuration: config.PingTimeoutDuration,
|
||||
OrgID: config.OrgID,
|
||||
OverrideDNS: config.OverrideDNS,
|
||||
DisableRelay: config.DisableRelay,
|
||||
PreferLocalRoutes: config.PreferLocalRoutes,
|
||||
SubnetRouter: config.SubnetRouter,
|
||||
DisableRoutesAndAliasesOnExitNode: config.DisableRoutesAndAliasesOnExitNode,
|
||||
EnableUAPI: true,
|
||||
}
|
||||
go olm.StartTunnel(tunnelConfig)
|
||||
} else {
|
||||
|
||||
+92
-63
@@ -15,8 +15,8 @@ import (
|
||||
"github.com/fosrl/newt/util"
|
||||
olmDevice "github.com/fosrl/olm/device"
|
||||
"github.com/fosrl/olm/dns"
|
||||
dnsOverride "github.com/fosrl/olm/dns/override"
|
||||
"github.com/fosrl/olm/peers"
|
||||
"github.com/fosrl/olm/subnetrouter"
|
||||
"github.com/fosrl/olm/websocket"
|
||||
"golang.zx2c4.com/wireguard/device"
|
||||
"golang.zx2c4.com/wireguard/tun"
|
||||
@@ -75,6 +75,12 @@ func (o *Olm) handleConnect(msg websocket.WSMessage) {
|
||||
return
|
||||
}
|
||||
|
||||
// A server-provided DNS config always overrides the client's own local
|
||||
// config (from CLI flags / API connect request) - see applyDNSConfigUpdate.
|
||||
if wgData.DNSConfig != nil {
|
||||
o.applyDNSConfigUpdate(*wgData.DNSConfig)
|
||||
}
|
||||
|
||||
// When handed an already-open FD (mobile/NetworkExtension platforms), the
|
||||
// TUN device's addresses and routes are owned and reconciled by the host
|
||||
// platform from NetworkSettings (e.g. Apple's NEPacketTunnelProvider via
|
||||
@@ -175,7 +181,12 @@ func (o *Olm) handleConnect(msg websocket.WSMessage) {
|
||||
logger.Warn("Failed to parse tunnel IP %q: %v", interfaceIP, err)
|
||||
}
|
||||
|
||||
// Create and start DNS proxy
|
||||
// The DNS proxy is always created - it resolves both exit node aliases
|
||||
// (unaffected by DisableRoutesAndAliasesOnExitNode - see connectExitNode)
|
||||
// and site resource aliases, which the peer manager created below adds
|
||||
// and removes dynamically as an exit node connects/disconnects (see
|
||||
// PeerManager.SetExitNode/ClearExitNode). o.dnsProxy is still nil-checked
|
||||
// everywhere it's used, in case creation itself fails.
|
||||
o.dnsProxy, err = dns.NewDNSProxy(o.middleDev, o.tunnelConfig.MTU, wgData.UtilitySubnet, o.tunnelConfig.UpstreamDNS, o.tunnelConfig.TunnelDNS, interfaceIP, o.tunnelConfig.MatchDomains, o.tunnelConfig.PublicDNS)
|
||||
if err != nil {
|
||||
logger.Error("Failed to create DNS proxy: %v", err)
|
||||
@@ -193,22 +204,42 @@ func (o *Olm) handleConnect(msg websocket.WSMessage) {
|
||||
logger.Error("Failed to o.tunnelConfigure interface: %v", err)
|
||||
}
|
||||
|
||||
if o.tunnelConfig.SubnetRouter {
|
||||
if err := subnetrouter.Enable(o.tunnelConfig.InterfaceName, o.primaryTunnelIP); err != nil {
|
||||
logger.Error("Failed to enable subnet router: %v", err)
|
||||
} else {
|
||||
logger.Info("Subnet router enabled on %s (SNAT to %s)", o.tunnelConfig.InterfaceName, o.primaryTunnelIP)
|
||||
}
|
||||
}
|
||||
|
||||
// The utility subnet route is what makes the DNS proxy (always created
|
||||
// above) reachable at all, so it's always added too, independent of
|
||||
// DisableRoutesAndAliasesOnExitNode - which only gates routes/aliases for
|
||||
// individual site resources, applied dynamically below as the peer
|
||||
// manager is told about the exit node's connection state.
|
||||
if err := network.AddRoutesWithSource([]string{wgData.UtilitySubnet}, o.tunnelConfig.InterfaceName, interfaceIP); err != nil { // also route the utility subnet
|
||||
logger.Error("Failed to add route for utility subnet: %v", err)
|
||||
}
|
||||
|
||||
// Create peer manager with integrated peer monitoring
|
||||
// Create peer manager with integrated peer monitoring. If
|
||||
// DisableRoutesAndAliasesOnExitNode is enabled, resource routes/aliases
|
||||
// are suppressed reactively once an exit node (ExitNodeConfig peer or
|
||||
// gateway mode) actually activates below/later - see
|
||||
// PeerManager.SetExitNode/SetGateway - rather than being pre-seeded here,
|
||||
// so a failed initial exit-node/gateway setup can never leave routes
|
||||
// stuck suppressed with nothing active to justify it.
|
||||
o.peerManager = peers.NewPeerManager(peers.PeerManagerConfig{
|
||||
Device: o.dev,
|
||||
DNSProxy: o.dnsProxy,
|
||||
InterfaceName: o.tunnelConfig.InterfaceName,
|
||||
PrivateKey: o.privateKey,
|
||||
MiddleDev: o.middleDev,
|
||||
LocalIP: interfaceIP,
|
||||
SharedBind: o.sharedBind,
|
||||
WSClient: o.websocket,
|
||||
APIServer: o.apiServer,
|
||||
PublicDNS: o.tunnelConfig.PublicDNS,
|
||||
Device: o.dev,
|
||||
DNSProxy: o.dnsProxy,
|
||||
InterfaceName: o.tunnelConfig.InterfaceName,
|
||||
PrivateKey: o.privateKey,
|
||||
MiddleDev: o.middleDev,
|
||||
LocalIP: interfaceIP,
|
||||
SharedBind: o.sharedBind,
|
||||
WSClient: o.websocket,
|
||||
APIServer: o.apiServer,
|
||||
PublicDNS: o.tunnelConfig.PublicDNS,
|
||||
DisableRoutesAndAliasesOnExitNode: o.tunnelConfig.DisableRoutesAndAliasesOnExitNode,
|
||||
})
|
||||
|
||||
for i := range wgData.Sites {
|
||||
@@ -237,62 +268,55 @@ func (o *Olm) handleConnect(msg websocket.WSMessage) {
|
||||
|
||||
o.peerManager.Start()
|
||||
|
||||
if err := o.dnsProxy.Start(); err != nil { // start DNS proxy first so there is no downtime
|
||||
logger.Error("Failed to start DNS proxy: %v", err)
|
||||
}
|
||||
// OnTokenUpdate (see olm.go) typically fires before this peer manager
|
||||
// existed - it runs during the initial token/auth fetch, well before this
|
||||
// "olm/wg/connect" message - so push in whatever hole-punch bypass
|
||||
// endpoints it already recorded now that there's somewhere to put them.
|
||||
o.flushPendingHolepunchBypassEndpoints()
|
||||
o.flushPendingDNSBypassEndpoints()
|
||||
|
||||
// Register JIT handler: when the DNS proxy resolves a local record, check whether
|
||||
// the owning site is already connected and, if not, initiate a JIT connection.
|
||||
o.dnsProxy.SetJITHandler(func(siteId int) {
|
||||
pm := o.getPeerManager()
|
||||
if pm == nil || o.websocket == nil {
|
||||
return
|
||||
if o.dnsProxy != nil {
|
||||
if err := o.dnsProxy.Start(); err != nil { // start DNS proxy first so there is no downtime
|
||||
logger.Error("Failed to start DNS proxy: %v", err)
|
||||
}
|
||||
|
||||
// Site already has an active peer connection - nothing to do.
|
||||
if _, exists := pm.GetPeer(siteId); exists {
|
||||
return
|
||||
}
|
||||
|
||||
o.peerSendMu.Lock()
|
||||
defer o.peerSendMu.Unlock()
|
||||
|
||||
// A JIT request for this site is already in-flight - avoid duplicate sends.
|
||||
if _, pending := o.jitPendingSites[siteId]; pending {
|
||||
return
|
||||
}
|
||||
|
||||
chainId := generateChainId()
|
||||
logger.Info("DNS-triggered JIT connect for site %d (chainId=%s)", siteId, chainId)
|
||||
stopFunc, _ := o.websocket.SendMessageInterval("olm/wg/server/peer/init", map[string]interface{}{
|
||||
"siteId": siteId,
|
||||
"chainId": chainId,
|
||||
}, 2*time.Second, 10)
|
||||
o.stopPeerInits[chainId] = stopFunc
|
||||
o.jitPendingSites[siteId] = chainId
|
||||
})
|
||||
|
||||
if o.tunnelConfig.OverrideDNS {
|
||||
// When the host platform already applies DNS natively (NEDNSSettings on
|
||||
// macOS/iOS, scoped to the tunnel session and auto-cleaned by the OS no
|
||||
// matter how the session ends), skip olm's own raw scutil-based override -
|
||||
// there is nothing for it to add and, unlike NEDNSSettings, it has no way
|
||||
// to guarantee cleanup if this process dies uncleanly. See NativeDNSManaged.
|
||||
if !o.tunnelConfig.NativeDNSManaged {
|
||||
// Set up DNS override to use our DNS proxy
|
||||
if err := dnsOverride.SetupDNSOverride(o.tunnelConfig.InterfaceName, o.dnsProxy.GetProxyIP()); err != nil {
|
||||
logger.Error("Failed to setup DNS override: %v", err)
|
||||
// Register JIT handler: when the DNS proxy resolves a local record, check whether
|
||||
// the owning site is already connected and, if not, initiate a JIT connection.
|
||||
o.dnsProxy.SetJITHandler(func(siteId int) {
|
||||
pm := o.getPeerManager()
|
||||
if pm == nil || o.websocket == nil {
|
||||
return
|
||||
}
|
||||
|
||||
// Start the external watchdog (if configured). The watchdog will
|
||||
// reset DNS if this process dies before it can call
|
||||
// RestoreDNSOverride. This is a no-op when no watchdog
|
||||
// subcommand has been configured on the OlmConfig.
|
||||
o.startDNSWatchdog(o.tunnelConfig.InterfaceName)
|
||||
}
|
||||
// Site already has an active peer connection - nothing to do.
|
||||
if _, exists := pm.GetPeer(siteId); exists {
|
||||
return
|
||||
}
|
||||
|
||||
network.SetDNSServers([]string{o.dnsProxy.GetProxyIP().String()})
|
||||
o.peerSendMu.Lock()
|
||||
defer o.peerSendMu.Unlock()
|
||||
|
||||
// A JIT request for this site is already in-flight - avoid duplicate sends.
|
||||
if _, pending := o.jitPendingSites[siteId]; pending {
|
||||
return
|
||||
}
|
||||
|
||||
chainId := generateChainId()
|
||||
logger.Info("DNS-triggered JIT connect for site %d (chainId=%s)", siteId, chainId)
|
||||
stopFunc, _ := o.websocket.SendMessageInterval("olm/wg/server/peer/init", map[string]interface{}{
|
||||
"siteId": siteId,
|
||||
"chainId": chainId,
|
||||
}, 2*time.Second, 10)
|
||||
o.stopPeerInits[chainId] = stopFunc
|
||||
o.jitPendingSites[siteId] = chainId
|
||||
})
|
||||
}
|
||||
|
||||
if o.tunnelConfig.OverrideDNS && o.dnsProxy != nil {
|
||||
if err := o.applyDNSOverride(true); err != nil {
|
||||
logger.Error("%v", err)
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
if wgData.ExitNode != nil && wgData.ExitNode.Connect {
|
||||
@@ -307,6 +331,10 @@ func (o *Olm) handleConnect(msg websocket.WSMessage) {
|
||||
|
||||
o.registered = true
|
||||
|
||||
if len(o.tunnelConfig.GatewaySiteIds) > 0 {
|
||||
o.applyPendingGatewayConfig(o.tunnelConfig.GatewaySiteResourceId, o.tunnelConfig.GatewaySiteIds)
|
||||
}
|
||||
|
||||
// Start ping monitor now that we are registered and connected
|
||||
o.websocket.StartPingMonitor()
|
||||
|
||||
@@ -372,7 +400,7 @@ func (o *Olm) handleTerminate(msg websocket.WSMessage) {
|
||||
logger.Info("Terminate reason (code: %s): %s", errorData.Code, errorData.Message)
|
||||
|
||||
if errorData.Code == "TERMINATED_INACTIVITY" {
|
||||
logger.Info("Ignoring...")
|
||||
logger.Debug("Ignoring TERMINATED_INACTIVITY message for now because we could have been sleeping...")
|
||||
return
|
||||
}
|
||||
|
||||
@@ -385,6 +413,7 @@ func (o *Olm) handleTerminate(msg websocket.WSMessage) {
|
||||
o.apiServer.SetConnectionStatus(false)
|
||||
o.apiServer.SetRegistered(false)
|
||||
o.apiServer.ClearPeerStatuses()
|
||||
o.apiServer.SetGatewayStatus(false, 0, nil)
|
||||
|
||||
network.ClearNetworkSettings()
|
||||
|
||||
|
||||
@@ -0,0 +1,122 @@
|
||||
package olm
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
|
||||
"github.com/fosrl/newt/logger"
|
||||
"github.com/fosrl/newt/network"
|
||||
dnsOverride "github.com/fosrl/olm/dns/override"
|
||||
"github.com/fosrl/olm/websocket"
|
||||
)
|
||||
|
||||
// applyDNSConfigUpdate merges a server-provided DNS config override into the
|
||||
// running tunnel config. Every set field always overrides the client's own
|
||||
// local config (from CLI flags / API connect request, or a previous update);
|
||||
// an unset field leaves the current value alone. Called both from the
|
||||
// initial "olm/wg/connect" message, before the DNS proxy exists yet (the
|
||||
// updated tunnelConfig feeds into its construction in handleConnect), and
|
||||
// from a later live "olm/wg/dns/update" push, where the running proxy is
|
||||
// updated directly.
|
||||
func (o *Olm) applyDNSConfigUpdate(cfg DNSConfigUpdate) {
|
||||
logger.Info("Applying DNS config from server: %+v", cfg)
|
||||
|
||||
if len(cfg.UpstreamDNS) > 0 {
|
||||
o.tunnelConfig.UpstreamDNS = cfg.UpstreamDNS
|
||||
if o.dnsProxy != nil {
|
||||
o.dnsProxy.SetUpstreamDNS(cfg.UpstreamDNS)
|
||||
}
|
||||
o.updateDNSBypassEndpoints(cfg.UpstreamDNS)
|
||||
}
|
||||
|
||||
if len(cfg.MatchDomains) > 0 {
|
||||
o.tunnelConfig.MatchDomains = cfg.MatchDomains
|
||||
if o.dnsProxy != nil {
|
||||
o.dnsProxy.SetMatchDomains(cfg.MatchDomains)
|
||||
}
|
||||
}
|
||||
|
||||
if cfg.TunnelDNS != nil {
|
||||
o.tunnelConfig.TunnelDNS = *cfg.TunnelDNS
|
||||
if o.dnsProxy != nil {
|
||||
o.dnsProxy.SetTunnelDNS(*cfg.TunnelDNS)
|
||||
}
|
||||
}
|
||||
|
||||
if cfg.OverrideDNS != nil && *cfg.OverrideDNS != o.tunnelConfig.OverrideDNS {
|
||||
if o.dnsProxy != nil {
|
||||
if err := o.applyDNSOverride(*cfg.OverrideDNS); err != nil {
|
||||
logger.Error("Failed to apply DNS override update: %v", err)
|
||||
return
|
||||
}
|
||||
}
|
||||
o.tunnelConfig.OverrideDNS = *cfg.OverrideDNS
|
||||
}
|
||||
}
|
||||
|
||||
// applyDNSOverride installs or removes olm's own system DNS override
|
||||
// (pointing the host resolver at the DNS proxy) and requires the DNS proxy
|
||||
// to already be running. When the host platform already manages DNS
|
||||
// natively (see NativeDNSManaged), this only updates the OS resolver list -
|
||||
// there's no raw override to add or remove. Used both at initial connect
|
||||
// (see handleConnect) and for a live "olm/wg/dns/update" toggle of
|
||||
// OverrideDNS via applyDNSConfigUpdate.
|
||||
func (o *Olm) applyDNSOverride(enable bool) error {
|
||||
if o.dnsProxy == nil {
|
||||
return fmt.Errorf("cannot toggle DNS override: DNS proxy is not running")
|
||||
}
|
||||
|
||||
if enable {
|
||||
if !o.tunnelConfig.NativeDNSManaged {
|
||||
if err := dnsOverride.SetupDNSOverride(o.tunnelConfig.InterfaceName, o.dnsProxy.GetProxyIP()); err != nil {
|
||||
return fmt.Errorf("failed to setup DNS override: %w", err)
|
||||
}
|
||||
|
||||
// Start the external watchdog (if configured). The watchdog will
|
||||
// reset DNS if this process dies before it can call
|
||||
// RestoreDNSOverride. This is a no-op when no watchdog
|
||||
// subcommand has been configured on the OlmConfig.
|
||||
o.startDNSWatchdog(o.tunnelConfig.InterfaceName)
|
||||
}
|
||||
|
||||
network.SetDNSServers([]string{o.dnsProxy.GetProxyIP().String()})
|
||||
} else {
|
||||
if !o.tunnelConfig.NativeDNSManaged {
|
||||
if err := dnsOverride.RestoreDNSOverride(); err != nil {
|
||||
return fmt.Errorf("failed to restore DNS: %w", err)
|
||||
}
|
||||
|
||||
o.stopDNSWatchdog()
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// handleDNSConfigUpdate handles a server-initiated request to change the
|
||||
// client's DNS configuration (upstream DNS, tunnel DNS, override DNS, match
|
||||
// domains) after it is already connected, without requiring a full
|
||||
// reconnect. Mirrors the DNSConfig field sent on the initial
|
||||
// "olm/wg/connect" message - see DNSConfigUpdate.
|
||||
func (o *Olm) handleDNSConfigUpdate(msg websocket.WSMessage) {
|
||||
logger.Debug("Received DNS config update message: %v", msg.Data)
|
||||
|
||||
if !o.tunnelRunning {
|
||||
logger.Debug("Tunnel stopped, ignoring DNS config update message")
|
||||
return
|
||||
}
|
||||
|
||||
jsonData, err := json.Marshal(msg.Data)
|
||||
if err != nil {
|
||||
logger.Error("Error marshaling DNS config update data: %v", err)
|
||||
return
|
||||
}
|
||||
|
||||
var update DNSConfigUpdate
|
||||
if err := json.Unmarshal(jsonData, &update); err != nil {
|
||||
logger.Error("Error unmarshaling DNS config update data: %v", err)
|
||||
return
|
||||
}
|
||||
|
||||
o.applyDNSConfigUpdate(update)
|
||||
}
|
||||
@@ -210,7 +210,14 @@ persistent_keepalive_interval=%d`, util.FixKey(cfg.PublicKey), allowedIP, resolv
|
||||
|
||||
if pm := o.getPeerManager(); pm != nil {
|
||||
pm.SetExitNode(strings.Split(cfg.ServerIP, "/")[0], strings.Split(cfg.TunnelIP, "/")[0])
|
||||
// Distinct from the hole-punch exit nodes registered in olm.go's
|
||||
// OnTokenUpdate handler: this is the exit node actually connected as a
|
||||
// WireGuard peer above. Its own traffic must stay off the gateway
|
||||
// route the same way a site peer's endpoint does, or it would loop
|
||||
// through the tunnel it's part of maintaining.
|
||||
pm.AddGatewayBypassEndpoint(resolvedEndpoint)
|
||||
}
|
||||
o.exitNodeResolvedEndpoint = resolvedEndpoint
|
||||
|
||||
logger.Info("Connected to exit node at %s", resolvedEndpoint)
|
||||
return nil
|
||||
@@ -232,9 +239,14 @@ func (o *Olm) removeExitNodePeerLocked() error {
|
||||
}
|
||||
cfg := o.exitNode
|
||||
o.exitNode = nil
|
||||
resolvedEndpoint := o.exitNodeResolvedEndpoint
|
||||
o.exitNodeResolvedEndpoint = ""
|
||||
|
||||
if pm := o.getPeerManager(); pm != nil {
|
||||
pm.ClearExitNode()
|
||||
if resolvedEndpoint != "" {
|
||||
pm.RemoveGatewayBypassEndpoint(resolvedEndpoint)
|
||||
}
|
||||
}
|
||||
|
||||
if o.dnsProxy != nil {
|
||||
|
||||
+288
@@ -0,0 +1,288 @@
|
||||
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, 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(siteResourceId, siteIds)
|
||||
}
|
||||
|
||||
// DisableGateway fully clears gateway state. Requires the tunnel to already
|
||||
// be registered/connected; rejects otherwise.
|
||||
func (o *Olm) DisableGateway() error {
|
||||
if !o.registered {
|
||||
return fmt.Errorf("cannot disable gateway: not registered/connected")
|
||||
}
|
||||
pm := o.getPeerManager()
|
||||
if pm == nil {
|
||||
return fmt.Errorf("cannot disable gateway: tunnel not running")
|
||||
}
|
||||
if err := pm.ClearGateway(); err != nil {
|
||||
return err
|
||||
}
|
||||
o.apiServer.SetGatewayStatus(false, 0, nil)
|
||||
return nil
|
||||
}
|
||||
|
||||
// applySelectGateway resolves the control-plane endpoint host and delegates
|
||||
// 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(siteResourceId int, siteIds []int) error {
|
||||
pm := o.getPeerManager()
|
||||
if pm == nil {
|
||||
return fmt.Errorf("cannot select gateway: tunnel not running")
|
||||
}
|
||||
if err := pm.SetGateway(siteResourceId, siteIds, extractControlEndpointHost(o.tunnelConfig.Endpoint)); err != nil {
|
||||
return err
|
||||
}
|
||||
o.apiServer.SetGatewayStatus(true, siteResourceId, siteIds)
|
||||
return nil
|
||||
}
|
||||
|
||||
// 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(siteResourceId int, requestedSiteIds []int) {
|
||||
pm := o.getPeerManager()
|
||||
if pm == nil {
|
||||
return
|
||||
}
|
||||
|
||||
var tracked []int
|
||||
for _, id := range requestedSiteIds {
|
||||
if _, ok := pm.GetPeer(id); ok {
|
||||
tracked = append(tracked, id)
|
||||
} else {
|
||||
logger.Warn("Gateway site %d requested at connect time was not found among tracked peers; skipping", id)
|
||||
}
|
||||
}
|
||||
|
||||
if len(tracked) == 0 {
|
||||
logger.Warn("None of the requested gateway site IDs were found among tracked peers; gateway not established")
|
||||
return
|
||||
}
|
||||
|
||||
if err := o.applySelectGateway(siteResourceId, tracked); err != nil {
|
||||
logger.Error("Failed to establish gateway from StartTunnel config: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// flushPendingHolepunchBypassEndpoints re-registers every currently-known
|
||||
// hole-punch bypass endpoint (see the OnTokenUpdate handler in olm.go) with
|
||||
// the peer manager. OnTokenUpdate typically fires before the peer manager
|
||||
// exists - it runs during the initial token/auth fetch in
|
||||
// websocket.Client.establishConnection, well before the server's
|
||||
// "olm/wg/connect" message creates the peer manager here in handleConnect -
|
||||
// so anything recorded into o.hpBypassEndpoints while pm was nil needs to be
|
||||
// pushed in once it becomes available. AddGatewayBypassEndpoint is
|
||||
// idempotent, so calling it again for an endpoint OnTokenUpdate already
|
||||
// managed to register directly (e.g. a later token refresh, once the peer
|
||||
// manager already existed) is harmless.
|
||||
func (o *Olm) flushPendingHolepunchBypassEndpoints() {
|
||||
pm := o.getPeerManager()
|
||||
if pm == nil {
|
||||
return
|
||||
}
|
||||
|
||||
o.hpBypassMu.Lock()
|
||||
defer o.hpBypassMu.Unlock()
|
||||
for hostport := range o.hpBypassEndpoints {
|
||||
pm.AddGatewayBypassEndpoint(hostport)
|
||||
}
|
||||
}
|
||||
|
||||
// updateDNSBypassEndpoints diffs servers (the DNS proxy's upstream/primary
|
||||
// and secondary DNS servers - see TunnelConfig.UpstreamDNS and
|
||||
// dns.DNSProxy.SetUpstreamDNS) against the currently-registered set and
|
||||
// adds/removes gateway bypass routes for the difference, via the same
|
||||
// AddGatewayBypassEndpoint/RemoveGatewayBypassEndpoint machinery used for
|
||||
// hole-punch endpoints above. This keeps the DNS proxy's own outbound queries
|
||||
// to its real upstream resolvers off the gateway default-route-equivalent, so
|
||||
// they reach the real servers directly instead of looping back through the
|
||||
// tunnel. Called from StartTunnel (initial value and dynamic system-DNS
|
||||
// updates, including SetSystemDNS pushes) and applyDNSConfigUpdate (live
|
||||
// server-pushed overrides). Safe to call before the peer manager exists (see
|
||||
// flushPendingDNSBypassEndpoints) and safe to call repeatedly with the same
|
||||
// servers (no-op).
|
||||
func (o *Olm) updateDNSBypassEndpoints(servers []string) {
|
||||
pm := o.getPeerManager()
|
||||
|
||||
newBypassEndpoints := make(map[string]bool, len(servers))
|
||||
for _, server := range servers {
|
||||
newBypassEndpoints[server] = true
|
||||
}
|
||||
|
||||
o.dnsBypassMu.Lock()
|
||||
defer o.dnsBypassMu.Unlock()
|
||||
if pm != nil {
|
||||
for server := range newBypassEndpoints {
|
||||
if !o.dnsBypassEndpoints[server] {
|
||||
pm.AddGatewayBypassEndpoint(server)
|
||||
}
|
||||
}
|
||||
for server := range o.dnsBypassEndpoints {
|
||||
if !newBypassEndpoints[server] {
|
||||
pm.RemoveGatewayBypassEndpoint(server)
|
||||
}
|
||||
}
|
||||
}
|
||||
o.dnsBypassEndpoints = newBypassEndpoints
|
||||
}
|
||||
|
||||
// flushPendingDNSBypassEndpoints re-registers every currently-known upstream
|
||||
// DNS bypass endpoint with the peer manager. Mirrors
|
||||
// flushPendingHolepunchBypassEndpoints: updateDNSBypassEndpoints typically
|
||||
// runs before the peer manager exists (the initial UpstreamDNS value is
|
||||
// applied in StartTunnel, and a DNS config override may arrive at the very
|
||||
// start of handleConnect - see olm/dns_config.go - both well before
|
||||
// handleConnect constructs the peer manager further down), so whatever was
|
||||
// recorded needs to be pushed in once it becomes available.
|
||||
// AddGatewayBypassEndpoint is idempotent.
|
||||
func (o *Olm) flushPendingDNSBypassEndpoints() {
|
||||
pm := o.getPeerManager()
|
||||
if pm == nil {
|
||||
return
|
||||
}
|
||||
|
||||
o.dnsBypassMu.Lock()
|
||||
defer o.dnsBypassMu.Unlock()
|
||||
for server := range o.dnsBypassEndpoints {
|
||||
pm.AddGatewayBypassEndpoint(server)
|
||||
}
|
||||
}
|
||||
|
||||
// extractControlEndpointHost returns the bare host (no scheme/port) of the
|
||||
// Pangolin server olm is registered against, for gateway bypass-route
|
||||
// purposes. Falls back to the raw endpoint string on parse failure -
|
||||
// resolveEndpointIPLocked/net.SplitHostPort tolerate a bare host - rather
|
||||
// than failing the whole gateway activation over a cosmetic parse issue.
|
||||
func extractControlEndpointHost(endpoint string) string {
|
||||
u, err := url.Parse(endpoint)
|
||||
if err != nil || u.Hostname() == "" {
|
||||
return endpoint
|
||||
}
|
||||
return u.Hostname()
|
||||
}
|
||||
+162
-25
@@ -8,10 +8,12 @@ import (
|
||||
"fmt"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/netip"
|
||||
_ "net/http/pprof"
|
||||
"net/netip"
|
||||
"os"
|
||||
"os/exec"
|
||||
"runtime"
|
||||
"strconv"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
@@ -27,6 +29,7 @@ import (
|
||||
"github.com/fosrl/olm/dns"
|
||||
dnsOverride "github.com/fosrl/olm/dns/override"
|
||||
"github.com/fosrl/olm/peers"
|
||||
"github.com/fosrl/olm/subnetrouter"
|
||||
"github.com/fosrl/olm/websocket"
|
||||
"golang.zx2c4.com/wireguard/device"
|
||||
"golang.zx2c4.com/wireguard/tun"
|
||||
@@ -63,6 +66,32 @@ type Olm struct {
|
||||
// secondary address on the same interface/WireGuard device as the site peers.
|
||||
exitNode *ExitNodeConfig
|
||||
exitNodeMu sync.Mutex
|
||||
// exitNodeResolvedEndpoint is the exit node's WireGuard endpoint (already
|
||||
// DNS-resolved to "ip:port") as of the last successful connectExitNode
|
||||
// call, kept alongside exitNode purely so removeExitNodePeerLocked can
|
||||
// unregister the exact same gateway bypass-route target it registered -
|
||||
// see connectExitNode's AddGatewayBypassEndpoint call. Guarded by exitNodeMu.
|
||||
exitNodeResolvedEndpoint string
|
||||
|
||||
// hpBypassEndpoints tracks the "host:relayPort" endpoints currently
|
||||
// registered as gateway bypass targets for hole-punch exit nodes (STUN-like
|
||||
// probing, distinct from a connected exit node's own WireGuard peer above) -
|
||||
// diffed against each OnTokenUpdate so stale entries are unregistered and
|
||||
// new ones protected while gateway mode is active.
|
||||
hpBypassEndpoints map[string]bool
|
||||
hpBypassMu sync.Mutex
|
||||
|
||||
// dnsBypassEndpoints tracks the "host:port" upstream DNS servers (see
|
||||
// TunnelConfig.UpstreamDNS / dns.DNSProxy's upstreamDNS - the DNS proxy's
|
||||
// primary/secondary real resolvers) currently registered as gateway bypass
|
||||
// targets, so the proxy's own outbound DNS queries aren't captured by the
|
||||
// gateway default-route-equivalent and end up looping back through the
|
||||
// tunnel. Diffed against every update - the initial value in StartTunnel, a
|
||||
// live server-pushed DNS config override, or system DNS detection/
|
||||
// SetSystemDNS - so stale entries are unregistered and new ones protected
|
||||
// while gateway mode is active. Mirrors hpBypassEndpoints.
|
||||
dnsBypassEndpoints map[string]bool
|
||||
dnsBypassMu sync.Mutex
|
||||
|
||||
// primaryTunnelIP is the site tunnel's own address (wgData.TunnelIP), set once
|
||||
// per connect in handleConnect. It's the interface's first/primary address -
|
||||
@@ -119,6 +148,30 @@ func (o *Olm) getPeerManager() *peers.PeerManager {
|
||||
return pm
|
||||
}
|
||||
|
||||
// sharedBindNetwork returns the network family to use for the shared UDP
|
||||
// socket (WireGuard + hole punch traffic). Everywhere but Windows this is a
|
||||
// real dual-stack "udp" wildcard bind, which is what lets the socket reach an
|
||||
// exit node/site over IPv6 when that's the only family the network path has
|
||||
// a route for (the tunnel payload itself is always IPv4 - see
|
||||
// network.ConfigureInterface - but the transport reaching the server/site
|
||||
// endpoint isn't restricted to IPv4).
|
||||
//
|
||||
// On Windows, a dual-stack wildcard bind is unreliable when another VPN's
|
||||
// virtual adapter is also active: Windows can select that adapter's IPv6
|
||||
// address as the implicit local address for a send to an IPv4 destination,
|
||||
// which the OS then rejects outright (WSAEINVAL, "The requested address is
|
||||
// not valid in its context"). RebindSocket already worked around this for
|
||||
// its own bind by using "udp4" explicitly; do the same here for the initial
|
||||
// bind so olm doesn't need a second VPN active to hit it. This does mean
|
||||
// Windows can't reach an IPv6-only exit node/site, same trade-off
|
||||
// RebindSocket already made. See https://github.com/fosrl/olm/issues/134.
|
||||
func sharedBindNetwork() string {
|
||||
if runtime.GOOS == "windows" {
|
||||
return "udp4"
|
||||
}
|
||||
return "udp"
|
||||
}
|
||||
|
||||
// initTunnelInfo creates the shared UDP socket and holepunch manager.
|
||||
// This is used during initial tunnel setup and when switching organizations.
|
||||
func (o *Olm) initTunnelInfo(clientID string) error {
|
||||
@@ -140,7 +193,7 @@ func (o *Olm) initTunnelInfo(clientID string) error {
|
||||
IP: net.IPv4zero,
|
||||
}
|
||||
|
||||
udpConn, err := net.ListenUDP("udp", localAddr)
|
||||
udpConn, err := net.ListenUDP(sharedBindNetwork(), localAddr)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to create UDP socket: %w", err)
|
||||
}
|
||||
@@ -161,6 +214,11 @@ func (o *Olm) initTunnelInfo(clientID string) error {
|
||||
// Create the holepunch manager
|
||||
o.holePunchManager = holepunch.NewManager(sharedBind, clientID, "olm", privateKey.PublicKey().String(), o.tunnelConfig.PublicDNS)
|
||||
|
||||
// A user-disabled hole punch must fully suppress outbound hole punch
|
||||
// packets, not just change what's reported to the server as "relay" (see
|
||||
// SetEnabled's doc comment and https://github.com/fosrl/olm/issues/134).
|
||||
o.holePunchManager.SetEnabled(o.tunnelConfig.Holepunch)
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -221,13 +279,14 @@ func Init(ctx context.Context, config OlmConfig) (*Olm, error) {
|
||||
apiServer.SetAgent(config.Agent)
|
||||
|
||||
newOlm := &Olm{
|
||||
logFile: logFile,
|
||||
olmCtx: ctx,
|
||||
apiServer: apiServer,
|
||||
olmConfig: config,
|
||||
stopPeerSends: make(map[string]func()),
|
||||
stopPeerInits: make(map[string]func()),
|
||||
jitPendingSites: make(map[int]string),
|
||||
logFile: logFile,
|
||||
olmCtx: ctx,
|
||||
apiServer: apiServer,
|
||||
olmConfig: config,
|
||||
stopPeerSends: make(map[string]func()),
|
||||
stopPeerInits: make(map[string]func()),
|
||||
jitPendingSites: make(map[int]string),
|
||||
hpBypassEndpoints: make(map[string]bool),
|
||||
}
|
||||
|
||||
newOlm.registerAPICallbacks()
|
||||
@@ -242,18 +301,22 @@ 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,
|
||||
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,
|
||||
DisableRoutesAndAliasesOnExitNode: req.DisableRoutesAndAliasesOnExitNode,
|
||||
GatewaySiteIds: req.GatewaySiteIds,
|
||||
GatewaySiteResourceId: req.GatewaySiteResourceId,
|
||||
}
|
||||
|
||||
var err error
|
||||
@@ -358,6 +421,16 @@ func (o *Olm) registerAPICallbacks() {
|
||||
|
||||
return nil
|
||||
},
|
||||
// onSelectGateway
|
||||
func(req api.GatewayRequest) error {
|
||||
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 {
|
||||
logger.Info("Received disable-gateway request via API")
|
||||
return o.DisableGateway()
|
||||
},
|
||||
)
|
||||
}
|
||||
|
||||
@@ -465,6 +538,7 @@ func (o *Olm) StartTunnel(config TunnelConfig) {
|
||||
if o.dnsProxy != nil {
|
||||
o.dnsProxy.SetUpstreamDNS(servers)
|
||||
}
|
||||
o.updateDNSBypassEndpoints(servers)
|
||||
} else {
|
||||
logger.Debug("Not updating UpstreamDNS: statically configured to %v", config.UpstreamDNS)
|
||||
}
|
||||
@@ -488,6 +562,13 @@ func (o *Olm) StartTunnel(config TunnelConfig) {
|
||||
if len(o.tunnelConfig.UpstreamDNS) == 0 {
|
||||
o.tunnelConfig.UpstreamDNS = []string{"8.8.8.8:53"}
|
||||
}
|
||||
// Register the startup upstream DNS servers as gateway bypass targets
|
||||
// regardless of how they were determined (statically configured, defaulted
|
||||
// above, or already applied from a dynamic detection callback a moment
|
||||
// ago) - the peer manager doesn't exist yet at this point, so this only
|
||||
// records intent; flushPendingDNSBypassEndpoints (called from
|
||||
// handleConnect) pushes it in once the peer manager is created.
|
||||
o.updateDNSBypassEndpoints(o.tunnelConfig.UpstreamDNS)
|
||||
|
||||
// Reset terminated status when tunnel starts
|
||||
o.apiServer.SetTerminated(false)
|
||||
@@ -577,6 +658,16 @@ 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.
|
||||
o.websocket.RegisterHandler("olm/wg/dns/update", o.handleDNSConfigUpdate)
|
||||
|
||||
o.websocket.RegisterHandler("olm/ping/exitNodes", func(msg websocket.WSMessage) {
|
||||
logger.Debug("Received exit node ping request")
|
||||
|
||||
@@ -728,6 +819,36 @@ func (o *Olm) StartTunnel(config TunnelConfig) {
|
||||
|
||||
logger.Debug("Updated hole punch exit nodes: %v", hpExitNodes)
|
||||
|
||||
// pm can be nil here: this callback runs from establishConnection, as
|
||||
// part of the initial token/auth fetch, which happens well before the
|
||||
// server's "olm/wg/connect" message creates the peer manager in
|
||||
// handleConnect - so on a fresh connect there usually isn't one yet.
|
||||
// o.hpBypassEndpoints is still updated unconditionally so it reflects
|
||||
// the current set regardless; flushPendingHolepunchBypassEndpoints
|
||||
// (called from handleConnect once the peer manager exists) pushes
|
||||
// whatever was recorded here into it.
|
||||
pm := o.getPeerManager()
|
||||
newBypassEndpoints := make(map[string]bool, len(hpExitNodes))
|
||||
for _, node := range hpExitNodes {
|
||||
newBypassEndpoints[net.JoinHostPort(node.Endpoint, strconv.Itoa(int(node.RelayPort)))] = true
|
||||
}
|
||||
|
||||
o.hpBypassMu.Lock()
|
||||
if pm != nil {
|
||||
for hostport := range newBypassEndpoints {
|
||||
if !o.hpBypassEndpoints[hostport] {
|
||||
pm.AddGatewayBypassEndpoint(hostport)
|
||||
}
|
||||
}
|
||||
for hostport := range o.hpBypassEndpoints {
|
||||
if !newBypassEndpoints[hostport] {
|
||||
pm.RemoveGatewayBypassEndpoint(hostport)
|
||||
}
|
||||
}
|
||||
}
|
||||
o.hpBypassEndpoints = newBypassEndpoints
|
||||
o.hpBypassMu.Unlock()
|
||||
|
||||
// Start hole punching using the manager
|
||||
logger.Info("Starting hole punch for %d exit nodes", len(exitNodes))
|
||||
if err := o.holePunchManager.StartMultipleExitNodes(hpExitNodes); err != nil {
|
||||
@@ -748,6 +869,7 @@ func (o *Olm) StartTunnel(config TunnelConfig) {
|
||||
o.apiServer.SetRegistered(false)
|
||||
o.apiServer.ClearOlmError()
|
||||
o.apiServer.ClearPeerStatuses()
|
||||
o.apiServer.SetGatewayStatus(false, 0, nil)
|
||||
network.ClearNetworkSettings()
|
||||
|
||||
o.Close()
|
||||
@@ -831,6 +953,12 @@ func (o *Olm) Close() {
|
||||
o.stopDNSWatchdog()
|
||||
}
|
||||
|
||||
if o.tunnelConfig.SubnetRouter {
|
||||
if err := subnetrouter.Disable(o.tunnelConfig.InterfaceName); err != nil {
|
||||
logger.Error("Failed to disable subnet router: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
if o.holePunchManager != nil {
|
||||
o.holePunchManager.Stop()
|
||||
o.holePunchManager = nil
|
||||
@@ -853,11 +981,20 @@ func (o *Olm) Close() {
|
||||
|
||||
// The WireGuard device and TUN interface are being torn down below, which takes
|
||||
// the exit node peer and its secondary address with them - just clear the
|
||||
// in-memory record so a stale config isn't reused on the next connect.
|
||||
// in-memory record so a stale config isn't reused on the next connect. The
|
||||
// peer manager (and its gateway bypass-route state, including anything
|
||||
// registered for this exit node or for hole-punch nodes below) was already
|
||||
// torn down above, so these resets are purely to avoid stale diffing state
|
||||
// carrying into the next connect, not for route cleanup.
|
||||
o.exitNodeMu.Lock()
|
||||
o.exitNode = nil
|
||||
o.exitNodeResolvedEndpoint = ""
|
||||
o.exitNodeMu.Unlock()
|
||||
|
||||
o.hpBypassMu.Lock()
|
||||
o.hpBypassEndpoints = make(map[string]bool)
|
||||
o.hpBypassMu.Unlock()
|
||||
|
||||
if o.uapiListener != nil {
|
||||
_ = o.uapiListener.Close()
|
||||
o.uapiListener = nil
|
||||
@@ -951,6 +1088,7 @@ func (o *Olm) StopTunnel() error {
|
||||
o.apiServer.SetConnectionStatus(false)
|
||||
o.apiServer.SetRegistered(false)
|
||||
o.apiServer.ClearOlmError()
|
||||
o.apiServer.SetGatewayStatus(false, 0, nil)
|
||||
|
||||
network.ClearNetworkSettings()
|
||||
o.apiServer.ClearPeerStatuses()
|
||||
@@ -1215,7 +1353,7 @@ func (o *Olm) RebindSocket() error {
|
||||
IP: net.IPv4zero,
|
||||
}
|
||||
|
||||
newConn, err = net.ListenUDP("udp4", localAddr)
|
||||
newConn, err = net.ListenUDP(sharedBindNetwork(), localAddr)
|
||||
if err != nil {
|
||||
// If we can't reuse the port, find a new one
|
||||
logger.Warn("Could not rebind to port %d, finding new port: %v", currentPort, err)
|
||||
@@ -1229,8 +1367,7 @@ func (o *Olm) RebindSocket() error {
|
||||
IP: net.IPv4zero,
|
||||
}
|
||||
|
||||
// Use udp4 explicitly to avoid IPv6 dual-stack issues
|
||||
newConn, err = net.ListenUDP("udp4", localAddr)
|
||||
newConn, err = net.ListenUDP(sharedBindNetwork(), localAddr)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to create new UDP socket: %w", err)
|
||||
}
|
||||
|
||||
+85
-27
@@ -228,28 +228,49 @@ func (o *Olm) handleWgPeerRelay(msg websocket.WSMessage) {
|
||||
|
||||
var relayData struct {
|
||||
peers.RelayPeerData
|
||||
ChainId string `json:"chainId"`
|
||||
ChainId string `json:"chainId"`
|
||||
ChainIds []string `json:"chainIds"`
|
||||
}
|
||||
if err := json.Unmarshal(jsonData, &relayData); err != nil {
|
||||
logger.Error("Error unmarshaling relay data: %v", err)
|
||||
return
|
||||
}
|
||||
|
||||
if monitor := pm.GetPeerMonitor(); monitor != nil {
|
||||
monitor.CancelRelaySend(relayData.ChainId)
|
||||
siteIds := relayData.SiteIds
|
||||
relayEndpoints := relayData.RelayEndpoints
|
||||
chainIds := relayData.ChainIds
|
||||
if len(siteIds) == 0 {
|
||||
siteIds = []int{relayData.SiteId}
|
||||
relayEndpoints = []string{relayData.RelayEndpoint}
|
||||
chainIds = []string{relayData.ChainId}
|
||||
}
|
||||
|
||||
primaryRelay, err := util.ResolveDomainUpstream(relayData.RelayEndpoint, o.tunnelConfig.PublicDNS)
|
||||
monitor := pm.GetPeerMonitor()
|
||||
|
||||
if err != nil {
|
||||
logger.Error("Failed to resolve primary relay endpoint: %v", err)
|
||||
return
|
||||
for i, siteId := range siteIds {
|
||||
var relayEndpoint, chainId string
|
||||
if i < len(relayEndpoints) {
|
||||
relayEndpoint = relayEndpoints[i]
|
||||
}
|
||||
if i < len(chainIds) {
|
||||
chainId = chainIds[i]
|
||||
}
|
||||
|
||||
if monitor != nil {
|
||||
monitor.CancelRelaySend(chainId)
|
||||
}
|
||||
|
||||
primaryRelay, err := util.ResolveDomainUpstream(relayEndpoint, o.tunnelConfig.PublicDNS)
|
||||
if err != nil {
|
||||
logger.Error("Failed to resolve primary relay endpoint for site %d: %v", siteId, err)
|
||||
continue
|
||||
}
|
||||
|
||||
// Update HTTP server to mark this peer as using relay
|
||||
o.apiServer.UpdatePeerRelayStatus(siteId, relayEndpoint, true)
|
||||
|
||||
pm.RelayPeer(siteId, primaryRelay, relayData.RelayPort)
|
||||
}
|
||||
|
||||
// Update HTTP server to mark this peer as using relay
|
||||
o.apiServer.UpdatePeerRelayStatus(relayData.SiteId, relayData.RelayEndpoint, true)
|
||||
|
||||
pm.RelayPeer(relayData.SiteId, primaryRelay, relayData.RelayPort)
|
||||
}
|
||||
|
||||
func (o *Olm) handleWgPeerUnrelay(msg websocket.WSMessage) {
|
||||
@@ -270,27 +291,48 @@ func (o *Olm) handleWgPeerUnrelay(msg websocket.WSMessage) {
|
||||
|
||||
var relayData struct {
|
||||
peers.UnRelayPeerData
|
||||
ChainId string `json:"chainId"`
|
||||
ChainId string `json:"chainId"`
|
||||
ChainIds []string `json:"chainIds"`
|
||||
}
|
||||
if err := json.Unmarshal(jsonData, &relayData); err != nil {
|
||||
logger.Error("Error unmarshaling relay data: %v", err)
|
||||
return
|
||||
}
|
||||
|
||||
if monitor := pm.GetPeerMonitor(); monitor != nil {
|
||||
monitor.CancelRelaySend(relayData.ChainId)
|
||||
siteIds := relayData.SiteIds
|
||||
endpoints := relayData.Endpoints
|
||||
chainIds := relayData.ChainIds
|
||||
if len(siteIds) == 0 {
|
||||
siteIds = []int{relayData.SiteId}
|
||||
endpoints = []string{relayData.Endpoint}
|
||||
chainIds = []string{relayData.ChainId}
|
||||
}
|
||||
|
||||
primaryRelay, err := util.ResolveDomainUpstream(relayData.Endpoint, o.tunnelConfig.PublicDNS)
|
||||
monitor := pm.GetPeerMonitor()
|
||||
|
||||
if err != nil {
|
||||
logger.Warn("Failed to resolve primary relay endpoint: %v", err)
|
||||
for i, siteId := range siteIds {
|
||||
var endpoint, chainId string
|
||||
if i < len(endpoints) {
|
||||
endpoint = endpoints[i]
|
||||
}
|
||||
if i < len(chainIds) {
|
||||
chainId = chainIds[i]
|
||||
}
|
||||
|
||||
if monitor != nil {
|
||||
monitor.CancelRelaySend(chainId)
|
||||
}
|
||||
|
||||
primaryRelay, err := util.ResolveDomainUpstream(endpoint, o.tunnelConfig.PublicDNS)
|
||||
if err != nil {
|
||||
logger.Warn("Failed to resolve primary relay endpoint for site %d: %v", siteId, err)
|
||||
}
|
||||
|
||||
// Update HTTP server to mark this peer as using relay
|
||||
o.apiServer.UpdatePeerRelayStatus(siteId, endpoint, false)
|
||||
|
||||
pm.UnRelayPeer(siteId, primaryRelay)
|
||||
}
|
||||
|
||||
// Update HTTP server to mark this peer as using relay
|
||||
o.apiServer.UpdatePeerRelayStatus(relayData.SiteId, relayData.Endpoint, false)
|
||||
|
||||
pm.UnRelayPeer(relayData.SiteId, primaryRelay)
|
||||
}
|
||||
|
||||
// handleWgPeerLocal handles the server's acknowledgement of an "olm/wg/local" message.
|
||||
@@ -314,15 +356,23 @@ func (o *Olm) handleWgPeerLocal(msg websocket.WSMessage) {
|
||||
|
||||
var localData struct {
|
||||
peers.LocalPeerAckData
|
||||
ChainId string `json:"chainId"`
|
||||
ChainId string `json:"chainId"`
|
||||
ChainIds []string `json:"chainIds"`
|
||||
}
|
||||
if err := json.Unmarshal(jsonData, &localData); err != nil {
|
||||
logger.Error("Error unmarshaling local ack data: %v", err)
|
||||
return
|
||||
}
|
||||
|
||||
chainIds := localData.ChainIds
|
||||
if len(chainIds) == 0 {
|
||||
chainIds = []string{localData.ChainId}
|
||||
}
|
||||
|
||||
if monitor := pm.GetPeerMonitor(); monitor != nil {
|
||||
monitor.CancelLocalSend(localData.ChainId)
|
||||
for _, chainId := range chainIds {
|
||||
monitor.CancelLocalSend(chainId)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -346,15 +396,23 @@ func (o *Olm) handleWgPeerUnlocal(msg websocket.WSMessage) {
|
||||
|
||||
var localData struct {
|
||||
peers.LocalPeerAckData
|
||||
ChainId string `json:"chainId"`
|
||||
ChainId string `json:"chainId"`
|
||||
ChainIds []string `json:"chainIds"`
|
||||
}
|
||||
if err := json.Unmarshal(jsonData, &localData); err != nil {
|
||||
logger.Error("Error unmarshaling unlocal ack data: %v", err)
|
||||
return
|
||||
}
|
||||
|
||||
chainIds := localData.ChainIds
|
||||
if len(chainIds) == 0 {
|
||||
chainIds = []string{localData.ChainId}
|
||||
}
|
||||
|
||||
if monitor := pm.GetPeerMonitor(); monitor != nil {
|
||||
monitor.CancelLocalSend(localData.ChainId)
|
||||
for _, chainId := range chainIds {
|
||||
monitor.CancelLocalSend(chainId)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -11,6 +11,22 @@ type WgData struct {
|
||||
TunnelIP string `json:"tunnelIP"`
|
||||
UtilitySubnet string `json:"utilitySubnet"` // this is for things like the DNS server, and alias addresses
|
||||
ExitNode *ExitNodeConfig `json:"exitNode,omitempty"`
|
||||
DNSConfig *DNSConfigUpdate `json:"dnsConfig,omitempty"`
|
||||
}
|
||||
|
||||
// DNSConfigUpdate describes a server-driven override of the olm client's DNS
|
||||
// configuration - the same settings that can otherwise only be set locally
|
||||
// (see TunnelConfig's UpstreamDNS/OverrideDNS/TunnelDNS/MatchDomains). It
|
||||
// arrives on the initial "olm/wg/connect" message and can also be sent later
|
||||
// via "olm/wg/dns/update" to change the running config without reconnecting.
|
||||
// Every field is optional/nil-able: an omitted field leaves the client's
|
||||
// current value (local config, or whatever a previous update set) unchanged,
|
||||
// while a present field always overrides it.
|
||||
type DNSConfigUpdate struct {
|
||||
UpstreamDNS []string `json:"upstreamDns,omitempty"`
|
||||
OverrideDNS *bool `json:"overrideDns,omitempty"`
|
||||
TunnelDNS *bool `json:"tunnelDns,omitempty"`
|
||||
MatchDomains []string `json:"matchDomains,omitempty"`
|
||||
}
|
||||
|
||||
// ExitNodeConfig describes an exit node the olm client can connect to for
|
||||
@@ -158,4 +174,42 @@ type TunnelConfig struct {
|
||||
// false, preserving the routing behavior from before this option was
|
||||
// introduced.
|
||||
PreferLocalRoutes bool
|
||||
|
||||
// SubnetRouter, when enabled, lets this client forward traffic from its
|
||||
// local network out over the tunnel: forwarded packets are NATed to the
|
||||
// client's own tunnel IP before being encrypted, since the server side
|
||||
// authorizes traffic by the client's tunnel identity, not by whatever
|
||||
// LAN address it originally arrived with. Linux only. Defaults to false.
|
||||
SubnetRouter bool
|
||||
|
||||
// DisableRoutesAndAliasesOnExitNode, when enabled, makes the exit node
|
||||
// take precedence over individual resources: for as long as an exit node
|
||||
// is connected, olm removes routes to the host's routing table for site
|
||||
// resources (server IPs, remote subnets) and their alias DNS records, and
|
||||
// restores them the moment the exit node disconnects. If the tunnel
|
||||
// starts with an exit node already selected, these routes/aliases are
|
||||
// never added in the first place. The exit node's own routes and aliases
|
||||
// are unaffected, and it can connect/disconnect at any time - via the
|
||||
// initial connect, a server push (olm/wg/exitnode/connect|disconnect), or
|
||||
// the local API - all of which converge through the same connect/
|
||||
// disconnect path (see PeerManager.SetExitNode/ClearExitNode).
|
||||
// WireGuard AllowedIPs are still configured throughout, so the tunnel can
|
||||
// still be used by anything that reaches it without the OS routing table
|
||||
// (e.g. a file descriptor/netstack consumer). Gateway routes (the
|
||||
// default-route-equivalent and its endpoint bypass routes) are unaffected
|
||||
// and are always installed. Defaults to false.
|
||||
DisableRoutesAndAliasesOnExitNode bool
|
||||
|
||||
// GatewaySiteIds, when non-empty, designates these site IDs as gateway
|
||||
// (full-tunnel/default-route) candidates from the moment the tunnel
|
||||
// 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,210 @@
|
||||
package peers
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/fosrl/newt/network"
|
||||
"golang.zx2c4.com/wireguard/conn"
|
||||
"golang.zx2c4.com/wireguard/device"
|
||||
"golang.zx2c4.com/wireguard/tun/tuntest"
|
||||
"golang.zx2c4.com/wireguard/wgctrl/wgtypes"
|
||||
)
|
||||
|
||||
// newTestDevice returns a real *device.Device backed by an in-memory channel
|
||||
// TUN (golang.zx2c4.com/wireguard/tun/tuntest) and a standard UDP bind - no
|
||||
// OS TUN interface or elevated privileges required, so this is safe to run
|
||||
// as a normal unit test. Used to exercise the real AddAllowedIP/
|
||||
// RemoveAllowedIP/ConfigurePeer IPC calls that suppressResourceRoutesLocked/
|
||||
// restoreResourceRoutesLocked make, which a hand-rolled PeerManager with
|
||||
// device left nil (see newGatewayTestManager/newExitNodeTestManager) can't
|
||||
// safely call.
|
||||
func newTestDevice(t *testing.T) (*device.Device, wgtypes.Key) {
|
||||
t.Helper()
|
||||
privateKey, err := wgtypes.GeneratePrivateKey()
|
||||
if err != nil {
|
||||
t.Fatalf("GeneratePrivateKey: %v", err)
|
||||
}
|
||||
|
||||
dev := device.NewDevice(tuntest.NewChannelTUN().TUN(), conn.NewDefaultBind(), device.NewLogger(device.LogLevelError, "test: "))
|
||||
t.Cleanup(dev.Close)
|
||||
|
||||
if err := dev.IpcSet("private_key=" + hexKey(privateKey) + "\n"); err != nil {
|
||||
t.Fatalf("failed to set device private key: %v", err)
|
||||
}
|
||||
|
||||
return dev, privateKey
|
||||
}
|
||||
|
||||
func hexKey(k wgtypes.Key) string {
|
||||
b := [32]byte(k)
|
||||
const hextable = "0123456789abcdef"
|
||||
out := make([]byte, 64)
|
||||
for i, c := range b {
|
||||
out[i*2] = hextable[c>>4]
|
||||
out[i*2+1] = hextable[c&0x0f]
|
||||
}
|
||||
return string(out)
|
||||
}
|
||||
|
||||
// TestSuppressRestoreResourceRoutesStripsWireGuardAllowedIPs is the
|
||||
// counterpart to TestSetGatewaySuppressesResourceRoutesInNetworkSettings,
|
||||
// but for WireGuard's own AllowedIPs rather than the OS routing table/
|
||||
// NetworkSettings: it verifies suppressResourceRoutesLocked/
|
||||
// restoreResourceRoutesLocked actually add/remove the site peer's resource
|
||||
// CIDRs (remote subnet, alias) from WireGuard itself - not just the system
|
||||
// route - so that nothing reaching the tunnel interface directly (e.g. a
|
||||
// mobile netstack/FD consumer bypassing the OS route table) can still reach
|
||||
// a suppressed resource via WireGuard's own crypto-key routing. The server
|
||||
// IP allowed-ip entry must survive throughout, since it's what keeps the
|
||||
// site's own control/monitoring traffic (pings, handshakes) alive while
|
||||
// suppressed.
|
||||
func TestSuppressRestoreResourceRoutesStripsWireGuardAllowedIPs(t *testing.T) {
|
||||
network.ClearNetworkSettings()
|
||||
defer network.ClearNetworkSettings()
|
||||
|
||||
dev, privKey := newTestDevice(t)
|
||||
peerKey, err := wgtypes.GeneratePrivateKey()
|
||||
if err != nil {
|
||||
t.Fatalf("GeneratePrivateKey (peer): %v", err)
|
||||
}
|
||||
peerPubKey := peerKey.PublicKey()
|
||||
|
||||
pm := &PeerManager{
|
||||
device: dev,
|
||||
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),
|
||||
privateKey: privKey,
|
||||
interfaceName: "fake0",
|
||||
localIP: "100.90.128.8",
|
||||
disableRoutesAndAliasesOnExitNode: true,
|
||||
}
|
||||
|
||||
site := SiteConfig{
|
||||
SiteId: 5,
|
||||
PublicKey: peerPubKey.String(),
|
||||
Endpoint: "127.0.0.1:1", // never dialed - IpcSet doesn't connect
|
||||
ServerIP: "100.90.128.1/20",
|
||||
RemoteSubnets: []string{"172.18.21.32/24"},
|
||||
}
|
||||
|
||||
// Directly claim ownership and push the peer's full AllowedIPs, mirroring
|
||||
// what AddPeer does, without needing the rest of AddPeer's machinery
|
||||
// (DNS proxy, peer monitor, holepunch test) that isn't relevant here.
|
||||
pm.peers[5] = site
|
||||
pm.claimAllowedIP(5, "172.18.21.32/24")
|
||||
wgConfig := site
|
||||
wgConfig.AllowedIps = []string{"172.18.21.32/24"}
|
||||
if err := ConfigurePeer(dev, wgConfig, privKey, false, 0, nil); err != nil {
|
||||
t.Fatalf("seed ConfigurePeer: %v", err)
|
||||
}
|
||||
|
||||
before, err := dev.IpcGet()
|
||||
if err != nil {
|
||||
t.Fatalf("IpcGet before suppress: %v", err)
|
||||
}
|
||||
if !strings.Contains(before, "172.18.21.0/24") {
|
||||
t.Fatalf("test setup broken: seeded allowed_ip missing before suppress:\n%s", before)
|
||||
}
|
||||
if !strings.Contains(before, "100.90.128.1/32") {
|
||||
t.Fatalf("test setup broken: server IP allowed_ip missing before suppress:\n%s", before)
|
||||
}
|
||||
|
||||
pm.mu.Lock()
|
||||
pm.suppressResourceRoutesLocked()
|
||||
pm.mu.Unlock()
|
||||
|
||||
after, err := dev.IpcGet()
|
||||
if err != nil {
|
||||
t.Fatalf("IpcGet after suppress: %v", err)
|
||||
}
|
||||
if strings.Contains(after, "172.18.21.0/24") {
|
||||
t.Fatalf("resource allowed_ip still present in WireGuard after suppression (exit node active):\n%s", after)
|
||||
}
|
||||
if !strings.Contains(after, "100.90.128.1/32") {
|
||||
t.Fatalf("server IP allowed_ip must survive suppression (needed for site monitoring/liveness):\n%s", after)
|
||||
}
|
||||
|
||||
pm.mu.Lock()
|
||||
pm.restoreResourceRoutesLocked()
|
||||
pm.mu.Unlock()
|
||||
|
||||
restored, err := dev.IpcGet()
|
||||
if err != nil {
|
||||
t.Fatalf("IpcGet after restore: %v", err)
|
||||
}
|
||||
if !strings.Contains(restored, "172.18.21.0/24") {
|
||||
t.Fatalf("resource allowed_ip not restored in WireGuard after restore:\n%s", restored)
|
||||
}
|
||||
if !strings.Contains(restored, "100.90.128.1/32") {
|
||||
t.Fatalf("server IP allowed_ip missing after restore:\n%s", restored)
|
||||
}
|
||||
}
|
||||
|
||||
// TestShouldPushAllowedIPLockedAndOwnershipFilteringWhileSuppressed covers
|
||||
// the "tunnel starts with an exit node already active" case at the unit
|
||||
// level: getOwnedAllowedIPs/getWireGuardAllowedIPs (what AddPeer's inline
|
||||
// ownership computation, addAllowedIp, the route optimizer, etc. all defer
|
||||
// to - see shouldPushAllowedIPLocked) must exclude a resource CIDR while
|
||||
// suppressed even though the underlying claim/ownership is registered
|
||||
// normally, and must always keep the gateway CIDR.
|
||||
func TestShouldPushAllowedIPLockedAndOwnershipFilteringWhileSuppressed(t *testing.T) {
|
||||
pm := &PeerManager{
|
||||
peers: make(map[int]SiteConfig),
|
||||
allowedIPOwners: make(map[string]int),
|
||||
allowedIPClaims: make(map[string]map[int]bool),
|
||||
disableRoutesAndAliasesOnExitNode: true,
|
||||
}
|
||||
pm.peers[5] = SiteConfig{SiteId: 5, ServerIP: "100.90.128.1/20"}
|
||||
|
||||
pm.claimAllowedIP(5, "172.18.21.0/24")
|
||||
pm.claimAllowedIP(5, gatewayCIDR)
|
||||
|
||||
if !pm.shouldPushAllowedIPLocked(gatewayCIDR) {
|
||||
t.Fatalf("gateway CIDR must always be pushable")
|
||||
}
|
||||
if !pm.shouldPushAllowedIPLocked("172.18.21.0/24") {
|
||||
t.Fatalf("resource CIDR must be pushable while not suppressed")
|
||||
}
|
||||
owned := pm.getOwnedAllowedIPs(5)
|
||||
if len(owned) != 2 {
|
||||
t.Fatalf("expected both claimed CIDRs owned before suppression, got %v", owned)
|
||||
}
|
||||
|
||||
pm.resourceRoutesSuppressed = true
|
||||
|
||||
if pm.shouldPushAllowedIPLocked("172.18.21.0/24") {
|
||||
t.Fatalf("resource CIDR must not be pushable while suppressed")
|
||||
}
|
||||
if !pm.shouldPushAllowedIPLocked(gatewayCIDR) {
|
||||
t.Fatalf("gateway CIDR must still be pushable while suppressed")
|
||||
}
|
||||
|
||||
owned = pm.getOwnedAllowedIPs(5)
|
||||
if len(owned) != 1 || owned[0] != gatewayCIDR {
|
||||
t.Fatalf("expected only the gateway CIDR owned while suppressed, got %v", owned)
|
||||
}
|
||||
|
||||
wgIPs := pm.getWireGuardAllowedIPs(5)
|
||||
want := map[string]bool{"100.90.128.1/32": true, gatewayCIDR: true}
|
||||
if len(wgIPs) != len(want) {
|
||||
t.Fatalf("expected server IP + gateway CIDR only while suppressed, got %v", wgIPs)
|
||||
}
|
||||
for _, ip := range wgIPs {
|
||||
if !want[ip] {
|
||||
t.Fatalf("unexpected allowed IP %q while suppressed: %v", ip, wgIPs)
|
||||
}
|
||||
}
|
||||
|
||||
pm.resourceRoutesSuppressed = false
|
||||
owned = pm.getOwnedAllowedIPs(5)
|
||||
if len(owned) != 2 {
|
||||
t.Fatalf("expected both claimed CIDRs owned again after restore, got %v", owned)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,106 @@
|
||||
package peers
|
||||
|
||||
import "testing"
|
||||
|
||||
// newExitNodeTestManager returns a PeerManager for exercising the
|
||||
// DisableRoutesAndAliasesOnExitNode state machine (resourceRoutesSuppressed,
|
||||
// exitNodeActive, gatewayActive) without touching the OS routing table, a
|
||||
// WireGuard device, or a DNS proxy - safe as long as pm.peers stays empty,
|
||||
// same constraint as newGatewayTestManager in gateway_test.go. Unlike
|
||||
// NewPeerManager, disableRoutesAndAliasesOnExitNode is the only state seeded
|
||||
// here; exitNodeActive/gatewayActive/resourceRoutesSuppressed always start
|
||||
// false, matching NewPeerManager's real (purely reactive, no pre-seeding)
|
||||
// behavior.
|
||||
func newExitNodeTestManager(disableRoutesAndAliasesOnExitNode bool) *PeerManager {
|
||||
return &PeerManager{
|
||||
peers: make(map[int]SiteConfig),
|
||||
allowedIPOwners: make(map[string]int),
|
||||
allowedIPClaims: make(map[string]map[int]bool),
|
||||
disableRoutesAndAliasesOnExitNode: disableRoutesAndAliasesOnExitNode,
|
||||
}
|
||||
}
|
||||
|
||||
func TestSetClearExitNodeSuppressesResourceRoutesWhenEnabled(t *testing.T) {
|
||||
pm := newExitNodeTestManager(true)
|
||||
|
||||
if pm.resourceRoutesSuppressed {
|
||||
t.Fatalf("must start unsuppressed")
|
||||
}
|
||||
|
||||
pm.SetExitNode("100.64.0.1", "100.64.0.2")
|
||||
if !pm.exitNodeActive {
|
||||
t.Fatalf("SetExitNode must mark the exit node active")
|
||||
}
|
||||
if !pm.resourceRoutesSuppressed {
|
||||
t.Fatalf("SetExitNode must suppress resource routes when the feature is enabled")
|
||||
}
|
||||
|
||||
pm.ClearExitNode()
|
||||
if pm.exitNodeActive {
|
||||
t.Fatalf("ClearExitNode must mark the exit node inactive")
|
||||
}
|
||||
if pm.resourceRoutesSuppressed {
|
||||
t.Fatalf("ClearExitNode must restore resource routes")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSetClearExitNodeNoopWhenDisabled(t *testing.T) {
|
||||
pm := newExitNodeTestManager(false)
|
||||
|
||||
pm.SetExitNode("100.64.0.1", "100.64.0.2")
|
||||
if pm.resourceRoutesSuppressed {
|
||||
t.Fatalf("resource routes must never be suppressed when the feature is disabled")
|
||||
}
|
||||
|
||||
pm.ClearExitNode()
|
||||
if pm.resourceRoutesSuppressed {
|
||||
t.Fatalf("resource routes must stay unsuppressed when the feature is disabled")
|
||||
}
|
||||
}
|
||||
|
||||
// TestClearExitNodeKeepsSuppressedWhileGatewayActive covers the two
|
||||
// independent "exit node" signals (see exitNodeOrGatewayActiveLocked):
|
||||
// disconnecting the ExitNodeConfig WireGuard peer must not restore resource
|
||||
// routes if gateway/full-tunnel mode - what client apps and the CLI actually
|
||||
// call "select exit node" - is still active. gatewayActive is set directly
|
||||
// here (bypassing SetGateway, which touches the OS routing table) purely to
|
||||
// exercise ClearExitNode's OR-check.
|
||||
func TestClearExitNodeKeepsSuppressedWhileGatewayActive(t *testing.T) {
|
||||
pm := newExitNodeTestManager(true)
|
||||
|
||||
pm.SetExitNode("100.64.0.1", "100.64.0.2")
|
||||
if !pm.resourceRoutesSuppressed {
|
||||
t.Fatalf("SetExitNode must suppress resource routes")
|
||||
}
|
||||
|
||||
pm.mu.Lock()
|
||||
pm.gatewayActive = true
|
||||
pm.mu.Unlock()
|
||||
|
||||
pm.ClearExitNode()
|
||||
if pm.exitNodeActive {
|
||||
t.Fatalf("ClearExitNode must mark the exit node inactive regardless of gateway state")
|
||||
}
|
||||
if !pm.resourceRoutesSuppressed {
|
||||
t.Fatalf("resource routes must stay suppressed while gateway mode is still active")
|
||||
}
|
||||
}
|
||||
|
||||
func TestExitNodeOrGatewayActiveLocked(t *testing.T) {
|
||||
pm := newExitNodeTestManager(true)
|
||||
|
||||
if pm.exitNodeOrGatewayActiveLocked() {
|
||||
t.Fatalf("neither signal is active yet")
|
||||
}
|
||||
|
||||
pm.exitNodeActive = true
|
||||
if !pm.exitNodeOrGatewayActiveLocked() {
|
||||
t.Fatalf("exit node signal alone must count as active")
|
||||
}
|
||||
pm.exitNodeActive = false
|
||||
|
||||
pm.gatewayActive = true
|
||||
if !pm.exitNodeOrGatewayActiveLocked() {
|
||||
t.Fatalf("gateway signal alone must count as active")
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
+878
-60
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,212 @@
|
||||
package monitor
|
||||
|
||||
import (
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/fosrl/newt/logger"
|
||||
"github.com/fosrl/olm/websocket"
|
||||
)
|
||||
|
||||
// Batch tuning: each new item added resets a batchDebounce trailing window, so a burst of
|
||||
// decisions that trickle in over a few hundred ms (e.g. 100 sites each independently
|
||||
// finishing their own rapid holepunch test after a network-wide blip) still collapses into
|
||||
// one flush instead of splitting across several. batchMaxWait bounds how long a steady
|
||||
// trickle of arrivals can keep postponing that flush, so an item is never held back more
|
||||
// than batchMaxWait past its own arrival. Once sent, items the server hasn't acknowledged
|
||||
// yet are retried together every batchInterval, up to batchMaxAttempts times, mirroring the
|
||||
// cadence the previous per-site SendMessageInterval senders used.
|
||||
const (
|
||||
batchDebounce = 200 * time.Millisecond
|
||||
batchMaxWait = 750 * time.Millisecond
|
||||
batchInterval = 2 * time.Second
|
||||
batchMaxAttempts = 10
|
||||
)
|
||||
|
||||
// batchSendItem is a single queued site decision awaiting acknowledgement.
|
||||
type batchSendItem struct {
|
||||
siteID int
|
||||
endpoint string // only populated for the local batcher; ignored otherwise
|
||||
attempts int
|
||||
}
|
||||
|
||||
// batchSender coalesces per-site "olm/wg/*" notifications of a single message type into
|
||||
// periodic batched websocket messages, retrying items the server hasn't acknowledged (via
|
||||
// cancel) until they succeed or exhaust batchMaxAttempts. When the connected server doesn't
|
||||
// understand the batched wire format (see websocket.Client.SupportsBatchedSiteMessages),
|
||||
// it falls back to sending each item as its own message in the pre-batching singular form.
|
||||
type batchSender struct {
|
||||
mu sync.Mutex
|
||||
items map[string]*batchSendItem // chainId -> item
|
||||
messageType string
|
||||
wsClient *websocket.Client
|
||||
debounceTimer *time.Timer
|
||||
burstStarted time.Time
|
||||
runOnce sync.Once
|
||||
stopChan chan struct{}
|
||||
}
|
||||
|
||||
func newBatchSender(wsClient *websocket.Client, messageType string) *batchSender {
|
||||
return &batchSender{
|
||||
items: make(map[string]*batchSendItem),
|
||||
messageType: messageType,
|
||||
wsClient: wsClient,
|
||||
stopChan: make(chan struct{}),
|
||||
}
|
||||
}
|
||||
|
||||
// add queues siteID for the next flush and returns the chainId that identifies it for
|
||||
// cancel/ack purposes. endpoint is only meaningful for the local batcher.
|
||||
func (b *batchSender) add(siteID int, endpoint string) string {
|
||||
chainId := generateChainId()
|
||||
|
||||
b.mu.Lock()
|
||||
b.items[chainId] = &batchSendItem{siteID: siteID, endpoint: endpoint}
|
||||
|
||||
now := time.Now()
|
||||
if b.debounceTimer == nil {
|
||||
// First item of a new burst: start the trailing window.
|
||||
b.burstStarted = now
|
||||
b.debounceTimer = time.AfterFunc(batchDebounce, b.flush)
|
||||
} else {
|
||||
// Extend the window for this new arrival, capped so a burst that keeps trickling
|
||||
// in still flushes within batchMaxWait of its first item.
|
||||
delay := batchDebounce
|
||||
if remaining := batchMaxWait - now.Sub(b.burstStarted); remaining < delay {
|
||||
if remaining < 0 {
|
||||
remaining = 0
|
||||
}
|
||||
delay = remaining
|
||||
}
|
||||
b.debounceTimer.Reset(delay)
|
||||
}
|
||||
b.mu.Unlock()
|
||||
|
||||
b.runOnce.Do(func() { go b.run() })
|
||||
|
||||
return chainId
|
||||
}
|
||||
|
||||
// cancel removes chainId from the pending set, e.g. once the server has acknowledged it.
|
||||
// Returns true if it was pending.
|
||||
func (b *batchSender) cancel(chainId string) bool {
|
||||
b.mu.Lock()
|
||||
defer b.mu.Unlock()
|
||||
if _, ok := b.items[chainId]; !ok {
|
||||
return false
|
||||
}
|
||||
delete(b.items, chainId)
|
||||
return true
|
||||
}
|
||||
|
||||
// cancelAll clears all pending items, e.g. on shutdown.
|
||||
func (b *batchSender) cancelAll() {
|
||||
b.mu.Lock()
|
||||
b.items = make(map[string]*batchSendItem)
|
||||
b.mu.Unlock()
|
||||
}
|
||||
|
||||
type readySend struct {
|
||||
chainId string
|
||||
siteID int
|
||||
endpoint string
|
||||
}
|
||||
|
||||
// flush sends everything currently pending, dropping items that have exhausted
|
||||
// batchMaxAttempts. If the server supports the batched wire format it goes out as a single
|
||||
// message; otherwise each item is sent individually in the pre-batching singular form.
|
||||
func (b *batchSender) flush() {
|
||||
b.mu.Lock()
|
||||
b.debounceTimer = nil
|
||||
|
||||
ready := make([]readySend, 0, len(b.items))
|
||||
for chainId, item := range b.items {
|
||||
item.attempts++
|
||||
if item.attempts > batchMaxAttempts {
|
||||
logger.Warn("olm: giving up on %s for site %d (chain %s) after %d attempts", b.messageType, item.siteID, chainId, batchMaxAttempts)
|
||||
delete(b.items, chainId)
|
||||
continue
|
||||
}
|
||||
ready = append(ready, readySend{chainId: chainId, siteID: item.siteID, endpoint: item.endpoint})
|
||||
}
|
||||
wsClient := b.wsClient
|
||||
messageType := b.messageType
|
||||
b.mu.Unlock()
|
||||
|
||||
if len(ready) == 0 || wsClient == nil {
|
||||
return
|
||||
}
|
||||
|
||||
if !wsClient.SupportsBatchedSiteMessages() {
|
||||
b.sendIndividually(wsClient, messageType, ready)
|
||||
return
|
||||
}
|
||||
|
||||
siteIds := make([]int, len(ready))
|
||||
chainIds := make([]string, len(ready))
|
||||
endpoints := make([]string, len(ready))
|
||||
hasEndpoints := false
|
||||
for i, item := range ready {
|
||||
siteIds[i] = item.siteID
|
||||
chainIds[i] = item.chainId
|
||||
endpoints[i] = item.endpoint
|
||||
if item.endpoint != "" {
|
||||
hasEndpoints = true
|
||||
}
|
||||
}
|
||||
|
||||
data := map[string]interface{}{
|
||||
"siteIds": siteIds,
|
||||
"chainIds": chainIds,
|
||||
}
|
||||
if hasEndpoints {
|
||||
data["endpoints"] = endpoints
|
||||
}
|
||||
|
||||
if err := wsClient.SendMessage(messageType, data); err != nil {
|
||||
logger.Error("olm: failed to send batched %s: %v", messageType, err)
|
||||
} else {
|
||||
logger.Info("olm: sent batched %s for %d site(s)", messageType, len(siteIds))
|
||||
}
|
||||
}
|
||||
|
||||
// sendIndividually sends each item as its own message in the singular siteId/chainId form,
|
||||
// for servers that predate the batched wire format.
|
||||
func (b *batchSender) sendIndividually(wsClient *websocket.Client, messageType string, ready []readySend) {
|
||||
for _, item := range ready {
|
||||
data := map[string]interface{}{
|
||||
"siteId": item.siteID,
|
||||
"chainId": item.chainId,
|
||||
}
|
||||
if item.endpoint != "" {
|
||||
data["endpoint"] = item.endpoint
|
||||
}
|
||||
if err := wsClient.SendMessage(messageType, data); err != nil {
|
||||
logger.Error("olm: failed to send %s for site %d: %v", messageType, item.siteID, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// run periodically resends any items still awaiting acknowledgement.
|
||||
func (b *batchSender) run() {
|
||||
ticker := time.NewTicker(batchInterval)
|
||||
defer ticker.Stop()
|
||||
for {
|
||||
select {
|
||||
case <-b.stopChan:
|
||||
return
|
||||
case <-ticker.C:
|
||||
b.flush()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// close stops the background retry loop. The sender must not be used again afterwards.
|
||||
func (b *batchSender) close() {
|
||||
select {
|
||||
case <-b.stopChan:
|
||||
// already closed
|
||||
default:
|
||||
close(b.stopChan)
|
||||
}
|
||||
}
|
||||
+43
-97
@@ -40,9 +40,10 @@ type PeerMonitor struct {
|
||||
wsClient *websocket.Client
|
||||
publicDNS []string
|
||||
|
||||
// Relay sender tracking
|
||||
relaySends map[string]func()
|
||||
relaySendMu sync.Mutex
|
||||
// Relay sender tracking — batched per message type so many relay/unrelay decisions
|
||||
// arriving together (e.g. after a network-wide blip) go out as one websocket message.
|
||||
relayBatch *batchSender
|
||||
unrelayBatch *batchSender
|
||||
|
||||
// Netstack fields
|
||||
middleDev *middleDevice.MiddleDevice
|
||||
@@ -82,9 +83,9 @@ type PeerMonitor struct {
|
||||
localSwitchCallback func(siteId int, endpoint string) // invoked when a local endpoint becomes active
|
||||
localFallbackCallback func(siteId int) // invoked when we fall back from a local endpoint
|
||||
|
||||
// Local connection sender tracking, keyed by chainId (informational messages only)
|
||||
localSends map[string]func()
|
||||
localSendMu sync.Mutex
|
||||
// Local connection sender tracking (informational messages only), batched the same way.
|
||||
localBatch *batchSender
|
||||
unlocalBatch *batchSender
|
||||
|
||||
// Exponential backoff fields for holepunch monitor
|
||||
defaultHolepunchMinInterval time.Duration // Minimum interval (initial)
|
||||
@@ -149,14 +150,16 @@ func NewPeerMonitor(wsClient *websocket.Client, middleDev *middleDevice.MiddleDe
|
||||
holepunchEndpoints: make(map[int]string),
|
||||
holepunchStatus: make(map[int]bool),
|
||||
relayedPeers: make(map[int]bool),
|
||||
relaySends: make(map[string]func()),
|
||||
relayBatch: newBatchSender(wsClient, "olm/wg/relay"),
|
||||
unrelayBatch: newBatchSender(wsClient, "olm/wg/unrelay"),
|
||||
holepunchMaxAttempts: 3, // Trigger relay after 3 consecutive failures
|
||||
holepunchFailures: make(map[int]int),
|
||||
localEndpoints: make(map[int][]string),
|
||||
localActiveEndpoint: make(map[int]string),
|
||||
localFailures: make(map[int]int),
|
||||
localTestTimeout: 300 * time.Millisecond, // local network round trips should be fast
|
||||
localSends: make(map[string]func()),
|
||||
localBatch: newBatchSender(wsClient, "olm/wg/local"),
|
||||
unlocalBatch: newBatchSender(wsClient, "olm/wg/unlocal"),
|
||||
// Rapid initial test settings: complete within ~1.5 seconds
|
||||
rapidTestInterval: 200 * time.Millisecond, // 200ms between attempts
|
||||
rapidTestTimeout: 400 * time.Millisecond, // 400ms timeout per attempt
|
||||
@@ -628,23 +631,16 @@ func (pm *PeerMonitor) handleConnectionStatusChange(siteID int, status Connectio
|
||||
}
|
||||
}
|
||||
|
||||
// sendRelay sends a relay message to the server with retry, keyed by chainId
|
||||
// sendRelay queues a relay message for the server, batched with any other relay decisions
|
||||
// made in the same short window, with retry keyed by chainId.
|
||||
func (pm *PeerMonitor) sendRelay(siteID int) error {
|
||||
if pm.wsClient == nil {
|
||||
return fmt.Errorf("websocket client is nil")
|
||||
}
|
||||
|
||||
chainId := generateChainId()
|
||||
stopFunc, _ := pm.wsClient.SendMessageInterval("olm/wg/relay", map[string]interface{}{
|
||||
"siteId": siteID,
|
||||
"chainId": chainId,
|
||||
}, 2*time.Second, 10)
|
||||
chainId := pm.relayBatch.add(siteID, "")
|
||||
|
||||
pm.relaySendMu.Lock()
|
||||
pm.relaySends[chainId] = stopFunc
|
||||
pm.relaySendMu.Unlock()
|
||||
|
||||
logger.Info("Sent relay message for site %d (chain %s)", siteID, chainId)
|
||||
logger.Info("Queued relay message for site %d (chain %s)", siteID, chainId)
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -654,23 +650,16 @@ func (pm *PeerMonitor) RequestRelay(siteID int) error {
|
||||
return pm.sendRelay(siteID)
|
||||
}
|
||||
|
||||
// sendUnRelay sends an unrelay message to the server with retry, keyed by chainId
|
||||
// sendUnRelay queues an unrelay message for the server, batched with any other unrelay
|
||||
// decisions made in the same short window, with retry keyed by chainId.
|
||||
func (pm *PeerMonitor) sendUnRelay(siteID int) error {
|
||||
if pm.wsClient == nil {
|
||||
return fmt.Errorf("websocket client is nil")
|
||||
}
|
||||
|
||||
chainId := generateChainId()
|
||||
stopFunc, _ := pm.wsClient.SendMessageInterval("olm/wg/unrelay", map[string]interface{}{
|
||||
"siteId": siteID,
|
||||
"chainId": chainId,
|
||||
}, 2*time.Second, 10)
|
||||
chainId := pm.unrelayBatch.add(siteID, "")
|
||||
|
||||
pm.relaySendMu.Lock()
|
||||
pm.relaySends[chainId] = stopFunc
|
||||
pm.relaySendMu.Unlock()
|
||||
|
||||
logger.Info("Sent unrelay message for site %d (chain %s)", siteID, chainId)
|
||||
logger.Info("Queued unrelay message for site %d (chain %s)", siteID, chainId)
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -683,18 +672,9 @@ func (pm *PeerMonitor) sendLocal(siteID int, endpoint string) {
|
||||
return
|
||||
}
|
||||
|
||||
chainId := generateChainId()
|
||||
stopFunc, _ := pm.wsClient.SendMessageInterval("olm/wg/local", map[string]interface{}{
|
||||
"siteId": siteID,
|
||||
"endpoint": endpoint,
|
||||
"chainId": chainId,
|
||||
}, 2*time.Second, 10)
|
||||
chainId := pm.localBatch.add(siteID, endpoint)
|
||||
|
||||
pm.localSendMu.Lock()
|
||||
pm.localSends[chainId] = stopFunc
|
||||
pm.localSendMu.Unlock()
|
||||
|
||||
logger.Info("Sent local-connection message for site %d (%s, chain %s)", siteID, endpoint, chainId)
|
||||
logger.Info("Queued local-connection message for site %d (%s, chain %s)", siteID, endpoint, chainId)
|
||||
}
|
||||
|
||||
// sendUnLocal notifies the server that this peer fell back from its local network endpoint,
|
||||
@@ -704,65 +684,39 @@ func (pm *PeerMonitor) sendUnLocal(siteID int) {
|
||||
return
|
||||
}
|
||||
|
||||
chainId := generateChainId()
|
||||
stopFunc, _ := pm.wsClient.SendMessageInterval("olm/wg/unlocal", map[string]interface{}{
|
||||
"siteId": siteID,
|
||||
"chainId": chainId,
|
||||
}, 2*time.Second, 10)
|
||||
chainId := pm.unlocalBatch.add(siteID, "")
|
||||
|
||||
pm.localSendMu.Lock()
|
||||
pm.localSends[chainId] = stopFunc
|
||||
pm.localSendMu.Unlock()
|
||||
|
||||
logger.Info("Sent unlocal-connection message for site %d (chain %s)", siteID, chainId)
|
||||
logger.Info("Queued unlocal-connection message for site %d (chain %s)", siteID, chainId)
|
||||
}
|
||||
|
||||
// CancelLocalSend stops the interval sender for the given chainId, if one exists.
|
||||
// If chainId is empty, all active local-connection senders are stopped.
|
||||
// CancelLocalSend removes chainId from the pending local/unlocal batches, if present.
|
||||
// If chainId is empty, all pending local-connection items are cleared.
|
||||
func (pm *PeerMonitor) CancelLocalSend(chainId string) {
|
||||
pm.localSendMu.Lock()
|
||||
defer pm.localSendMu.Unlock()
|
||||
|
||||
if chainId == "" {
|
||||
for id, stop := range pm.localSends {
|
||||
if stop != nil {
|
||||
stop()
|
||||
}
|
||||
delete(pm.localSends, id)
|
||||
}
|
||||
pm.localBatch.cancelAll()
|
||||
pm.unlocalBatch.cancelAll()
|
||||
logger.Info("Cancelled all local-connection senders")
|
||||
return
|
||||
}
|
||||
|
||||
if stop, ok := pm.localSends[chainId]; ok {
|
||||
stop()
|
||||
delete(pm.localSends, chainId)
|
||||
if pm.localBatch.cancel(chainId) || pm.unlocalBatch.cancel(chainId) {
|
||||
logger.Info("Cancelled local-connection sender for chain %s", chainId)
|
||||
} else {
|
||||
logger.Warn("CancelLocalSend: no active sender for chain %s", chainId)
|
||||
}
|
||||
}
|
||||
|
||||
// CancelRelaySend stops the interval sender for the given chainId, if one exists.
|
||||
// If chainId is empty, all active relay senders are stopped.
|
||||
// CancelRelaySend removes chainId from the pending relay/unrelay batches, if present.
|
||||
// If chainId is empty, all pending relay items are cleared.
|
||||
func (pm *PeerMonitor) CancelRelaySend(chainId string) {
|
||||
pm.relaySendMu.Lock()
|
||||
defer pm.relaySendMu.Unlock()
|
||||
|
||||
if chainId == "" {
|
||||
for id, stop := range pm.relaySends {
|
||||
if stop != nil {
|
||||
stop()
|
||||
}
|
||||
delete(pm.relaySends, id)
|
||||
}
|
||||
pm.relayBatch.cancelAll()
|
||||
pm.unrelayBatch.cancelAll()
|
||||
logger.Info("Cancelled all relay senders")
|
||||
return
|
||||
}
|
||||
|
||||
if stop, ok := pm.relaySends[chainId]; ok {
|
||||
stop()
|
||||
delete(pm.relaySends, chainId)
|
||||
if pm.relayBatch.cancel(chainId) || pm.unrelayBatch.cancel(chainId) {
|
||||
logger.Info("Cancelled relay sender for chain %s", chainId)
|
||||
} else {
|
||||
logger.Warn("CancelRelaySend: no active sender for chain %s", chainId)
|
||||
@@ -1177,25 +1131,17 @@ func (pm *PeerMonitor) Close() {
|
||||
}
|
||||
pm.exitNodeMu.Unlock()
|
||||
|
||||
// Stop all pending relay senders
|
||||
pm.relaySendMu.Lock()
|
||||
for chainId, stop := range pm.relaySends {
|
||||
if stop != nil {
|
||||
stop()
|
||||
}
|
||||
delete(pm.relaySends, chainId)
|
||||
}
|
||||
pm.relaySendMu.Unlock()
|
||||
// Stop all pending relay/unrelay batch senders
|
||||
pm.relayBatch.cancelAll()
|
||||
pm.relayBatch.close()
|
||||
pm.unrelayBatch.cancelAll()
|
||||
pm.unrelayBatch.close()
|
||||
|
||||
// Stop all pending local-connection senders
|
||||
pm.localSendMu.Lock()
|
||||
for chainId, stop := range pm.localSends {
|
||||
if stop != nil {
|
||||
stop()
|
||||
}
|
||||
delete(pm.localSends, chainId)
|
||||
}
|
||||
pm.localSendMu.Unlock()
|
||||
// Stop all pending local-connection batch senders
|
||||
pm.localBatch.cancelAll()
|
||||
pm.localBatch.close()
|
||||
pm.unlocalBatch.cancelAll()
|
||||
pm.unlocalBatch.close()
|
||||
|
||||
pm.mutex.Lock()
|
||||
defer pm.mutex.Unlock()
|
||||
|
||||
+15
-4
@@ -35,15 +35,26 @@ type PeerRemove struct {
|
||||
SiteId int `json:"siteId"`
|
||||
}
|
||||
|
||||
// RelayPeerData represents the server's acknowledgement of an "olm/wg/relay" message. The
|
||||
// server replies in kind: SiteId/RelayEndpoint for a single-site request, or the parallel
|
||||
// SiteIds/RelayEndpoints arrays for a batched one (RelayPort is shared by the whole batch).
|
||||
type RelayPeerData struct {
|
||||
SiteId int `json:"siteId"`
|
||||
RelayEndpoint string `json:"relayEndpoint"`
|
||||
SiteId int `json:"siteId,omitempty"`
|
||||
RelayEndpoint string `json:"relayEndpoint,omitempty"`
|
||||
RelayPort uint16 `json:"relayPort"`
|
||||
|
||||
SiteIds []int `json:"siteIds,omitempty"`
|
||||
RelayEndpoints []string `json:"relayEndpoints,omitempty"`
|
||||
}
|
||||
|
||||
// UnRelayPeerData represents the server's acknowledgement of an "olm/wg/unrelay" message,
|
||||
// in either its single-site or batched (parallel-array) form. See RelayPeerData.
|
||||
type UnRelayPeerData struct {
|
||||
SiteId int `json:"siteId"`
|
||||
Endpoint string `json:"endpoint"`
|
||||
SiteId int `json:"siteId,omitempty"`
|
||||
Endpoint string `json:"endpoint,omitempty"`
|
||||
|
||||
SiteIds []int `json:"siteIds,omitempty"`
|
||||
Endpoints []string `json:"endpoints,omitempty"`
|
||||
}
|
||||
|
||||
// LocalPeerAckData represents the server's acknowledgement of an "olm/wg/local" or
|
||||
|
||||
@@ -0,0 +1,222 @@
|
||||
//go:build linux
|
||||
|
||||
// Package subnetrouter lets this client forward LAN traffic out over its own
|
||||
// WireGuard tunnel, source-NAT'd to the tunnel's own IP. Pangolin's
|
||||
// server-side routing/ACLs are keyed on the client's tunnel IP as its
|
||||
// identity, so traffic merely forwarded from the LAN (which arrives with the
|
||||
// LAN device's own source address) would not be recognized - it must be
|
||||
// rewritten to look like it came from this client before it goes out over
|
||||
// the tunnel, the same way a NAT router masquerades LAN traffic behind its
|
||||
// WAN IP.
|
||||
package subnetrouter
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"net/netip"
|
||||
"os"
|
||||
"strings"
|
||||
"sync"
|
||||
|
||||
"github.com/fosrl/newt/logger"
|
||||
"github.com/google/nftables"
|
||||
"github.com/google/nftables/expr"
|
||||
"golang.org/x/sys/unix"
|
||||
)
|
||||
|
||||
const (
|
||||
tableName = "olm_subnet_router"
|
||||
ipForwardSys = "/proc/sys/net/ipv4/ip_forward"
|
||||
)
|
||||
|
||||
// ipForwardMu guards the two package-level fields below, which record
|
||||
// whether Enable had to flip ip_forward on itself, so Disable only ever
|
||||
// restores a value it actually changed - mirroring dns/override's
|
||||
// instance-free save/restore convention.
|
||||
var (
|
||||
ipForwardMu sync.Mutex
|
||||
weEnabledForward bool
|
||||
)
|
||||
|
||||
// Enable turns this host into a subnet router: it enables IPv4 forwarding
|
||||
// (if not already on) and installs an nftables table that SNATs anything
|
||||
// leaving interfaceName whose source isn't already tunnelIP, and accepts
|
||||
// forwarding to/from interfaceName so a default-deny FORWARD policy
|
||||
// elsewhere on the host doesn't drop it.
|
||||
//
|
||||
// It is idempotent: any table left behind by a previous run (e.g. after a
|
||||
// crash) is torn down first, so repeated Enable/Disable cycles across
|
||||
// reconnects never conflict with stale state.
|
||||
func Enable(interfaceName string, tunnelIP netip.Addr) error {
|
||||
if !tunnelIP.Is4() {
|
||||
return fmt.Errorf("subnet router requires an IPv4 tunnel address, got %v", tunnelIP)
|
||||
}
|
||||
|
||||
// Best-effort cleanup of anything left over from a previous run.
|
||||
if err := Disable(interfaceName); err != nil {
|
||||
logger.Debug("subnetrouter: pre-enable cleanup: %v", err)
|
||||
}
|
||||
|
||||
if err := enableIPForward(); err != nil {
|
||||
return fmt.Errorf("failed to enable IPv4 forwarding: %w", err)
|
||||
}
|
||||
|
||||
conn := &nftables.Conn{}
|
||||
|
||||
table := conn.AddTable(&nftables.Table{
|
||||
Family: nftables.TableFamilyIPv4,
|
||||
Name: tableName,
|
||||
})
|
||||
|
||||
postrouting := conn.AddChain(&nftables.Chain{
|
||||
Name: "postrouting",
|
||||
Table: table,
|
||||
Type: nftables.ChainTypeNAT,
|
||||
Hooknum: nftables.ChainHookPostrouting,
|
||||
Priority: nftables.ChainPriorityNATSource,
|
||||
})
|
||||
|
||||
addr := tunnelIP.As4()
|
||||
conn.AddRule(&nftables.Rule{
|
||||
Table: table,
|
||||
Chain: postrouting,
|
||||
Exprs: []expr.Any{
|
||||
// oifname == interfaceName
|
||||
&expr.Meta{Key: expr.MetaKeyOIFNAME, Register: 1},
|
||||
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: ifname(interfaceName)},
|
||||
// ip saddr != tunnelIP
|
||||
&expr.Payload{
|
||||
DestRegister: 1,
|
||||
Base: expr.PayloadBaseNetworkHeader,
|
||||
Offset: 12, // IPv4 source address offset
|
||||
Len: 4,
|
||||
},
|
||||
&expr.Cmp{Op: expr.CmpOpNeq, Register: 1, Data: addr[:]},
|
||||
// snat to tunnelIP
|
||||
&expr.Immediate{Register: 1, Data: addr[:]},
|
||||
&expr.NAT{
|
||||
Type: expr.NATTypeSourceNAT,
|
||||
Family: unix.NFPROTO_IPV4,
|
||||
RegAddrMin: 1,
|
||||
RegAddrMax: 1,
|
||||
},
|
||||
},
|
||||
})
|
||||
|
||||
forward := conn.AddChain(&nftables.Chain{
|
||||
Name: "forward",
|
||||
Table: table,
|
||||
Type: nftables.ChainTypeFilter,
|
||||
Hooknum: nftables.ChainHookForward,
|
||||
Priority: nftables.ChainPriorityFilter,
|
||||
})
|
||||
|
||||
for _, key := range []expr.MetaKey{expr.MetaKeyIIFNAME, expr.MetaKeyOIFNAME} {
|
||||
conn.AddRule(&nftables.Rule{
|
||||
Table: table,
|
||||
Chain: forward,
|
||||
Exprs: []expr.Any{
|
||||
&expr.Meta{Key: key, Register: 1},
|
||||
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: ifname(interfaceName)},
|
||||
&expr.Verdict{Kind: expr.VerdictAccept},
|
||||
},
|
||||
})
|
||||
}
|
||||
|
||||
if err := conn.Flush(); err != nil {
|
||||
// Roll back the forwarding sysctl change too, so a failed Enable
|
||||
// doesn't leave the host with forwarding on and no NAT rules.
|
||||
_ = disableIPForwardIfWeEnabledIt()
|
||||
return fmt.Errorf("failed to apply nftables rules: %w", err)
|
||||
}
|
||||
|
||||
logger.Debug("subnetrouter: enabled on %s (snat to %s)", interfaceName, tunnelIP)
|
||||
return nil
|
||||
}
|
||||
|
||||
// Disable removes the nftables table added by Enable (a no-op if it doesn't
|
||||
// exist) and restores ip_forward to whatever it was before Enable, but only
|
||||
// if Enable is what changed it.
|
||||
func Disable(interfaceName string) error {
|
||||
conn := &nftables.Conn{}
|
||||
|
||||
tables, err := conn.ListTables()
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to list nftables tables: %w", err)
|
||||
}
|
||||
|
||||
var found bool
|
||||
for _, t := range tables {
|
||||
if t.Name == tableName && t.Family == nftables.TableFamilyIPv4 {
|
||||
conn.DelTable(t)
|
||||
found = true
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
var flushErr error
|
||||
if found {
|
||||
flushErr = conn.Flush()
|
||||
}
|
||||
|
||||
forwardErr := disableIPForwardIfWeEnabledIt()
|
||||
|
||||
if flushErr != nil {
|
||||
return fmt.Errorf("failed to remove nftables table: %w", flushErr)
|
||||
}
|
||||
return forwardErr
|
||||
}
|
||||
|
||||
// enableIPForward turns on IPv4 forwarding if it isn't already on, recording
|
||||
// whether this call is the one that changed it.
|
||||
func enableIPForward() error {
|
||||
ipForwardMu.Lock()
|
||||
defer ipForwardMu.Unlock()
|
||||
|
||||
current, err := readIPForward()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if current {
|
||||
weEnabledForward = false
|
||||
return nil
|
||||
}
|
||||
|
||||
if err := os.WriteFile(ipForwardSys, []byte("1\n"), 0644); err != nil {
|
||||
return err
|
||||
}
|
||||
weEnabledForward = true
|
||||
return nil
|
||||
}
|
||||
|
||||
// disableIPForwardIfWeEnabledIt restores ip_forward to 0, but only if a
|
||||
// prior enableIPForward call is what turned it on.
|
||||
func disableIPForwardIfWeEnabledIt() error {
|
||||
ipForwardMu.Lock()
|
||||
defer ipForwardMu.Unlock()
|
||||
|
||||
if !weEnabledForward {
|
||||
return nil
|
||||
}
|
||||
|
||||
if err := os.WriteFile(ipForwardSys, []byte("0\n"), 0644); err != nil {
|
||||
return err
|
||||
}
|
||||
weEnabledForward = false
|
||||
return nil
|
||||
}
|
||||
|
||||
func readIPForward() (bool, error) {
|
||||
data, err := os.ReadFile(ipForwardSys)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
return strings.TrimSpace(string(data)) == "1", nil
|
||||
}
|
||||
|
||||
// ifname encodes an interface name the way nftables expects it: NUL-padded
|
||||
// to IFNAMSIZ (16) bytes.
|
||||
func ifname(name string) []byte {
|
||||
b := make([]byte, 16)
|
||||
copy(b, name)
|
||||
return b
|
||||
}
|
||||
@@ -0,0 +1,21 @@
|
||||
//go:build !linux
|
||||
|
||||
package subnetrouter
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"net/netip"
|
||||
)
|
||||
|
||||
// Enable always fails on non-Linux platforms: there is no nftables/netfilter
|
||||
// to install SNAT rules into. Callers should log this as a warning, not
|
||||
// treat it as fatal.
|
||||
func Enable(interfaceName string, tunnelIP netip.Addr) error {
|
||||
return fmt.Errorf("subnet router is only supported on Linux")
|
||||
}
|
||||
|
||||
// Disable is a no-op on non-Linux platforms, since Enable never succeeds
|
||||
// there and so never leaves anything to clean up.
|
||||
func Disable(interfaceName string) error {
|
||||
return nil
|
||||
}
|
||||
+16
-14
@@ -104,7 +104,8 @@ type Client struct {
|
||||
configVersionMux sync.RWMutex
|
||||
token string // Cached authentication token
|
||||
exitNodes []ExitNode // Cached exit nodes from token response
|
||||
tokenMux sync.RWMutex // Protects token and exitNodes
|
||||
serverVersion string // Server version from the last token response
|
||||
tokenMux sync.RWMutex // Protects token, exitNodes and serverVersion
|
||||
forceNewToken bool // Flag to force fetching a new token on next connection
|
||||
processingMessage bool // Flag to track if a message is currently being processed
|
||||
processingMux sync.RWMutex // Protects processingMessage
|
||||
@@ -423,11 +424,11 @@ func (c *Client) RegisterHandler(messageType string, handler MessageHandler) {
|
||||
c.handlers[messageType] = handler
|
||||
}
|
||||
|
||||
func (c *Client) getToken() (string, []ExitNode, error) {
|
||||
func (c *Client) getToken() (string, []ExitNode, string, error) {
|
||||
// Parse the base URL to ensure we have the correct hostname
|
||||
baseURL, err := url.Parse(c.baseURL)
|
||||
if err != nil {
|
||||
return "", nil, fmt.Errorf("failed to parse base URL: %w", err)
|
||||
return "", nil, "", fmt.Errorf("failed to parse base URL: %w", err)
|
||||
}
|
||||
|
||||
// Ensure we have the base URL without trailing slashes
|
||||
@@ -439,7 +440,7 @@ func (c *Client) getToken() (string, []ExitNode, error) {
|
||||
if c.tlsConfig.ClientCertFile != "" || c.tlsConfig.ClientKeyFile != "" || len(c.tlsConfig.CAFiles) > 0 || c.tlsConfig.PKCS12File != "" {
|
||||
tlsConfig, err = c.setupTLS()
|
||||
if err != nil {
|
||||
return "", nil, fmt.Errorf("failed to setup TLS configuration: %w", err)
|
||||
return "", nil, "", fmt.Errorf("failed to setup TLS configuration: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -461,7 +462,7 @@ func (c *Client) getToken() (string, []ExitNode, error) {
|
||||
jsonData, err := json.Marshal(tokenData)
|
||||
|
||||
if err != nil {
|
||||
return "", nil, fmt.Errorf("failed to marshal token request data: %w", err)
|
||||
return "", nil, "", fmt.Errorf("failed to marshal token request data: %w", err)
|
||||
}
|
||||
|
||||
// Create a new request
|
||||
@@ -471,7 +472,7 @@ func (c *Client) getToken() (string, []ExitNode, error) {
|
||||
bytes.NewBuffer(jsonData),
|
||||
)
|
||||
if err != nil {
|
||||
return "", nil, fmt.Errorf("failed to create request: %w", err)
|
||||
return "", nil, "", fmt.Errorf("failed to create request: %w", err)
|
||||
}
|
||||
|
||||
// Set headers
|
||||
@@ -490,7 +491,7 @@ func (c *Client) getToken() (string, []ExitNode, error) {
|
||||
}
|
||||
resp, err := client.Do(req)
|
||||
if err != nil {
|
||||
return "", nil, fmt.Errorf("failed to request new token: %w", err)
|
||||
return "", nil, "", fmt.Errorf("failed to request new token: %w", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
@@ -500,33 +501,33 @@ func (c *Client) getToken() (string, []ExitNode, error) {
|
||||
|
||||
// Return AuthError for 401/403 status codes
|
||||
if resp.StatusCode == http.StatusUnauthorized || resp.StatusCode == http.StatusForbidden {
|
||||
return "", nil, &AuthError{
|
||||
return "", nil, "", &AuthError{
|
||||
StatusCode: resp.StatusCode,
|
||||
Message: string(body),
|
||||
}
|
||||
}
|
||||
|
||||
// For other errors (5xx, network issues, etc.), return regular error
|
||||
return "", nil, fmt.Errorf("failed to get token with status code: %d, body: %s", resp.StatusCode, string(body))
|
||||
return "", nil, "", fmt.Errorf("failed to get token with status code: %d, body: %s", resp.StatusCode, string(body))
|
||||
}
|
||||
|
||||
var tokenResp TokenResponse
|
||||
if err := json.NewDecoder(resp.Body).Decode(&tokenResp); err != nil {
|
||||
logger.Error("websocket: Failed to decode token response.")
|
||||
return "", nil, fmt.Errorf("failed to decode token response: %w", err)
|
||||
return "", nil, "", fmt.Errorf("failed to decode token response: %w", err)
|
||||
}
|
||||
|
||||
if !tokenResp.Success {
|
||||
return "", nil, fmt.Errorf("failed to get token: %s", tokenResp.Message)
|
||||
return "", nil, "", fmt.Errorf("failed to get token: %s", tokenResp.Message)
|
||||
}
|
||||
|
||||
if tokenResp.Data.Token == "" {
|
||||
return "", nil, fmt.Errorf("received empty token from server")
|
||||
return "", nil, "", fmt.Errorf("received empty token from server")
|
||||
}
|
||||
|
||||
logger.Debug("websocket: Received token: %s", tokenResp.Data.Token)
|
||||
|
||||
return tokenResp.Data.Token, tokenResp.Data.ExitNodes, nil
|
||||
return tokenResp.Data.Token, tokenResp.Data.ExitNodes, tokenResp.Data.ServerVersion, nil
|
||||
}
|
||||
|
||||
func (c *Client) connectWithRetry() {
|
||||
@@ -564,13 +565,14 @@ func (c *Client) establishConnection() error {
|
||||
c.tokenMux.Lock()
|
||||
needNewToken := c.token == "" || c.forceNewToken
|
||||
if needNewToken {
|
||||
token, exitNodes, err := c.getToken()
|
||||
token, exitNodes, serverVersion, err := c.getToken()
|
||||
if err != nil {
|
||||
c.tokenMux.Unlock()
|
||||
return fmt.Errorf("failed to get token: %w", err)
|
||||
}
|
||||
c.token = token
|
||||
c.exitNodes = exitNodes
|
||||
c.serverVersion = serverVersion
|
||||
c.forceNewToken = false
|
||||
|
||||
if c.onTokenUpdate != nil {
|
||||
|
||||
@@ -0,0 +1,75 @@
|
||||
package websocket
|
||||
|
||||
import (
|
||||
"strconv"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// minBatchedSiteMessagesVersion is the first pangolin server version that understands the
|
||||
// batched siteIds/chainIds (and endpoints/relayEndpoints) form of the olm/wg/relay,
|
||||
// olm/wg/unrelay, olm/wg/local and olm/wg/unlocal messages. Servers older than this only
|
||||
// understand the singular siteId/chainId form, so olm must fall back to sending those
|
||||
// messages one at a time.
|
||||
const minBatchedSiteMessagesVersion = "1.24.0"
|
||||
|
||||
// SupportsBatchedSiteMessages reports whether the server we're connected to is new enough
|
||||
// to understand batched relay/unrelay/local/unlocal messages, based on the serverVersion
|
||||
// returned in the last token exchange.
|
||||
func (c *Client) SupportsBatchedSiteMessages() bool {
|
||||
return supportsBatchedSiteMessages(c.ServerVersion())
|
||||
}
|
||||
|
||||
// ServerVersion returns the version reported by the server during the last token exchange,
|
||||
// or "" if unknown (e.g. no successful connection has completed yet).
|
||||
func (c *Client) ServerVersion() string {
|
||||
c.tokenMux.RLock()
|
||||
defer c.tokenMux.RUnlock()
|
||||
return c.serverVersion
|
||||
}
|
||||
|
||||
func supportsBatchedSiteMessages(serverVersion string) bool {
|
||||
if serverVersion == "" {
|
||||
// Unknown server version: assume it predates batching rather than risk sending a
|
||||
// format an old server can't parse.
|
||||
return false
|
||||
}
|
||||
return compareVersions(baseVersion(serverVersion), minBatchedSiteMessagesVersion) >= 0
|
||||
}
|
||||
|
||||
// baseVersion strips any "-suffix" build metadata (e.g. "1.24.0-s.5" -> "1.24.0").
|
||||
func baseVersion(v string) string {
|
||||
if i := strings.IndexByte(v, '-'); i >= 0 {
|
||||
return v[:i]
|
||||
}
|
||||
return v
|
||||
}
|
||||
|
||||
// compareVersions compares two dotted numeric version strings (e.g. "1.24.0"), returning
|
||||
// -1 if a < b, 0 if equal, or 1 if a > b. Missing or non-numeric components are treated as
|
||||
// 0, so a partial or malformed version degrades gracefully instead of panicking.
|
||||
func compareVersions(a, b string) int {
|
||||
as := strings.Split(a, ".")
|
||||
bs := strings.Split(b, ".")
|
||||
|
||||
max := len(as)
|
||||
if len(bs) > max {
|
||||
max = len(bs)
|
||||
}
|
||||
|
||||
for i := 0; i < max; i++ {
|
||||
var an, bn int
|
||||
if i < len(as) {
|
||||
an, _ = strconv.Atoi(as[i])
|
||||
}
|
||||
if i < len(bs) {
|
||||
bn, _ = strconv.Atoi(bs[i])
|
||||
}
|
||||
if an != bn {
|
||||
if an < bn {
|
||||
return -1
|
||||
}
|
||||
return 1
|
||||
}
|
||||
}
|
||||
return 0
|
||||
}
|
||||
@@ -0,0 +1,47 @@
|
||||
package websocket
|
||||
|
||||
import "testing"
|
||||
|
||||
func TestSupportsBatchedSiteMessages(t *testing.T) {
|
||||
cases := []struct {
|
||||
version string
|
||||
want bool
|
||||
}{
|
||||
{"", false},
|
||||
{"1.23.0", false},
|
||||
{"1.23.9", false},
|
||||
{"1.24.0", true},
|
||||
{"1.24.1", true},
|
||||
{"1.25.0", true},
|
||||
{"2.0.0", true},
|
||||
{"1.24.0-s.5", true}, // cloud build of a supported base version
|
||||
{"1.23.0-s.99", false}, // cloud build of an unsupported base version
|
||||
{"not-a-version", false},
|
||||
}
|
||||
|
||||
for _, c := range cases {
|
||||
if got := supportsBatchedSiteMessages(c.version); got != c.want {
|
||||
t.Errorf("supportsBatchedSiteMessages(%q) = %v, want %v", c.version, got, c.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestCompareVersions(t *testing.T) {
|
||||
cases := []struct {
|
||||
a, b string
|
||||
want int
|
||||
}{
|
||||
{"1.24.0", "1.24.0", 0},
|
||||
{"1.23.0", "1.24.0", -1},
|
||||
{"1.24.0", "1.23.0", 1},
|
||||
{"1.24", "1.24.0", 0},
|
||||
{"1.24.1", "1.24", 1},
|
||||
{"2.0.0", "1.99.99", 1},
|
||||
}
|
||||
|
||||
for _, c := range cases {
|
||||
if got := compareVersions(c.a, c.b); got != c.want {
|
||||
t.Errorf("compareVersions(%q, %q) = %d, want %d", c.a, c.b, got, c.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user