Compare commits

..
92 Commits
Author SHA1 Message Date
Owen SchwartzandGitHub 4f54e27b22 Merge pull request #133 from fosrl/dev
1.18.2
2026-08-03 15:50:07 -04:00
Owen 3542fee459 Dont rely on newt 2026-08-03 15:47:21 -04:00
Owen 60c7703f07 Merge branch 'main' into dev 2026-08-03 12:14:52 -04:00
Owen d8714df81f filter out magic packets in the middle device to prevent flapping 2026-08-03 12:14:43 -04:00
Owen 6faa3a0273 Attempt to fix disconnecting 2026-07-31 10:22:48 -04:00
Owen fbc4fb2827 Add a write deadline on the websocket 2026-07-31 09:45:31 -04:00
Owen SchwartzandGitHub 96e1d0f98c Merge pull request #131 from fosrl/dev
Prevent flapping on route optimizer
2026-07-29 17:51:59 -04:00
Owen 29663cdb81 Prevent flapping on route optimizer 2026-07-29 17:13:59 -04:00
Owen SchwartzandGitHub 23597835d8 Merge pull request #130 from fosrl/dev
1.8.0
2026-07-18 21:27:18 -04:00
Owen 8ef735185f Update go mod 2026-07-18 21:24:56 -04:00
Owen 99d249db2d add PreferLocalRoutes option 2026-07-18 17:28:56 -04:00
Owen 9c3c04c728 Match domains dns rename 2026-07-17 17:57:26 -04:00
Owen 39aaaafa17 Dont remove non controlled routes 2026-07-17 17:39:24 -04:00
Owen 3d88b321e8 Rapid test again when we fail local 2026-07-17 15:55:31 -04:00
Owen e1214e21cd Reflect local status in the api 2026-07-17 13:58:20 -04:00
Owen a2f0f64c2f Cancel local send from chainId 2026-07-16 11:29:55 -04:00
Owen 0e06fb0152 Rapid test the local endpoints as well 2026-07-16 11:11:02 -04:00
Owen 7929ce9cf9 Test local connections and choose those first 2026-07-16 10:57:42 -04:00
Owen 9513433c07 Add match domains config 2026-07-15 15:51:13 -04:00
Owen SchwartzandGitHub 1319354914 Merge pull request #128 from fosrl/dev
Pull system dns
2026-07-07 20:46:49 -04:00
Owen b39be4f5b0 Update newt 2026-07-07 20:45:08 -04:00
Owen 929f183ed0 Improve linux to use networkmanager correctly 2026-07-07 20:38:06 -04:00
Owen 0cf5eb2ad0 Adjust darwin to use sysctl 2026-07-07 17:08:35 -04:00
Owen 1920f52699 Also check network manager on linux 2026-07-07 16:41:01 -04:00
Owen 7b07745650 Fix windows only pulling dhcp 2026-07-07 16:23:25 -04:00
Owen 1f301db892 Mod tidy 2026-07-07 15:49:51 -04:00
Owen 4c600bab15 Add some logging 2026-07-06 17:51:49 -04:00
Owen f356aed39b Add ios stub 2026-07-06 16:56:37 -04:00
Owen 7a88fac395 Test the dns server first then fall back 2026-07-06 16:08:10 -04:00
Owen efc012b43d Groundwork for monitoring and updating the dns 2026-06-25 20:21:25 -04:00
Owen SchwartzandGitHub 83678dedcb Merge pull request #127 from fosrl/dev
Update newt
2026-06-25 06:49:55 -07:00
Owen e8e004e5e1 Update newt 2026-06-25 09:49:30 -04:00
Owen SchwartzandGitHub b08994e019 Merge pull request #126 from fosrl/dev
Handle remove on a per site level correctly
2026-06-25 06:38:55 -07:00
Owen f87b804051 Merge branch 'main' into dev 2026-06-25 09:37:40 -04:00
Owen f02045d936 Aliases are by site 2026-06-24 18:36:07 -04:00
Owen 55377980b8 Dont remove alias if used by other peers 2026-06-24 16:42:29 -04:00
Owen SchwartzandGitHub 403e14327d Merge pull request #125 from fosrl/dev
1.6.0
2026-06-10 11:19:52 -07:00
Owen 31f6d699a8 Add dns watchdog and override fixer 2026-06-09 21:17:13 -07:00
Owen 6e41eb99ab Merge branch 'main' into dev 2026-06-09 12:03:59 -07:00
Owen 12c1918e26 Include chainId for warning the user with an error 2026-06-09 12:03:50 -07:00
Owen 0d6503c03f Add comment 2026-06-09 12:03:50 -07:00
miloschwartz b789c9af75 update issue template 2026-05-25 21:41:23 -07:00
Owen SchwartzandGitHub 8cc854e313 Merge pull request #121 from fosrl/dev
Support JIT sync message
2026-05-13 15:19:06 -07:00
Owen 2cd16e24a1 Handle jit peers with alias 2026-05-13 15:09:45 -07:00
Owen SchwartzandGitHub b6b35b4581 Merge pull request #120 from LaurenceJJones/investigate/106-handleWgPeerUpdate-alias-merge
fix(peers): merge Aliases in handleWgPeerUpdate
2026-05-11 10:12:56 -07:00
LaurenceandCursor 3f64336b0b peers: merge Aliases in handleWgPeerUpdate
WireGuard update-peer messages now copy Aliases from the payload into
the merged SiteConfig so UpdatePeer can refresh alias DNS records.

Fixes fosrl/olm#106

Co-authored-by: Cursor <cursoragent@cursor.com>
2026-05-11 10:45:08 +01:00
Owen SchwartzandGitHub e94c13f601 Merge pull request #119 from fosrl/dev
Guard add peer with is registered`
2026-05-06 22:08:22 -07:00
Owen fff806c53d Guard add peer with is registered` 2026-05-06 22:06:03 -07:00
Owen SchwartzandGitHub 84d7e8d926 Merge pull request #118 from fosrl/dev
Update newt
2026-04-27 20:13:48 -07:00
Owen fde70dd15b Update newt 2026-04-27 20:13:20 -07:00
Owen 7bf6da1729 Update newt 2026-04-22 12:22:39 -07:00
Owen SchwartzandGitHub 9f68f171ba Merge pull request #117 from fosrl/dev
1.5.0
2026-04-22 12:13:09 -07:00
Owen 7a8f4ab049 Merge branch 'main' into dev 2026-04-22 12:09:50 -07:00
Owen 334ea156b6 Merge branch 'private-site-ha' into dev 2026-04-21 15:08:11 -07:00
Owen aa838fec61 Mention the cli 2026-04-19 15:49:08 -07:00
Owen 6eaf8c1475 Basic route selector working 2026-04-13 18:24:11 -07:00
Owen SchwartzandGitHub df6a84648b Merge pull request #100 from LaurenceJJones/fix/issue-38-stale-dns-cleanup
feat(DNS): Add static cleanup funcs
2026-04-07 21:26:38 -04:00
Owen 5ef6b21a6e Add CODEOWNERS 2026-04-07 11:34:38 -04:00
Owen 7d83518951 Get peer manager 2026-03-20 17:28:39 -07:00
Owen 964532777a Increase attempts 2026-03-19 17:24:41 -07:00
Owen SchwartzandGitHub 703fe4fe5d Merge pull request #105 from fosrl/dev
Fix nil pointer deference
2026-03-19 16:16:06 -07:00
Owen 42ef1f5ee3 Fix nil pointer deference 2026-03-19 15:21:50 -07:00
Owen SchwartzandGitHub 31eed74933 Merge pull request #103 from fosrl/dev
Update dockerfile for new version
2026-03-17 11:29:04 -07:00
Owen ac5c11dff0 Update dockerfile for new version 2026-03-17 11:27:44 -07:00
Owen SchwartzandGitHub 4d0c43fc3e Merge pull request #102 from fosrl/dev
Update cicd
2026-03-16 17:53:13 -07:00
Owen 815997d7ce Update cicd 2026-03-16 17:52:40 -07:00
Owen SchwartzandGitHub c77c162bae Merge pull request #101 from fosrl/dev
1.4.3
2026-03-16 16:44:08 -07:00
Owen 703c606af5 Handle no chainId case 2026-03-16 14:31:16 -07:00
Owen 4bc0508c7d Remove redundant info 2026-03-16 13:50:21 -07:00
Owen 3de8dc9fc2 Add optional compression 2026-03-12 17:49:12 -07:00
Owen c2b5ef96a4 Jit of aliases working 2026-03-12 17:26:46 -07:00
Owen e326da3d3e Merge branch 'dev' into jit 2026-03-12 16:53:16 -07:00
Owen 53def4e2f6 Merge branch 'main' into dev 2026-03-12 16:51:06 -07:00
Owen e85fd9d71e Bump newt version 2026-03-12 16:50:41 -07:00
Owen 98a24960f5 Remove extra restore function 2026-03-12 16:50:41 -07:00
Owen e82387d515 Actually pull the upstream from the dns var 2026-03-12 16:50:41 -07:00
Owen b3cb3e1c92 Add hardcoded public dns 2026-03-12 16:50:41 -07:00
Laurence f250702177 feat(DNS): Add static cleanup funcs
To aid CLI in cleaning up configuration we expose static functions that know how to handle each provider and platform linked to https://github.com/fosrl/cli/issues/38
2026-03-12 12:26:03 +00:00
Owen 22cd02ae15 Alias jit handler 2026-03-11 15:56:51 -07:00
André GilersonandOwen Schwartz 3f258d3500 Fix crash when peer has nil publicKey in site config
Skip sites with empty/nil publicKey instead of passing them to the
WireGuard UAPI layer, which expects a valid 64-char hex string. A nil
key occurs when a Newt site has never connected. Previously this caused
all sites to fail with "hex string does not fit the slice".
2026-03-07 20:44:25 -08:00
Owen e2690bcc03 Store site id 2026-03-06 16:19:00 -08:00
Owen f2d0e6a14c Merge branch 'dev' into jit 2026-03-06 16:08:24 -08:00
LaurenceandOwen Schwartz ae88766d85 test(dns): add dns test cases for nodata 2026-03-06 16:08:01 -08:00
LaurenceandOwen Schwartz 9ae49e36d5 refactor(dns): simplify DNSRecordStore from trie to map
Replace trie-based domain lookup with simple map for O(1) lookups.
  Add exists boolean to GetRecords for proper NODATA vs NXDOMAIN responses.
2026-03-06 16:08:01 -08:00
LaurenceandOwen Schwartz 5ca4825800 refactor(dns): trie + unified record set for DNSRecordStore
- Replace four maps (aRecords, aaaaRecords, aWildcards, aaaaWildcards) with a label trie for exact lookups and a single wildcards map
- Store one recordSet (A + AAAA) per domain/pattern instead of separate A and AAAA maps
- Exact lookups O(labels); PTR unchanged (map); API and behaviour unchanged
2026-03-06 16:08:01 -08:00
Owen 809dbe77de Make chainId in relay message bckwd compat 2026-03-06 15:27:03 -08:00
Owen c67c2a60a1 Handle canceling sends for relay 2026-03-06 15:15:31 -08:00
Owen 051c0fdfd8 Working jit with chain ids 2026-03-04 17:51:48 -08:00
Owen e7507e0837 Add api endpoints to jit 2026-03-04 17:01:17 -08:00
Laurence 8549dc8746 enhance(dns): expose stale cleanup functionality
When the tunnel is forced close an integration may want to manually call cleanup function to fix stale issues without having the knowledge of which configuration to cleanup
2026-02-26 11:30:12 +00:00
Owen 21b66fbb34 Update iss 2026-02-25 14:57:56 -08:00
Owen 9c0e37eddb Send token 2026-02-24 19:47:30 -08:00
52 changed files with 5463 additions and 897 deletions
+1
View File
@@ -0,0 +1 @@
* @oschwartz10612 @miloschwartz
+3 -2
View File
@@ -14,12 +14,13 @@ body:
label: Environment
description: Please fill out the relevant details below for your environment.
value: |
- OS Type & Version: (e.g., Ubuntu 22.04)
- OS Type & Version:
- Pangolin Version:
- Edition (Community or Enterprise):
- Gerbil Version:
- Traefik Version:
- Newt Version:
- Olm Version: (if applicable)
- Client Version:
validations:
required: true
+671 -350
View File
File diff suppressed because it is too large Load Diff
+9 -10
View File
@@ -1,4 +1,8 @@
FROM golang:1.25-alpine AS builder
# FROM golang:1.25-alpine AS builder
FROM public.ecr.aws/docker/library/golang:1.25-alpine AS builder
# Install git and ca-certificates
RUN apk --no-cache add ca-certificates git tzdata
# Set the working directory inside the container
WORKDIR /app
@@ -13,21 +17,16 @@ RUN go mod download
COPY . .
# Build the application
RUN CGO_ENABLED=0 GOOS=linux go build -o /olm
ARG VERSION=dev
RUN CGO_ENABLED=0 GOOS=linux go build -ldflags="-s -w -X main.olmVersion=${VERSION}" -o /olm
# Start a new stage from scratch
FROM alpine:3.23 AS runner
FROM public.ecr.aws/docker/library/alpine:3.23 AS runner
RUN apk --no-cache add ca-certificates
RUN apk --no-cache add ca-certificates tzdata iputils
# Copy the pre-built binary file from the previous stage and the entrypoint script
COPY --from=builder /olm /usr/local/bin/
COPY entrypoint.sh /
RUN chmod +x /entrypoint.sh
# Copy the entrypoint script
ENTRYPOINT ["/entrypoint.sh"]
# Command to run the executable
CMD ["olm"]
+11 -8
View File
@@ -2,6 +2,9 @@
all: local
VERSION ?= dev
LDFLAGS = -X main.olmVersion=$(VERSION)
local:
CGO_ENABLED=0 go build -o ./bin/olm
@@ -43,25 +46,25 @@ go-build-release: \
go-build-release-windows-amd64 \
go-build-release-linux-arm64:
CGO_ENABLED=0 GOOS=linux GOARCH=arm64 go build -o bin/olm_linux_arm64
CGO_ENABLED=0 GOOS=linux GOARCH=arm64 go build -ldflags "$(LDFLAGS)" -o bin/olm_linux_arm64
go-build-release-linux-arm32-v7:
CGO_ENABLED=0 GOOS=linux GOARCH=arm GOARM=7 go build -o bin/olm_linux_arm32
CGO_ENABLED=0 GOOS=linux GOARCH=arm GOARM=7 go build -ldflags "$(LDFLAGS)" -o bin/olm_linux_arm32
go-build-release-linux-arm32-v6:
CGO_ENABLED=0 GOOS=linux GOARCH=arm GOARM=6 go build -o bin/olm_linux_arm32v6
CGO_ENABLED=0 GOOS=linux GOARCH=arm GOARM=6 go build -ldflags "$(LDFLAGS)" -o bin/olm_linux_arm32v6
go-build-release-linux-amd64:
CGO_ENABLED=0 GOOS=linux GOARCH=amd64 go build -o bin/olm_linux_amd64
CGO_ENABLED=0 GOOS=linux GOARCH=amd64 go build -ldflags "$(LDFLAGS)" -o bin/olm_linux_amd64
go-build-release-linux-riscv64:
CGO_ENABLED=0 GOOS=linux GOARCH=riscv64 go build -o bin/olm_linux_riscv64
CGO_ENABLED=0 GOOS=linux GOARCH=riscv64 go build -ldflags "$(LDFLAGS)" -o bin/olm_linux_riscv64
go-build-release-darwin-arm64:
CGO_ENABLED=0 GOOS=darwin GOARCH=arm64 go build -o bin/olm_darwin_arm64
CGO_ENABLED=0 GOOS=darwin GOARCH=arm64 go build -ldflags "$(LDFLAGS)" -o bin/olm_darwin_arm64
go-build-release-darwin-amd64:
CGO_ENABLED=0 GOOS=darwin GOARCH=amd64 go build -o bin/olm_darwin_amd64
CGO_ENABLED=0 GOOS=darwin GOARCH=amd64 go build -ldflags "$(LDFLAGS)" -o bin/olm_darwin_amd64
go-build-release-windows-amd64:
CGO_ENABLED=0 GOOS=windows GOARCH=amd64 go build -o bin/olm_windows_amd64.exe
CGO_ENABLED=0 GOOS=windows GOARCH=amd64 go build -ldflags "$(LDFLAGS)" -o bin/olm_windows_amd64.exe
+1
View File
@@ -1,4 +1,5 @@
# Olm
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 a [WireGuard](https://www.wireguard.com/) tunnel client designed to securely connect your computer to Newt sites running on remote networks.
+91 -3
View File
@@ -29,6 +29,7 @@ type ConnectionRequest struct {
PingInterval string `json:"pingInterval,omitempty"`
PingTimeout string `json:"pingTimeout,omitempty"`
OrgID string `json:"orgId,omitempty"`
MatchDomains []string `json:"matchDomains,omitempty"`
}
// SwitchOrgRequest defines the structure for switching organizations
@@ -50,6 +51,7 @@ type PeerStatus struct {
LastSeen time.Time `json:"lastSeen"`
Endpoint string `json:"endpoint,omitempty"`
IsRelay bool `json:"isRelay"`
IsLocal bool `json:"isLocal"` // true when connected via a local network endpoint, bypassing both the public endpoint and relay
PeerIP string `json:"peerAddress,omitempty"`
HolepunchConnected bool `json:"holepunchConnected"`
}
@@ -78,6 +80,13 @@ type MetadataChangeRequest struct {
Postures map[string]any `json:"postures"`
}
// JITConnectionRequest defines the structure for a dynamic Just-In-Time connection request.
// Either SiteID or ResourceID must be provided (but not necessarily both).
type JITConnectionRequest struct {
Site string `json:"site,omitempty"`
Resource string `json:"resource,omitempty"`
}
// API represents the HTTP server and its state
type API struct {
addr string
@@ -92,6 +101,7 @@ type API struct {
onExit func() error
onRebind func() error
onPowerMode func(PowerModeRequest) error
onJITConnect func(JITConnectionRequest) error
statusMu sync.RWMutex
peerStatuses map[int]*PeerStatus
@@ -143,6 +153,7 @@ func (s *API) SetHandlers(
onExit func() error,
onRebind func() error,
onPowerMode func(PowerModeRequest) error,
onJITConnect func(JITConnectionRequest) error,
) {
s.onConnect = onConnect
s.onSwitchOrg = onSwitchOrg
@@ -151,6 +162,7 @@ func (s *API) SetHandlers(
s.onExit = onExit
s.onRebind = onRebind
s.onPowerMode = onPowerMode
s.onJITConnect = onJITConnect
}
// Start starts the HTTP server
@@ -169,6 +181,7 @@ func (s *API) Start() error {
mux.HandleFunc("/health", s.handleHealth)
mux.HandleFunc("/rebind", s.handleRebind)
mux.HandleFunc("/power-mode", s.handlePowerMode)
mux.HandleFunc("/jit-connect", s.handleJITConnect)
s.server = &http.Server{
Handler: mux,
@@ -217,7 +230,7 @@ func (s *API) Stop() error {
return nil
}
func (s *API) AddPeerStatus(siteID int, siteName string, connected bool, rtt time.Duration, endpoint string, isRelay bool) {
func (s *API) AddPeerStatus(siteID int, siteName string, connected bool, rtt time.Duration, endpoint string, isRelay bool, isLocal bool) {
s.statusMu.Lock()
defer s.statusMu.Unlock()
@@ -235,10 +248,11 @@ func (s *API) AddPeerStatus(siteID int, siteName string, connected bool, rtt tim
status.LastSeen = time.Now()
status.Endpoint = endpoint
status.IsRelay = isRelay
status.IsLocal = isLocal
}
// UpdatePeerStatus updates the status of a peer including endpoint and relay info
func (s *API) UpdatePeerStatus(siteID int, connected bool, rtt time.Duration, endpoint string, isRelay bool) {
// UpdatePeerStatus updates the status of a peer including endpoint, relay, and local info
func (s *API) UpdatePeerStatus(siteID int, connected bool, rtt time.Duration, endpoint string, isRelay bool, isLocal bool) {
s.statusMu.Lock()
defer s.statusMu.Unlock()
@@ -255,6 +269,7 @@ func (s *API) UpdatePeerStatus(siteID int, connected bool, rtt time.Duration, en
status.LastSeen = time.Now()
status.Endpoint = endpoint
status.IsRelay = isRelay
status.IsLocal = isLocal
}
func (s *API) RemovePeerStatus(siteID int) { // remove the peer from the status map
@@ -351,6 +366,31 @@ func (s *API) UpdatePeerRelayStatus(siteID int, endpoint string, isRelay bool) {
status.Endpoint = endpoint
status.IsRelay = isRelay
if isRelay {
// Relay and local are mutually exclusive; local always wins when viable.
status.IsLocal = false
}
}
// UpdatePeerLocalStatus updates only the local-connection status of a peer. A peer using a
// local connection is never simultaneously relayed.
func (s *API) UpdatePeerLocalStatus(siteID int, endpoint string, isLocal bool) {
s.statusMu.Lock()
defer s.statusMu.Unlock()
status, exists := s.peerStatuses[siteID]
if !exists {
status = &PeerStatus{
SiteID: siteID,
}
s.peerStatuses[siteID] = status
}
status.Endpoint = endpoint
status.IsLocal = isLocal
if isLocal {
status.IsRelay = false
}
}
// UpdatePeerHolepunchStatus updates the holepunch connection status of a peer
@@ -633,6 +673,54 @@ func (s *API) handleRebind(w http.ResponseWriter, r *http.Request) {
})
}
// handleJITConnect handles the /jit-connect endpoint.
// It initiates a dynamic Just-In-Time connection to a site identified by either
// a site or a resource. Exactly one of the two must be provided.
func (s *API) handleJITConnect(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
http.Error(w, "Method not allowed", http.StatusMethodNotAllowed)
return
}
var req JITConnectionRequest
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
http.Error(w, fmt.Sprintf("Invalid request body: %v", err), http.StatusBadRequest)
return
}
// Validate that exactly one of site or resource is provided
if req.Site == "" && req.Resource == "" {
http.Error(w, "Missing required field: either site or resource must be provided", http.StatusBadRequest)
return
}
if req.Site != "" && req.Resource != "" {
http.Error(w, "Ambiguous request: provide either site or resource, not both", http.StatusBadRequest)
return
}
if req.Site != "" {
logger.Info("Received JIT connection request via API: site=%s", req.Site)
} else {
logger.Info("Received JIT connection request via API: resource=%s", req.Resource)
}
if s.onJITConnect != nil {
if err := s.onJITConnect(req); err != nil {
http.Error(w, fmt.Sprintf("JIT connection failed: %v", err), http.StatusInternalServerError)
return
}
} else {
http.Error(w, "JIT connect 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": "JIT connection request accepted",
})
}
// 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) {
+68 -24
View File
@@ -27,6 +27,13 @@ type OlmConfig struct {
UpstreamDNS []string `json:"upstreamDNS"`
InterfaceName string `json:"interface"`
// MatchDomains lists FQDN wildcard patterns (using * and ? wildcards, e.g.
// "*.proxy.internal") that olm should check against local records / resolve
// via UpstreamDNS. Queries for domains that don't match any pattern are sent
// directly to the host's own system DNS servers instead. Empty means match
// every domain (i.e. the feature is disabled).
MatchDomains []string `json:"matchDomainsDNS"`
// Logging
LogLevel string `json:"logLevel"`
@@ -40,11 +47,12 @@ 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"`
DisableHolepunch bool `json:"disableHolepunch"`
TlsClientCert string `json:"tlsClientCert"`
OverrideDNS bool `json:"overrideDNS"`
TunnelDNS bool `json:"tunnelDNS"`
DisableRelay bool `json:"disableRelay"`
PreferLocalRoutes bool `json:"preferLocalRoutes"`
// DoNotCreateNewClient bool `json:"doNotCreateNewClient"`
// Parsed values (not in JSON)
@@ -99,6 +107,7 @@ func DefaultConfig() *OlmConfig {
config.sources["mtu"] = string(SourceDefault)
config.sources["dns"] = string(SourceDefault)
config.sources["upstreamDNS"] = string(SourceDefault)
config.sources["matchDomains"] = string(SourceDefault)
config.sources["logLevel"] = string(SourceDefault)
config.sources["interface"] = string(SourceDefault)
config.sources["enableApi"] = string(SourceDefault)
@@ -110,6 +119,7 @@ func DefaultConfig() *OlmConfig {
config.sources["overrideDNS"] = string(SourceDefault)
config.sources["tunnelDNS"] = string(SourceDefault)
config.sources["disableRelay"] = string(SourceDefault)
config.sources["preferLocalRoutes"] = string(SourceDefault)
// config.sources["doNotCreateNewClient"] = string(SourceDefault)
return config
@@ -229,6 +239,10 @@ func loadConfigFromEnv(config *OlmConfig) {
config.UpstreamDNS = []string{val}
config.sources["upstreamDNS"] = string(SourceEnv)
}
if val := os.Getenv("MATCH_DOMAINS_DNS"); val != "" {
config.MatchDomains = splitComma(val)
config.sources["matchDomains"] = string(SourceEnv)
}
if val := os.Getenv("LOG_LEVEL"); val != "" {
config.LogLevel = val
config.sources["logLevel"] = string(SourceEnv)
@@ -269,6 +283,10 @@ func loadConfigFromEnv(config *OlmConfig) {
config.DisableRelay = true
config.sources["disableRelay"] = string(SourceEnv)
}
if val := os.Getenv("PREFER_LOCAL_ROUTES"); val == "true" {
config.PreferLocalRoutes = true
config.sources["preferLocalRoutes"] = string(SourceEnv)
}
if val := os.Getenv("TUNNEL_DNS"); val == "true" {
config.TunnelDNS = true
config.sources["tunnelDNS"] = string(SourceEnv)
@@ -285,25 +303,27 @@ 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),
"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,
"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,
// "doNotCreateNewClient": config.DoNotCreateNewClient,
}
@@ -317,6 +337,8 @@ func loadConfigFromCLI(config *OlmConfig, args []string) (bool, bool, error) {
serviceFlags.StringVar(&config.DNS, "dns", config.DNS, "DNS server to use")
var upstreamDNSFlag string
serviceFlags.StringVar(&upstreamDNSFlag, "upstream-dns", "", "Upstream DNS server(s) (comma-separated, default: 8.8.8.8:53)")
var matchDomainsFlag string
serviceFlags.StringVar(&matchDomainsFlag, "match-domains-dns", "", "FQDN wildcard patterns (comma-separated, e.g. '*.proxy.internal,*.host-0?.autoco.internal') to check against local records/upstream DNS; queries for non-matching domains are sent directly to the system's DNS servers (default: match all domains)")
serviceFlags.StringVar(&config.LogLevel, "log-level", config.LogLevel, "Log level (DEBUG, INFO, WARN, ERROR, FATAL)")
serviceFlags.StringVar(&config.InterfaceName, "interface", config.InterfaceName, "Name of the WireGuard interface")
serviceFlags.StringVar(&config.HTTPAddr, "http-addr", config.HTTPAddr, "HTTP server address (e.g., ':9452')")
@@ -327,6 +349,7 @@ func loadConfigFromCLI(config *OlmConfig, args []string) (bool, bool, error) {
serviceFlags.BoolVar(&config.DisableHolepunch, "disable-holepunch", config.DisableHolepunch, "Disable hole punching")
serviceFlags.BoolVar(&config.OverrideDNS, "override-dns", config.OverrideDNS, "When enabled, the client uses custom DNS servers to resolve internal resources and aliases. This overrides your system's default DNS settings. Queries that cannot be resolved as a Pangolin resource will be forwarded to your configured Upstream DNS Server. (default false)")
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.DoNotCreateNewClient, "do-not-create-new-client", config.DoNotCreateNewClient, "Do not create new client")
@@ -348,6 +371,11 @@ func loadConfigFromCLI(config *OlmConfig, args []string) (bool, bool, error) {
}
}
// Parse match domains flag if provided
if matchDomainsFlag != "" {
config.MatchDomains = splitComma(matchDomainsFlag)
}
// Track which values were changed by CLI args
if config.Endpoint != origValues["endpoint"].(string) {
config.sources["endpoint"] = string(SourceCLI)
@@ -373,6 +401,9 @@ func loadConfigFromCLI(config *OlmConfig, args []string) (bool, bool, error) {
if fmt.Sprintf("%v", config.UpstreamDNS) != origValues["upstreamDNS"].(string) {
config.sources["upstreamDNS"] = string(SourceCLI)
}
if fmt.Sprintf("%v", config.MatchDomains) != origValues["matchDomains"].(string) {
config.sources["matchDomains"] = string(SourceCLI)
}
if config.LogLevel != origValues["logLevel"].(string) {
config.sources["logLevel"] = string(SourceCLI)
}
@@ -403,6 +434,9 @@ func loadConfigFromCLI(config *OlmConfig, args []string) (bool, bool, error) {
if config.DisableRelay != origValues["disableRelay"].(bool) {
config.sources["disableRelay"] = string(SourceCLI)
}
if config.PreferLocalRoutes != origValues["preferLocalRoutes"].(bool) {
config.sources["preferLocalRoutes"] = string(SourceCLI)
}
if config.TunnelDNS != origValues["tunnelDNS"].(bool) {
config.sources["tunnelDNS"] = string(SourceCLI)
}
@@ -481,6 +515,10 @@ func mergeConfigs(dest, src *OlmConfig) {
dest.UpstreamDNS = src.UpstreamDNS
dest.sources["upstreamDNS"] = string(SourceFile)
}
if len(src.MatchDomains) > 0 {
dest.MatchDomains = src.MatchDomains
dest.sources["matchDomains"] = string(SourceFile)
}
if src.LogLevel != "" && src.LogLevel != "INFO" {
dest.LogLevel = src.LogLevel
dest.sources["logLevel"] = string(SourceFile)
@@ -530,6 +568,10 @@ func mergeConfigs(dest, src *OlmConfig) {
dest.DisableRelay = src.DisableRelay
dest.sources["disableRelay"] = string(SourceFile)
}
if src.PreferLocalRoutes {
dest.PreferLocalRoutes = src.PreferLocalRoutes
dest.sources["preferLocalRoutes"] = string(SourceFile)
}
// if src.DoNotCreateNewClient {
// dest.DoNotCreateNewClient = src.DoNotCreateNewClient
// dest.sources["doNotCreateNewClient"] = string(SourceFile)
@@ -598,6 +640,7 @@ func (c *OlmConfig) ShowConfig() {
fmt.Printf(" mtu = %d [%s]\n", c.MTU, getSource("mtu"))
fmt.Printf(" dns = %s [%s]\n", c.DNS, getSource("dns"))
fmt.Printf(" upstream-dns = %v [%s]\n", c.UpstreamDNS, getSource("upstreamDNS"))
fmt.Printf(" match-domains-dns = %v [%s]\n", c.MatchDomains, getSource("matchDomains"))
fmt.Printf(" interface = %s [%s]\n", c.InterfaceName, getSource("interface"))
// Logging
@@ -621,6 +664,7 @@ func (c *OlmConfig) ShowConfig() {
fmt.Printf(" override-dns = %v [%s]\n", c.OverrideDNS, getSource("overrideDNS"))
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(" 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"))
+108 -40
View File
@@ -1,6 +1,7 @@
package device
import (
"bytes"
"io"
"net/netip"
"os"
@@ -8,6 +9,7 @@ import (
"sync/atomic"
"time"
"github.com/fosrl/newt/bind"
"github.com/fosrl/newt/logger"
"golang.zx2c4.com/wireguard/tun"
)
@@ -24,7 +26,7 @@ type FilterRule struct {
// closeAwareDevice wraps a tun.Device along with a flag
// indicating whether its Close method was called.
type closeAwareDevice struct {
isClosed atomic.Bool
isClosed atomic.Bool
tun.Device
closeEventCh chan struct{}
wg sync.WaitGroup
@@ -423,6 +425,33 @@ func extractDestIP(packet []byte) (netip.Addr, bool) {
return netip.Addr{}, false
}
// extractUDPPayload returns the UDP payload of packet, if packet is a well-formed
// IPv4 or IPv6 UDP datagram (ignoring IPv6 extension headers).
func extractUDPPayload(packet []byte) ([]byte, bool) {
if len(packet) < 20 {
return nil, false
}
const udpProtocol = 17
switch packet[0] >> 4 {
case 4:
ihl := int(packet[0]&0x0f) * 4
if ihl < 20 || len(packet) < ihl+8 || packet[9] != udpProtocol {
return nil, false
}
return packet[ihl+8:], true
case 6:
const ipv6HeaderLen = 40
if len(packet) < ipv6HeaderLen+8 || packet[6] != udpProtocol {
return nil, false
}
return packet[ipv6HeaderLen+8:], true
}
return nil, false
}
// Read intercepts packets going UP from the TUN device (towards WireGuard)
func (d *MiddleDevice) Read(bufs [][]byte, sizes []int, offset int) (n int, err error) {
for {
@@ -497,17 +526,19 @@ func (d *MiddleDevice) Read(bufs [][]byte, sizes []int, offset int) (n int, err
rules := d.rules
d.rulesMutex.RUnlock()
if len(rules) == 0 {
return n, nil
}
// Process packets and filter out handled ones
// Process packets and filter out handled ones. This always runs (even with
// no per-IP rules registered) so magic connectivity-test packets can be
// dropped before they reach WireGuard - see isLeakedMagicPacket.
writeIdx := 0
for readIdx := 0; readIdx < n; readIdx++ {
packet := bufs[readIdx][offset : offset+sizes[readIdx]]
if isLeakedMagicPacket(packet) {
continue
}
destIP, ok := extractDestIP(packet)
if !ok {
if !ok || len(rules) == 0 {
if writeIdx != readIdx {
bufs[writeIdx] = bufs[readIdx]
sizes[writeIdx] = sizes[readIdx]
@@ -539,6 +570,74 @@ func (d *MiddleDevice) Read(bufs [][]byte, sizes []int, offset int) (n int, err
}
}
// isLeakedMagicPacket reports whether packet carries one of our UDP connectivity-test
// magic payloads (see bind.IsMagicPacket). These packets are sent directly between
// physical UDP sockets by the local-endpoint holepunch tester and must never be
// encapsulated by WireGuard: if OS routing sends one into this TUN interface instead
// of out the real network interface (e.g. because the destination falls inside a
// routed tunnel subnet), tunneling and echoing it back would make a LAN-local
// endpoint falsely appear directly reachable. Dropping it here makes the test
// correctly time out instead.
func isLeakedMagicPacket(packet []byte) bool {
payload, ok := extractUDPPayload(packet)
return ok && isMagicPacket(payload)
}
// IsMagicPacket reports whether payload is one of our connectivity-test magic
// packets (a MagicTestRequest or MagicTestResponse). These packets are meant to
// travel directly between physical UDP sockets and must never be encapsulated by
// WireGuard - e.g. if OS routing mistakenly sends one into a WireGuard TUN
// interface (because the destination falls inside a routed tunnel subnet), it
// should be dropped there rather than tunneled, which would otherwise make a
// LAN-local endpoint test falsely appear to succeed over the tunnel.
func isMagicPacket(payload []byte) bool {
if len(payload) >= bind.MagicTestRequestLen && bytes.HasPrefix(payload, bind.MagicTestRequest) {
return true
}
if len(payload) >= bind.MagicTestResponseLen && bytes.HasPrefix(payload, bind.MagicTestResponse) {
return true
}
return false
}
// filterDownstreamBufs drops packets going DOWN to the TUN device (from WireGuard)
// that are handled by a per-IP rule or are a leaked magic connectivity-test packet
// (see isLeakedMagicPacket) - always checked, even with no rules registered. It
// returns bufs unchanged (no allocation) unless a packet actually needs to be
// dropped, at which point it switches to an owned copy of the buffers kept so far.
func filterDownstreamBufs(bufs [][]byte, rules []FilterRule, offset int) [][]byte {
filtered := bufs
for i, buf := range bufs {
drop := len(buf) <= offset
if !drop {
packet := buf[offset:]
if isLeakedMagicPacket(packet) {
drop = true
} else if destIP, ok := extractDestIP(packet); ok && len(rules) > 0 {
for _, rule := range rules {
if rule.DestIP == destIP && rule.Handler(packet) {
drop = true
break
}
}
}
}
if drop {
if len(filtered) == len(bufs) {
// First drop: switch to an owned, growable copy of everything kept so far.
filtered = append([][]byte(nil), bufs[:i]...)
}
continue
}
if len(filtered) != len(bufs) {
filtered = append(filtered, buf)
}
}
return filtered
}
// Write intercepts packets going DOWN to the TUN device (from WireGuard)
func (d *MiddleDevice) Write(bufs [][]byte, offset int) (int, error) {
for {
@@ -558,38 +657,7 @@ func (d *MiddleDevice) Write(bufs [][]byte, offset int) (int, error) {
rules := d.rules
d.rulesMutex.RUnlock()
var filteredBufs [][]byte
if len(rules) == 0 {
filteredBufs = bufs
} else {
filteredBufs = make([][]byte, 0, len(bufs))
for _, buf := range bufs {
if len(buf) <= offset {
continue
}
packet := buf[offset:]
destIP, ok := extractDestIP(packet)
if !ok {
filteredBufs = append(filteredBufs, buf)
continue
}
handled := false
for _, rule := range rules {
if rule.DestIP == destIP {
if rule.Handler(packet) {
handled = true
break
}
}
}
if !handled {
filteredBufs = append(filteredBufs, buf)
}
}
}
filteredBufs := filterDownstreamBufs(bufs, rules, offset)
if len(filteredBufs) == 0 {
return len(bufs), nil
@@ -660,4 +728,4 @@ func (d *MiddleDevice) WriteToTun(bufs [][]byte, offset int) (int, error) {
return n, err
}
}
}
+114
View File
@@ -4,9 +4,22 @@ import (
"net/netip"
"testing"
"github.com/fosrl/newt/bind"
"github.com/fosrl/newt/util"
)
// buildIPv4UDPPacket builds a minimal IPv4/UDP packet (no options) carrying payload.
func buildIPv4UDPPacket(payload []byte) []byte {
const ipHeaderLen = 20
const udpHeaderLen = 8
packet := make([]byte, ipHeaderLen+udpHeaderLen+len(payload))
packet[0] = 0x45 // version 4, IHL 5
packet[9] = 17 // protocol: UDP
copy(packet[ipHeaderLen+udpHeaderLen:], payload)
return packet
}
func TestExtractDestIP(t *testing.T) {
tests := []struct {
name string
@@ -88,6 +101,49 @@ func TestGetProtocol(t *testing.T) {
}
}
func TestIsLeakedMagicPacket(t *testing.T) {
request := make([]byte, bind.MagicTestRequestLen)
copy(request, bind.MagicTestRequest)
response := make([]byte, bind.MagicTestResponseLen)
copy(response, bind.MagicTestResponse)
tests := []struct {
name string
packet []byte
want bool
}{
{
name: "magic test request leaked into tunnel",
packet: buildIPv4UDPPacket(request),
want: true,
},
{
name: "magic test response leaked into tunnel",
packet: buildIPv4UDPPacket(response),
want: true,
},
{
name: "ordinary UDP payload",
packet: buildIPv4UDPPacket([]byte("just some ordinary application data")),
want: false,
},
{
name: "too short to be a packet",
packet: []byte{0x45, 0x00},
want: false,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if got := isLeakedMagicPacket(tt.packet); got != tt.want {
t.Errorf("isLeakedMagicPacket() = %v, want %v", got, tt.want)
}
})
}
}
func BenchmarkExtractDestIP(b *testing.B) {
packet := []byte{
0x45, 0x00, 0x00, 0x54, 0x00, 0x00, 0x40, 0x00,
@@ -100,3 +156,61 @@ func BenchmarkExtractDestIP(b *testing.B) {
extractDestIP(packet)
}
}
func TestFilterDownstreamBufsNoDropIsAllocFree(t *testing.T) {
bufs := make([][]byte, 128)
for i := range bufs {
bufs[i] = buildIPv4UDPPacket(make([]byte, 1372))
}
allocs := testing.AllocsPerRun(1000, func() {
out := filterDownstreamBufs(bufs, nil, 0)
if len(out) != len(bufs) {
t.Fatalf("expected no packets dropped, got %d/%d", len(out), len(bufs))
}
})
if allocs != 0 {
t.Errorf("filterDownstreamBufs() with nothing to drop allocated %v times per call, want 0", allocs)
}
}
func TestFilterDownstreamBufsDropsMagicPacket(t *testing.T) {
request := make([]byte, bind.MagicTestRequestLen)
copy(request, bind.MagicTestRequest)
bufs := [][]byte{
buildIPv4UDPPacket([]byte("ordinary payload one")),
buildIPv4UDPPacket(request),
buildIPv4UDPPacket([]byte("ordinary payload two")),
}
out := filterDownstreamBufs(bufs, nil, 0)
if len(out) != 2 {
t.Fatalf("expected 1 packet dropped, got %d remaining", len(out))
}
}
func BenchmarkFilterDownstreamBufsNoDrop(b *testing.B) {
bufs := make([][]byte, 128)
for i := range bufs {
bufs[i] = buildIPv4UDPPacket(make([]byte, 1372))
}
b.ResetTimer()
b.ReportAllocs()
for i := 0; i < b.N; i++ {
filterDownstreamBufs(bufs, nil, 0)
}
}
func BenchmarkIsLeakedMagicPacket(b *testing.B) {
// A typical ~1400 byte ordinary application payload (the common case on the
// hot path - almost every real packet should look like this).
ordinary := buildIPv4UDPPacket(make([]byte, 1372))
b.ResetTimer()
for i := 0; i < b.N; i++ {
isLeakedMagicPacket(ordinary)
}
}
+167 -10
View File
@@ -5,6 +5,7 @@ import (
"fmt"
"net"
"net/netip"
"strings"
"sync"
"time"
@@ -38,6 +39,20 @@ type DNSProxy struct {
middleDevice *device.MiddleDevice // Reference to MiddleDevice for packet filtering and TUN writes
recordStore *DNSRecordStore // Local DNS records
// matchDomains lists the FQDN wildcard patterns (using * and ? wildcards, see
// matchWildcard) that this proxy is responsible for. Queries whose name matches
// one of these patterns are checked against local records and, failing that,
// forwarded to upstreamDNS. Queries that match none of the patterns are sent
// directly to localDNS instead, bypassing local records and upstreamDNS
// entirely. An empty matchDomains means "match everything" (i.e. behave as if
// this feature were not configured).
matchDomains []string
// localDNS holds the host's own system DNS servers (as reported by
// SystemDNSMonitor / PublicDNS), used to resolve queries that don't match
// matchDomains rather than sending them upstream or through the tunnel.
localDNS []string
matchMu sync.RWMutex
// Tunnel DNS fields - for sending queries over WireGuard
tunnelIP netip.Addr // WireGuard interface IP (source for tunneled queries)
tunnelStack *stack.Stack // Separate netstack for outbound tunnel queries
@@ -45,13 +60,24 @@ type DNSProxy struct {
tunnelActivePorts map[uint16]bool
tunnelPortsLock sync.Mutex
// jitHandler is called when a local record is resolved for a site that may not be
// connected yet, giving the caller a chance to initiate a JIT connection.
// It is invoked asynchronously so it never blocks DNS resolution.
jitHandler func(siteId int)
ctx context.Context
cancel context.CancelFunc
wg sync.WaitGroup
}
// NewDNSProxy creates a new DNS proxy
func NewDNSProxy(middleDevice *device.MiddleDevice, mtu int, utilitySubnet string, upstreamDns []string, tunnelDns bool, tunnelIP string) (*DNSProxy, error) {
// NewDNSProxy creates a new DNS proxy.
//
// matchDomains, if non-empty, restricts local-record lookup and upstream
// forwarding to queries whose name matches one of the given wildcard patterns
// (see matchWildcard). Queries that match none of the patterns are instead
// forwarded directly to localDNS (the host's own system DNS servers). Pass an
// empty matchDomains to match every query, preserving prior behavior.
func NewDNSProxy(middleDevice *device.MiddleDevice, mtu int, utilitySubnet string, upstreamDns []string, tunnelDns bool, tunnelIP string, matchDomains []string, localDNS []string) (*DNSProxy, error) {
proxyIP, err := PickIPFromSubnet(utilitySubnet)
if err != nil {
return nil, fmt.Errorf("failed to pick DNS proxy IP from subnet: %v", err)
@@ -71,6 +97,8 @@ func NewDNSProxy(middleDevice *device.MiddleDevice, mtu int, utilitySubnet strin
tunnelDNS: tunnelDns,
recordStore: NewDNSRecordStore(),
tunnelActivePorts: make(map[uint16]bool),
matchDomains: matchDomains,
localDNS: localDNS,
ctx: ctx,
cancel: cancel,
}
@@ -378,12 +406,43 @@ func (p *DNSProxy) handleDNSQuery(udpConn *gonet.UDPConn, queryData []byte, clie
question := msg.Question[0]
logger.Debug("DNS query for %s (type %s)", question.Name, dns.TypeToString[question.Qtype])
// If matchDomains is configured and this query's name doesn't match any of
// the configured patterns, skip local records and upstream entirely and
// send it straight to the host's own system DNS servers.
if !p.matchesConfiguredDomains(question.Name) {
logger.Debug("Query for %s does not match configured domains, forwarding to local DNS %v", question.Name, p.getLocalDNS())
response := p.forwardToLocalDNS(msg)
if response == nil {
logger.Error("Failed to get DNS response for %s from local DNS", question.Name)
return
}
responseData, err := response.Pack()
if err != nil {
logger.Error("Failed to pack DNS response: %v", err)
return
}
if _, err := udpConn.WriteTo(responseData, clientAddr); err != nil {
logger.Error("Failed to send DNS response: %v", err)
}
return
}
// Check if we have local records for this query
var response *dns.Msg
if question.Qtype == dns.TypeA || question.Qtype == dns.TypeAAAA || question.Qtype == dns.TypePTR {
response = p.checkLocalRecords(msg, question)
}
// If a local A/AAAA record was found, notify the JIT handler so that the owning
// site can be connected on-demand if it is not yet active.
if response != nil && p.jitHandler != nil &&
(question.Qtype == dns.TypeA || question.Qtype == dns.TypeAAAA) {
if siteId, ok := p.recordStore.GetSiteIdForDomain(question.Name); ok && siteId != 0 {
handler := p.jitHandler
go handler(siteId)
}
}
// If no local records, forward to upstream
if response == nil {
logger.Debug("No local record for %s, forwarding upstream to %v", question.Name, p.upstreamDNS)
@@ -447,19 +506,20 @@ func (p *DNSProxy) checkLocalRecords(query *dns.Msg, question dns.Question) *dns
return nil
}
ips := p.recordStore.GetRecords(question.Name, recordType)
if len(ips) == 0 {
ips, exists := p.recordStore.GetRecords(question.Name, recordType)
if !exists {
// Domain not found in local records, forward to upstream
return nil
}
logger.Debug("Found %d local record(s) for %s", len(ips), question.Name)
// Create response message
// Create response message (NODATA if no records, otherwise with answers)
response := new(dns.Msg)
response.SetReply(query)
response.Authoritative = true
// Add answer records
// Add answer records (loop is a no-op if ips is empty)
for _, ip := range ips {
var rr dns.RR
if question.Qtype == dns.TypeA {
@@ -489,6 +549,77 @@ func (p *DNSProxy) checkLocalRecords(query *dns.Msg, question dns.Question) *dns
return response
}
// matchesConfiguredDomains reports whether name matches one of the configured
// matchDomains wildcard patterns. If matchDomains is empty, every name is
// considered a match (i.e. the feature is disabled).
func (p *DNSProxy) matchesConfiguredDomains(name string) bool {
p.matchMu.RLock()
patterns := p.matchDomains
p.matchMu.RUnlock()
if len(patterns) == 0 {
return true
}
name = strings.ToLower(dns.Fqdn(name))
for _, pattern := range patterns {
pattern = strings.ToLower(dns.Fqdn(pattern))
if matchWildcard(pattern, name) {
return true
}
}
return false
}
// getLocalDNS returns the currently configured local (system) DNS servers.
func (p *DNSProxy) getLocalDNS() []string {
p.matchMu.RLock()
defer p.matchMu.RUnlock()
return p.localDNS
}
// forwardToLocalDNS forwards a DNS query directly to the host's own system DNS
// servers (localDNS), always using host networking regardless of tunnelDNS -
// these queries are for domains the caller has explicitly excluded from
// Pangolin resolution, so they should never traverse the tunnel.
func (p *DNSProxy) forwardToLocalDNS(query *dns.Msg) *dns.Msg {
servers := p.getLocalDNS()
if len(servers) == 0 {
logger.Warn("No local DNS servers configured, dropping query for %s", query.Question[0].Name)
return nil
}
var lastErr error
for _, server := range servers {
response, err := p.queryUpstreamDirect(server, query, 2*time.Second)
if err == nil {
return response
}
lastErr = err
}
logger.Error("All local DNS servers failed: %v", lastErr)
return nil
}
// SetMatchDomains replaces the list of wildcard domain patterns (see
// matchWildcard) that this proxy checks against local records / upstream DNS.
// Queries not matching any pattern are sent to localDNS instead. Pass an
// empty slice to match every query (i.e. disable filtering).
func (p *DNSProxy) SetMatchDomains(patterns []string) {
p.matchMu.Lock()
defer p.matchMu.Unlock()
p.matchDomains = patterns
}
// SetLocalDNS replaces the list of local (host system) DNS servers used to
// resolve queries that don't match matchDomains. Servers must be in
// "host:port" format (e.g. "192.168.1.1:53").
func (p *DNSProxy) SetLocalDNS(servers []string) {
p.matchMu.Lock()
defer p.matchMu.Unlock()
p.localDNS = servers
}
// forwardToUpstream forwards a DNS query to upstream DNS servers
func (p *DNSProxy) forwardToUpstream(query *dns.Msg) *dns.Msg {
// Try primary DNS server
@@ -717,11 +848,30 @@ func (p *DNSProxy) runPacketSender() {
}
}
// SetJITHandler registers a callback that is invoked whenever a local DNS record is
// resolved for an A or AAAA query. The siteId identifies which site owns the record.
// The handler is called in its own goroutine so it must be safe to call concurrently.
// Pass nil to disable JIT notifications.
func (p *DNSProxy) SetJITHandler(handler func(siteId int)) {
p.jitHandler = handler
}
// SetUpstreamDNS replaces the list of upstream DNS servers used to forward
// queries that are not served by local records. The servers must be in
// "host:port" format (e.g. "8.8.8.8:53").
func (p *DNSProxy) SetUpstreamDNS(servers []string) {
if len(servers) == 0 {
return
}
p.upstreamDNS = servers
}
// 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
func (p *DNSProxy) AddDNSRecord(domain string, ip net.IP) error {
return p.recordStore.AddRecord(domain, ip)
func (p *DNSProxy) AddDNSRecord(domain string, ip net.IP, siteId int) error {
logger.Debug("Adding dns record for domain %s with IP %s (siteId=%d)", domain, ip.String(), siteId)
return p.recordStore.AddRecord(domain, ip, siteId)
}
// RemoveDNSRecord removes a DNS record from the local store
@@ -730,8 +880,15 @@ func (p *DNSProxy) RemoveDNSRecord(domain string, ip net.IP) {
p.recordStore.RemoveRecord(domain, ip)
}
// GetDNSRecords returns all IP addresses for a domain and record type
func (p *DNSProxy) GetDNSRecords(domain string, recordType RecordType) []net.IP {
// RemoveDNSRecordForSite removes DNS records for a domain that are owned by a specific site.
// If ip is nil, removes all records for the domain that are owned by that site.
func (p *DNSProxy) RemoveDNSRecordForSite(domain string, ip net.IP, siteId int) {
p.recordStore.RemoveRecordForSite(domain, ip, siteId)
}
// GetDNSRecords returns all IP addresses for a domain and record type.
// The second return value indicates whether the domain exists.
func (p *DNSProxy) GetDNSRecords(domain string, recordType RecordType) ([]net.IP, bool) {
return p.recordStore.GetRecords(domain, recordType)
}
+178
View File
@@ -0,0 +1,178 @@
package dns
import (
"net"
"testing"
"github.com/miekg/dns"
)
func TestCheckLocalRecordsNODATAForAAAA(t *testing.T) {
proxy := &DNSProxy{
recordStore: NewDNSRecordStore(),
}
// Add an A record for a domain (no AAAA record)
ip := net.ParseIP("10.0.0.1")
err := proxy.recordStore.AddRecord("myservice.internal", ip, 0)
if err != nil {
t.Fatalf("Failed to add A record: %v", err)
}
// Query AAAA for domain with only A record - should return NODATA
query := new(dns.Msg)
query.SetQuestion("myservice.internal.", dns.TypeAAAA)
response := proxy.checkLocalRecords(query, query.Question[0])
if response == nil {
t.Fatal("Expected NODATA response, got nil (would forward to upstream)")
}
if response.Rcode != dns.RcodeSuccess {
t.Errorf("Expected Rcode NOERROR (0), got %d", response.Rcode)
}
if len(response.Answer) != 0 {
t.Errorf("Expected empty answer section for NODATA, got %d answers", len(response.Answer))
}
if !response.Authoritative {
t.Error("Expected response to be authoritative")
}
// Query A for same domain - should return the record
query = new(dns.Msg)
query.SetQuestion("myservice.internal.", dns.TypeA)
response = proxy.checkLocalRecords(query, query.Question[0])
if response == nil {
t.Fatal("Expected response with A record, got nil")
}
if len(response.Answer) != 1 {
t.Fatalf("Expected 1 answer, got %d", len(response.Answer))
}
aRecord, ok := response.Answer[0].(*dns.A)
if !ok {
t.Fatal("Expected A record in answer")
}
if !aRecord.A.Equal(ip.To4()) {
t.Errorf("Expected IP %v, got %v", ip.To4(), aRecord.A)
}
}
func TestCheckLocalRecordsNODATAForA(t *testing.T) {
proxy := &DNSProxy{
recordStore: NewDNSRecordStore(),
}
// Add an AAAA record for a domain (no A record)
ip := net.ParseIP("2001:db8::1")
err := proxy.recordStore.AddRecord("ipv6only.internal", ip, 0)
if err != nil {
t.Fatalf("Failed to add AAAA record: %v", err)
}
// Query A for domain with only AAAA record - should return NODATA
query := new(dns.Msg)
query.SetQuestion("ipv6only.internal.", dns.TypeA)
response := proxy.checkLocalRecords(query, query.Question[0])
if response == nil {
t.Fatal("Expected NODATA response, got nil")
}
if response.Rcode != dns.RcodeSuccess {
t.Errorf("Expected Rcode NOERROR (0), got %d", response.Rcode)
}
if len(response.Answer) != 0 {
t.Errorf("Expected empty answer section, got %d answers", len(response.Answer))
}
if !response.Authoritative {
t.Error("Expected response to be authoritative")
}
// Query AAAA for same domain - should return the record
query = new(dns.Msg)
query.SetQuestion("ipv6only.internal.", dns.TypeAAAA)
response = proxy.checkLocalRecords(query, query.Question[0])
if response == nil {
t.Fatal("Expected response with AAAA record, got nil")
}
if len(response.Answer) != 1 {
t.Fatalf("Expected 1 answer, got %d", len(response.Answer))
}
aaaaRecord, ok := response.Answer[0].(*dns.AAAA)
if !ok {
t.Fatal("Expected AAAA record in answer")
}
if !aaaaRecord.AAAA.Equal(ip) {
t.Errorf("Expected IP %v, got %v", ip, aaaaRecord.AAAA)
}
}
func TestCheckLocalRecordsNonExistentDomain(t *testing.T) {
proxy := &DNSProxy{
recordStore: NewDNSRecordStore(),
}
// Add a record so the store isn't empty
err := proxy.recordStore.AddRecord("exists.internal", net.ParseIP("10.0.0.1"), 0)
if err != nil {
t.Fatalf("Failed to add record: %v", err)
}
// Query A for non-existent domain - should return nil (forward to upstream)
query := new(dns.Msg)
query.SetQuestion("unknown.internal.", dns.TypeA)
response := proxy.checkLocalRecords(query, query.Question[0])
if response != nil {
t.Error("Expected nil for non-existent domain, got response")
}
// Query AAAA for non-existent domain - should also return nil
query = new(dns.Msg)
query.SetQuestion("unknown.internal.", dns.TypeAAAA)
response = proxy.checkLocalRecords(query, query.Question[0])
if response != nil {
t.Error("Expected nil for non-existent domain, got response")
}
}
func TestCheckLocalRecordsNODATAWildcard(t *testing.T) {
proxy := &DNSProxy{
recordStore: NewDNSRecordStore(),
}
// Add a wildcard A record (no AAAA)
ip := net.ParseIP("10.0.0.1")
err := proxy.recordStore.AddRecord("*.wildcard.internal", ip, 0)
if err != nil {
t.Fatalf("Failed to add wildcard A record: %v", err)
}
// Query AAAA for wildcard-matched domain - should return NODATA
query := new(dns.Msg)
query.SetQuestion("host.wildcard.internal.", dns.TypeAAAA)
response := proxy.checkLocalRecords(query, query.Question[0])
if response == nil {
t.Fatal("Expected NODATA response for wildcard match, got nil")
}
if response.Rcode != dns.RcodeSuccess {
t.Errorf("Expected Rcode NOERROR (0), got %d", response.Rcode)
}
if len(response.Answer) != 0 {
t.Errorf("Expected empty answer section, got %d answers", len(response.Answer))
}
// Query A for wildcard-matched domain - should return the record
query = new(dns.Msg)
query.SetQuestion("host.wildcard.internal.", dns.TypeA)
response = proxy.checkLocalRecords(query, query.Question[0])
if response == nil {
t.Fatal("Expected response with A record, got nil")
}
if len(response.Answer) != 1 {
t.Fatalf("Expected 1 answer, got %d", len(response.Answer))
}
}
+244 -173
View File
@@ -18,24 +18,29 @@ const (
RecordTypePTR RecordType = RecordType(dns.TypePTR)
)
// DNSRecordStore manages local DNS records for A, AAAA, and PTR queries
// recordSet holds A and AAAA records for a single domain or wildcard pattern
type recordSet struct {
A []net.IP
AAAA []net.IP
SiteId int
owners map[string]map[int]bool // IP string -> owning site IDs
}
// DNSRecordStore manages local DNS records for A, AAAA, and PTR queries.
// Exact domains are stored in a map; wildcard patterns are in a separate map.
type DNSRecordStore struct {
mu sync.RWMutex
aRecords map[string][]net.IP // domain -> list of IPv4 addresses
aaaaRecords map[string][]net.IP // domain -> list of IPv6 addresses
aWildcards map[string][]net.IP // wildcard pattern -> list of IPv4 addresses
aaaaWildcards map[string][]net.IP // wildcard pattern -> list of IPv6 addresses
ptrRecords map[string]string // IP address string -> domain name
mu sync.RWMutex
exact map[string]*recordSet // normalized FQDN -> A/AAAA records
wildcards map[string]*recordSet // wildcard pattern -> A/AAAA records
ptrRecords map[string]string // IP address string -> domain name
}
// NewDNSRecordStore creates a new DNS record store
func NewDNSRecordStore() *DNSRecordStore {
return &DNSRecordStore{
aRecords: make(map[string][]net.IP),
aaaaRecords: make(map[string][]net.IP),
aWildcards: make(map[string][]net.IP),
aaaaWildcards: make(map[string][]net.IP),
ptrRecords: make(map[string]string),
exact: make(map[string]*recordSet),
wildcards: make(map[string]*recordSet),
ptrRecords: make(map[string]string),
}
}
@@ -43,44 +48,61 @@ func NewDNSRecordStore() *DNSRecordStore {
// domain should be in FQDN format (e.g., "example.com.")
// domain can contain wildcards: * (0+ chars) and ? (exactly 1 char)
// ip should be a valid IPv4 or IPv6 address
// siteId is the site that owns this alias/domain
// Automatically adds a corresponding PTR record for non-wildcard domains
func (s *DNSRecordStore) AddRecord(domain string, ip net.IP) error {
func (s *DNSRecordStore) AddRecord(domain string, ip net.IP, siteId int) error {
s.mu.Lock()
defer s.mu.Unlock()
// Ensure domain ends with a dot (FQDN format)
if len(domain) == 0 || domain[len(domain)-1] != '.' {
domain = domain + "."
}
// Normalize domain to lowercase FQDN
domain = strings.ToLower(dns.Fqdn(domain))
// Check if domain contains wildcards
isWildcard := strings.ContainsAny(domain, "*?")
if ip.To4() != nil {
// IPv4 address
if isWildcard {
s.aWildcards[domain] = append(s.aWildcards[domain], ip)
} else {
s.aRecords[domain] = append(s.aRecords[domain], ip)
// Automatically add PTR record for non-wildcard domains
s.ptrRecords[ip.String()] = domain
}
} else if ip.To16() != nil {
// IPv6 address
if isWildcard {
s.aaaaWildcards[domain] = append(s.aaaaWildcards[domain], ip)
} else {
s.aaaaRecords[domain] = append(s.aaaaRecords[domain], ip)
// Automatically add PTR record for non-wildcard domains
s.ptrRecords[ip.String()] = domain
}
} else {
isV4 := ip.To4() != nil
if !isV4 && ip.To16() == nil {
return &net.ParseError{Type: "IP address", Text: ip.String()}
}
// Choose the appropriate map based on whether this is a wildcard
m := s.exact
if isWildcard {
m = s.wildcards
}
if m[domain] == nil {
m[domain] = &recordSet{SiteId: siteId, owners: make(map[string]map[int]bool)}
}
rs := m[domain]
if rs.owners == nil {
rs.owners = make(map[string]map[int]bool)
}
ipKey := ip.String()
if rs.owners[ipKey] == nil {
rs.owners[ipKey] = make(map[int]bool)
}
rs.owners[ipKey][siteId] = true
if isV4 {
for _, existing := range rs.A {
if existing.Equal(ip) {
return nil
}
}
rs.A = append(rs.A, ip)
} else {
for _, existing := range rs.AAAA {
if existing.Equal(ip) {
return nil
}
}
rs.AAAA = append(rs.AAAA, ip)
}
// Add PTR record for non-wildcard domains
if !isWildcard {
s.ptrRecords[ip.String()] = domain
}
return nil
}
@@ -109,92 +131,128 @@ func (s *DNSRecordStore) AddPTRRecord(ip net.IP, domain string) error {
// If ip is nil, removes all records for the domain (including wildcards)
// Automatically removes corresponding PTR records for non-wildcard domains
func (s *DNSRecordStore) RemoveRecord(domain string, ip net.IP) {
s.removeRecord(domain, ip, 0, false)
}
// RemoveRecordForSite removes DNS records owned by a specific site.
// If ip is nil, it removes all records for the domain owned by that site.
func (s *DNSRecordStore) RemoveRecordForSite(domain string, ip net.IP, siteId int) {
s.removeRecord(domain, ip, siteId, true)
}
func (s *DNSRecordStore) removeRecord(domain string, ip net.IP, siteId int, bySite bool) {
s.mu.Lock()
defer s.mu.Unlock()
// Ensure domain ends with a dot (FQDN format)
if len(domain) == 0 || domain[len(domain)-1] != '.' {
domain = domain + "."
}
// Normalize domain to lowercase FQDN
domain = strings.ToLower(dns.Fqdn(domain))
// Check if domain contains wildcards
isWildcard := strings.ContainsAny(domain, "*?")
// Choose the appropriate map
m := s.exact
if isWildcard {
m = s.wildcards
}
rs := m[domain]
if rs == nil {
return
}
if rs.owners == nil {
rs.owners = make(map[string]map[int]bool)
}
if ip == nil {
// Remove all records for this domain
if isWildcard {
delete(s.aWildcards, domain)
delete(s.aaaaWildcards, domain)
} else {
// For non-wildcard domains, remove PTR records for all IPs
if ips, ok := s.aRecords[domain]; ok {
for _, ipAddr := range ips {
// Only remove PTR if it points to this domain
if ptrDomain, exists := s.ptrRecords[ipAddr.String()]; exists && ptrDomain == domain {
delete(s.ptrRecords, ipAddr.String())
}
}
if bySite {
rs.A = s.removeOwnedIPs(rs, rs.A, siteId, !isWildcard, domain)
rs.AAAA = s.removeOwnedIPs(rs, rs.AAAA, siteId, !isWildcard, domain)
if len(rs.A) == 0 && len(rs.AAAA) == 0 {
delete(m, domain)
}
if ips, ok := s.aaaaRecords[domain]; ok {
for _, ipAddr := range ips {
// Only remove PTR if it points to this domain
if ptrDomain, exists := s.ptrRecords[ipAddr.String()]; exists && ptrDomain == domain {
delete(s.ptrRecords, ipAddr.String())
}
}
}
delete(s.aRecords, domain)
delete(s.aaaaRecords, domain)
return
}
// Remove all records for this domain
if !isWildcard {
for _, ipAddr := range rs.A {
if ptrDomain, exists := s.ptrRecords[ipAddr.String()]; exists && ptrDomain == domain {
delete(s.ptrRecords, ipAddr.String())
}
}
for _, ipAddr := range rs.AAAA {
if ptrDomain, exists := s.ptrRecords[ipAddr.String()]; exists && ptrDomain == domain {
delete(s.ptrRecords, ipAddr.String())
}
}
}
delete(m, domain)
return
}
// Remove specific IP
ipKey := ip.String()
if bySite {
owners := rs.owners[ipKey]
if len(owners) == 0 {
return
}
delete(owners, siteId)
if len(owners) > 0 {
return
}
delete(rs.owners, ipKey)
}
if ip.To4() != nil {
// Remove specific IPv4 address
if isWildcard {
if ips, ok := s.aWildcards[domain]; ok {
s.aWildcards[domain] = removeIP(ips, ip)
if len(s.aWildcards[domain]) == 0 {
delete(s.aWildcards, domain)
}
}
} else {
if ips, ok := s.aRecords[domain]; ok {
s.aRecords[domain] = removeIP(ips, ip)
if len(s.aRecords[domain]) == 0 {
delete(s.aRecords, domain)
}
// Automatically remove PTR record if it points to this domain
if ptrDomain, exists := s.ptrRecords[ip.String()]; exists && ptrDomain == domain {
delete(s.ptrRecords, ip.String())
}
rs.A = removeIP(rs.A, ip)
if !isWildcard {
if ptrDomain, exists := s.ptrRecords[ip.String()]; exists && ptrDomain == domain {
delete(s.ptrRecords, ip.String())
}
}
} else if ip.To16() != nil {
// Remove specific IPv6 address
if isWildcard {
if ips, ok := s.aaaaWildcards[domain]; ok {
s.aaaaWildcards[domain] = removeIP(ips, ip)
if len(s.aaaaWildcards[domain]) == 0 {
delete(s.aaaaWildcards, domain)
}
}
} else {
if ips, ok := s.aaaaRecords[domain]; ok {
s.aaaaRecords[domain] = removeIP(ips, ip)
if len(s.aaaaRecords[domain]) == 0 {
delete(s.aaaaRecords, domain)
}
// Automatically remove PTR record if it points to this domain
if ptrDomain, exists := s.ptrRecords[ip.String()]; exists && ptrDomain == domain {
delete(s.ptrRecords, ip.String())
}
} else {
rs.AAAA = removeIP(rs.AAAA, ip)
if !isWildcard {
if ptrDomain, exists := s.ptrRecords[ip.String()]; exists && ptrDomain == domain {
delete(s.ptrRecords, ip.String())
}
}
}
delete(rs.owners, ipKey)
// Clean up empty record sets
if len(rs.A) == 0 && len(rs.AAAA) == 0 {
delete(m, domain)
}
}
func (s *DNSRecordStore) removeOwnedIPs(rs *recordSet, ips []net.IP, siteId int, removePTR bool, domain string) []net.IP {
kept := make([]net.IP, 0, len(ips))
for _, ipAddr := range ips {
ipKey := ipAddr.String()
owners := rs.owners[ipKey]
if len(owners) == 0 {
kept = append(kept, ipAddr)
continue
}
delete(owners, siteId)
if len(owners) > 0 {
kept = append(kept, ipAddr)
continue
}
delete(rs.owners, ipKey)
if removePTR {
if ptrDomain, exists := s.ptrRecords[ipKey]; exists && ptrDomain == domain {
delete(s.ptrRecords, ipKey)
}
}
}
return kept
}
// RemovePTRRecord removes a PTR record for an IP address
@@ -205,61 +263,80 @@ func (s *DNSRecordStore) RemovePTRRecord(ip net.IP) {
delete(s.ptrRecords, ip.String())
}
// GetRecords returns all IP addresses for a domain and record type
// First checks for exact matches, then checks wildcard patterns
func (s *DNSRecordStore) GetRecords(domain string, recordType RecordType) []net.IP {
// GetSiteIdForDomain returns the siteId associated with the given domain.
// It checks exact matches first, then wildcard patterns.
// The second return value is false if the domain is not found in local records.
func (s *DNSRecordStore) GetSiteIdForDomain(domain string) (int, bool) {
s.mu.RLock()
defer s.mu.RUnlock()
// Normalize domain to lowercase FQDN
domain = strings.ToLower(dns.Fqdn(domain))
var records []net.IP
switch recordType {
case RecordTypeA:
// Check exact match first
if ips, ok := s.aRecords[domain]; ok {
// Return a copy to prevent external modifications
records = make([]net.IP, len(ips))
copy(records, ips)
return records
}
// Check wildcard patterns
for pattern, ips := range s.aWildcards {
if matchWildcard(pattern, domain) {
records = append(records, ips...)
}
}
if len(records) > 0 {
// Return a copy
result := make([]net.IP, len(records))
copy(result, records)
return result
}
// Check exact match first
if rs, exists := s.exact[domain]; exists {
return rs.SiteId, true
}
case RecordTypeAAAA:
// Check exact match first
if ips, ok := s.aaaaRecords[domain]; ok {
// Return a copy to prevent external modifications
records = make([]net.IP, len(ips))
copy(records, ips)
return records
}
// Check wildcard patterns
for pattern, ips := range s.aaaaWildcards {
if matchWildcard(pattern, domain) {
records = append(records, ips...)
}
}
if len(records) > 0 {
// Return a copy
result := make([]net.IP, len(records))
copy(result, records)
return result
// Check wildcard matches
for pattern, rs := range s.wildcards {
if matchWildcard(pattern, domain) {
return rs.SiteId, true
}
}
return records
return 0, false
}
// GetRecords returns all IP addresses for a domain and record type.
// The second return value indicates whether the domain exists at all
// (true = domain exists, use NODATA if no records; false = NXDOMAIN).
func (s *DNSRecordStore) GetRecords(domain string, recordType RecordType) ([]net.IP, bool) {
s.mu.RLock()
defer s.mu.RUnlock()
domain = strings.ToLower(dns.Fqdn(domain))
// Check exact match first
if rs, exists := s.exact[domain]; exists {
var ips []net.IP
if recordType == RecordTypeA {
ips = rs.A
} else {
ips = rs.AAAA
}
if len(ips) > 0 {
out := make([]net.IP, len(ips))
copy(out, ips)
return out, true
}
// Domain exists but no records of this type
return nil, true
}
// Check wildcard matches
var records []net.IP
matched := false
for pattern, rs := range s.wildcards {
if !matchWildcard(pattern, domain) {
continue
}
matched = true
if recordType == RecordTypeA {
records = append(records, rs.A...)
} else {
records = append(records, rs.AAAA...)
}
}
if !matched {
return nil, false
}
if len(records) == 0 {
return nil, true
}
out := make([]net.IP, len(records))
copy(out, records)
return out, true
}
// GetPTRRecord returns the domain name for a PTR record query
@@ -288,34 +365,30 @@ func (s *DNSRecordStore) HasRecord(domain string, recordType RecordType) bool {
s.mu.RLock()
defer s.mu.RUnlock()
// Normalize domain to lowercase FQDN
domain = strings.ToLower(dns.Fqdn(domain))
switch recordType {
case RecordTypeA:
// Check exact match
if _, ok := s.aRecords[domain]; ok {
// Check exact match
if rs, exists := s.exact[domain]; exists {
if recordType == RecordTypeA && len(rs.A) > 0 {
return true
}
// Check wildcard patterns
for pattern := range s.aWildcards {
if matchWildcard(pattern, domain) {
return true
}
}
case RecordTypeAAAA:
// Check exact match
if _, ok := s.aaaaRecords[domain]; ok {
if recordType == RecordTypeAAAA && len(rs.AAAA) > 0 {
return true
}
// Check wildcard patterns
for pattern := range s.aaaaWildcards {
if matchWildcard(pattern, domain) {
return true
}
}
}
// Check wildcard matches
for pattern, rs := range s.wildcards {
if !matchWildcard(pattern, domain) {
continue
}
if recordType == RecordTypeA && len(rs.A) > 0 {
return true
}
if recordType == RecordTypeAAAA && len(rs.AAAA) > 0 {
return true
}
}
return false
}
@@ -339,10 +412,8 @@ func (s *DNSRecordStore) Clear() {
s.mu.Lock()
defer s.mu.Unlock()
s.aRecords = make(map[string][]net.IP)
s.aaaaRecords = make(map[string][]net.IP)
s.aWildcards = make(map[string][]net.IP)
s.aaaaWildcards = make(map[string][]net.IP)
s.exact = make(map[string]*recordSet)
s.wildcards = make(map[string]*recordSet)
s.ptrRecords = make(map[string]string)
}
@@ -494,4 +565,4 @@ func IPToReverseDNS(ip net.IP) string {
}
return ""
}
}
+93 -38
View File
@@ -170,38 +170,47 @@ func TestDNSRecordStoreWildcard(t *testing.T) {
// Add wildcard records
wildcardIP := net.ParseIP("10.0.0.1")
err := store.AddRecord("*.autoco.internal", wildcardIP)
err := store.AddRecord("*.autoco.internal", wildcardIP, 0)
if err != nil {
t.Fatalf("Failed to add wildcard record: %v", err)
}
// Add exact record
exactIP := net.ParseIP("10.0.0.2")
err = store.AddRecord("exact.autoco.internal", exactIP)
err = store.AddRecord("exact.autoco.internal", exactIP, 0)
if err != nil {
t.Fatalf("Failed to add exact record: %v", err)
}
// Test exact match takes precedence
ips := store.GetRecords("exact.autoco.internal.", RecordTypeA)
ips, exists := store.GetRecords("exact.autoco.internal.", RecordTypeA)
if !exists {
t.Error("Expected domain to exist")
}
if len(ips) != 1 {
t.Errorf("Expected 1 IP for exact match, got %d", len(ips))
}
if !ips[0].Equal(exactIP) {
if len(ips) > 0 && !ips[0].Equal(exactIP) {
t.Errorf("Expected exact IP %v, got %v", exactIP, ips[0])
}
// Test wildcard match
ips = store.GetRecords("host.autoco.internal.", RecordTypeA)
ips, exists = store.GetRecords("host.autoco.internal.", RecordTypeA)
if !exists {
t.Error("Expected wildcard match to exist")
}
if len(ips) != 1 {
t.Errorf("Expected 1 IP for wildcard match, got %d", len(ips))
}
if !ips[0].Equal(wildcardIP) {
if len(ips) > 0 && !ips[0].Equal(wildcardIP) {
t.Errorf("Expected wildcard IP %v, got %v", wildcardIP, ips[0])
}
// Test non-match (base domain)
ips = store.GetRecords("autoco.internal.", RecordTypeA)
ips, exists = store.GetRecords("autoco.internal.", RecordTypeA)
if exists {
t.Error("Expected base domain to not exist")
}
if len(ips) != 0 {
t.Errorf("Expected 0 IPs for base domain, got %d", len(ips))
}
@@ -212,13 +221,16 @@ func TestDNSRecordStoreComplexWildcard(t *testing.T) {
// Add complex wildcard pattern
ip1 := net.ParseIP("10.0.0.1")
err := store.AddRecord("*.host-0?.autoco.internal", ip1)
err := store.AddRecord("*.host-0?.autoco.internal", ip1, 0)
if err != nil {
t.Fatalf("Failed to add wildcard record: %v", err)
}
// Test matching domain
ips := store.GetRecords("sub.host-01.autoco.internal.", RecordTypeA)
ips, exists := store.GetRecords("sub.host-01.autoco.internal.", RecordTypeA)
if !exists {
t.Error("Expected complex wildcard match to exist")
}
if len(ips) != 1 {
t.Errorf("Expected 1 IP for complex wildcard match, got %d", len(ips))
}
@@ -227,13 +239,19 @@ func TestDNSRecordStoreComplexWildcard(t *testing.T) {
}
// Test non-matching domain (missing prefix)
ips = store.GetRecords("host-01.autoco.internal.", RecordTypeA)
ips, exists = store.GetRecords("host-01.autoco.internal.", RecordTypeA)
if exists {
t.Error("Expected domain without prefix to not exist")
}
if len(ips) != 0 {
t.Errorf("Expected 0 IPs for domain without prefix, got %d", len(ips))
}
// Test non-matching domain (wrong ? position)
ips = store.GetRecords("sub.host-012.autoco.internal.", RecordTypeA)
ips, exists = store.GetRecords("sub.host-012.autoco.internal.", RecordTypeA)
if exists {
t.Error("Expected domain with wrong ? match to not exist")
}
if len(ips) != 0 {
t.Errorf("Expected 0 IPs for domain with wrong ? match, got %d", len(ips))
}
@@ -244,13 +262,16 @@ func TestDNSRecordStoreRemoveWildcard(t *testing.T) {
// Add wildcard record
ip := net.ParseIP("10.0.0.1")
err := store.AddRecord("*.autoco.internal", ip)
err := store.AddRecord("*.autoco.internal", ip, 0)
if err != nil {
t.Fatalf("Failed to add wildcard record: %v", err)
}
// Verify it exists
ips := store.GetRecords("host.autoco.internal.", RecordTypeA)
ips, exists := store.GetRecords("host.autoco.internal.", RecordTypeA)
if !exists {
t.Error("Expected domain to exist before removal")
}
if len(ips) != 1 {
t.Errorf("Expected 1 IP before removal, got %d", len(ips))
}
@@ -259,7 +280,10 @@ func TestDNSRecordStoreRemoveWildcard(t *testing.T) {
store.RemoveRecord("*.autoco.internal", nil)
// Verify it's gone
ips = store.GetRecords("host.autoco.internal.", RecordTypeA)
ips, exists = store.GetRecords("host.autoco.internal.", RecordTypeA)
if exists {
t.Error("Expected domain to not exist after removal")
}
if len(ips) != 0 {
t.Errorf("Expected 0 IPs after removal, got %d", len(ips))
}
@@ -273,36 +297,36 @@ func TestDNSRecordStoreMultipleWildcards(t *testing.T) {
ip2 := net.ParseIP("10.0.0.2")
ip3 := net.ParseIP("10.0.0.3")
err := store.AddRecord("*.prod.autoco.internal", ip1)
err := store.AddRecord("*.prod.autoco.internal", ip1, 0)
if err != nil {
t.Fatalf("Failed to add first wildcard: %v", err)
}
err = store.AddRecord("*.dev.autoco.internal", ip2)
err = store.AddRecord("*.dev.autoco.internal", ip2, 0)
if err != nil {
t.Fatalf("Failed to add second wildcard: %v", err)
}
// Add a broader wildcard that matches both
err = store.AddRecord("*.autoco.internal", ip3)
err = store.AddRecord("*.autoco.internal", ip3, 0)
if err != nil {
t.Fatalf("Failed to add third wildcard: %v", err)
}
// Test domain matching only the prod pattern and the broad pattern
ips := store.GetRecords("host.prod.autoco.internal.", RecordTypeA)
ips, _ := store.GetRecords("host.prod.autoco.internal.", RecordTypeA)
if len(ips) != 2 {
t.Errorf("Expected 2 IPs (prod + broad), got %d", len(ips))
}
// Test domain matching only the dev pattern and the broad pattern
ips = store.GetRecords("service.dev.autoco.internal.", RecordTypeA)
ips, _ = store.GetRecords("service.dev.autoco.internal.", RecordTypeA)
if len(ips) != 2 {
t.Errorf("Expected 2 IPs (dev + broad), got %d", len(ips))
}
// Test domain matching only the broad pattern
ips = store.GetRecords("host.test.autoco.internal.", RecordTypeA)
ips, _ = store.GetRecords("host.test.autoco.internal.", RecordTypeA)
if len(ips) != 1 {
t.Errorf("Expected 1 IP (broad only), got %d", len(ips))
}
@@ -313,13 +337,13 @@ func TestDNSRecordStoreIPv6Wildcard(t *testing.T) {
// Add IPv6 wildcard record
ip := net.ParseIP("2001:db8::1")
err := store.AddRecord("*.autoco.internal", ip)
err := store.AddRecord("*.autoco.internal", ip, 0)
if err != nil {
t.Fatalf("Failed to add IPv6 wildcard record: %v", err)
}
// Test wildcard match for IPv6
ips := store.GetRecords("host.autoco.internal.", RecordTypeAAAA)
ips, _ := store.GetRecords("host.autoco.internal.", RecordTypeAAAA)
if len(ips) != 1 {
t.Errorf("Expected 1 IPv6 for wildcard match, got %d", len(ips))
}
@@ -333,7 +357,7 @@ func TestHasRecordWildcard(t *testing.T) {
// Add wildcard record
ip := net.ParseIP("10.0.0.1")
err := store.AddRecord("*.autoco.internal", ip)
err := store.AddRecord("*.autoco.internal", ip, 0)
if err != nil {
t.Fatalf("Failed to add wildcard record: %v", err)
}
@@ -354,7 +378,7 @@ func TestDNSRecordStoreCaseInsensitive(t *testing.T) {
// Add record with mixed case
ip := net.ParseIP("10.0.0.1")
err := store.AddRecord("MyHost.AutoCo.Internal", ip)
err := store.AddRecord("MyHost.AutoCo.Internal", ip, 0)
if err != nil {
t.Fatalf("Failed to add mixed case record: %v", err)
}
@@ -368,7 +392,7 @@ func TestDNSRecordStoreCaseInsensitive(t *testing.T) {
}
for _, domain := range testCases {
ips := store.GetRecords(domain, RecordTypeA)
ips, _ := store.GetRecords(domain, RecordTypeA)
if len(ips) != 1 {
t.Errorf("Expected 1 IP for domain %q, got %d", domain, len(ips))
}
@@ -379,7 +403,7 @@ func TestDNSRecordStoreCaseInsensitive(t *testing.T) {
// Test wildcard with mixed case
wildcardIP := net.ParseIP("10.0.0.2")
err = store.AddRecord("*.Example.Com", wildcardIP)
err = store.AddRecord("*.Example.Com", wildcardIP, 0)
if err != nil {
t.Fatalf("Failed to add mixed case wildcard: %v", err)
}
@@ -392,7 +416,7 @@ func TestDNSRecordStoreCaseInsensitive(t *testing.T) {
}
for _, domain := range wildcardTestCases {
ips := store.GetRecords(domain, RecordTypeA)
ips, _ := store.GetRecords(domain, RecordTypeA)
if len(ips) != 1 {
t.Errorf("Expected 1 IP for wildcard domain %q, got %d", domain, len(ips))
}
@@ -403,7 +427,7 @@ func TestDNSRecordStoreCaseInsensitive(t *testing.T) {
// Test removal with different case
store.RemoveRecord("MYHOST.AUTOCO.INTERNAL", nil)
ips := store.GetRecords("myhost.autoco.internal.", RecordTypeA)
ips, _ := store.GetRecords("myhost.autoco.internal.", RecordTypeA)
if len(ips) != 0 {
t.Errorf("Expected 0 IPs after removal, got %d", len(ips))
}
@@ -665,7 +689,7 @@ func TestClearPTRRecords(t *testing.T) {
store.AddPTRRecord(ip2, "host2.example.com.")
// Add some A records too
store.AddRecord("test.example.com.", net.ParseIP("10.0.0.1"))
store.AddRecord("test.example.com.", net.ParseIP("10.0.0.1"), 0)
// Verify PTR records exist
if !store.HasPTRRecord("1.1.168.192.in-addr.arpa.") {
@@ -695,7 +719,7 @@ func TestAutomaticPTRRecordOnAdd(t *testing.T) {
// Add an A record - should automatically add PTR record
domain := "host.example.com."
ip := net.ParseIP("192.168.1.100")
err := store.AddRecord(domain, ip)
err := store.AddRecord(domain, ip, 0)
if err != nil {
t.Fatalf("Failed to add A record: %v", err)
}
@@ -713,7 +737,7 @@ func TestAutomaticPTRRecordOnAdd(t *testing.T) {
// Add AAAA record - should also automatically add PTR record
domain6 := "ipv6host.example.com."
ip6 := net.ParseIP("2001:db8::1")
err = store.AddRecord(domain6, ip6)
err = store.AddRecord(domain6, ip6, 0)
if err != nil {
t.Fatalf("Failed to add AAAA record: %v", err)
}
@@ -735,7 +759,7 @@ func TestAutomaticPTRRecordOnRemove(t *testing.T) {
// Add an A record (with automatic PTR)
domain := "host.example.com."
ip := net.ParseIP("192.168.1.100")
store.AddRecord(domain, ip)
store.AddRecord(domain, ip, 0)
// Verify PTR exists
reverseDomain := "100.1.168.192.in-addr.arpa."
@@ -752,12 +776,43 @@ func TestAutomaticPTRRecordOnRemove(t *testing.T) {
}
// Verify A record is also gone
ips := store.GetRecords(domain, RecordTypeA)
ips, _ := store.GetRecords(domain, RecordTypeA)
if len(ips) != 0 {
t.Errorf("Expected A record to be removed, got %d records", len(ips))
}
}
func TestRemoveRecordForSiteKeepsSharedAliasIP(t *testing.T) {
store := NewDNSRecordStore()
domain := "shared.example.com."
ip := net.ParseIP("192.168.1.100")
if err := store.AddRecord(domain, ip, 10); err != nil {
t.Fatalf("Failed to add record for site 10: %v", err)
}
if err := store.AddRecord(domain, ip, 20); err != nil {
t.Fatalf("Failed to add record for site 20: %v", err)
}
store.RemoveRecordForSite(domain, ip, 10)
ips, exists := store.GetRecords(domain, RecordTypeA)
if !exists {
t.Fatal("Expected shared record to still exist after removing one site owner")
}
if len(ips) != 1 || !ips[0].Equal(ip) {
t.Fatalf("Expected shared IP to remain after first owner removal, got %v", ips)
}
store.RemoveRecordForSite(domain, ip, 20)
ips, exists = store.GetRecords(domain, RecordTypeA)
if exists {
t.Fatalf("Expected domain to be removed after last owner removal, got %v", ips)
}
}
func TestAutomaticPTRRecordOnRemoveAll(t *testing.T) {
store := NewDNSRecordStore()
@@ -765,8 +820,8 @@ func TestAutomaticPTRRecordOnRemoveAll(t *testing.T) {
domain := "host.example.com."
ip1 := net.ParseIP("192.168.1.100")
ip2 := net.ParseIP("192.168.1.101")
store.AddRecord(domain, ip1)
store.AddRecord(domain, ip2)
store.AddRecord(domain, ip1, 0)
store.AddRecord(domain, ip2, 0)
// Verify both PTR records exist
reverseDomain1 := "100.1.168.192.in-addr.arpa."
@@ -796,7 +851,7 @@ func TestNoPTRForWildcardRecords(t *testing.T) {
// Add wildcard record - should NOT create PTR record
domain := "*.example.com."
ip := net.ParseIP("192.168.1.100")
err := store.AddRecord(domain, ip)
err := store.AddRecord(domain, ip, 0)
if err != nil {
t.Fatalf("Failed to add wildcard record: %v", err)
}
@@ -820,7 +875,7 @@ func TestPTRRecordOverwrite(t *testing.T) {
// Add first domain with IP
domain1 := "host1.example.com."
ip := net.ParseIP("192.168.1.100")
store.AddRecord(domain1, ip)
store.AddRecord(domain1, ip, 0)
// Verify PTR points to first domain
reverseDomain := "100.1.168.192.in-addr.arpa."
@@ -834,7 +889,7 @@ func TestPTRRecordOverwrite(t *testing.T) {
// Add second domain with same IP - should overwrite PTR
domain2 := "host2.example.com."
store.AddRecord(domain2, ip)
store.AddRecord(domain2, ip, 0)
// Verify PTR now points to second domain (last one added)
result, ok = store.GetPTRRecord(reverseDomain)
+13 -1
View File
@@ -13,4 +13,16 @@ func SetupDNSOverride(interfaceName string, proxyIp netip.Addr) error {
// RestoreDNSOverride is a no-op on Android
func RestoreDNSOverride() error {
return nil
}
}
// CleanupStaleState is a no-op on Android as DNS configuration is handled by the VpnService API
func CleanupStaleState(interfaceName string) error {
_ = interfaceName
return nil
}
// ForceResetDNS is a no-op on Android.
func ForceResetDNS(interfaceName string) error {
_ = interfaceName
return nil
}
+52
View File
@@ -15,6 +15,13 @@ var configurator platform.DNSConfigurator
// SetupDNSOverride configures the system DNS to use the DNS proxy on macOS
// Uses scutil for DNS configuration
func SetupDNSOverride(interfaceName string, proxyIp netip.Addr) error {
// Defensively clear any stale DNS state from a previous unclean shutdown
// before installing the new override. This makes a second tunnel start
// safe even if the previous client crashed without restoring DNS.
if err := CleanupStaleState(interfaceName); err != nil {
logger.Warn("Pre-setup stale DNS cleanup failed (continuing): %v", err)
}
var err error
configurator, err = platform.NewDarwinDNSConfigurator()
if err != nil {
@@ -61,3 +68,48 @@ func RestoreDNSOverride() error {
logger.Info("DNS configuration restored successfully")
return nil
}
// CleanupStaleState removes any stale DNS configuration left over from a previous
// unclean shutdown (e.g., system crash, power loss while tunnel was active).
// This function should be called early during startup, before any network operations,
// to ensure DNS is working properly.
//
// On macOS, this cleans up any scutil DNS keys that were created but not removed.
func CleanupStaleState(interfaceName string) error {
_ = interfaceName
if err := platform.CleanupStaleDarwinDNS(); err != nil {
logger.Warn("Failed to cleanup stale Darwin DNS config: %v", err)
return fmt.Errorf("Darwin DNS cleanup: %w", err)
}
logger.Info("Stale DNS state cleanup completed successfully")
return nil
}
// ForceResetDNS forcibly clears any DNS override state, whether or not the
// current process installed it. This is intended for the "reset-dns" CLI
// command and for the watchdog process to recover from a stuck override
// left behind by a crashed client.
func ForceResetDNS(interfaceName string) error {
logger.Info("Forcing DNS reset on Darwin (interface=%s)", interfaceName)
// First clean up any persisted state from a previous session.
cleanupErr := CleanupStaleState(interfaceName)
// Then, if the current process happens to hold a live configurator,
// instruct it to restore DNS as well so in-memory state is consistent.
if configurator != nil {
if err := configurator.RestoreDNS(); err != nil {
logger.Warn("ForceResetDNS: in-memory restore failed: %v", err)
}
configurator = nil
}
// As a last-resort defense, sweep any scutil keys matching our naming
// convention even if no state file exists.
if err := platform.SweepOlmScutilKeys(); err != nil {
logger.Warn("ForceResetDNS: scutil sweep failed: %v", err)
}
return cleanupErr
}
+13 -1
View File
@@ -12,4 +12,16 @@ func SetupDNSOverride(interfaceName string, proxyIp netip.Addr) error {
// RestoreDNSOverride is a no-op on iOS as DNS configuration is handled by the system
func RestoreDNSOverride() error {
return nil
}
}
// CleanupStaleState is a no-op on iOS as DNS configuration is handled by the system
func CleanupStaleState(interfaceName string) error {
_ = interfaceName
return nil
}
// ForceResetDNS is a no-op on iOS.
func ForceResetDNS(interfaceName string) error {
_ = interfaceName
return nil
}
+75
View File
@@ -15,6 +15,13 @@ var configurator platform.DNSConfigurator
// SetupDNSOverride configures the system DNS to use the DNS proxy on Linux/FreeBSD
// Detects the DNS manager by reading /etc/resolv.conf and verifying runtime availability
func SetupDNSOverride(interfaceName string, proxyIp netip.Addr) error {
// Defensively clear any stale DNS state from a previous unclean shutdown
// before installing the new override. This makes a second tunnel start
// safe even if the previous client crashed without restoring DNS.
if err := CleanupStaleState(interfaceName); err != nil {
logger.Warn("Pre-setup stale DNS cleanup failed (continuing): %v", err)
}
var err error
// Detect which DNS manager is in use by checking /etc/resolv.conf and runtime availability
@@ -98,3 +105,71 @@ func RestoreDNSOverride() error {
logger.Info("DNS configuration restored successfully")
return nil
}
// CleanupStaleState removes any stale DNS configuration left over from a previous
// unclean shutdown (e.g., system crash, power loss while tunnel was active).
// This function should be called early during startup, before any network operations,
// to ensure DNS is working properly.
//
// It checks and cleans up stale state from all supported DNS managers:
// - NetworkManager: removes /etc/NetworkManager/conf.d/olm-dns.conf
// - resolvconf: removes entry for the provided interface
// - File-based: restores /etc/resolv.conf from backup if it exists
//
// This is safe to call even if no stale state exists.
func CleanupStaleState(interfaceName string) error {
var errs []error
// Clean up NetworkManager stale config
if err := platform.CleanupStaleNetworkManagerDNS(); err != nil {
logger.Warn("Failed to cleanup stale NetworkManager DNS config: %v", err)
errs = append(errs, fmt.Errorf("NetworkManager cleanup: %w", err))
} else {
logger.Debug("NetworkManager DNS cleanup completed")
}
// Clean up resolvconf stale entries for the provided interface
if err := platform.CleanupStaleResolvconfDNS(interfaceName); err != nil {
logger.Warn("Failed to cleanup stale resolvconf DNS config: %v", err)
errs = append(errs, fmt.Errorf("resolvconf cleanup: %w", err))
} else {
logger.Debug("resolvconf DNS cleanup completed")
}
// Clean up file-based stale backup
if err := platform.CleanupStaleFileDNS(); err != nil {
logger.Warn("Failed to cleanup stale file-based DNS config: %v", err)
errs = append(errs, fmt.Errorf("file DNS cleanup: %w", err))
} else {
logger.Debug("File-based DNS cleanup completed")
}
if len(errs) > 0 {
return fmt.Errorf("some DNS cleanup operations failed: %v", errs)
}
logger.Info("Stale DNS state cleanup completed successfully")
return nil
}
// ForceResetDNS forcibly clears any DNS override state, whether or not the
// current process installed it. This is intended for the "reset-dns" CLI
// command and for the watchdog process to recover from a stuck override
// left behind by a crashed client.
func ForceResetDNS(interfaceName string) error {
logger.Info("Forcing DNS reset on Linux/FreeBSD (interface=%s)", interfaceName)
// First clean up any persisted state from a previous session.
cleanupErr := CleanupStaleState(interfaceName)
// Then, if the current process happens to hold a live configurator,
// instruct it to restore DNS as well so in-memory state is consistent.
if configurator != nil {
if err := configurator.RestoreDNS(); err != nil {
logger.Warn("ForceResetDNS: in-memory restore failed: %v", err)
}
configurator = nil
}
return cleanupErr
}
+36
View File
@@ -15,6 +15,12 @@ var configurator platform.DNSConfigurator
// SetupDNSOverride configures the system DNS to use the DNS proxy on Windows
// Uses registry-based configuration (automatically extracts interface GUID)
func SetupDNSOverride(interfaceName string, proxyIp netip.Addr) error {
// Defensively clear any stale DNS state from a previous unclean shutdown
// before installing the new override.
if err := CleanupStaleState(interfaceName); err != nil {
logger.Warn("Pre-setup stale DNS cleanup failed (continuing): %v", err)
}
var err error
configurator, err = platform.NewWindowsDNSConfigurator(interfaceName)
if err != nil {
@@ -61,3 +67,33 @@ func RestoreDNSOverride() error {
logger.Info("DNS configuration restored successfully")
return nil
}
// CleanupStaleState removes any stale DNS configuration left over from a previous
// unclean shutdown (e.g., system crash, power loss while tunnel was active).
// This function should be called early during startup, before any network operations,
// to ensure DNS is working properly.
//
// On Windows, DNS configuration is tied to the interface GUID. When the WireGuard
// interface is recreated, it gets a new GUID, so there's no stale state to clean up.
func CleanupStaleState(interfaceName string) error {
// Windows DNS configuration via registry is interface-specific.
// When the WireGuard interface is recreated, it gets a new GUID,
// so there's no leftover state to clean up from previous sessions.
_ = interfaceName
logger.Debug("Windows DNS cleanup: no stale state to clean (interface-specific)")
return nil
}
// ForceResetDNS forcibly clears any DNS override state. On Windows this is
// largely a no-op because the registry override is tied to the interface
// GUID and is reclaimed when the interface is torn down.
func ForceResetDNS(interfaceName string) error {
logger.Info("Forcing DNS reset on Windows (interface=%s)", interfaceName)
if configurator != nil {
if err := configurator.RestoreDNS(); err != nil {
logger.Warn("ForceResetDNS: in-memory restore failed: %v", err)
}
configurator = nil
}
return CleanupStaleState(interfaceName)
}
+136
View File
@@ -0,0 +1,136 @@
package olm
import (
"context"
"fmt"
"net"
"net/http"
"os"
"time"
"github.com/fosrl/newt/logger"
)
// WatchdogConfig configures the DNS override watchdog. The watchdog runs as
// an external process (spawned via SpawnWatchdog) and monitors a parent olm
// process. When the parent appears to have died without restoring DNS, the
// watchdog forcibly resets the system DNS configuration.
type WatchdogConfig struct {
// ParentPID is the PID of the olm process that installed the DNS
// override. The watchdog exits when this PID is no longer alive.
ParentPID int
// SocketPath is the path to the olm Unix domain socket (or named pipe
// on Windows). The watchdog uses it as a secondary liveness signal.
// May be empty if no socket-based API is enabled.
SocketPath string
// InterfaceName is the name of the WireGuard interface whose DNS
// override should be reset on parent death.
InterfaceName string
// CheckInterval is how often to poll the parent's liveness.
// Defaults to 5 seconds when zero.
CheckInterval time.Duration
// FailureThreshold is the number of consecutive failed liveness checks
// before the watchdog declares the parent dead and resets DNS.
// Defaults to 3 when zero.
FailureThreshold int
}
// RunWatchdog runs the watchdog loop in the current process until either
// (a) the parent dies and DNS is reset, or (b) ctx is cancelled.
func RunWatchdog(ctx context.Context, cfg WatchdogConfig) error {
if cfg.ParentPID <= 0 {
return fmt.Errorf("watchdog: invalid parent PID %d", cfg.ParentPID)
}
if cfg.CheckInterval <= 0 {
cfg.CheckInterval = 5 * time.Second
}
if cfg.FailureThreshold <= 0 {
cfg.FailureThreshold = 3
}
logger.Info("DNS watchdog started: parent=%d interval=%s threshold=%d socket=%q interface=%q",
cfg.ParentPID, cfg.CheckInterval, cfg.FailureThreshold, cfg.SocketPath, cfg.InterfaceName)
ticker := time.NewTicker(cfg.CheckInterval)
defer ticker.Stop()
consecutiveFailures := 0
for {
select {
case <-ctx.Done():
logger.Info("DNS watchdog context cancelled, exiting cleanly")
return ctx.Err()
case <-ticker.C:
}
alive := isParentAlive(cfg.ParentPID, cfg.SocketPath)
if alive {
if consecutiveFailures > 0 {
logger.Debug("DNS watchdog: parent recovered after %d failures", consecutiveFailures)
}
consecutiveFailures = 0
continue
}
consecutiveFailures++
logger.Warn("DNS watchdog: parent liveness check failed (%d/%d)",
consecutiveFailures, cfg.FailureThreshold)
if consecutiveFailures >= cfg.FailureThreshold {
logger.Warn("DNS watchdog: parent declared dead, forcing DNS reset")
if err := ForceResetDNS(cfg.InterfaceName); err != nil {
logger.Error("DNS watchdog: ForceResetDNS failed: %v", err)
return err
}
logger.Info("DNS watchdog: DNS reset complete, exiting")
return nil
}
}
}
// isParentAlive returns true if the parent process appears to be alive. It
// considers the parent alive if EITHER the PID is still running OR the
// socket-based health endpoint responds. This dual check avoids false
// positives where one signal is flaky (e.g., socket blocked but process
// still recovering).
func isParentAlive(pid int, socketPath string) bool {
if pidAlive(pid) {
return true
}
// Process is gone; double-check via socket to avoid races where PID
// recycling or signal-0 quirks lie to us. Socket should already be
// gone too.
if socketPath != "" && socketHealthy(socketPath) {
return true
}
return false
}
// socketHealthy attempts a fast /health request over the unix socket.
func socketHealthy(socketPath string) bool {
if _, err := os.Stat(socketPath); err != nil {
return false
}
client := &http.Client{
Timeout: 2 * time.Second,
Transport: &http.Transport{
DialContext: func(ctx context.Context, _, _ string) (net.Conn, error) {
d := net.Dialer{Timeout: 2 * time.Second}
return d.DialContext(ctx, "unix", socketPath)
},
},
}
resp, err := client.Get("http://localhost/health")
if err != nil {
return false
}
defer resp.Body.Close()
return resp.StatusCode == http.StatusOK
}
+119
View File
@@ -0,0 +1,119 @@
//go:build !windows
package olm
import (
"fmt"
"os"
"os/exec"
"strconv"
"syscall"
"github.com/fosrl/newt/logger"
)
// SpawnWatchdogConfig captures the inputs needed to launch the external
// watchdog subprocess that monitors the calling olm process and forces a
// DNS reset if the parent dies before restoring DNS.
type SpawnWatchdogConfig struct {
// Executable is the path to the binary that will host the watchdog
// (typically os.Executable()). The binary must understand the
// watchdog subcommand layout described below.
Executable string
// Subcommand is the argv prefix the binary uses to enter watchdog
// mode (e.g., []string{"watchdog"} or []string{"dns", "watchdog"}).
Subcommand []string
// InterfaceName is the WireGuard interface whose DNS override should
// be reset if the parent dies.
InterfaceName string
// SocketPath is the parent's olm API socket path (may be empty).
SocketPath string
// LogFile, if non-empty, is the path the watchdog writes its stdout
// and stderr to. If empty, /dev/null is used.
LogFile string
}
// SpawnWatchdog launches the watchdog subprocess in a detached process group
// so that it survives the death of the parent. The returned *exec.Cmd is the
// handle the parent should call StopWatchdog on during clean shutdown.
//
// The spawned process is invoked as:
//
// <Executable> <Subcommand...> --parent-pid=<ppid> \
// --interface=<InterfaceName> [--socket=<SocketPath>]
//
// Both pangolin (cli) and olm should map their watchdog subcommand to
// RunWatchdog.
func SpawnWatchdog(cfg SpawnWatchdogConfig) (*exec.Cmd, error) {
if cfg.Executable == "" {
return nil, fmt.Errorf("watchdog: executable is required")
}
if len(cfg.Subcommand) == 0 {
return nil, fmt.Errorf("watchdog: subcommand is required")
}
args := append([]string{}, cfg.Subcommand...)
args = append(args,
"--parent-pid="+strconv.Itoa(os.Getpid()),
"--interface="+cfg.InterfaceName,
)
if cfg.SocketPath != "" {
args = append(args, "--socket="+cfg.SocketPath)
}
cmd := exec.Command(cfg.Executable, args...)
// Detach: new session so the watchdog is not killed by a signal
// delivered to the parent's process group.
cmd.SysProcAttr = &syscall.SysProcAttr{
Setsid: true,
}
// Direct watchdog output to a log file or /dev/null so it doesn't
// share file descriptors with the parent's TTY.
logTarget := cfg.LogFile
if logTarget == "" {
logTarget = os.DevNull
}
logFile, err := os.OpenFile(logTarget, os.O_CREATE|os.O_WRONLY|os.O_APPEND, 0o644)
if err != nil {
return nil, fmt.Errorf("watchdog: open log file: %w", err)
}
cmd.Stdin = nil
cmd.Stdout = logFile
cmd.Stderr = logFile
if err := cmd.Start(); err != nil {
_ = logFile.Close()
return nil, fmt.Errorf("watchdog: start: %w", err)
}
// We don't need our handle on the log file after the subprocess
// inherits it.
_ = logFile.Close()
logger.Info("DNS watchdog spawned (pid=%d, exe=%s)", cmd.Process.Pid, cfg.Executable)
return cmd, nil
}
// StopWatchdog asks the watchdog to exit cleanly via SIGTERM and reaps it.
// Safe to call with a nil cmd.
func StopWatchdog(cmd *exec.Cmd) {
if cmd == nil || cmd.Process == nil {
return
}
pid := cmd.Process.Pid
if err := cmd.Process.Signal(syscall.SIGTERM); err != nil {
logger.Debug("DNS watchdog stop signal failed (pid=%d): %v", pid, err)
}
// Reap in the background; we don't want to block shutdown if the
// watchdog is wedged.
go func() {
_ = cmd.Wait()
logger.Debug("DNS watchdog (pid=%d) reaped", pid)
}()
}
+29
View File
@@ -0,0 +1,29 @@
//go:build windows
package olm
import (
"os/exec"
)
// SpawnWatchdogConfig is provided on Windows for API symmetry but the
// watchdog itself is effectively a no-op there (see watchdog_windows.go).
type SpawnWatchdogConfig struct {
Executable string
Subcommand []string
InterfaceName string
SocketPath string
LogFile string
}
// SpawnWatchdog is a no-op on Windows; DNS overrides are interface-GUID
// scoped and reclaimed when the interface is removed.
func SpawnWatchdog(cfg SpawnWatchdogConfig) (*exec.Cmd, error) {
_ = cfg
return nil, nil
}
// StopWatchdog is a no-op on Windows.
func StopWatchdog(cmd *exec.Cmd) {
_ = cmd
}
+25
View File
@@ -0,0 +1,25 @@
//go:build !windows
package olm
import (
"os"
"syscall"
)
// pidAlive returns true if the process with the given PID is still alive.
// On Unix-like systems we use signal 0, which performs error checking but
// does not deliver an actual signal.
func pidAlive(pid int) bool {
if pid <= 0 {
return false
}
proc, err := os.FindProcess(pid)
if err != nil {
return false
}
if err := proc.Signal(syscall.Signal(0)); err != nil {
return false
}
return true
}
+15
View File
@@ -0,0 +1,15 @@
//go:build windows
package olm
// pidAlive on Windows. Reliable PID probing on Windows requires syscall
// OpenProcess with PROCESS_QUERY_LIMITED_INFORMATION followed by
// GetExitCodeProcess, which is non-trivial. Since DNS override on Windows
// is interface-GUID-scoped and is naturally cleaned up when the WireGuard
// interface goes away, the watchdog is effectively a no-op on Windows.
// We always report the parent as alive so the watchdog never tears down
// DNS based on PID checks.
func pidAlive(pid int) bool {
_ = pid
return true
}
+123
View File
@@ -417,3 +417,126 @@ func (d *DarwinDNSConfigurator) clearState() error {
logger.Debug("Cleared DNS state file")
return nil
}
// CleanupStaleDarwinDNS removes any stale DNS configuration left by the Darwin
// configurator from a previous unclean shutdown. This is a static function that can be
// called without creating a configurator instance, useful for cleanup before network operations.
func CleanupStaleDarwinDNS() error {
// Always sweep orphaned Olm scutil keys regardless of whether a state
// file exists. This protects against cases where the state file was
// lost (e.g., user home wiped, write failed) but DNS keys are still
// installed in the running scutil session.
defer func() {
_ = SweepOlmScutilKeys()
// Flush DNS cache after any sweep so changes take effect.
_ = exec.Command(dscacheutilPath, "-flushcache").Run()
_ = exec.Command("killall", "-HUP", "mDNSResponder").Run()
}()
stateFilePath := getDNSStateFilePath()
// Check if state file exists
data, err := os.ReadFile(stateFilePath)
if err != nil {
if os.IsNotExist(err) {
// No state file, nothing to clean up
return nil
}
return fmt.Errorf("read state file: %w", err)
}
var state DNSPersistentState
if err := json.Unmarshal(data, &state); err != nil {
// Invalid state file, remove it
os.Remove(stateFilePath)
return nil
}
if len(state.CreatedKeys) == 0 {
// No keys to clean up
return nil
}
logger.Info("Found DNS state from previous session, cleaning up %d keys", len(state.CreatedKeys))
// Remove all keys from previous session using scutil directly
for _, key := range state.CreatedKeys {
logger.Debug("Removing leftover DNS key: %s", key)
cmd := fmt.Sprintf("open\nremove %s\nquit\n", key)
scutilCmd := exec.Command(scutilPath)
scutilCmd.Stdin = strings.NewReader(cmd)
if err := scutilCmd.Run(); err != nil {
logger.Warn("Failed to remove DNS key %s: %v", key, err)
}
}
// Clear state file
if err := os.Remove(stateFilePath); err != nil && !os.IsNotExist(err) {
logger.Warn("Failed to clear DNS state file: %v", err)
}
// Flush DNS cache after cleanup
cacheCmd := exec.Command(dscacheutilPath, "-flushcache")
_ = cacheCmd.Run()
killCmd := exec.Command("killall", "-HUP", "mDNSResponder")
_ = killCmd.Run()
return nil
}
// SweepOlmScutilKeys enumerates scutil State:/Network/Service/Olm-* keys and
// removes any that are present. This is a best-effort safety net used when
// state files have been lost or never written.
func SweepOlmScutilKeys() error {
// list scutil keys matching our naming convention
listOutput, err := runScutilOnce("list State:/Network/Service/Olm-.*/DNS\n")
if err != nil {
return fmt.Errorf("scutil list: %w", err)
}
var keys []string
scanner := bufio.NewScanner(bytes.NewReader(listOutput))
for scanner.Scan() {
line := strings.TrimSpace(scanner.Text())
// scutil output format: subKey [0] = State:/Network/Service/Olm-Override/DNS
idx := strings.Index(line, "State:/Network/Service/Olm-")
if idx < 0 {
continue
}
key := strings.TrimSpace(line[idx:])
if key != "" {
keys = append(keys, key)
}
}
if len(keys) == 0 {
return nil
}
logger.Info("Sweeping %d orphaned Olm scutil DNS keys", len(keys))
var commands strings.Builder
for _, key := range keys {
commands.WriteString(fmt.Sprintf("remove %s\n", key))
}
if _, err := runScutilOnce(commands.String()); err != nil {
return fmt.Errorf("scutil sweep remove: %w", err)
}
return nil
}
// runScutilOnce runs a one-shot scutil command sequence wrapped with open/quit
// without requiring a configurator instance.
func runScutilOnce(commands string) ([]byte, error) {
wrapped := fmt.Sprintf("open\n%squit\n", commands)
cmd := exec.Command(scutilPath)
cmd.Stdin = strings.NewReader(wrapped)
output, err := cmd.CombinedOutput()
if err != nil {
return nil, fmt.Errorf("scutil command failed: %w, output: %s", err, output)
}
return output, nil
}
+24
View File
@@ -218,3 +218,27 @@ func copyFile(src, dst string) error {
return nil
}
// CleanupStaleFileDNS removes any stale DNS configuration left by the file-based
// configurator from a previous unclean shutdown. This is a static function that can be
// called without creating a configurator instance, useful for cleanup before network operations.
func CleanupStaleFileDNS() error {
// Check if backup file exists from a previous session
if _, err := os.Stat(resolvConfBackupPath); os.IsNotExist(err) {
// No backup file, nothing to clean up
return nil
}
// A backup exists, which means we crashed while DNS was configured
// Restore the original resolv.conf
if err := copyFile(resolvConfBackupPath, resolvConfPath); err != nil {
return fmt.Errorf("restore from backup during cleanup: %w", err)
}
// Remove backup file
if err := os.Remove(resolvConfBackupPath); err != nil {
return fmt.Errorf("remove backup file during cleanup: %w", err)
}
return nil
}
+272 -6
View File
@@ -4,6 +4,7 @@ package dns
import (
"context"
"encoding/binary"
"errors"
"fmt"
"net/netip"
@@ -16,12 +17,25 @@ import (
const (
// NetworkManager D-Bus constants
networkManagerDest = "org.freedesktop.NetworkManager"
networkManagerDbusObjectNode = "/org/freedesktop/NetworkManager"
networkManagerDbusDNSManagerInterface = "org.freedesktop.NetworkManager.DnsManager"
networkManagerDbusDNSManagerObjectNode = networkManagerDbusObjectNode + "/DnsManager"
networkManagerDbusDNSManagerModeProperty = networkManagerDbusDNSManagerInterface + ".Mode"
networkManagerDbusVersionProperty = "org.freedesktop.NetworkManager.Version"
networkManagerDest = "org.freedesktop.NetworkManager"
networkManagerDbusObjectNode = "/org/freedesktop/NetworkManager"
networkManagerDbusDNSManagerInterface = "org.freedesktop.NetworkManager.DnsManager"
networkManagerDbusDNSManagerObjectNode = networkManagerDbusObjectNode + "/DnsManager"
networkManagerDbusDNSManagerModeProperty = networkManagerDbusDNSManagerInterface + ".Mode"
networkManagerDbusVersionProperty = "org.freedesktop.NetworkManager.Version"
networkManagerDbusActiveConnsProperty = networkManagerDest + ".ActiveConnections"
networkManagerDbusActiveInterface = "org.freedesktop.NetworkManager.Connection.Active"
networkManagerDbusActiveIP4ConfigProperty = networkManagerDbusActiveInterface + ".Ip4Config"
networkManagerDbusActiveIP6ConfigProperty = networkManagerDbusActiveInterface + ".Ip6Config"
networkManagerDbusActiveDevicesProperty = networkManagerDbusActiveInterface + ".Devices"
networkManagerDbusIP4ConfigInterface = "org.freedesktop.NetworkManager.IP4Config"
networkManagerDbusIP6ConfigInterface = "org.freedesktop.NetworkManager.IP6Config"
networkManagerDbusDeviceInterface = "org.freedesktop.NetworkManager.Device"
networkManagerDbusDeviceDhcp4ConfigProp = networkManagerDbusDeviceInterface + ".Dhcp4Config"
networkManagerDbusDeviceDhcp6ConfigProp = networkManagerDbusDeviceInterface + ".Dhcp6Config"
networkManagerDbusDhcp4ConfigInterface = "org.freedesktop.NetworkManager.DHCP4Config"
networkManagerDbusDhcp6ConfigInterface = "org.freedesktop.NetworkManager.DHCP6Config"
networkManagerDbusGetAppliedConnMethod = networkManagerDbusDeviceInterface + ".GetAppliedConnection"
// NetworkManager dispatcher script path
networkManagerDispatcherDir = "/etc/NetworkManager/dispatcher.d"
@@ -301,6 +315,220 @@ func GetNetworkManagerDNSMode() (string, error) {
return mode, nil
}
// GetNetworkManagerNameservers returns the DNS servers NetworkManager knows
// about for every active connection, read live via D-Bus.
//
// olm's own NetworkManager DNS override (see NetworkManagerDNSConfigurator)
// works by writing a [global-dns-domain-*] section to
// /etc/NetworkManager/conf.d/olm-dns.conf and reloading NetworkManager. That
// is NetworkManager's global DNS override mechanism: it replaces the DNS
// servers NetworkManager's DnsManager computes as "effective" system-wide,
// for every connection - not just what gets written to /etc/resolv.conf. So
// once olm's override is active, even each connection's merged
// IP4Config/IP6Config.NameserverData (the previous, sole source used here)
// can end up reporting olm's own proxy address instead of the real network
// DNS.
//
// To recover the real DNS regardless, this also reads two further sources
// that NetworkManager's DNS merging - and therefore olm's global-dns override
// - never touches, since both are populated independently of it:
// - Dhcp4Config/Dhcp6Config.Options["*name_servers"]: the raw nameserver
// list straight from the DHCP lease.
// - Device.GetAppliedConnection()'s ipv4.dns/ipv6.dns: the DNS servers
// explicitly configured on the connection profile itself, e.g. a static
// DNS override set by the user directly in NetworkManager (the
// NetworkManager equivalent of a manually-set Windows adapter DNS).
//
// IP4Config/IP6Config.NameserverData is still queried too, as a fallback for
// setups the other two don't cover. Any of olm's own address that leaks
// through any of these sources is expected to be dropped by the caller via
// SystemDNSMonitor.SetExcludeIP.
func GetNetworkManagerNameservers() ([]netip.Addr, error) {
conn, err := dbus.SystemBus()
if err != nil {
return nil, fmt.Errorf("connect to system bus: %w", err)
}
defer conn.Close()
nm := conn.Object(networkManagerDest, networkManagerDbusObjectNode)
activeVariant, err := nm.GetProperty(networkManagerDbusActiveConnsProperty)
if err != nil {
return nil, fmt.Errorf("get active connections: %w", err)
}
activePaths, ok := activeVariant.Value().([]dbus.ObjectPath)
if !ok {
return nil, errors.New("ActiveConnections is not a list of object paths")
}
ipConfigSources := []struct {
activeProperty string
configIface string
}{
{networkManagerDbusActiveIP4ConfigProperty, networkManagerDbusIP4ConfigInterface},
{networkManagerDbusActiveIP6ConfigProperty, networkManagerDbusIP6ConfigInterface},
}
seen := make(map[netip.Addr]bool)
var servers []netip.Addr
add := func(addr netip.Addr) {
addr = addr.Unmap()
if !addr.IsValid() || addr.IsLoopback() || addr.IsLinkLocalUnicast() {
return
}
if !seen[addr] {
seen[addr] = true
servers = append(servers, addr)
}
}
for _, activePath := range activePaths {
active := conn.Object(networkManagerDest, activePath)
for _, src := range ipConfigSources {
cfgVariant, err := active.GetProperty(src.activeProperty)
if err != nil {
continue
}
cfgPath, ok := cfgVariant.Value().(dbus.ObjectPath)
if !ok || cfgPath == "" || cfgPath == "/" {
continue
}
nsVariant, err := conn.Object(networkManagerDest, cfgPath).GetProperty(src.configIface + ".NameserverData")
if err != nil {
continue
}
entries, ok := nsVariant.Value().([]map[string]dbus.Variant)
if !ok {
continue
}
for _, entry := range entries {
addrVariant, ok := entry["address"]
if !ok {
continue
}
addrStr, ok := addrVariant.Value().(string)
if !ok {
continue
}
if addr, err := netip.ParseAddr(addrStr); err == nil {
add(addr)
}
}
}
devicesVariant, err := active.GetProperty(networkManagerDbusActiveDevicesProperty)
if err != nil {
continue
}
devicePaths, ok := devicesVariant.Value().([]dbus.ObjectPath)
if !ok {
continue
}
for _, devicePath := range devicePaths {
device := conn.Object(networkManagerDest, devicePath)
for _, addr := range dhcpLeaseNameservers(conn, device, networkManagerDbusDeviceDhcp4ConfigProp, networkManagerDbusDhcp4ConfigInterface, "domain_name_servers") {
add(addr)
}
for _, addr := range dhcpLeaseNameservers(conn, device, networkManagerDbusDeviceDhcp6ConfigProp, networkManagerDbusDhcp6ConfigInterface, "dhcp6_name_servers") {
add(addr)
}
for _, addr := range appliedConnectionNameservers(device) {
add(addr)
}
}
}
return servers, nil
}
// dhcpLeaseNameservers reads a space-separated nameserver list out of a
// device's Dhcp4Config/Dhcp6Config Options, straight from the DHCP lease -
// data NetworkManager's DNS merging (and therefore olm's own global-dns
// override) never touches.
func dhcpLeaseNameservers(conn *dbus.Conn, device dbus.BusObject, configProperty, configIface, optionsKey string) []netip.Addr {
cfgVariant, err := device.GetProperty(configProperty)
if err != nil {
return nil
}
cfgPath, ok := cfgVariant.Value().(dbus.ObjectPath)
if !ok || cfgPath == "" || cfgPath == "/" {
return nil
}
optsVariant, err := conn.Object(networkManagerDest, cfgPath).GetProperty(configIface + ".Options")
if err != nil {
return nil
}
opts, ok := optsVariant.Value().(map[string]dbus.Variant)
if !ok {
return nil
}
raw, ok := opts[optionsKey]
if !ok {
return nil
}
str, ok := raw.Value().(string)
if !ok {
return nil
}
var addrs []netip.Addr
for _, field := range strings.Fields(str) {
if addr, err := netip.ParseAddr(field); err == nil {
addrs = append(addrs, addr)
}
}
return addrs
}
// appliedConnectionNameservers reads the ipv4.dns/ipv6.dns servers configured
// on the device's currently-applied connection profile - e.g. a static DNS
// override set by the user directly in NetworkManager - independent of DHCP
// and of olm's own global-dns override.
func appliedConnectionNameservers(device dbus.BusObject) []netip.Addr {
var settings map[string]map[string]dbus.Variant
var versionID uint64
if err := device.Call(networkManagerDbusGetAppliedConnMethod, 0, uint32(0)).Store(&settings, &versionID); err != nil {
return nil
}
var addrs []netip.Addr
if ipv4, ok := settings["ipv4"]; ok {
if dnsVariant, ok := ipv4["dns"]; ok {
if raw, ok := dnsVariant.Value().([]uint32); ok {
for _, v := range raw {
var b [4]byte
// NetworkManager encodes IPv4 addresses in this setting as
// network-byte-order bytes reinterpreted as a native uint32.
binary.LittleEndian.PutUint32(b[:], v)
addrs = append(addrs, netip.AddrFrom4(b))
}
}
}
}
if ipv6, ok := settings["ipv6"]; ok {
if dnsVariant, ok := ipv6["dns"]; ok {
if raw, ok := dnsVariant.Value().([][]byte); ok {
for _, b := range raw {
if len(b) == 16 {
var arr [16]byte
copy(arr[:], b)
addrs = append(addrs, netip.AddrFrom16(arr))
}
}
}
}
}
return addrs
}
// GetNetworkManagerVersion returns the version of NetworkManager
func GetNetworkManagerVersion() (string, error) {
conn, err := dbus.SystemBus()
@@ -323,3 +551,41 @@ func GetNetworkManagerVersion() (string, error) {
return version, nil
}
// CleanupStaleNetworkManagerDNS removes any stale DNS configuration left by NetworkManager
// configurator from a previous unclean shutdown. This is a static function that can be called
// without creating a configurator instance, useful for cleanup before network operations.
func CleanupStaleNetworkManagerDNS() error {
confPath := networkManagerConfDir + "/" + networkManagerDNSConfFile
// Check if our config file exists from a previous session
if _, err := os.Stat(confPath); os.IsNotExist(err) {
// No config file, nothing to clean up
return nil
}
// Remove the stale configuration file
if err := os.Remove(confPath); err != nil && !os.IsNotExist(err) {
return fmt.Errorf("remove stale DNS config file: %w", err)
}
// Try to reload NetworkManager if it's available
if IsNetworkManagerAvailable() {
conn, err := dbus.SystemBus()
if err != nil {
return fmt.Errorf("connect to system bus for reload: %w", err)
}
defer conn.Close()
obj := conn.Object(networkManagerDest, networkManagerDbusObjectNode)
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
defer cancel()
if err := obj.CallWithContext(ctx, networkManagerDest+".Reload", 0, uint32(0)).Store(); err != nil {
return fmt.Errorf("reload NetworkManager after cleanup: %w", err)
}
}
return nil
}
+34
View File
@@ -219,3 +219,37 @@ func IsResolvconfAvailable() bool {
cmd := exec.Command(resolvconfCommand, "--version")
return cmd.Run() == nil
}
// CleanupStaleResolvconfDNS removes any stale DNS configuration left by the resolvconf
// configurator from a previous unclean shutdown. This is a static function that can be
// called without creating a configurator instance, useful for cleanup before network operations.
// The interfaceName parameter specifies which interface entry to clean up (typically "olm").
func CleanupStaleResolvconfDNS(interfaceName string) error {
if !IsResolvconfAvailable() {
// resolvconf not available, nothing to clean up
return nil
}
// Detect resolvconf implementation type
implType, err := detectResolvconfType()
if err != nil {
// Can't detect type, try default
implType = "resolvconf"
}
// Try to delete any existing entry for this interface
// This is idempotent - if no entry exists, resolvconf will just return success
var cmd *exec.Cmd
switch implType {
case "openresolv":
cmd = exec.Command(resolvconfCommand, "-f", "-d", interfaceName)
default:
cmd = exec.Command(resolvconfCommand, "-d", interfaceName)
}
// Ignore errors - the entry may not exist, which is fine
_ = cmd.Run()
return nil
}
+274
View File
@@ -0,0 +1,274 @@
package dns
import (
"context"
"net"
"net/netip"
"sort"
"sync"
"time"
"github.com/fosrl/newt/logger"
"github.com/miekg/dns"
)
const defaultPollInterval = 30 * time.Second
// dnsHealthCheckTimeout bounds how long we wait for a candidate DNS server to
// answer a health-check query before considering it unusable.
const dnsHealthCheckTimeout = 2 * time.Second
// SystemDNSMonitor monitors the host system's DNS configuration and notifies
// callers when it changes. The reported servers are in "host:port" format
// (e.g. "8.8.8.8:53") and can be used directly as UpstreamDNS and PublicDNS.
//
// Platform behaviour:
// - Linux: reads /run/systemd/resolve/resolv.conf when present (updated by
// systemd-resolved on every DHCP change), then falls back to
// /etc/resolv.conf.olm.backup (written before olm overrides DNS), and
// finally /etc/resolv.conf.
// - macOS: reads the unscoped resolvers from `scutil --dns`, falling back
// to /etc/resolv.conf if scutil is unavailable. This includes olm's own
// supplemental scutil DNS override entry, which is expected to be
// filtered out via SetExcludeIP.
// - Windows: enumerates every network adapter's effective DNS servers
// (static if set, else DHCP-assigned) from the registry.
// - Other platforms: returns an empty list (no-op monitor).
type SystemDNSMonitor struct {
mu sync.RWMutex
current []string // last health-checked, applied server list
lastRaw []string // last raw (exclude-filtered but unvalidated) candidate list seen
onChange func(servers []string)
interval time.Duration
stopCh chan struct{}
excludeMu sync.RWMutex
excludeIPs map[netip.Addr]bool
}
// NewSystemDNSMonitor creates a new monitor. onChange is called with the new
// server list whenever a change is detected; it is also called once from Start
// with the initial values. A zero interval uses the 30-second default.
func NewSystemDNSMonitor(interval time.Duration, onChange func(servers []string)) *SystemDNSMonitor {
if interval <= 0 {
interval = defaultPollInterval
}
return &SystemDNSMonitor{
interval: interval,
onChange: onChange,
stopCh: make(chan struct{}),
excludeIPs: make(map[netip.Addr]bool),
}
}
// SetExcludeIP registers an IP address that must never appear in the reported
// DNS server list. Call this after olm's DNS proxy is created to prevent the
// proxy's own IP from being returned as an upstream server when the OS DNS has
// been overridden to point at the proxy.
func (m *SystemDNSMonitor) SetExcludeIP(ip netip.Addr) {
m.excludeMu.Lock()
m.excludeIPs[ip.Unmap()] = true
m.excludeMu.Unlock()
}
// Start reads the current system DNS immediately, fires onChange, then polls
// in the background until Stop is called or ctx is cancelled.
func (m *SystemDNSMonitor) Start(ctx context.Context) {
m.applyCandidates(m.readFiltered())
go m.run(ctx)
}
// Stop halts the background polling goroutine.
func (m *SystemDNSMonitor) Stop() {
select {
case <-m.stopCh:
default:
close(m.stopCh)
}
}
// Current returns the most recently observed system DNS servers.
func (m *SystemDNSMonitor) Current() []string {
m.mu.RLock()
defer m.mu.RUnlock()
out := make([]string, len(m.current))
copy(out, m.current)
return out
}
// readFiltered calls the platform-specific readSystemDNS and removes any
// addresses that have been excluded via SetExcludeIP. If all addresses are
// excluded the function returns nil so the caller can retain the last
// known-good value.
func (m *SystemDNSMonitor) readFiltered() []string {
return m.filterExcluded(readSystemDNS())
}
// filterExcluded removes any addresses that have been excluded via
// SetExcludeIP from servers. Used both for the internally-polled server list
// (readFiltered) and for server lists reported externally (ReportExternal) by
// platforms - Android, iOS - where olm cannot read the OS's DNS configuration
// itself.
func (m *SystemDNSMonitor) filterExcluded(servers []string) []string {
m.excludeMu.RLock()
excludeIPs := m.excludeIPs
m.excludeMu.RUnlock()
if len(excludeIPs) == 0 {
return servers
}
var filtered []string
for _, s := range servers {
host, _, err := net.SplitHostPort(s)
if err != nil {
filtered = append(filtered, s)
continue
}
addr, err := netip.ParseAddr(host)
if err != nil || excludeIPs[addr.Unmap()] {
continue
}
filtered = append(filtered, s)
}
return filtered
}
// ReportExternal applies an externally-observed DNS server list (e.g. from
// Android's ConnectivityManager or iOS's SCDynamicStore, where the platform
// itself - not olm - must detect the OS's real DNS configuration) through the
// same exclude-IP filtering, health-check validation, and change-detection as
// the internal poll loop, firing onChange if the result differs from the last
// known value.
func (m *SystemDNSMonitor) ReportExternal(servers []string) {
m.applyCandidates(m.filterExcluded(servers))
}
func (m *SystemDNSMonitor) run(ctx context.Context) {
ticker := time.NewTicker(m.interval)
defer ticker.Stop()
for {
select {
case <-ctx.Done():
return
case <-m.stopCh:
return
case <-ticker.C:
m.applyCandidates(m.readFiltered())
}
}
}
// applyCandidates takes an exclude-filtered (but not yet health-checked) list
// of candidate DNS servers - from either the internal poll loop or
// ReportExternal - and, only if it differs from the last raw list seen (to
// avoid re-running network health checks on every 30-second poll tick when
// nothing has actually changed), health-checks it via filterUnreachable and
// applies whatever passes, firing onChange if the result changed.
//
// If none of the candidates pass the health check, the previous known-good
// value is retained rather than clobbered - this is what protects against
// e.g. a carrier reporting a DNS server (such as T-Mobile's internal ULA
// DNS64 resolvers) that is technically "the system DNS" but not actually
// reachable/usable from wherever queries are sent.
func (m *SystemDNSMonitor) applyCandidates(raw []string) {
if len(raw) == 0 {
return
}
m.mu.Lock()
if dnsSlicesEqual(m.lastRaw, raw) {
m.mu.Unlock()
logger.Debug("System DNS candidates unchanged, skipping health check: %v", raw)
return
}
m.lastRaw = raw
m.mu.Unlock()
logger.Debug("System DNS candidates changed, health-checking: %v", raw)
validated := filterUnreachable(raw)
if len(validated) == 0 {
logger.Warn("None of the detected DNS servers answered a health-check query, keeping previous value: %v", raw)
return
}
m.mu.Lock()
changed := !dnsSlicesEqual(m.current, validated)
if changed {
m.current = validated
}
m.mu.Unlock()
if changed && m.onChange != nil {
logger.Info("System DNS changed: %v", validated)
m.onChange(validated)
}
}
// dnsServerReachable is a seam for tests; production code always uses probeDNSServerErr.
var dnsServerReachable = probeDNSServerErr
// filterUnreachable validates that each candidate server actually answers a
// DNS query before it's trusted, rather than statically guessing from the
// address (e.g. rejecting all private/ULA addresses, which would also reject
// a perfectly valid home router forwarding to a real resolver). Checks run
// concurrently so multiple candidates don't serialize the timeout.
func filterUnreachable(servers []string) []string {
if len(servers) == 0 {
return servers
}
reachable := make([]bool, len(servers))
errs := make([]error, len(servers))
var wg sync.WaitGroup
for i, server := range servers {
wg.Add(1)
go func(i int, server string) {
defer wg.Done()
reachable[i], errs[i] = dnsServerReachable(server)
}(i, server)
}
wg.Wait()
var result []string
for i, server := range servers {
if reachable[i] {
result = append(result, server)
} else {
logger.Debug("Discarding DNS server %s: failed health check: %v", server, errs[i])
}
}
return result
}
// probeDNSServerErr sends a minimal root NS query to confirm a candidate server
// actually answers, without depending on any specific external hostname being
// reachable (which could itself be blocked/filtered independently of whether
// the resolver works). The returned error is kept (rather than just a bool) so
// callers can log why a candidate was rejected (unreachable route, timeout, etc.).
func probeDNSServerErr(server string) (bool, error) {
client := &dns.Client{Timeout: dnsHealthCheckTimeout}
msg := new(dns.Msg)
msg.SetQuestion(".", dns.TypeNS)
_, _, err := client.Exchange(msg, server)
return err == nil, err
}
// dnsSlicesEqual reports whether two server lists are equal regardless of order.
func dnsSlicesEqual(a, b []string) bool {
if len(a) != len(b) {
return false
}
ac := make([]string, len(a))
bc := make([]string, len(b))
copy(ac, a)
copy(bc, b)
sort.Strings(ac)
sort.Strings(bc)
for i := range ac {
if ac[i] != bc[i] {
return false
}
}
return true
}
+18
View File
@@ -0,0 +1,18 @@
//go:build android
package dns
// readSystemDNS returns nil on Android: olm cannot read the OS's DNS
// configuration itself here, so the app detects it (via ConnectivityManager)
// and pushes it in through Olm.SetSystemDNS instead (see SystemDnsMonitor.java).
//
// This is a dedicated file (rather than falling through the general
// sysresolver_stub.go catch-all) because wireguard-android's build passes
// "-tags linux" to share Linux netlink code with Android, and that custom tag
// makes "!linux" evaluate to false even on a real GOOS=android build, which
// would otherwise make sysresolver_stub.go stop applying and leave
// readSystemDNS undefined. An explicit "android" constraint isn't affected by
// that, since nothing passes a conflicting "-tags android".
func readSystemDNS() []string {
return nil
}
+116
View File
@@ -0,0 +1,116 @@
//go:build darwin && !ios && !nosysresolver
package dns
import (
"bufio"
"net"
"net/netip"
"os"
"os/exec"
"strings"
)
// scutilPath is the well-known location of scutil on macOS.
const scutilPath = "/usr/sbin/scutil"
// readSystemDNS returns the current system DNS servers in "host:53" format.
//
// olm's own DNS override is itself a scutil supplemental resolver (see
// dns/platform/darwin.go), and macOS gives supplemental resolvers priority
// over the primary network service's resolver when generating the merged
// configuration - which is also what gets mirrored into /etc/resolv.conf. So
// once olm's override is active, /etc/resolv.conf (and a naive read of just
// the top of "scutil --dns") reflects olm's own proxy address, not the
// physical network's real DNS.
//
// Instead this reads every resolver in the unscoped "DNS configuration"
// section of `scutil --dns` (the "(for scoped queries)" section that follows
// only duplicates per-interface resolvers and is skipped), which includes
// both the real physical-network resolver and olm's own supplemental one.
// olm's own address is expected to be filtered out by the caller via
// SystemDNSMonitor.SetExcludeIP, the same mechanism used on Windows to drop
// olm's own adapter DNS entry.
//
// /etc/resolv.conf is kept as a fallback for when scutil is unavailable.
func readSystemDNS() []string {
if out, err := exec.Command(scutilPath, "--dns").Output(); err == nil {
if servers := parseScutilDNS(string(out)); len(servers) > 0 {
return servers
}
}
return parseMacResolvConf("/etc/resolv.conf")
}
// parseScutilDNS extracts nameserver addresses from the unscoped "DNS
// configuration" section at the top of `scutil --dns` output, stopping at
// the "DNS configuration (for scoped queries)" section that follows it.
func parseScutilDNS(output string) []string {
var result []string
seen := make(map[string]bool)
scanner := bufio.NewScanner(strings.NewReader(output))
for scanner.Scan() {
line := strings.TrimSpace(scanner.Text())
if strings.HasPrefix(line, "DNS configuration (for scoped queries)") {
break
}
if !strings.HasPrefix(line, "nameserver[") {
continue
}
parts := strings.SplitN(line, ":", 2)
if len(parts) != 2 {
continue
}
addr, err := netip.ParseAddr(strings.TrimSpace(parts[1]))
if err != nil {
continue
}
if addr.IsLoopback() || addr.IsLinkLocalUnicast() {
continue
}
hp := net.JoinHostPort(addr.String(), "53")
if !seen[hp] {
seen[hp] = true
result = append(result, hp)
}
}
return result
}
func parseMacResolvConf(path string) []string {
f, err := os.Open(path)
if err != nil {
return nil
}
defer f.Close()
var result []string
seen := make(map[string]bool)
scanner := bufio.NewScanner(f)
for scanner.Scan() {
line := strings.TrimSpace(scanner.Text())
if !strings.HasPrefix(line, "nameserver") {
continue
}
fields := strings.Fields(line)
if len(fields) < 2 {
continue
}
addr, err := netip.ParseAddr(fields[1])
if err != nil {
continue
}
if addr.IsLoopback() || addr.IsLinkLocalUnicast() {
continue
}
s := net.JoinHostPort(addr.String(), "53")
if !seen[s] {
seen[s] = true
result = append(result, s)
}
}
return result
}
+17
View File
@@ -0,0 +1,17 @@
//go:build darwin && !ios && nosysresolver
package dns
// readSystemDNS is disabled by the nosysresolver build tag. This is used for
// the macOS app build: unlike the CLI, the app's PacketTunnel system
// extension additionally applies NEDNSSettings (see apple/PacketTunnel),
// which can become the system's primary resolver and make /etc/resolv.conf
// reflect olm's own proxy IP instead of the real upstream DNS. Rather than
// have olm poll a value that may be self-referential, the app pushes the
// real DNS servers in via SetSystemDNS (detected in Swift via
// SCDynamicStore) exactly like Android and iOS. The CLI keeps the real
// implementation in sysresolver_darwin.go, since it has no such override
// mechanism and /etc/resolv.conf always reflects the physical network there.
func readSystemDNS() []string {
return nil
}
+18
View File
@@ -0,0 +1,18 @@
//go:build ios
package dns
// readSystemDNS returns nil on iOS: olm cannot read the OS's DNS
// configuration itself here, so the app must detect it and push it in
// through the equivalent of Olm.SetSystemDNS instead.
//
// This is a dedicated file (rather than falling through the general
// sysresolver_stub.go catch-all) because Go's build constraint evaluator
// treats GOOS=ios as implicitly satisfying the "darwin" tag as well as
// "ios". sysresolver_stub.go excludes with "!darwin", which is false for an
// iOS build, so the stub silently stops applying and would leave
// readSystemDNS undefined. An explicit "ios" constraint isn't affected by
// that ambiguity.
func readSystemDNS() []string {
return nil
}
+129
View File
@@ -0,0 +1,129 @@
//go:build linux && !android
package dns
import (
"bufio"
"net"
"net/netip"
"os"
"strings"
platform "github.com/fosrl/olm/dns/platform"
)
// readSystemDNS returns the current system DNS servers in "host:53" format.
//
// Resolution order:
// 1. /run/systemd/resolve/resolv.conf, but only when systemd-resolved is
// actually running (checked live via D-Bus) — the file lives in /run and
// can linger there, frozen at whatever it last contained, long after the
// service that maintained it has stopped (e.g. it ran earlier in the
// boot and was since disabled). Trusting its mere existence would report
// that stale snapshot forever instead of falling through to a live
// source. When the service is actually up the file is maintained with
// the real per-link DNS servers, updated on every DHCP change and never
// touched by olm's D-Bus DNS override.
// 2. NetworkManager, queried live over D-Bus — NetworkManager's own view of
// each active connection's DNS servers, independent of what is currently
// written to /etc/resolv.conf. This covers NetworkManager's "dnsmasq" and
// "unbound" DNS modes, where /etc/resolv.conf only contains a loopback
// stub address, and stays accurate even if olm's own override has
// directly overwritten /etc/resolv.conf, without going stale the way a
// one-time backup snapshot would if the real DNS changes mid-override
// (e.g. the user switches WiFi networks). olm's own NetworkManager
// override is itself a NetworkManager-level global DNS override (see
// platform.GetNetworkManagerNameservers), so this also reads each
// device's raw DHCP lease and applied-connection settings, which that
// override does not touch, to recover the real servers.
// 3. /etc/resolv.conf.olm.backup — written by olm before it overrides
// /etc/resolv.conf on non-systemd systems, for when NetworkManager isn't
// in use at all.
// 4. /etc/resolv.conf — plain fallback.
//
// Loopback and link-local addresses (e.g. 127.0.0.53, ::1) are excluded
// because they are stub resolver addresses, not real upstream servers.
func readSystemDNS() []string {
// Prefer systemd-resolved's resolved (non-stub) resolv.conf, but only if
// systemd-resolved is actually alive right now - see resolution order
// note above on why the file's existence alone isn't enough.
if platform.IsSystemdResolvedAvailable() {
if servers := parseResolvConf("/run/systemd/resolve/resolv.conf"); len(servers) > 0 {
return servers
}
}
if servers := readNetworkManagerDNS(); len(servers) > 0 {
return servers
}
// If olm has already overridden /etc/resolv.conf the backup holds the
// original pre-override DNS servers.
if _, err := os.Stat("/etc/resolv.conf.olm.backup"); err == nil {
if servers := parseResolvConf("/etc/resolv.conf.olm.backup"); len(servers) > 0 {
return servers
}
}
return parseResolvConf("/etc/resolv.conf")
}
// readNetworkManagerDNS returns the DNS servers NetworkManager reports over
// D-Bus for its active connections, in "host:53" format. Returns nil if
// NetworkManager isn't running or reports nothing usable.
func readNetworkManagerDNS() []string {
addrs, err := platform.GetNetworkManagerNameservers()
if err != nil || len(addrs) == 0 {
return nil
}
result := make([]string, 0, len(addrs))
for _, addr := range addrs {
result = append(result, addrToHostPort(addr))
}
return result
}
// parseResolvConf reads nameserver lines from a resolv.conf-style file,
// skipping loopback and link-local addresses.
func parseResolvConf(path string) []string {
f, err := os.Open(path)
if err != nil {
return nil
}
defer f.Close()
var result []string
seen := make(map[string]bool)
scanner := bufio.NewScanner(f)
for scanner.Scan() {
line := strings.TrimSpace(scanner.Text())
if !strings.HasPrefix(line, "nameserver") {
continue
}
fields := strings.Fields(line)
if len(fields) < 2 {
continue
}
addr, err := netip.ParseAddr(fields[1])
if err != nil {
continue
}
if addr.IsLoopback() || addr.IsLinkLocalUnicast() {
continue
}
s := addrToHostPort(addr)
if !seen[s] {
seen[s] = true
result = append(result, s)
}
}
return result
}
// addrToHostPort converts a netip.Addr to "addr:53" format, wrapping IPv6
// addresses in brackets as required by net.JoinHostPort.
func addrToHostPort(addr netip.Addr) string {
return net.JoinHostPort(addr.String(), "53")
}
+23
View File
@@ -0,0 +1,23 @@
//go:build !linux && !darwin && !windows && !android
package dns
// readSystemDNS returns nil on platforms where automatic DNS discovery is not
// implemented (freebsd, etc.). Callers should fall back to a statically
// configured DNS server.
//
// android and ios are excluded from this constraint (and have their own
// sysresolver_android.go / sysresolver_ios.go with an explicit "android" /
// "ios" tag) rather than falling through the "!linux" / "!darwin" catch-all
// here:
// - wireguard-android's build passes "-tags linux" to share Linux netlink code
// (a deliberate, long-standing convention, since Android's kernel is Linux), and Go's
// build constraint evaluator can't distinguish a custom "-tags linux" from the real
// GOOS=linux - so with that tag set, "!linux" is false even though GOOS is actually
// android, and this file would silently stop applying, leaving readSystemDNS undefined.
// - Go's build constraint evaluator treats GOOS=ios as implicitly satisfying the
// "darwin" tag as well as "ios", so "!darwin" is false on an iOS build too, which
// would otherwise leave readSystemDNS undefined there as well.
func readSystemDNS() []string {
return nil
}
+149
View File
@@ -0,0 +1,149 @@
package dns
import (
"net/netip"
"reflect"
"testing"
)
// stubReachable overrides dnsServerReachable for the duration of the test so
// tests don't depend on real network access, restoring the original on
// cleanup.
func stubReachable(t *testing.T, fn func(server string) bool) {
t.Helper()
orig := dnsServerReachable
dnsServerReachable = func(server string) (bool, error) { return fn(server), nil }
t.Cleanup(func() { dnsServerReachable = orig })
}
func allReachable(t *testing.T) {
stubReachable(t, func(string) bool { return true })
}
func TestReportExternalFiltersExcludedIP(t *testing.T) {
allReachable(t)
var got []string
m := &SystemDNSMonitor{
excludeIPs: make(map[netip.Addr]bool),
onChange: func(servers []string) {
got = servers
},
}
m.SetExcludeIP(netip.MustParseAddr("10.0.0.1"))
m.ReportExternal([]string{"10.0.0.1:53", "8.8.8.8:53"})
want := []string{"8.8.8.8:53"}
if !reflect.DeepEqual(got, want) {
t.Fatalf("onChange servers = %v, want %v", got, want)
}
if !reflect.DeepEqual(m.Current(), want) {
t.Fatalf("Current() = %v, want %v", m.Current(), want)
}
}
func TestReportExternalAllExcludedIsNoop(t *testing.T) {
allReachable(t)
called := false
m := &SystemDNSMonitor{
excludeIPs: make(map[netip.Addr]bool),
current: []string{"1.1.1.1:53"},
onChange: func(servers []string) {
called = true
},
}
m.SetExcludeIP(netip.MustParseAddr("10.0.0.1"))
m.ReportExternal([]string{"10.0.0.1:53"})
if called {
t.Fatal("onChange should not fire when all reported servers are excluded")
}
want := []string{"1.1.1.1:53"}
if !reflect.DeepEqual(m.Current(), want) {
t.Fatalf("Current() = %v, want unchanged %v", m.Current(), want)
}
}
func TestReportExternalOnlyFiresOnChange(t *testing.T) {
allReachable(t)
calls := 0
m := &SystemDNSMonitor{
excludeIPs: make(map[netip.Addr]bool),
onChange: func(servers []string) {
calls++
},
}
m.ReportExternal([]string{"8.8.8.8:53"})
m.ReportExternal([]string{"8.8.8.8:53"})
if calls != 1 {
t.Fatalf("onChange fired %d times, want 1", calls)
}
}
func TestReportExternalDropsUnreachableServer(t *testing.T) {
// Simulates e.g. T-Mobile's private ULA DNS64 resolver: technically "the
// system DNS" per the OS, but doesn't actually answer queries.
stubReachable(t, func(server string) bool {
return server != "[fd00:976a::9]:53"
})
var got []string
m := &SystemDNSMonitor{
excludeIPs: make(map[netip.Addr]bool),
onChange: func(servers []string) {
got = servers
},
}
m.ReportExternal([]string{"[fd00:976a::9]:53", "8.8.8.8:53"})
want := []string{"8.8.8.8:53"}
if !reflect.DeepEqual(got, want) {
t.Fatalf("onChange servers = %v, want %v", got, want)
}
}
func TestReportExternalKeepsPreviousValueWhenAllUnreachable(t *testing.T) {
stubReachable(t, func(string) bool { return true })
called := false
m := &SystemDNSMonitor{
excludeIPs: make(map[netip.Addr]bool),
current: []string{"1.1.1.1:53"},
onChange: func(servers []string) {
called = true
},
}
// Now simulate every candidate failing the health check (e.g. moved to a
// network where none of the reported servers actually respond).
stubReachable(t, func(string) bool { return false })
m.ReportExternal([]string{"[fd00:976a::9]:53", "[fd00:976a::10]:53"})
if called {
t.Fatal("onChange should not fire when no candidate passes the health check")
}
want := []string{"1.1.1.1:53"}
if !reflect.DeepEqual(m.Current(), want) {
t.Fatalf("Current() = %v, want unchanged %v", m.Current(), want)
}
}
func TestFilterUnreachable(t *testing.T) {
stubReachable(t, func(server string) bool {
return server == "8.8.8.8:53"
})
got := filterUnreachable([]string{"10.0.0.1:53", "8.8.8.8:53", "9.9.9.9:53"})
want := []string{"8.8.8.8:53"}
if !reflect.DeepEqual(got, want) {
t.Fatalf("filterUnreachable() = %v, want %v", got, want)
}
}
+110
View File
@@ -0,0 +1,110 @@
//go:build windows
package dns
import (
"fmt"
"net"
"net/netip"
"golang.org/x/sys/windows/registry"
)
const (
tcpipInterfacesPath = `SYSTEM\CurrentControlSet\Services\Tcpip\Parameters\Interfaces`
dhcpNameServerKey = "DhcpNameServer"
staticNameServerKey = "NameServer"
)
// readSystemDNS returns the current system DNS servers in "host:53" format by
// enumerating every network adapter in the Windows registry.
//
// For each adapter olm reads the effective DNS servers: static (NameServer)
// if set, since a static entry overrides DHCP for that adapter and is what
// the OS resolver actually uses, otherwise falling back to the DHCP-assigned
// servers (DhcpNameServer). This also picks up olm's own WireGuard adapter,
// which olm points at its local DNS proxy via a static NameServer entry; that
// address is expected to be filtered out by the caller via
// SystemDNSMonitor.SetExcludeIP. Loopback and link-local addresses are
// excluded.
func readSystemDNS() []string {
key, err := registry.OpenKey(registry.LOCAL_MACHINE, tcpipInterfacesPath, registry.ENUMERATE_SUB_KEYS)
if err != nil {
return nil
}
defer key.Close()
subkeys, err := key.ReadSubKeyNames(-1)
if err != nil {
return nil
}
seen := make(map[string]bool)
var result []string
for _, guid := range subkeys {
path := fmt.Sprintf(`%s\%s`, tcpipInterfacesPath, guid)
iKey, err := registry.OpenKey(registry.LOCAL_MACHINE, path, registry.QUERY_VALUE)
if err != nil {
continue
}
servers, _, err := iKey.GetStringValue(staticNameServerKey)
if err != nil || servers == "" {
servers, _, err = iKey.GetStringValue(dhcpNameServerKey)
}
iKey.Close()
if err != nil || servers == "" {
continue
}
for _, s := range splitWinDNSList(servers) {
addr, err := netip.ParseAddr(s)
if err != nil {
continue
}
if addr.IsLoopback() || addr.IsLinkLocalUnicast() {
continue
}
hp := net.JoinHostPort(addr.String(), "53")
if !seen[hp] {
seen[hp] = true
result = append(result, hp)
}
}
}
return result
}
// splitWinDNSList splits a Windows DNS server list that may be comma- or
// space-separated.
func splitWinDNSList(s string) []string {
var out []string
for _, part := range splitByRunes(s, []rune{',', ' '}) {
if part != "" {
out = append(out, part)
}
}
return out
}
func splitByRunes(s string, delims []rune) []string {
var result []string
start := 0
for i, r := range s {
for _, d := range delims {
if r == d {
if i > start {
result = append(result, s[start:i])
}
start = i + len(string(r))
break
}
}
}
if start < len(s) {
result = append(result, s[start:])
}
return result
}
+68
View File
@@ -0,0 +1,68 @@
package main
import (
"context"
"flag"
"fmt"
"time"
"github.com/fosrl/newt/logger"
dnsOverride "github.com/fosrl/olm/dns/override"
)
const (
defaultWatchdogInterval = 5 * time.Second
defaultWatchdogThreshold = 3
)
// runDNSWatchdogCommand handles the `olm watchdog` subcommand. The watchdog
// is meant to be spawned by an olm process after it installs a DNS
// override, and forcibly resets the system DNS if that parent dies before
// restoring it.
func runDNSWatchdogCommand(ctx context.Context, args []string) error {
fs := flag.NewFlagSet("watchdog", flag.ContinueOnError)
parentPID := fs.Int("parent-pid", 0, "PID of the olm process to monitor")
socketPath := fs.String("socket", "", "Path to the olm API unix socket (optional)")
interfaceName := fs.String("interface", "", "WireGuard interface name (used for cleanup)")
interval := fs.Duration("interval", defaultWatchdogInterval, "Liveness check interval")
threshold := fs.Int("threshold", defaultWatchdogThreshold, "Consecutive failures before DNS reset")
if err := fs.Parse(args); err != nil {
return err
}
if *parentPID <= 0 {
return fmt.Errorf("--parent-pid is required and must be positive")
}
// Ensure logger is initialised for the watchdog process.
logger.Init(nil)
return dnsOverride.RunWatchdog(ctx, dnsOverride.WatchdogConfig{
ParentPID: *parentPID,
SocketPath: *socketPath,
InterfaceName: *interfaceName,
CheckInterval: *interval,
FailureThreshold: *threshold,
})
}
// runResetDNSCommand handles the `olm reset-dns` subcommand. It forcibly
// removes any DNS override state left behind on the system.
func runResetDNSCommand(args []string) error {
fs := flag.NewFlagSet("reset-dns", flag.ContinueOnError)
interfaceName := fs.String("interface", "olm", "WireGuard interface name")
if err := fs.Parse(args); err != nil {
return err
}
logger.Init(nil)
if err := dnsOverride.ForceResetDNS(*interfaceName); err != nil {
return err
}
fmt.Println("DNS reset complete")
return nil
}
+10 -10
View File
@@ -1,18 +1,18 @@
module github.com/fosrl/olm
go 1.25
go 1.25.0
require (
github.com/Microsoft/go-winio v0.6.2
github.com/fosrl/newt v1.9.0
github.com/fosrl/newt v1.15.0
github.com/godbus/dbus/v5 v5.2.2
github.com/gorilla/websocket v1.5.3
github.com/miekg/dns v1.1.70
golang.org/x/sys v0.40.0
golang.org/x/sys v0.46.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
software.sslmate.com/src/go-pkcs12 v0.7.0
software.sslmate.com/src/go-pkcs12 v0.7.3
)
require (
@@ -20,15 +20,15 @@ require (
github.com/google/go-cmp v0.7.0 // indirect
github.com/vishvananda/netlink v1.3.1 // indirect
github.com/vishvananda/netns v0.0.5 // indirect
golang.org/x/crypto v0.46.0 // indirect
golang.org/x/crypto v0.53.0 // indirect
golang.org/x/exp v0.0.0-20251113190631-e25ba8c21ef6 // indirect
golang.org/x/mod v0.31.0 // indirect
golang.org/x/net v0.48.0 // indirect
golang.org/x/sync v0.19.0 // indirect
golang.org/x/mod v0.34.0 // indirect
golang.org/x/net v0.56.0 // indirect
golang.org/x/sync v0.20.0 // indirect
golang.org/x/time v0.12.0 // indirect
golang.org/x/tools v0.40.0 // indirect
golang.org/x/tools v0.43.0 // indirect
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2 // indirect
golang.zx2c4.com/wireguard/windows v0.5.3 // indirect
golang.zx2c4.com/wireguard/windows v1.0.1 // indirect
)
// To be used ONLY for local development
+18 -18
View File
@@ -1,7 +1,7 @@
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.9.0 h1:66eJMo6fA+YcBTbddxTfNJXNQo1WWKzmn6zPRP5kSDE=
github.com/fosrl/newt v1.9.0/go.mod h1:d1+yYMnKqg4oLqAM9zdbjthjj2FQEVouiACjqU468ck=
github.com/fosrl/newt v1.15.0 h1:WpL0whZM1FMjUe2Vy5jSH1bgbxm1O9k1qCyF/mqZT+s=
github.com/fosrl/newt v1.15.0/go.mod h1:l6kWoZPSaXT+ZRUjiyPgwflRqZWYaXpUj9oQ0sOPh4o=
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=
@@ -16,33 +16,33 @@ github.com/vishvananda/netlink v1.3.1 h1:3AEMt62VKqz90r0tmNhog0r/PpWKmrEShJU0wJW
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.46.0 h1:cKRW/pmt1pKAfetfu+RCEvjvZkA9RimPbh7bhFjGVBU=
golang.org/x/crypto v0.46.0/go.mod h1:Evb/oLKmMraqjZ2iQTwDwvCtJkczlDuTmdJXoZVzqU0=
golang.org/x/crypto v0.53.0 h1:QZ4Muo8THX6CizN2vPPd5fBGHyogrdK9fG4wLPFUsto=
golang.org/x/crypto v0.53.0/go.mod h1:DNLU434OwVakk9PzuwV8w62mAJpRJL3vsgcfp4Qnsio=
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.31.0 h1:HaW9xtz0+kOcWKwli0ZXy79Ix+UW/vOfmWI5QVd2tgI=
golang.org/x/mod v0.31.0/go.mod h1:43JraMp9cGx1Rx3AqioxrbrhNsLl2l/iNAvuBkrezpg=
golang.org/x/net v0.48.0 h1:zyQRTTrjc33Lhh0fBgT/H3oZq9WuvRR5gPC70xpDiQU=
golang.org/x/net v0.48.0/go.mod h1:+ndRgGjkh8FGtu1w1FGbEC31if4VrNVMuKTgcAAnQRY=
golang.org/x/sync v0.19.0 h1:vV+1eWNmZ5geRlYjzm2adRgW2/mcpevXNg50YZtPCE4=
golang.org/x/sync v0.19.0/go.mod h1:9KTHXmSnoGruLpwFjVSX0lNNA75CykiMECbovNTZqGI=
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/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.40.0 h1:DBZZqJ2Rkml6QMQsZywtnjnnGvHza6BTfYFWY9kjEWQ=
golang.org/x/sys v0.40.0/go.mod h1:OgkHotnGiDImocRcuBABYBEXf8A9a87e/uXjp9XT3ks=
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/time v0.12.0 h1:ScB/8o8olJvc+CQPWrK3fPZNfh7qgwCrY0zJmoEQLSE=
golang.org/x/time v0.12.0/go.mod h1:CDIdPxbZBQxdj6cxyCIdrNogrJKMJ7pr37NYpMcMDSg=
golang.org/x/tools v0.40.0 h1:yLkxfA+Qnul4cs9QA3KnlFu0lVmd8JJfoq+E41uSutA=
golang.org/x/tools v0.40.0/go.mod h1:Ik/tzLRlbscWpqqMRjyWYDisX8bG13FrdXp3o4Sr9lc=
golang.org/x/tools v0.43.0 h1:12BdW9CeB3Z+J/I/wj34VMl8X+fEXBxVR90JeMX5E7s=
golang.org/x/tools v0.43.0/go.mod h1:uHkMso649BX2cZK6+RpuIPXS3ho2hZo4FVwfoy1vIk0=
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2 h1:B82qJJgjvYKsXS9jeunTOisW56dUokqW/FOteYJJ/yg=
golang.zx2c4.com/wintun v0.0.0-20230126152724-0fa3db229ce2/go.mod h1:deeaetjYA+DHMHg+sMSMI58GrEteJUUzzw7en6TJQcI=
golang.zx2c4.com/wireguard v0.0.0-20250521234502-f333402bd9cb h1:whnFRlWMcXI9d+ZbWg+4sHnLp52d5yiIPUxMBSt4X9A=
golang.zx2c4.com/wireguard v0.0.0-20250521234502-f333402bd9cb/go.mod h1:rpwXGsirqLqN2L0JDJQlwOboGHmptD5ZD6T2VmcqhTw=
golang.zx2c4.com/wireguard/wgctrl v0.0.0-20241231184526-a9ab2273dd10 h1:3GDAcqdIg1ozBNLgPy4SLT84nfcBjr6rhGtXYtrkWLU=
golang.zx2c4.com/wireguard/wgctrl v0.0.0-20241231184526-a9ab2273dd10/go.mod h1:T97yPqesLiNrOYxkwmhMI0ZIlJDm+p0PMR8eRVeR5tQ=
golang.zx2c4.com/wireguard/windows v0.5.3 h1:On6j2Rpn3OEMXqBq00QEDC7bWSZrPIHKIus8eIuExIE=
golang.zx2c4.com/wireguard/windows v0.5.3/go.mod h1:9TEe8TJmtwyQebdFwAkEWOPr3prrtqm+REGFifP60hI=
golang.zx2c4.com/wireguard/windows v1.0.1 h1:eOxiDVbywPC+ZQqvdCK7x+ZwWXKbYv50TtH8ysFIbw8=
golang.zx2c4.com/wireguard/windows v1.0.1/go.mod h1:+fbT3FFdX4zzYDLwJh5+HPEcNN/3HyNdzhNSVsQM+zs=
gvisor.dev/gvisor v0.0.0-20250503011706-39ed1f5ac29c h1:m/r7OM+Y2Ty1sgBQ7Qb27VgIMBW8ZZhT4gLnUyDIhzI=
gvisor.dev/gvisor v0.0.0-20250503011706-39ed1f5ac29c/go.mod h1:3r5CMtNQMKIvBlrmM9xWUNamjKBYPOWyXOjmg5Kts3g=
software.sslmate.com/src/go-pkcs12 v0.7.0 h1:Db8W44cB54TWD7stUFFSWxdfpdn6fZVcDl0w3R4RVM0=
software.sslmate.com/src/go-pkcs12 v0.7.0/go.mod h1:Qiz0EyvDRJjjxGyUQa2cCNZn/wMyzrRJ/qcDXOQazLI=
software.sslmate.com/src/go-pkcs12 v0.7.3 h1:JBQD3FDqYjTeyDAeZQklj2ar88ykBLtALloPJHyAauU=
software.sslmate.com/src/go-pkcs12 v0.7.3/go.mod h1:Qiz0EyvDRJjjxGyUQa2cCNZn/wMyzrRJ/qcDXOQazLI=
+26 -1
View File
@@ -13,6 +13,8 @@ import (
olmpkg "github.com/fosrl/olm/olm"
)
var olmVersion = "version_replaceme"
func main() {
// Check if we're running as a Windows service
if isWindowsService() {
@@ -162,6 +164,26 @@ func main() {
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
// Internal DNS subcommands. These are handled before normal flag
// parsing because they have their own argument layouts and need to
// run without setting up the full olm runtime.
if len(os.Args) > 1 {
switch os.Args[1] {
case "watchdog", "dns-watchdog":
if err := runDNSWatchdogCommand(signalCtx, os.Args[2:]); err != nil {
fmt.Fprintf(os.Stderr, "watchdog failed: %v\n", err)
os.Exit(1)
}
return
case "reset-dns":
if err := runResetDNSCommand(os.Args[2:]); err != nil {
fmt.Fprintf(os.Stderr, "reset-dns failed: %v\n", err)
os.Exit(1)
}
return
}
}
// Run in console mode
runOlmMainWithArgs(ctx, cancel, signalCtx, os.Args[1:])
}
@@ -190,7 +212,6 @@ func runOlmMainWithArgs(ctx context.Context, cancel context.CancelFunc, signalCt
os.Exit(0)
}
olmVersion := "version_replaceme"
if showVersion {
fmt.Println("Olm version " + olmVersion)
os.Exit(0)
@@ -220,6 +241,8 @@ func runOlmMainWithArgs(ctx context.Context, cancel context.CancelFunc, signalCt
OnExit: cancel, // Pass cancel function directly to trigger shutdown
OnTerminated: cancel,
PprofAddr: ":4444", // TODO: REMOVE OR MAKE CONFIGURABLE
// Re-invoke this binary in watchdog mode to clean up DNS if we die.
WatchdogSubcommand: []string{"watchdog"},
}
olm, err := olmpkg.Init(ctx, olmConfig)
@@ -240,6 +263,7 @@ func runOlmMainWithArgs(ctx context.Context, cancel context.CancelFunc, signalCt
MTU: config.MTU,
DNS: config.DNS,
UpstreamDNS: config.UpstreamDNS,
MatchDomains: config.MatchDomains,
InterfaceName: config.InterfaceName,
Holepunch: !config.DisableHolepunch,
TlsClientCert: config.TlsClientCert,
@@ -248,6 +272,7 @@ func runOlmMainWithArgs(ctx context.Context, cancel context.CancelFunc, signalCt
OrgID: config.OrgID,
OverrideDNS: config.OverrideDNS,
DisableRelay: config.DisableRelay,
PreferLocalRoutes: config.PreferLocalRoutes,
EnableUAPI: true,
}
go olm.StartTunnel(tunnelConfig)
+7 -7
View File
@@ -32,7 +32,7 @@ DefaultGroupName={#MyAppName}
DisableProgramGroupPage=yes
; Uncomment the following line to run in non administrative install mode (install for current user only).
;PrivilegesRequired=lowest
OutputBaseFilename=mysetup
OutputBaseFilename=olm_windows_installer
SolidCompression=yes
WizardStyle=modern
; Add this to ensure PATH changes are applied and the system is prompted for a restart if needed
@@ -78,7 +78,7 @@ begin
Result := True;
exit;
end;
// Perform a case-insensitive check to see if the path is already present.
// We add semicolons to prevent partial matches (e.g., matching C:\App in C:\App2).
if Pos(';' + UpperCase(Path) + ';', ';' + UpperCase(OrigPath) + ';') > 0 then
@@ -109,7 +109,7 @@ begin
PathList.Delimiter := ';';
PathList.StrictDelimiter := True;
PathList.DelimitedText := OrigPath;
// Find and remove the matching entry (case-insensitive)
for I := PathList.Count - 1 downto 0 do
begin
@@ -119,10 +119,10 @@ begin
PathList.Delete(I);
end;
end;
// Reconstruct the PATH
NewPath := PathList.DelimitedText;
// Write the new PATH back to the registry
if RegWriteExpandStringValue(HKEY_LOCAL_MACHINE,
'SYSTEM\CurrentControlSet\Control\Session Manager\Environment',
@@ -145,8 +145,8 @@ begin
// Get the application installation path
AppPath := ExpandConstant('{app}');
Log('Removing PATH entry for: ' + AppPath);
// Remove only our path entry from the system PATH
RemovePathEntry(AppPath);
end;
end;
end;
+62 -11
View File
@@ -7,6 +7,7 @@ import (
"runtime"
"strconv"
"strings"
"time"
"github.com/fosrl/newt/logger"
"github.com/fosrl/newt/network"
@@ -144,11 +145,19 @@ func (o *Olm) handleConnect(msg websocket.WSMessage) {
}
// Create and start DNS proxy
o.dnsProxy, err = dns.NewDNSProxy(o.middleDev, o.tunnelConfig.MTU, wgData.UtilitySubnet, o.tunnelConfig.UpstreamDNS, o.tunnelConfig.TunnelDNS, interfaceIP)
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)
}
// Tell the system DNS monitor to exclude the proxy IP so that subsequent
// polls never mistake the proxy for a real upstream server (on Linux the OS
// DNS is overridden to point at this IP, which would otherwise feed back
// into UpstreamDNS or PublicDNS on the next poll).
if o.dnsMonitor != nil && o.dnsProxy != nil {
o.dnsMonitor.SetExcludeIP(o.dnsProxy.GetProxyIP())
}
if err = network.ConfigureInterface(o.tunnelConfig.InterfaceName, wgData.TunnelIP, o.tunnelConfig.MTU); err != nil {
logger.Error("Failed to o.tunnelConfigure interface: %v", err)
}
@@ -168,20 +177,25 @@ func (o *Olm) handleConnect(msg websocket.WSMessage) {
SharedBind: o.sharedBind,
WSClient: o.websocket,
APIServer: o.apiServer,
PublicDNS: o.tunnelConfig.PublicDNS,
})
for i := range wgData.Sites {
site := wgData.Sites[i]
var siteEndpoint string
// here we are going to take the relay endpoint if it exists which means we requested a relay for this peer
if site.RelayEndpoint != "" {
siteEndpoint = site.RelayEndpoint
} else {
siteEndpoint = site.Endpoint
if site.PublicKey != "" {
var siteEndpoint string
// here we are going to take the relay endpoint if it exists which means we requested a relay for this peer
if site.RelayEndpoint != "" {
siteEndpoint = site.RelayEndpoint
} else {
siteEndpoint = site.Endpoint
}
o.apiServer.AddPeerStatus(site.SiteId, site.Name, false, 0, siteEndpoint, false, false)
}
o.apiServer.AddPeerStatus(site.SiteId, site.Name, false, 0, siteEndpoint, false)
// we still call this to add the aliases for jit lookup but we just do that then pass inside. need to skip the above so we dont add to the api
if err := o.peerManager.AddPeer(site); err != nil {
logger.Error("Failed to add peer: %v", err)
return
@@ -196,6 +210,37 @@ func (o *Olm) handleConnect(msg websocket.WSMessage) {
logger.Error("Failed to start DNS proxy: %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
}
// 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 {
// Set up DNS override to use our DNS proxy
if err := dnsOverride.SetupDNSOverride(o.tunnelConfig.InterfaceName, o.dnsProxy.GetProxyIP()); err != nil {
@@ -203,6 +248,12 @@ func (o *Olm) handleConnect(msg websocket.WSMessage) {
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)
network.SetDNSServers([]string{o.dnsProxy.GetProxyIP().String()})
}
@@ -273,12 +324,12 @@ func (o *Olm) handleTerminate(msg websocket.WSMessage) {
logger.Error("Error unmarshaling terminate error data: %v", err)
} else {
logger.Info("Terminate reason (code: %s): %s", errorData.Code, errorData.Message)
if errorData.Code == "TERMINATED_INACTIVITY" {
logger.Info("Ignoring...")
return
}
// Set the olm error in the API server so it can be exposed via status
o.apiServer.SetOlmError(errorData.Code, errorData.Message)
}
+57 -25
View File
@@ -2,6 +2,7 @@ package olm
import (
"encoding/json"
"fmt"
"time"
"github.com/fosrl/newt/holepunch"
@@ -31,21 +32,27 @@ func (o *Olm) handleWgPeerAddData(msg websocket.WSMessage) {
return
}
if _, exists := o.peerManager.GetPeer(addSubnetsData.SiteId); !exists {
pm := o.getPeerManager()
if pm == nil {
logger.Debug("Ignoring add-remote-subnets-aliases message: peerManager is nil (shutdown in progress)")
return
}
if _, exists := pm.GetPeer(addSubnetsData.SiteId); !exists {
logger.Debug("Peer %d not found for removing remote subnets and aliases", addSubnetsData.SiteId)
return
}
// Add new subnets
for _, subnet := range addSubnetsData.RemoteSubnets {
if err := o.peerManager.AddRemoteSubnet(addSubnetsData.SiteId, subnet); err != nil {
if err := pm.AddRemoteSubnet(addSubnetsData.SiteId, subnet); err != nil {
logger.Error("Failed to add allowed IP %s: %v", subnet, err)
}
}
// Add new aliases
for _, alias := range addSubnetsData.Aliases {
if err := o.peerManager.AddAlias(addSubnetsData.SiteId, alias); err != nil {
if err := pm.AddAlias(addSubnetsData.SiteId, alias); err != nil {
logger.Error("Failed to add alias %s: %v", alias.Alias, err)
}
}
@@ -72,21 +79,27 @@ func (o *Olm) handleWgPeerRemoveData(msg websocket.WSMessage) {
return
}
if _, exists := o.peerManager.GetPeer(removeSubnetsData.SiteId); !exists {
pm := o.getPeerManager()
if pm == nil {
logger.Debug("Ignoring remove-remote-subnets-aliases message: peerManager is nil (shutdown in progress)")
return
}
if _, exists := pm.GetPeer(removeSubnetsData.SiteId); !exists {
logger.Debug("Peer %d not found for removing remote subnets and aliases", removeSubnetsData.SiteId)
return
}
// Remove subnets
for _, subnet := range removeSubnetsData.RemoteSubnets {
if err := o.peerManager.RemoveRemoteSubnet(removeSubnetsData.SiteId, subnet); err != nil {
if err := pm.RemoveRemoteSubnet(removeSubnetsData.SiteId, subnet); err != nil {
logger.Error("Failed to remove allowed IP %s: %v", subnet, err)
}
}
// Remove aliases
for _, alias := range removeSubnetsData.Aliases {
if err := o.peerManager.RemoveAlias(removeSubnetsData.SiteId, alias.Alias); err != nil {
if err := pm.RemoveAlias(removeSubnetsData.SiteId, alias.Alias); err != nil {
logger.Error("Failed to remove alias %s: %v", alias.Alias, err)
}
}
@@ -113,7 +126,13 @@ func (o *Olm) handleWgPeerUpdateData(msg websocket.WSMessage) {
return
}
if _, exists := o.peerManager.GetPeer(updateSubnetsData.SiteId); !exists {
pm := o.getPeerManager()
if pm == nil {
logger.Debug("Ignoring update-remote-subnets-aliases message: peerManager is nil (shutdown in progress)")
return
}
if _, exists := pm.GetPeer(updateSubnetsData.SiteId); !exists {
logger.Debug("Peer %d not found for updating remote subnets and aliases", updateSubnetsData.SiteId)
return
}
@@ -122,14 +141,14 @@ func (o *Olm) handleWgPeerUpdateData(msg websocket.WSMessage) {
// This ensures that if an old and new subnet are the same on different peers,
// the route won't be temporarily removed
for _, subnet := range updateSubnetsData.NewRemoteSubnets {
if err := o.peerManager.AddRemoteSubnet(updateSubnetsData.SiteId, subnet); err != nil {
if err := pm.AddRemoteSubnet(updateSubnetsData.SiteId, subnet); err != nil {
logger.Error("Failed to add allowed IP %s: %v", subnet, err)
}
}
// Remove old subnets after new ones are added
for _, subnet := range updateSubnetsData.OldRemoteSubnets {
if err := o.peerManager.RemoveRemoteSubnet(updateSubnetsData.SiteId, subnet); err != nil {
if err := pm.RemoveRemoteSubnet(updateSubnetsData.SiteId, subnet); err != nil {
logger.Error("Failed to remove allowed IP %s: %v", subnet, err)
}
}
@@ -138,14 +157,14 @@ func (o *Olm) handleWgPeerUpdateData(msg websocket.WSMessage) {
// This ensures that if an old and new alias share the same IP, the IP won't be
// temporarily removed from the allowed IPs list
for _, alias := range updateSubnetsData.NewAliases {
if err := o.peerManager.AddAlias(updateSubnetsData.SiteId, alias); err != nil {
if err := pm.AddAlias(updateSubnetsData.SiteId, alias); err != nil {
logger.Error("Failed to add alias %s: %v", alias.Alias, err)
}
}
// Remove old aliases after new ones are added
for _, alias := range updateSubnetsData.OldAliases {
if err := o.peerManager.RemoveAlias(updateSubnetsData.SiteId, alias.Alias); err != nil {
if err := pm.RemoveAlias(updateSubnetsData.SiteId, alias.Alias); err != nil {
logger.Error("Failed to remove alias %s: %v", alias.Alias, err)
}
}
@@ -162,7 +181,8 @@ func (o *Olm) handleSync(msg websocket.WSMessage) {
return
}
if o.peerManager == nil {
pm := o.getPeerManager()
if pm == nil {
logger.Warn("Peer manager not initialized, ignoring sync request")
return
}
@@ -189,7 +209,7 @@ func (o *Olm) handleSync(msg websocket.WSMessage) {
}
// Get all current peers
currentPeers := o.peerManager.GetAllPeers()
currentPeers := pm.GetAllPeers()
currentPeerMap := make(map[int]peers.SiteConfig)
for _, peer := range currentPeers {
currentPeerMap[peer.SiteId] = peer
@@ -199,7 +219,7 @@ func (o *Olm) handleSync(msg websocket.WSMessage) {
for siteId := range currentPeerMap {
if _, exists := expectedPeers[siteId]; !exists {
logger.Info("Sync: Removing peer for site %d (no longer in expected config)", siteId)
if err := o.peerManager.RemovePeer(siteId); err != nil {
if err := pm.RemovePeer(siteId); err != nil {
logger.Error("Sync: Failed to remove peer %d: %v", siteId, err)
} else {
// Remove any exit nodes associated with this peer from hole punching
@@ -216,23 +236,35 @@ func (o *Olm) handleSync(msg websocket.WSMessage) {
// Find peers to add (in expected but not in current) and peers to update
for siteId, expectedSite := range expectedPeers {
if _, exists := currentPeerMap[siteId]; !exists {
// Only trigger add if this is NOT a JIT-only config (i.e., has more than just siteId and aliases)
jitOnly := expectedSite.PublicKey == ""
if jitOnly {
logger.Debug("Sync: Registering aliases for JIT-only site %d", siteId)
if err := pm.AddPeer(expectedSite); err != nil {
logger.Error("Sync: Failed to register aliases for JIT site %d: %v", siteId, err)
}
continue
}
// New peer - add it using the add flow (with holepunch)
logger.Info("Sync: Adding new peer for site %d", siteId)
o.holePunchManager.TriggerHolePunch()
// // TODO: do we need to send the message to the cloud to add the peer that way?
// if err := o.peerManager.AddPeer(expectedSite); err != nil {
// logger.Error("Sync: Failed to add peer %d: %v", siteId, err)
// } else {
// logger.Info("Sync: Successfully added peer for site %d", siteId)
// }
o.holePunchManager.ResetServerHolepunchInterval() // start sending immediately again so we fill in the endpoint on the cloud
// add the peer via the server
// this is important because newt needs to get triggered as well to add the peer once the hp is complete
o.stopPeerSend, _ = o.websocket.SendMessageInterval("olm/wg/server/peer/add", map[string]interface{}{
"siteId": expectedSite.SiteId,
}, 1*time.Second, 10)
chainId := fmt.Sprintf("sync-%d", expectedSite.SiteId)
o.peerSendMu.Lock()
if stop, ok := o.stopPeerSends[chainId]; ok {
stop()
}
stopFunc, _ := o.websocket.SendMessageInterval("olm/wg/server/peer/add", map[string]interface{}{
"siteId": expectedSite.SiteId,
"chainId": chainId,
}, 2*time.Second, 10)
o.stopPeerSends[chainId] = stopFunc
o.peerSendMu.Unlock()
} else {
// Existing peer - check if update is needed
@@ -291,7 +323,7 @@ func (o *Olm) handleSync(msg websocket.WSMessage) {
siteConfig.Aliases = expectedSite.Aliases
}
if err := o.peerManager.UpdatePeer(siteConfig); err != nil {
if err := pm.UpdatePeer(siteConfig); err != nil {
logger.Error("Sync: Failed to update peer %d: %v", siteId, err)
} else {
// If the endpoint changed, trigger holepunch to refresh NAT mappings
+259 -25
View File
@@ -2,11 +2,14 @@ package olm
import (
"context"
"crypto/rand"
"encoding/hex"
"fmt"
"net"
"net/http"
_ "net/http/pprof"
"os"
"os/exec"
"sync"
"time"
@@ -31,7 +34,7 @@ type Olm struct {
privateKey wgtypes.Key
logFile *os.File
registered bool
registered bool
tunnelRunning bool
uapiListener net.Listener
@@ -40,11 +43,18 @@ type Olm struct {
middleDev *olmDevice.MiddleDevice
sharedBind *bind.SharedBind
dnsProxy *dns.DNSProxy
apiServer *api.API
websocket *websocket.Client
holePunchManager *holepunch.Manager
peerManager *peers.PeerManager
dnsProxy *dns.DNSProxy
dnsMonitor *dns.SystemDNSMonitor
// pendingSystemDNS holds a SetSystemDNS report received before dnsMonitor exists
// (e.g. Android/iOS push a value while the tunnel is still starting up), so it
// isn't silently dropped. Drained into dnsMonitor as soon as StartTunnel creates it.
pendingSystemDNSMu sync.Mutex
pendingSystemDNS []string
apiServer *api.API
websocket *websocket.Client
holePunchManager *holepunch.Manager
peerManager *peers.PeerManager
peerManagerMu sync.RWMutex
// Power mode management
currentPowerMode string
powerModeMu sync.Mutex
@@ -65,10 +75,26 @@ type Olm struct {
stopRegister func()
updateRegister func(newData any)
stopPeerSend func()
stopPeerSends map[string]func()
stopPeerInits map[string]func()
jitPendingSites map[int]string // siteId -> chainId for in-flight JIT requests
peerSendMu sync.Mutex
// WaitGroup to track tunnel lifecycle
tunnelWg sync.WaitGroup
// External DNS watchdog process (spawned after DNS override is installed).
// nil when no watchdog is running.
dnsWatchdogCmd *exec.Cmd
}
// getPeerManager safely returns the current peerManager under a read-lock.
// Callers must check the returned value for nil before using it.
func (o *Olm) getPeerManager() *peers.PeerManager {
o.peerManagerMu.RLock()
pm := o.peerManager
o.peerManagerMu.RUnlock()
return pm
}
// initTunnelInfo creates the shared UDP socket and holepunch manager.
@@ -111,11 +137,18 @@ func (o *Olm) initTunnelInfo(clientID string) error {
logger.Info("Created shared UDP socket on port %d (refcount: %d)", sourcePort, sharedBind.GetRefCount())
// Create the holepunch manager
o.holePunchManager = holepunch.NewManager(sharedBind, clientID, "olm", privateKey.PublicKey().String())
o.holePunchManager = holepunch.NewManager(sharedBind, clientID, "olm", privateKey.PublicKey().String(), o.tunnelConfig.PublicDNS)
return nil
}
// generateChainId generates a random chain ID for tracking peer sender lifecycles.
func generateChainId() string {
b := make([]byte, 8)
_, _ = rand.Read(b)
return hex.EncodeToString(b)
}
func Init(ctx context.Context, config OlmConfig) (*Olm, error) {
logger.GetLogger().SetLevel(util.ParseLogLevel(config.LogLevel))
@@ -166,10 +199,13 @@ func Init(ctx context.Context, config OlmConfig) (*Olm, error) {
apiServer.SetAgent(config.Agent)
newOlm := &Olm{
logFile: logFile,
olmCtx: ctx,
apiServer: apiServer,
olmConfig: config,
logFile: logFile,
olmCtx: ctx,
apiServer: apiServer,
olmConfig: config,
stopPeerSends: make(map[string]func()),
stopPeerInits: make(map[string]func()),
jitPendingSites: make(map[int]string),
}
newOlm.registerAPICallbacks()
@@ -195,6 +231,7 @@ func (o *Olm) registerAPICallbacks() {
Holepunch: req.Holepunch,
TlsClientCert: req.TlsClientCert,
OrgID: req.OrgID,
MatchDomains: req.MatchDomains,
}
var err error
@@ -222,7 +259,7 @@ func (o *Olm) registerAPICallbacks() {
tunnelConfig.MTU = 1420
}
if req.DNS == "" {
tunnelConfig.DNS = "9.9.9.9"
tunnelConfig.DNS = "8.8.8.8"
}
// DNSProxyIP has no default - it must be provided if DNS proxy is desired
// UpstreamDNS defaults to 8.8.8.8 if not provided
@@ -284,24 +321,155 @@ func (o *Olm) registerAPICallbacks() {
logger.Info("Processing power mode change request via API: mode=%s", req.Mode)
return o.SetPowerMode(req.Mode)
},
func(req api.JITConnectionRequest) error {
logger.Info("Processing JIT connect request via API: site=%s resource=%s", req.Site, req.Resource)
chainId := generateChainId()
o.peerSendMu.Lock()
stopFunc, _ := o.websocket.SendMessageInterval("olm/wg/server/peer/init", map[string]interface{}{
"siteId": req.Site,
"resourceId": req.Resource,
"chainId": chainId,
}, 2*time.Second, 10)
o.stopPeerInits[chainId] = stopFunc
o.peerSendMu.Unlock()
return nil
},
)
}
// startDNSWatchdog launches an external watchdog process that will reset
// system DNS if this olm process dies before it can call
// RestoreDNSOverride. It is a no-op when the OlmConfig has no
// WatchdogSubcommand configured, or when the watchdog has already been
// started for this Olm instance.
func (o *Olm) startDNSWatchdog(interfaceName string) {
if o.dnsWatchdogCmd != nil {
return
}
if len(o.olmConfig.WatchdogSubcommand) == 0 {
logger.Debug("DNS watchdog disabled (no WatchdogSubcommand configured)")
return
}
executable := o.olmConfig.WatchdogExecutable
if executable == "" {
exe, err := os.Executable()
if err != nil {
logger.Warn("DNS watchdog: failed to resolve executable: %v", err)
return
}
executable = exe
}
cmd, err := dnsOverride.SpawnWatchdog(dnsOverride.SpawnWatchdogConfig{
Executable: executable,
Subcommand: o.olmConfig.WatchdogSubcommand,
InterfaceName: interfaceName,
SocketPath: o.olmConfig.SocketPath,
LogFile: o.olmConfig.WatchdogLogFile,
})
if err != nil {
logger.Warn("DNS watchdog: spawn failed: %v", err)
return
}
o.dnsWatchdogCmd = cmd
}
// stopDNSWatchdog stops any previously spawned DNS watchdog process.
// Safe to call when no watchdog was started.
func (o *Olm) stopDNSWatchdog() {
if o.dnsWatchdogCmd == nil {
return
}
dnsOverride.StopWatchdog(o.dnsWatchdogCmd)
o.dnsWatchdogCmd = nil
}
func (o *Olm) StartTunnel(config TunnelConfig) {
if o.tunnelRunning {
logger.Info("Tunnel already running")
return
}
// debug print out the whole config
logger.Debug("Starting tunnel with config: %+v", config)
o.tunnelRunning = true // Also set it here in case it is called externally
o.tunnelConfig = config
network.PreferLocalRoutes = config.PreferLocalRoutes
// Determine whether the system DNS monitor should also manage UpstreamDNS.
// If the caller did not provide an explicit UpstreamDNS (it was defaulted to
// 8.8.8.8:53 by the API handler), we want the monitor to keep it updated
// with whatever DNS the host network is currently using.
upstreamFromConfig := len(config.UpstreamDNS) > 0 &&
!(len(config.UpstreamDNS) == 1 && config.UpstreamDNS[0] == "8.8.8.8:53")
if upstreamFromConfig {
logger.Info("UpstreamDNS is statically configured (%v); automatic system DNS detection will only update PublicDNS, DNS forwarding will keep using the configured value even if it becomes unreachable on a new network", config.UpstreamDNS)
}
// Start the system DNS monitor. The callback fires synchronously once with
// the initial values so that PublicDNS (and optionally UpstreamDNS) are
// populated before the tunnel goroutine proceeds.
o.dnsMonitor = dns.NewSystemDNSMonitor(0, func(servers []string) {
if len(servers) == 0 {
return
}
logger.Info("Applying system DNS: %v", servers)
// PublicDNS must always reflect the physical-network DNS so that
// WireGuard endpoint hostnames and hole-punch targets can be resolved
// even after the system resolver has been overridden by olm.
o.tunnelConfig.PublicDNS = servers
if o.holePunchManager != nil {
o.holePunchManager.SetPublicDNS(servers)
}
if pm := o.getPeerManager(); pm != nil {
pm.SetPublicDNS(servers)
}
// Keep the DNS proxy's local-DNS fallback (used for MatchDomains
// misses) in sync with the host's real system DNS servers.
if o.dnsProxy != nil {
o.dnsProxy.SetLocalDNS(servers)
}
// UpstreamDNS is updated only when the caller did not supply an
// explicit value; dynamic updates keep the proxy forwarding to the
// network's real resolver as the host moves between networks.
if !upstreamFromConfig {
o.tunnelConfig.UpstreamDNS = servers
if o.dnsProxy != nil {
o.dnsProxy.SetUpstreamDNS(servers)
}
} else {
logger.Debug("Not updating UpstreamDNS: statically configured to %v", config.UpstreamDNS)
}
})
o.dnsMonitor.Start(o.olmCtx)
// Apply any SetSystemDNS report that arrived before dnsMonitor existed (e.g. an
// Android/iOS push that raced ahead of this goroutine).
if pending := o.takePendingSystemDNS(); len(pending) > 0 {
o.dnsMonitor.ReportExternal(pending)
}
// Fall back to hardcoded DNS if the system monitor could not detect any.
if len(o.tunnelConfig.PublicDNS) == 0 {
if o.tunnelConfig.DNS != "" {
o.tunnelConfig.PublicDNS = []string{o.tunnelConfig.DNS + ":53"}
} else {
o.tunnelConfig.PublicDNS = []string{"8.8.8.8:53"}
}
}
if len(o.tunnelConfig.UpstreamDNS) == 0 {
o.tunnelConfig.UpstreamDNS = []string{"8.8.8.8:53"}
}
// Reset terminated status when tunnel starts
o.apiServer.SetTerminated(false)
fingerprint := config.InitialFingerprint
if fingerprint == nil {
fingerprint = make(map[string]any)
@@ -313,7 +481,7 @@ func (o *Olm) StartTunnel(config TunnelConfig) {
}
o.SetFingerprint(fingerprint)
o.SetPostures(postures)
o.SetPostures(postures)
// Create a cancellable context for this tunnel process
tunnelCtx, cancel := context.WithCancel(o.olmCtx)
@@ -338,7 +506,6 @@ func (o *Olm) StartTunnel(config TunnelConfig) {
config.OrgID,
config.Endpoint,
30*time.Second, // 30 seconds
config.PingTimeoutDuration,
websocket.WithPingDataProvider(func() map[string]any {
o.metaMu.Lock()
defer o.metaMu.Unlock()
@@ -370,6 +537,8 @@ func (o *Olm) StartTunnel(config TunnelConfig) {
o.websocket.RegisterHandler("olm/wg/peer/update", o.handleWgPeerUpdate)
o.websocket.RegisterHandler("olm/wg/peer/relay", o.handleWgPeerRelay)
o.websocket.RegisterHandler("olm/wg/peer/unrelay", o.handleWgPeerUnrelay)
o.websocket.RegisterHandler("olm/wg/peer/local", o.handleWgPeerLocal)
o.websocket.RegisterHandler("olm/wg/peer/unlocal", o.handleWgPeerUnlocal)
// Handlers for managing remote subnets to a peer
o.websocket.RegisterHandler("olm/wg/peer/data/add", o.handleWgPeerAddData)
@@ -378,6 +547,7 @@ func (o *Olm) StartTunnel(config TunnelConfig) {
// Handler for peer handshake - adds exit node to holepunch rotation and notifies server
o.websocket.RegisterHandler("olm/wg/peer/holepunch/site/add", o.handleWgPeerHolepunchAddSite)
o.websocket.RegisterHandler("olm/wg/peer/chain/cancel", o.handleCancelChain)
o.websocket.RegisterHandler("olm/sync", o.handleSync)
o.websocket.OnConnect(func() error {
@@ -387,7 +557,7 @@ func (o *Olm) StartTunnel(config TunnelConfig) {
if o.registered {
o.websocket.StartPingMonitor()
logger.Debug("Already registered, skipping registration")
return nil
}
@@ -420,7 +590,8 @@ func (o *Olm) StartTunnel(config TunnelConfig) {
"userToken": userToken,
"fingerprint": o.fingerprint,
"postures": o.postures,
}, 1*time.Second, 10)
"chainId": generateChainId(), // use a random chainId for registration updates - it won't be used for cancellation since registration is a one-time message but for tracking the session
}, 2*time.Second, 20) // after 18 tries on the server side we send the error so dont change this without changing that
// Invoke onRegistered callback if configured
if o.olmConfig.OnRegistered != nil {
@@ -517,6 +688,23 @@ func (o *Olm) Close() {
o.stopRegister = nil
}
// Stop all pending peer init and send senders before closing websocket
o.peerSendMu.Lock()
for _, stop := range o.stopPeerInits {
if stop != nil {
stop()
}
}
o.stopPeerInits = make(map[string]func())
for _, stop := range o.stopPeerSends {
if stop != nil {
stop()
}
}
o.stopPeerSends = make(map[string]func())
o.jitPendingSites = make(map[int]string)
o.peerSendMu.Unlock()
// send a disconnect message to the cloud to show disconnected
if o.websocket != nil {
o.websocket.SendMessage("olm/disconnecting", map[string]any{})
@@ -531,16 +719,30 @@ func (o *Olm) Close() {
logger.Error("Failed to restore DNS: %v", err)
}
// Stop the watchdog *after* a successful DNS restore so that if we
// somehow crash mid-restore the watchdog still has a chance to clean
// up. The watchdog itself is a no-op if it was never spawned.
o.stopDNSWatchdog()
if o.holePunchManager != nil {
o.holePunchManager.Stop()
o.holePunchManager = nil
}
// Stop the system DNS monitor after hole punch is stopped (it feeds
// publicDNS into the hole punch manager).
if o.dnsMonitor != nil {
o.dnsMonitor.Stop()
o.dnsMonitor = nil
}
// Close() also calls Stop() internally
o.peerManagerMu.Lock()
if o.peerManager != nil {
o.peerManager.Close()
o.peerManager = nil
}
o.peerManagerMu.Unlock()
if o.uapiListener != nil {
_ = o.uapiListener.Close()
@@ -702,6 +904,38 @@ func (o *Olm) SetPostures(data map[string]any) {
o.postures = data
}
// SetSystemDNS reports DNS servers observed by platform-native code. On
// Android and iOS olm cannot read the OS's DNS configuration itself (unlike
// Linux/macOS/Windows, see dns.readSystemDNS), so the app/extension detects
// the real pre-override DNS servers and pushes them here as the network
// changes. The list is applied through the same exclude-IP filtering and
// change detection as the internally-polled SystemDNSMonitor.
func (o *Olm) SetSystemDNS(servers []string) {
logger.Info("SetSystemDNS called with: %v", servers)
if o.dnsMonitor == nil {
// StartTunnel hasn't created the monitor yet (mobile platforms may push a
// value the moment they start observing, before the tunnel goroutine has
// gotten far enough to construct it). Stash it so StartTunnel can apply it
// instead of falling back to a hardcoded default DNS server.
o.pendingSystemDNSMu.Lock()
o.pendingSystemDNS = servers
o.pendingSystemDNSMu.Unlock()
logger.Debug("dnsMonitor not yet started, queued SetSystemDNS value")
return
}
o.dnsMonitor.ReportExternal(servers)
}
// takePendingSystemDNS returns and clears any SetSystemDNS value reported before
// dnsMonitor existed.
func (o *Olm) takePendingSystemDNS() []string {
o.pendingSystemDNSMu.Lock()
defer o.pendingSystemDNSMu.Unlock()
pending := o.pendingSystemDNS
o.pendingSystemDNS = nil
return pending
}
// SetPowerMode switches between normal and low power modes
// In low power mode: websocket is closed (stopping pings) and monitoring intervals are set to 10 minutes
// In normal power mode: websocket is reconnected (restarting pings) and monitoring intervals are restored
@@ -752,14 +986,14 @@ func (o *Olm) SetPowerMode(mode string) error {
lowPowerInterval := 10 * time.Minute
if o.peerManager != nil {
peerMonitor := o.peerManager.GetPeerMonitor()
if pm := o.getPeerManager(); pm != nil {
peerMonitor := pm.GetPeerMonitor()
if peerMonitor != nil {
peerMonitor.SetPeerInterval(lowPowerInterval, lowPowerInterval)
peerMonitor.SetPeerHolepunchInterval(lowPowerInterval, lowPowerInterval)
logger.Info("Set monitoring intervals to 10 minutes for low power mode")
}
o.peerManager.UpdateAllPeersPersistentKeepalive(0) // disable
pm.UpdateAllPeersPersistentKeepalive(0) // disable
}
if o.holePunchManager != nil {
@@ -804,14 +1038,14 @@ func (o *Olm) SetPowerMode(mode string) error {
}
// Restore intervals and reconnect websocket
if o.peerManager != nil {
peerMonitor := o.peerManager.GetPeerMonitor()
if pm := o.getPeerManager(); pm != nil {
peerMonitor := pm.GetPeerMonitor()
if peerMonitor != nil {
peerMonitor.ResetPeerHolepunchInterval()
peerMonitor.ResetPeerInterval()
}
o.peerManager.UpdateAllPeersPersistentKeepalive(5)
pm.UpdateAllPeersPersistentKeepalive(5)
}
if o.holePunchManager != nil {
+259 -26
View File
@@ -20,9 +20,16 @@ func (o *Olm) handleWgPeerAdd(msg websocket.WSMessage) {
return
}
if o.stopPeerSend != nil {
o.stopPeerSend()
o.stopPeerSend = nil
// Check if connection setup is complete
if !o.registered {
logger.Warn("Not connected, ignoring add-peer message")
return
}
pm := o.getPeerManager()
if pm == nil {
logger.Debug("Ignoring add-peer message: peerManager is nil (shutdown in progress)")
return
}
jsonData, err := json.Marshal(msg.Data)
@@ -31,20 +38,45 @@ func (o *Olm) handleWgPeerAdd(msg websocket.WSMessage) {
return
}
var siteConfig peers.SiteConfig
if err := json.Unmarshal(jsonData, &siteConfig); err != nil {
var siteConfigMsg struct {
peers.SiteConfig
ChainId string `json:"chainId"`
}
if err := json.Unmarshal(jsonData, &siteConfigMsg); err != nil {
logger.Error("Error unmarshaling add data: %v", err)
return
}
if siteConfigMsg.ChainId != "" {
o.peerSendMu.Lock()
if stop, ok := o.stopPeerSends[siteConfigMsg.ChainId]; ok {
stop()
delete(o.stopPeerSends, siteConfigMsg.ChainId)
}
o.peerSendMu.Unlock()
} else {
// stop all of the stopPeerSends
o.peerSendMu.Lock()
for _, stop := range o.stopPeerSends {
stop()
}
o.stopPeerSends = make(map[string]func())
o.peerSendMu.Unlock()
}
if siteConfigMsg.PublicKey == "" {
logger.Warn("Skipping add-peer for site %d (%s): no public key available (site may not be connected)", siteConfigMsg.SiteId, siteConfigMsg.Name)
return
}
_ = o.holePunchManager.TriggerHolePunch() // Trigger immediate hole punch attempt so that if the peer decides to relay we have already punched close to when we need it
if err := o.peerManager.AddPeer(siteConfig); err != nil {
if err := pm.AddPeer(siteConfigMsg.SiteConfig); err != nil {
logger.Error("Failed to add peer: %v", err)
return
}
logger.Info("Successfully added peer for site %d", siteConfig.SiteId)
logger.Info("Successfully added peer for site %d", siteConfigMsg.SiteId)
}
func (o *Olm) handleWgPeerRemove(msg websocket.WSMessage) {
@@ -56,6 +88,18 @@ func (o *Olm) handleWgPeerRemove(msg websocket.WSMessage) {
return
}
// Check if connection setup is complete
if !o.registered {
logger.Warn("Not connected, ignoring remove-peer message")
return
}
pm := o.getPeerManager()
if pm == nil {
logger.Debug("Ignoring remove-peer message: peerManager is nil (shutdown in progress)")
return
}
jsonData, err := json.Marshal(msg.Data)
if err != nil {
logger.Error("Error marshaling data: %v", err)
@@ -68,7 +112,7 @@ func (o *Olm) handleWgPeerRemove(msg websocket.WSMessage) {
return
}
if err := o.peerManager.RemovePeer(removeData.SiteId); err != nil {
if err := pm.RemovePeer(removeData.SiteId); err != nil {
logger.Error("Failed to remove peer: %v", err)
return
}
@@ -93,6 +137,18 @@ func (o *Olm) handleWgPeerUpdate(msg websocket.WSMessage) {
return
}
// Check if connection setup is complete
if !o.registered {
logger.Warn("Not connected, ignoring update-peer message")
return
}
pm := o.getPeerManager()
if pm == nil {
logger.Debug("Ignoring update-peer message: peerManager is nil (shutdown in progress)")
return
}
jsonData, err := json.Marshal(msg.Data)
if err != nil {
logger.Error("Error marshaling data: %v", err)
@@ -106,7 +162,7 @@ func (o *Olm) handleWgPeerUpdate(msg websocket.WSMessage) {
}
// Get existing peer from PeerManager
existingPeer, exists := o.peerManager.GetPeer(updateData.SiteId)
existingPeer, exists := pm.GetPeer(updateData.SiteId)
if !exists {
logger.Warn("Peer with site ID %d not found", updateData.SiteId)
return
@@ -133,8 +189,11 @@ func (o *Olm) handleWgPeerUpdate(msg websocket.WSMessage) {
if updateData.RemoteSubnets != nil {
siteConfig.RemoteSubnets = updateData.RemoteSubnets
}
if updateData.Aliases != nil {
siteConfig.Aliases = updateData.Aliases
}
if err := o.peerManager.UpdatePeer(siteConfig); err != nil {
if err := pm.UpdatePeer(siteConfig); err != nil {
logger.Error("Failed to update peer: %v", err)
return
}
@@ -142,8 +201,10 @@ func (o *Olm) handleWgPeerUpdate(msg websocket.WSMessage) {
// If the endpoint changed, trigger holepunch to refresh NAT mappings
if updateData.Endpoint != "" && updateData.Endpoint != existingPeer.Endpoint {
logger.Info("Endpoint changed for site %d, triggering holepunch to refresh NAT mappings", updateData.SiteId)
_ = o.holePunchManager.TriggerHolePunch()
o.holePunchManager.ResetServerHolepunchInterval()
if o.holePunchManager != nil {
_ = o.holePunchManager.TriggerHolePunch()
o.holePunchManager.ResetServerHolepunchInterval()
}
}
logger.Info("Successfully updated peer for site %d", updateData.SiteId)
@@ -153,7 +214,8 @@ func (o *Olm) handleWgPeerRelay(msg websocket.WSMessage) {
logger.Debug("Received relay-peer message: %v", msg.Data)
// Check if peerManager is still valid (may be nil during shutdown)
if o.peerManager == nil {
pm := o.getPeerManager()
if pm == nil {
logger.Debug("Ignoring relay message: peerManager is nil (shutdown in progress)")
return
}
@@ -164,13 +226,21 @@ func (o *Olm) handleWgPeerRelay(msg websocket.WSMessage) {
return
}
var relayData peers.RelayPeerData
var relayData struct {
peers.RelayPeerData
ChainId string `json:"chainId"`
}
if err := json.Unmarshal(jsonData, &relayData); err != nil {
logger.Error("Error unmarshaling relay data: %v", err)
return
}
primaryRelay, err := util.ResolveDomain(relayData.RelayEndpoint)
if monitor := pm.GetPeerMonitor(); monitor != nil {
monitor.CancelRelaySend(relayData.ChainId)
}
primaryRelay, err := util.ResolveDomainUpstream(relayData.RelayEndpoint, o.tunnelConfig.PublicDNS)
if err != nil {
logger.Error("Failed to resolve primary relay endpoint: %v", err)
return
@@ -179,14 +249,15 @@ func (o *Olm) handleWgPeerRelay(msg websocket.WSMessage) {
// Update HTTP server to mark this peer as using relay
o.apiServer.UpdatePeerRelayStatus(relayData.SiteId, relayData.RelayEndpoint, true)
o.peerManager.RelayPeer(relayData.SiteId, primaryRelay, relayData.RelayPort)
pm.RelayPeer(relayData.SiteId, primaryRelay, relayData.RelayPort)
}
func (o *Olm) handleWgPeerUnrelay(msg websocket.WSMessage) {
logger.Debug("Received unrelay-peer message: %v", msg.Data)
// Check if peerManager is still valid (may be nil during shutdown)
if o.peerManager == nil {
pm := o.getPeerManager()
if pm == nil {
logger.Debug("Ignoring unrelay message: peerManager is nil (shutdown in progress)")
return
}
@@ -197,13 +268,21 @@ func (o *Olm) handleWgPeerUnrelay(msg websocket.WSMessage) {
return
}
var relayData peers.UnRelayPeerData
var relayData struct {
peers.UnRelayPeerData
ChainId string `json:"chainId"`
}
if err := json.Unmarshal(jsonData, &relayData); err != nil {
logger.Error("Error unmarshaling relay data: %v", err)
return
}
primaryRelay, err := util.ResolveDomain(relayData.Endpoint)
if monitor := pm.GetPeerMonitor(); monitor != nil {
monitor.CancelRelaySend(relayData.ChainId)
}
primaryRelay, err := util.ResolveDomainUpstream(relayData.Endpoint, o.tunnelConfig.PublicDNS)
if err != nil {
logger.Warn("Failed to resolve primary relay endpoint: %v", err)
}
@@ -211,7 +290,72 @@ func (o *Olm) handleWgPeerUnrelay(msg websocket.WSMessage) {
// Update HTTP server to mark this peer as using relay
o.apiServer.UpdatePeerRelayStatus(relayData.SiteId, relayData.Endpoint, false)
o.peerManager.UnRelayPeer(relayData.SiteId, primaryRelay)
pm.UnRelayPeer(relayData.SiteId, primaryRelay)
}
// handleWgPeerLocal handles the server's acknowledgement of an "olm/wg/local" message.
// olm already switched the peer to the local endpoint before sending that message (it
// doesn't wait for permission, unlike relay), so all this needs to do is stop the retry
// sender for the given chain.
func (o *Olm) handleWgPeerLocal(msg websocket.WSMessage) {
logger.Debug("Received local-peer ack message: %v", msg.Data)
pm := o.getPeerManager()
if pm == nil {
logger.Debug("Ignoring local ack message: peerManager is nil (shutdown in progress)")
return
}
jsonData, err := json.Marshal(msg.Data)
if err != nil {
logger.Error("Error marshaling data: %v", err)
return
}
var localData struct {
peers.LocalPeerAckData
ChainId string `json:"chainId"`
}
if err := json.Unmarshal(jsonData, &localData); err != nil {
logger.Error("Error unmarshaling local ack data: %v", err)
return
}
if monitor := pm.GetPeerMonitor(); monitor != nil {
monitor.CancelLocalSend(localData.ChainId)
}
}
// handleWgPeerUnlocal handles the server's acknowledgement of an "olm/wg/unlocal" message.
// Same as handleWgPeerLocal, olm has already fallen back from the local endpoint by the time
// it sends the notification, so this just stops the retry sender.
func (o *Olm) handleWgPeerUnlocal(msg websocket.WSMessage) {
logger.Debug("Received unlocal-peer ack message: %v", msg.Data)
pm := o.getPeerManager()
if pm == nil {
logger.Debug("Ignoring unlocal ack message: peerManager is nil (shutdown in progress)")
return
}
jsonData, err := json.Marshal(msg.Data)
if err != nil {
logger.Error("Error marshaling data: %v", err)
return
}
var localData struct {
peers.LocalPeerAckData
ChainId string `json:"chainId"`
}
if err := json.Unmarshal(jsonData, &localData); err != nil {
logger.Error("Error unmarshaling unlocal ack data: %v", err)
return
}
if monitor := pm.GetPeerMonitor(); monitor != nil {
monitor.CancelLocalSend(localData.ChainId)
}
}
func (o *Olm) handleWgPeerHolepunchAddSite(msg websocket.WSMessage) {
@@ -230,7 +374,8 @@ func (o *Olm) handleWgPeerHolepunchAddSite(msg websocket.WSMessage) {
}
var handshakeData struct {
SiteId int `json:"siteId"`
SiteId int `json:"siteId"`
ChainId string `json:"chainId"`
ExitNode struct {
PublicKey string `json:"publicKey"`
Endpoint string `json:"endpoint"`
@@ -243,8 +388,34 @@ func (o *Olm) handleWgPeerHolepunchAddSite(msg websocket.WSMessage) {
return
}
// Stop the peer init sender for this chain, if any
if handshakeData.ChainId != "" {
o.peerSendMu.Lock()
if stop, ok := o.stopPeerInits[handshakeData.ChainId]; ok {
stop()
delete(o.stopPeerInits, handshakeData.ChainId)
}
// If this chain was initiated by a DNS-triggered JIT request, clear the
// pending entry so the site can be re-triggered if needed in the future.
delete(o.jitPendingSites, handshakeData.SiteId)
o.peerSendMu.Unlock()
} else {
// Stop all of the stopPeerInits
o.peerSendMu.Lock()
for _, stop := range o.stopPeerInits {
stop()
}
o.stopPeerInits = make(map[string]func())
o.peerSendMu.Unlock()
}
// Get existing peer from PeerManager
_, exists := o.peerManager.GetPeer(handshakeData.SiteId)
pm := o.getPeerManager()
if pm == nil {
logger.Debug("Ignoring peer-handshake message: peerManager is nil (shutdown in progress)")
return
}
_, exists := pm.GetPeer(handshakeData.SiteId)
if exists {
logger.Warn("Peer with site ID %d already added", handshakeData.SiteId)
return
@@ -273,10 +444,72 @@ func (o *Olm) handleWgPeerHolepunchAddSite(msg websocket.WSMessage) {
o.holePunchManager.TriggerHolePunch() // Trigger immediate hole punch attempt
o.holePunchManager.ResetServerHolepunchInterval() // start sending immediately again so we fill in the endpoint on the cloud
// Send handshake acknowledgment back to server with retry
o.stopPeerSend, _ = o.websocket.SendMessageInterval("olm/wg/server/peer/add", map[string]interface{}{
"siteId": handshakeData.SiteId,
}, 1*time.Second, 10)
// Send handshake acknowledgment back to server with retry, keyed by chainId
chainId := handshakeData.ChainId
if chainId == "" {
chainId = generateChainId()
}
o.peerSendMu.Lock()
stopFunc, _ := o.websocket.SendMessageInterval("olm/wg/server/peer/add", map[string]interface{}{
"siteId": handshakeData.SiteId,
"chainId": chainId,
}, 2*time.Second, 10)
o.stopPeerSends[chainId] = stopFunc
o.peerSendMu.Unlock()
logger.Info("Initiated handshake for site %d with exit node %s", handshakeData.SiteId, handshakeData.ExitNode.Endpoint)
}
func (o *Olm) handleCancelChain(msg websocket.WSMessage) {
logger.Debug("Received cancel-chain message: %v", msg.Data)
jsonData, err := json.Marshal(msg.Data)
if err != nil {
logger.Error("Error marshaling cancel-chain data: %v", err)
return
}
var cancelData struct {
ChainId string `json:"chainId"`
}
if err := json.Unmarshal(jsonData, &cancelData); err != nil {
logger.Error("Error unmarshaling cancel-chain data: %v", err)
return
}
if cancelData.ChainId == "" {
logger.Warn("Received cancel-chain message with no chainId")
return
}
o.peerSendMu.Lock()
defer o.peerSendMu.Unlock()
found := false
if stop, ok := o.stopPeerInits[cancelData.ChainId]; ok {
stop()
delete(o.stopPeerInits, cancelData.ChainId)
found = true
}
// If this chain was a DNS-triggered JIT request, clear the pending entry so
// the site can be re-triggered on the next DNS lookup.
for siteId, chainId := range o.jitPendingSites {
if chainId == cancelData.ChainId {
delete(o.jitPendingSites, siteId)
break
}
}
if stop, ok := o.stopPeerSends[cancelData.ChainId]; ok {
stop()
delete(o.stopPeerSends, cancelData.ChainId)
found = true
}
if found {
logger.Info("Cancelled chain %s", cancelData.ChainId)
} else {
logger.Warn("Cancel-chain: no active sender found for chain %s", cancelData.ChainId)
}
}
+31
View File
@@ -48,6 +48,21 @@ type OlmConfig struct {
OnAuthError func(statusCode int, message string) // Called when auth fails (401/403)
OnOlmError func(code string, message string) // Called when registration fails
OnExit func() // Called when exit is requested via API
// DNS watchdog (optional). When WatchdogSubcommand is non-empty, the
// olm package will spawn an external watchdog subprocess after a DNS
// override is installed. The watchdog will reset the system DNS if
// this process dies before it can call RestoreDNSOverride.
//
// The watchdog is launched as:
// <WatchdogExecutable> <WatchdogSubcommand...> \
// --parent-pid=<pid> --interface=<name> [--socket=<path>]
//
// When WatchdogExecutable is empty, os.Executable() of the calling
// process is used. WatchdogLogFile defaults to /dev/null.
WatchdogExecutable string
WatchdogSubcommand []string
WatchdogLogFile string
}
type TunnelConfig struct {
@@ -61,8 +76,16 @@ type TunnelConfig struct {
MTU int
DNS string
UpstreamDNS []string
PublicDNS []string
InterfaceName string
// MatchDomains lists FQDN wildcard patterns (using * and ? wildcards) that
// olm should check against local records / resolve via UpstreamDNS. Queries
// that don't match any pattern are sent directly to the host's own system
// DNS servers (PublicDNS) instead of being handled by the DNS proxy at all.
// An empty MatchDomains matches every query, preserving prior behavior.
MatchDomains []string
// Advanced
Holepunch bool
TlsClientCert string
@@ -86,4 +109,12 @@ type TunnelConfig struct {
InitialPostures map[string]any
DisableRelay bool
// PreferLocalRoutes, when enabled, adds tunnel routes with an explicit
// high metric/priority so that an overlapping local/connected route to
// the same destination always takes precedence over the VPN route,
// rather than the two racing based on insertion order. Defaults to
// false, preserving the routing behavior from before this option was
// introduced.
PreferLocalRoutes bool
}
+433 -44
View File
@@ -6,6 +6,7 @@ import (
"strconv"
"strings"
"sync"
"time"
"github.com/fosrl/newt/bind"
"github.com/fosrl/newt/logger"
@@ -33,6 +34,7 @@ type PeerManagerConfig struct {
// WSClient is optional - if nil, relay messages won't be sent
WSClient *websocket.Client
APIServer *api.API
PublicDNS []string
}
type PeerManager struct {
@@ -50,10 +52,35 @@ type PeerManager struct {
// key is the CIDR string, value is a set of siteIds that want this IP
allowedIPClaims map[string]map[int]bool
APIServer *api.API
publicDNS []string
PersistentKeepalive int
routeOptimizerStop chan struct{}
optimizerTrigger chan struct{}
// lastOwnerChange tracks, per allowed-IP CIDR, when ownership was last transferred.
// Used to enforce a cooldown so routes don't flap between two similarly-performing sites.
lastOwnerChange map[string]time.Time
}
const (
// routeSwitchRTTMargin requires a candidate site's RTT to be at least this much
// better (as a fraction) than the current owner's before we consider it worth
// switching, so two similarly-performing sites don't flap back and forth.
routeSwitchRTTMargin = 0.20 // candidate must be >=20% faster
// routeSwitchMinAbsMargin is a floor on the RTT improvement required, so the
// percentage margin above doesn't become meaningless at very low RTTs (e.g. a
// 1ms vs 0.8ms "20% improvement" shouldn't trigger a switch).
routeSwitchMinAbsMargin = 5 * time.Millisecond
// routeSwitchCooldown is the minimum time to wait after transferring ownership
// of a route before it can be transferred again, unless the current owner's
// connection quality degrades (disconnects or falls back to relay).
routeSwitchCooldown = 30 * time.Second
)
// NewPeerManager creates a new PeerManager with an internal PeerMonitor
func NewPeerManager(config PeerManagerConfig) *PeerManager {
pm := &PeerManager{
@@ -65,6 +92,8 @@ func NewPeerManager(config PeerManagerConfig) *PeerManager {
allowedIPOwners: make(map[string]int),
allowedIPClaims: make(map[string]map[int]bool),
APIServer: config.APIServer,
publicDNS: config.PublicDNS,
lastOwnerChange: make(map[string]time.Time),
}
// Create the peer monitor
@@ -74,8 +103,13 @@ func NewPeerManager(config PeerManagerConfig) *PeerManager {
config.LocalIP,
config.SharedBind,
config.APIServer,
config.PublicDNS,
)
pm.optimizerTrigger = make(chan struct{}, 1)
pm.peerMonitor.SetLocalConnectionCallbacks(pm.LocalPeer, pm.UnLocalPeer)
return pm
}
@@ -93,6 +127,21 @@ func (pm *PeerManager) GetPeerMonitor() *monitor.PeerMonitor {
return pm.peerMonitor
}
// SetPublicDNS replaces the DNS servers used to resolve WireGuard peer
// endpoints and hole-punch targets. The servers must be in "host:port" format
// (e.g. "8.8.8.8:53"). The change takes effect for all future peer
// configuration calls; existing WireGuard peers are not re-resolved.
func (pm *PeerManager) SetPublicDNS(servers []string) {
pm.mu.Lock()
pm.publicDNS = servers
mon := pm.peerMonitor
pm.mu.Unlock()
if mon != nil {
mon.SetPublicDNS(servers)
}
}
func (pm *PeerManager) GetAllPeers() []SiteConfig {
pm.mu.RLock()
defer pm.mu.RUnlock()
@@ -107,6 +156,19 @@ func (pm *PeerManager) AddPeer(siteConfig SiteConfig) error {
pm.mu.Lock()
defer pm.mu.Unlock()
for _, alias := range siteConfig.Aliases {
address := net.ParseIP(alias.AliasAddress)
if address == nil {
continue
}
pm.dnsProxy.AddDNSRecord(alias.Alias, address, siteConfig.SiteId)
}
if siteConfig.PublicKey == "" {
logger.Debug("Skip adding site %d because no pub key", siteConfig.SiteId)
return nil
}
// build the allowed IPs list from the remote subnets and aliases and add them to the peer
allowedIPs := make([]string, 0, len(siteConfig.RemoteSubnets)+len(siteConfig.Aliases))
allowedIPs = append(allowedIPs, siteConfig.RemoteSubnets...)
@@ -129,7 +191,7 @@ func (pm *PeerManager) AddPeer(siteConfig SiteConfig) error {
wgConfig := siteConfig
wgConfig.AllowedIps = ownedIPs
if err := ConfigurePeer(pm.device, wgConfig, pm.privateKey, pm.peerMonitor.IsPeerRelayed(siteConfig.SiteId), pm.PersistentKeepalive); err != nil {
if err := ConfigurePeer(pm.device, wgConfig, pm.privateKey, pm.peerMonitor.IsPeerRelayed(siteConfig.SiteId), pm.PersistentKeepalive, pm.publicDNS); err != nil {
return err
}
@@ -139,18 +201,11 @@ func (pm *PeerManager) AddPeer(siteConfig SiteConfig) error {
if err := network.AddRoutes(siteConfig.RemoteSubnets, pm.interfaceName); err != nil {
logger.Error("Failed to add routes for remote subnets: %v", err)
}
for _, alias := range siteConfig.Aliases {
address := net.ParseIP(alias.AliasAddress)
if address == nil {
continue
}
pm.dnsProxy.AddDNSRecord(alias.Alias, address)
}
monitorAddress := strings.Split(siteConfig.ServerIP, "/")[0]
monitorPeer := net.JoinHostPort(monitorAddress, strconv.Itoa(int(siteConfig.ServerPort+1))) // +1 for the monitor port
err := pm.peerMonitor.AddPeer(siteConfig.SiteId, monitorPeer, siteConfig.Endpoint) // always use the real site endpoint for hole punch monitoring
err := pm.peerMonitor.AddPeer(siteConfig.SiteId, monitorPeer, siteConfig.Endpoint, siteConfig.LocalEndpoints) // always use the real site endpoint for hole punch monitoring
if err != nil {
logger.Warn("Failed to setup monitoring for site %d: %v", siteConfig.SiteId, err)
} else {
@@ -159,11 +214,11 @@ func (pm *PeerManager) AddPeer(siteConfig SiteConfig) error {
pm.peers[siteConfig.SiteId] = siteConfig
pm.APIServer.AddPeerStatus(siteConfig.SiteId, siteConfig.Name, false, 0, siteConfig.Endpoint, false)
pm.APIServer.AddPeerStatus(siteConfig.SiteId, siteConfig.Name, false, 0, siteConfig.Endpoint, false, false)
// Perform rapid initial holepunch test (outside of lock to avoid blocking)
// This quickly determines if holepunch is viable and triggers relay if not
go pm.performRapidInitialTest(siteConfig.SiteId, siteConfig.Endpoint)
go pm.performRapidInitialTest(siteConfig.SiteId, siteConfig.Endpoint, siteConfig.LocalEndpoints)
return nil
}
@@ -173,7 +228,7 @@ func (pm *PeerManager) AddPeer(siteConfig SiteConfig) error {
func (pm *PeerManager) UpdateAllPeersPersistentKeepalive(interval int) map[int]error {
pm.mu.RLock()
defer pm.mu.RUnlock()
pm.PersistentKeepalive = interval
errors := make(map[int]error)
@@ -226,7 +281,7 @@ func (pm *PeerManager) RemovePeer(siteId int) error {
}
}
if !subnetStillInUse {
if err := network.RemoveRoutes([]string{subnet}); err != nil {
if err := network.RemoveRoutes([]string{subnet}, pm.interfaceName); err != nil {
logger.Error("Failed to remove route for remote subnet %s: %v", subnet, err)
}
}
@@ -270,7 +325,7 @@ func (pm *PeerManager) RemovePeer(siteId int) error {
ownedIPs := pm.getOwnedAllowedIPs(promotedPeerId)
wgConfig := promotedPeer
wgConfig.AllowedIps = ownedIPs
if err := ConfigurePeer(pm.device, wgConfig, pm.privateKey, pm.peerMonitor.IsPeerRelayed(promotedPeerId), pm.PersistentKeepalive); err != nil {
if err := ConfigurePeer(pm.device, wgConfig, pm.privateKey, pm.peerMonitor.IsPeerRelayed(promotedPeerId), pm.PersistentKeepalive, pm.publicDNS); err != nil {
logger.Error("Failed to update promoted peer %d: %v", promotedPeerId, err)
}
}
@@ -295,6 +350,33 @@ func (pm *PeerManager) UpdatePeer(siteConfig SiteConfig) error {
return fmt.Errorf("peer with site ID %d not found", siteConfig.SiteId)
}
// Preserve the currently active local endpoint (if any) across updates so an in-progress
// local connection isn't disrupted by an unrelated site update.
siteConfig.ActiveLocalEndpoint = oldPeer.ActiveLocalEndpoint
// Update aliases
// Remove old aliases
for _, alias := range oldPeer.Aliases {
address := net.ParseIP(alias.AliasAddress)
if address == nil {
continue
}
pm.dnsProxy.RemoveDNSRecord(alias.Alias, address)
}
// Add new aliases
for _, alias := range siteConfig.Aliases {
address := net.ParseIP(alias.AliasAddress)
if address == nil {
continue
}
pm.dnsProxy.AddDNSRecord(alias.Alias, address, siteConfig.SiteId)
}
if siteConfig.PublicKey == "" {
logger.Debug("Skip updating site %d because no pub key", siteConfig.SiteId)
return nil
}
// If public key changed, remove old peer first
if siteConfig.PublicKey != oldPeer.PublicKey {
if err := RemovePeer(pm.device, siteConfig.SiteId, oldPeer.PublicKey); err != nil {
@@ -346,7 +428,7 @@ func (pm *PeerManager) UpdatePeer(siteConfig SiteConfig) error {
wgConfig := siteConfig
wgConfig.AllowedIps = ownedIPs
if err := ConfigurePeer(pm.device, wgConfig, pm.privateKey, pm.peerMonitor.IsPeerRelayed(siteConfig.SiteId), pm.PersistentKeepalive); err != nil {
if err := ConfigurePeer(pm.device, wgConfig, pm.privateKey, pm.peerMonitor.IsPeerRelayed(siteConfig.SiteId), pm.PersistentKeepalive, pm.publicDNS); err != nil {
return err
}
@@ -356,7 +438,7 @@ func (pm *PeerManager) UpdatePeer(siteConfig SiteConfig) error {
promotedOwnedIPs := pm.getOwnedAllowedIPs(promotedPeerId)
promotedWgConfig := promotedPeer
promotedWgConfig.AllowedIps = promotedOwnedIPs
if err := ConfigurePeer(pm.device, promotedWgConfig, pm.privateKey, pm.peerMonitor.IsPeerRelayed(promotedPeerId), pm.PersistentKeepalive); err != nil {
if err := ConfigurePeer(pm.device, promotedWgConfig, pm.privateKey, pm.peerMonitor.IsPeerRelayed(promotedPeerId), pm.PersistentKeepalive, pm.publicDNS); err != nil {
logger.Error("Failed to update promoted peer %d: %v", promotedPeerId, err)
}
}
@@ -405,7 +487,7 @@ func (pm *PeerManager) UpdatePeer(siteConfig SiteConfig) error {
}
}
if !subnetStillInUse {
if err := network.RemoveRoutes([]string{subnet}); err != nil {
if err := network.RemoveRoutes([]string{subnet}, pm.interfaceName); err != nil {
logger.Error("Failed to remove route for subnet %s: %v", subnet, err)
}
}
@@ -418,25 +500,8 @@ func (pm *PeerManager) UpdatePeer(siteConfig SiteConfig) error {
}
}
// Update aliases
// Remove old aliases
for _, alias := range oldPeer.Aliases {
address := net.ParseIP(alias.AliasAddress)
if address == nil {
continue
}
pm.dnsProxy.RemoveDNSRecord(alias.Alias, address)
}
// Add new aliases
for _, alias := range siteConfig.Aliases {
address := net.ParseIP(alias.AliasAddress)
if address == nil {
continue
}
pm.dnsProxy.AddDNSRecord(alias.Alias, address)
}
pm.peerMonitor.UpdateHolepunchEndpoint(siteConfig.SiteId, siteConfig.Endpoint)
pm.peerMonitor.UpdateLocalEndpoints(siteConfig.SiteId, siteConfig.LocalEndpoints)
monitorAddress := strings.Split(siteConfig.ServerIP, "/")[0]
monitorPeer := net.JoinHostPort(monitorAddress, strconv.Itoa(int(siteConfig.ServerPort+1))) // +1 for the monitor port
@@ -472,6 +537,7 @@ func (pm *PeerManager) releaseAllowedIP(siteId int, cidr string) (newOwner int,
delete(claims, siteId)
if len(claims) == 0 {
delete(pm.allowedIPClaims, cidr)
delete(pm.lastOwnerChange, cidr)
}
}
@@ -690,7 +756,7 @@ func (pm *PeerManager) RemoveRemoteSubnet(siteId int, ip string) error {
// Only remove route if no other peer needs it
if !subnetStillInUse {
if err := network.RemoveRoutes([]string{ip}); err != nil {
if err := network.RemoveRoutes([]string{ip}, pm.interfaceName); err != nil {
return err
}
}
@@ -713,7 +779,7 @@ func (pm *PeerManager) AddAlias(siteId int, alias Alias) error {
address := net.ParseIP(alias.AliasAddress)
if address != nil {
pm.dnsProxy.AddDNSRecord(alias.Alias, address)
pm.dnsProxy.AddDNSRecord(alias.Alias, address, siteId)
}
// Add an allowed IP for the alias
@@ -747,7 +813,7 @@ func (pm *PeerManager) RemoveAlias(siteId int, aliasName string) error {
if aliasToRemove != nil {
address := net.ParseIP(aliasToRemove.AliasAddress)
if address != nil {
pm.dnsProxy.RemoveDNSRecord(aliasName, address)
pm.dnsProxy.RemoveDNSRecordForSite(aliasName, address, siteId)
}
}
@@ -778,6 +844,11 @@ func (pm *PeerManager) RemoveAlias(siteId int, aliasName string) error {
func (pm *PeerManager) RelayPeer(siteId int, relayEndpoint string, relayPort uint16) {
pm.mu.Lock()
peer, exists := pm.peers[siteId]
if exists && peer.ActiveLocalEndpoint != "" {
pm.mu.Unlock()
logger.Info("Ignoring relay request for site %d: local connection is active", siteId)
return
}
if exists {
// Store the relay endpoint
peer.RelayEndpoint = relayEndpoint
@@ -820,15 +891,43 @@ endpoint=%s:%d`, util.FixKey(peer.PublicKey), formattedEndpoint, relayPort)
}
// performRapidInitialTest performs a rapid holepunch test for a newly added peer.
// If the test fails, it immediately requests relay to minimize connection delay.
// This runs in a goroutine to avoid blocking AddPeer.
func (pm *PeerManager) performRapidInitialTest(siteId int, endpoint string) {
// It races a test of the public endpoint against a test of the local candidate endpoints
// (if any) and waits for both to finish before acting, so we never request relay only to
// have it immediately superseded by a local connection (or vice versa). Local wins if it's
// viable at all; otherwise relay is requested only if the public endpoint isn't viable.
// This runs in a goroutine to avoid blocking AddPeer - the peer already starts out pointed
// at the public endpoint (set synchronously in AddPeer), so this just settles the peer onto
// its steady-state connection within ~1-2 seconds.
func (pm *PeerManager) performRapidInitialTest(siteId int, endpoint string, localEndpoints []string) {
if pm.peerMonitor == nil {
return
}
// Perform rapid test - this takes ~1-2 seconds max
holepunchViable := pm.peerMonitor.RapidTestPeer(siteId, endpoint)
var wg sync.WaitGroup
var localWinner string
var holepunchViable bool
if len(localEndpoints) > 0 {
wg.Add(1)
go func() {
defer wg.Done()
localWinner = pm.peerMonitor.RapidTestLocalEndpoints(siteId, localEndpoints)
}()
}
wg.Add(1)
go func() {
defer wg.Done()
holepunchViable = pm.peerMonitor.RapidTestPeer(siteId, endpoint)
}()
wg.Wait()
if localWinner != "" {
logger.Info("Rapid test: local connection viable for site %d, switching to %s", siteId, localWinner)
pm.LocalPeer(siteId, localWinner)
return
}
if !holepunchViable {
// Holepunch failed rapid test, request relay immediately
@@ -846,10 +945,12 @@ func (pm *PeerManager) Start() {
if pm.peerMonitor != nil {
pm.peerMonitor.Start()
}
pm.startRouteOptimizer()
}
// Stop stops the peer monitor
func (pm *PeerManager) Stop() {
pm.stopRouteOptimizer()
if pm.peerMonitor != nil {
pm.peerMonitor.Stop()
}
@@ -857,6 +958,7 @@ func (pm *PeerManager) Stop() {
// Close stops the peer monitor and cleans up resources
func (pm *PeerManager) Close() {
pm.stopRouteOptimizer()
if pm.peerMonitor != nil {
pm.peerMonitor.Close()
pm.peerMonitor = nil
@@ -887,6 +989,11 @@ func (pm *PeerManager) MarkPeerRelayed(siteID int, relayed bool) {
func (pm *PeerManager) UnRelayPeer(siteId int, endpoint string) error {
pm.mu.Lock()
peer, exists := pm.peers[siteId]
if exists && peer.ActiveLocalEndpoint != "" {
pm.mu.Unlock()
logger.Info("Ignoring unrelay request for site %d: local connection is active", siteId)
return nil
}
if exists {
// Store the relay endpoint
peer.Endpoint = endpoint
@@ -918,3 +1025,285 @@ endpoint=%s`, util.FixKey(peer.PublicKey), endpoint)
logger.Info("Switched peer %d back to direct connection at %s", siteId, endpoint)
return nil
}
// LocalPeer switches a peer to a local network endpoint discovered by the peer monitor.
// Local endpoints take priority over both the public endpoint and the relay, so this
// bypasses relay/public-endpoint bookkeeping entirely and just updates the WireGuard
// endpoint directly.
func (pm *PeerManager) LocalPeer(siteId int, localEndpoint string) {
pm.mu.Lock()
peer, exists := pm.peers[siteId]
if exists {
peer.ActiveLocalEndpoint = localEndpoint
pm.peers[siteId] = peer
}
pm.mu.Unlock()
if !exists {
logger.Error("Cannot switch to local connection: peer with site ID %d not found", siteId)
return
}
// Update only the endpoint for this peer (update_only preserves other settings)
wgConfig := fmt.Sprintf(`public_key=%s
update_only=true
endpoint=%s`, util.FixKey(peer.PublicKey), localEndpoint)
if err := pm.device.IpcSet(wgConfig); err != nil {
logger.Error("Failed to switch peer %d to local connection: %v", siteId, err)
return
}
if pm.APIServer != nil {
pm.APIServer.UpdatePeerLocalStatus(siteId, localEndpoint, true)
}
logger.Info("Switched peer %d to local connection at %s", siteId, localEndpoint)
}
// UnLocalPeer switches a peer away from its active local endpoint back to the public
// endpoint, resuming the normal public/relay monitoring logic from scratch (which will
// re-trigger relay on its own if the public endpoint also turns out to be unreachable).
func (pm *PeerManager) UnLocalPeer(siteId int) {
pm.mu.Lock()
peer, exists := pm.peers[siteId]
publicDNS := pm.publicDNS
if exists {
peer.ActiveLocalEndpoint = ""
pm.peers[siteId] = peer
}
pm.mu.Unlock()
if !exists {
logger.Error("Cannot fall back from local connection: peer with site ID %d not found", siteId)
return
}
resolved, err := util.ResolveDomainUpstream(formatEndpoint(peer.Endpoint), publicDNS)
if err != nil {
logger.Error("Failed to resolve fallback endpoint for peer %d: %v", siteId, err)
return
}
if err := pm.UnRelayPeer(siteId, resolved); err != nil {
logger.Error("Failed to fall back peer %d from local connection: %v", siteId, err)
return
}
if pm.APIServer != nil {
pm.APIServer.UpdatePeerLocalStatus(siteId, resolved, false)
}
}
// isBetterConnection returns true if connection quality (a) is better than (b).
// Priority: connected > disconnected, then direct > relayed, then lower RTT.
func isBetterConnection(aConn bool, aRelay bool, aRTT time.Duration,
bConn bool, bRelay bool, bRTT time.Duration) bool {
if aConn != bConn {
return aConn // connected beats disconnected
}
if !aConn {
return false // both offline, no preference
}
if aRelay != bRelay {
return !aRelay // direct beats relayed
}
// Same connectivity class: prefer lower RTT
if aRTT == 0 {
return false // unknown RTT, don't displace
}
if bRTT == 0 {
return true // current has no RTT data, prefer known
}
return aRTT < bRTT
}
// selectBestOwner returns the siteId of the best site to own the given IP,
// based on connection quality. Must be called with pm.mu held.
func (pm *PeerManager) selectBestOwner(claims map[int]bool) int {
bestSiteId := -1
var bestConn, bestRelay bool
var bestRTT time.Duration
for siteId := range claims {
conn, relay, rtt := pm.peerMonitor.GetConnectionQuality(siteId)
if bestSiteId < 0 || isBetterConnection(conn, relay, rtt, bestConn, bestRelay, bestRTT) {
bestSiteId = siteId
bestConn = conn
bestRelay = relay
bestRTT = rtt
}
}
return bestSiteId
}
// shouldSwitchOwner decides whether ownership of cidr should move from the current
// owner to the candidate. It applies hysteresis so two sites with roughly equal
// performance don't flap back and forth:
// - A switch driven by connectivity class (connected vs not, direct vs relayed) is
// always allowed immediately - these are correctness issues, not noise.
// - A switch driven purely by RTT requires both a minimum improvement margin and
// that the cooldown since the last switch of this route has elapsed.
//
// Must be called with pm.mu held.
func (pm *PeerManager) shouldSwitchOwner(cidr string, currentSiteId, candidateSiteId int) bool {
curConn, curRelay, curRTT := pm.peerMonitor.GetConnectionQuality(currentSiteId)
candConn, candRelay, candRTT := pm.peerMonitor.GetConnectionQuality(candidateSiteId)
// Connectivity-class differences (up/down, direct/relayed) are not subject to
// hysteresis - always act on them so we don't stay stuck on a broken route.
if curConn != candConn || curRelay != candRelay {
return true
}
if !curConn {
return false // both down, nothing to do
}
// Same connectivity class: only switch on a meaningful, sustained RTT win.
if candRTT == 0 || curRTT == 0 {
return false
}
minImprovement := time.Duration(float64(curRTT) * routeSwitchRTTMargin)
if minImprovement < routeSwitchMinAbsMargin {
minImprovement = routeSwitchMinAbsMargin
}
if candRTT > curRTT-minImprovement {
return false // not enough of an improvement to be worth switching
}
if lastChange, ok := pm.lastOwnerChange[cidr]; ok {
if time.Since(lastChange) < routeSwitchCooldown {
return false // switched too recently, avoid flapping
}
}
return true
}
// getWireGuardAllowedIPs returns the full set of IPs that should be in WireGuard
// for a peer: server IP /32 plus all shared IPs it currently owns.
// Must be called with pm.mu held.
func (pm *PeerManager) getWireGuardAllowedIPs(siteId int) []string {
peer, exists := pm.peers[siteId]
if !exists {
return nil
}
serverIP := strings.Split(peer.ServerIP, "/")[0] + "/32"
ips := []string{serverIP}
for cidr, owner := range pm.allowedIPOwners {
if owner == siteId {
ips = append(ips, cidr)
}
}
return ips
}
// transferOwnership moves WireGuard ownership of cidr from fromSiteId to toSiteId.
// Must be called with pm.mu held.
func (pm *PeerManager) transferOwnership(cidr string, fromSiteId int, toSiteId int) error {
// Update owner map first
pm.allowedIPOwners[cidr] = toSiteId
// Remove cidr from old owner's WireGuard allowed IPs
if fromPeer, exists := pm.peers[fromSiteId]; exists {
remaining := pm.getWireGuardAllowedIPs(fromSiteId) // cidr is no longer in owners, so it won't appear here
if err := RemoveAllowedIP(pm.device, fromPeer.PublicKey, remaining); err != nil {
// Revert
pm.allowedIPOwners[cidr] = fromSiteId
return fmt.Errorf("remove IP %s from site %d: %v", cidr, fromSiteId, err)
}
}
// Add cidr to new owner's WireGuard allowed IPs
if toPeer, exists := pm.peers[toSiteId]; exists {
if err := AddAllowedIP(pm.device, toPeer.PublicKey, cidr); err != nil {
return fmt.Errorf("add IP %s to site %d: %v", cidr, toSiteId, err)
}
}
return nil
}
// optimizeRoutes evaluates all shared IPs and reassigns ownership to the best site.
func (pm *PeerManager) optimizeRoutes() {
pm.mu.Lock()
defer pm.mu.Unlock()
for cidr, claims := range pm.allowedIPClaims {
if len(claims) <= 1 {
continue // No competition, nothing to optimize
}
currentOwner, hasOwner := pm.allowedIPOwners[cidr]
bestOwner := pm.selectBestOwner(claims)
if bestOwner < 0 {
continue
}
if hasOwner && currentOwner == bestOwner {
continue // Already on the best site
}
if !hasOwner {
// No current owner, just assign
pm.allowedIPOwners[cidr] = bestOwner
pm.lastOwnerChange[cidr] = time.Now()
if toPeer, exists := pm.peers[bestOwner]; exists {
if err := AddAllowedIP(pm.device, toPeer.PublicKey, cidr); err != nil {
logger.Error("Failed to assign IP %s to site %d: %v", cidr, bestOwner, err)
}
}
continue
}
if !pm.shouldSwitchOwner(cidr, currentOwner, bestOwner) {
continue // Not a big enough or sustained enough improvement, avoid flapping
}
logger.Info("Route optimizer: moving %s from site %d to site %d", cidr, currentOwner, bestOwner)
if err := pm.transferOwnership(cidr, currentOwner, bestOwner); err != nil {
logger.Error("Failed to transfer ownership of %s from site %d to site %d: %v",
cidr, currentOwner, bestOwner, err)
} else {
pm.lastOwnerChange[cidr] = time.Now()
}
}
}
// startRouteOptimizer registers the status-change callback and launches the optimizer goroutine.
func (pm *PeerManager) startRouteOptimizer() {
pm.routeOptimizerStop = make(chan struct{})
// Trigger optimization whenever any peer's connection status changes
if pm.peerMonitor != nil {
pm.peerMonitor.SetStatusChangeCallback(func(_ int) {
select {
case pm.optimizerTrigger <- struct{}{}:
default:
}
})
}
go func() {
ticker := time.NewTicker(5 * time.Second)
defer ticker.Stop()
for {
select {
case <-pm.routeOptimizerStop:
return
case <-pm.optimizerTrigger:
pm.optimizeRoutes()
case <-ticker.C:
pm.optimizeRoutes()
}
}
}()
}
// stopRouteOptimizer stops the route optimizer goroutine if it is running.
func (pm *PeerManager) stopRouteOptimizer() {
if pm.routeOptimizerStop != nil {
close(pm.routeOptimizerStop)
pm.routeOptimizerStop = nil
}
}
+531 -36
View File
@@ -2,6 +2,8 @@ package monitor
import (
"context"
"crypto/rand"
"encoding/hex"
"fmt"
"net"
"net/netip"
@@ -31,9 +33,14 @@ type PeerMonitor struct {
monitors map[int]*Client
mutex sync.Mutex
running bool
timeout time.Duration
timeout time.Duration
maxAttempts int
wsClient *websocket.Client
publicDNS []string
// Relay sender tracking
relaySends map[string]func()
relaySendMu sync.Mutex
// Netstack fields
middleDev *middleDevice.MiddleDevice
@@ -47,19 +54,36 @@ type PeerMonitor struct {
nsWg sync.WaitGroup
// Holepunch testing fields
sharedBind *bind.SharedBind
holepunchTester *holepunch.HolepunchTester
holepunchTimeout time.Duration
holepunchEndpoints map[int]string // siteID -> endpoint for holepunch testing
holepunchStatus map[int]bool // siteID -> connected status
holepunchStopChan chan struct{}
holepunchUpdateChan chan struct{}
sharedBind *bind.SharedBind
holepunchTester *holepunch.HolepunchTester
holepunchTimeout time.Duration
holepunchEndpoints map[int]string // siteID -> endpoint for holepunch testing
holepunchStatus map[int]bool // siteID -> connected status
holepunchStopChan chan struct{}
holepunchUpdateChan chan struct{}
// Relay tracking fields
relayedPeers map[int]bool // siteID -> whether the peer is currently relayed
holepunchMaxAttempts int // max consecutive failures before triggering relay
holepunchFailures map[int]int // siteID -> consecutive failure count
// Local endpoint testing fields. Local endpoints are ip:port addresses on the
// site host's local network interfaces (ordered best-to-worst by the server).
// When one is reachable it takes priority over both the public endpoint and
// the relay.
localEndpoints map[int][]string // siteID -> ordered candidate local endpoints
localActiveEndpoint map[int]string // siteID -> currently active local endpoint ("" = not using local)
localFailures map[int]int // siteID -> consecutive failures of the active local endpoint
localTestTimeout time.Duration // timeout for each local endpoint probe
// Local connection switch callbacks, set by the PeerManager
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
// Exponential backoff fields for holepunch monitor
defaultHolepunchMinInterval time.Duration // Minimum interval (initial)
defaultHolepunchMaxInterval time.Duration
@@ -78,11 +102,19 @@ type PeerMonitor struct {
apiServer *api.API
// WG connection status tracking
wgConnectionStatus map[int]bool // siteID -> WG connected status
wgConnectionStatus map[int]bool // siteID -> WG connected status
wgConnectionRTT map[int]time.Duration // siteID -> last known RTT
statusChangeCallback func(siteId int) // called when any peer's connection status changes
}
// NewPeerMonitor creates a new peer monitor with the given callback
func NewPeerMonitor(wsClient *websocket.Client, middleDev *middleDevice.MiddleDevice, localIP string, sharedBind *bind.SharedBind, apiServer *api.API) *PeerMonitor {
func generateChainId() string {
b := make([]byte, 8)
_, _ = rand.Read(b)
return hex.EncodeToString(b)
}
func NewPeerMonitor(wsClient *websocket.Client, middleDev *middleDevice.MiddleDevice, localIP string, sharedBind *bind.SharedBind, apiServer *api.API, publicDNS []string) *PeerMonitor {
ctx, cancel := context.WithCancel(context.Background())
pm := &PeerMonitor{
monitors: make(map[int]*Client),
@@ -91,6 +123,7 @@ func NewPeerMonitor(wsClient *websocket.Client, middleDev *middleDevice.MiddleDe
wsClient: wsClient,
middleDev: middleDev,
localIP: localIP,
publicDNS: publicDNS,
activePorts: make(map[uint16]bool),
nsCtx: ctx,
nsCancel: cancel,
@@ -99,14 +132,21 @@ 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()),
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()),
// Rapid initial test settings: complete within ~1.5 seconds
rapidTestInterval: 200 * time.Millisecond, // 200ms between attempts
rapidTestTimeout: 400 * time.Millisecond, // 400ms timeout per attempt
rapidTestMaxAttempts: 5, // 5 attempts = ~1-1.5 seconds total
apiServer: apiServer,
wgConnectionStatus: make(map[int]bool),
wgConnectionRTT: make(map[int]time.Duration),
// Exponential backoff settings for holepunch monitor
defaultHolepunchMinInterval: 2 * time.Second,
defaultHolepunchMaxInterval: 30 * time.Second,
@@ -124,12 +164,25 @@ func NewPeerMonitor(wsClient *websocket.Client, middleDev *middleDevice.MiddleDe
// Initialize holepunch tester if sharedBind is available
if sharedBind != nil {
pm.holepunchTester = holepunch.NewHolepunchTester(sharedBind)
pm.holepunchTester = holepunch.NewHolepunchTester(sharedBind, publicDNS)
}
return pm
}
// SetPublicDNS replaces the DNS servers used to resolve peer endpoints and
// hole-punch exit nodes. The servers must be in "host:port" format.
func (pm *PeerMonitor) SetPublicDNS(servers []string) {
pm.mutex.Lock()
pm.publicDNS = servers
tester := pm.holepunchTester
pm.mutex.Unlock()
if tester != nil {
tester.SetPublicDNS(servers)
}
}
// SetInterval changes how frequently peers are checked
func (pm *PeerMonitor) SetPeerInterval(minInterval, maxInterval time.Duration) {
pm.mutex.Lock()
@@ -204,7 +257,7 @@ func (pm *PeerMonitor) ResetPeerHolepunchInterval() {
}
// AddPeer adds a new peer to monitor
func (pm *PeerMonitor) AddPeer(siteID int, endpoint string, holepunchEndpoint string) error {
func (pm *PeerMonitor) AddPeer(siteID int, endpoint string, holepunchEndpoint string, localEndpoints []string) error {
pm.mutex.Lock()
defer pm.mutex.Unlock()
@@ -222,6 +275,9 @@ func (pm *PeerMonitor) AddPeer(siteID int, endpoint string, holepunchEndpoint st
pm.holepunchEndpoints[siteID] = holepunchEndpoint
pm.holepunchStatus[siteID] = false // Initially unknown/disconnected
pm.localEndpoints[siteID] = localEndpoints
pm.localActiveEndpoint[siteID] = ""
pm.localFailures[siteID] = 0
if pm.running {
if err := client.StartMonitor(func(status ConnectionStatus) {
@@ -244,6 +300,25 @@ func (pm *PeerMonitor) UpdateHolepunchEndpoint(siteID int, endpoint string) {
logger.Debug("Updated holepunch endpoint for site %d to %s", siteID, endpoint)
}
// UpdateLocalEndpoints updates the candidate local endpoints for a peer
func (pm *PeerMonitor) UpdateLocalEndpoints(siteID int, localEndpoints []string) {
pm.mutex.Lock()
defer pm.mutex.Unlock()
pm.localEndpoints[siteID] = localEndpoints
logger.Debug("Updated local endpoints for site %d: %v", siteID, localEndpoints)
}
// SetLocalConnectionCallbacks registers the callbacks invoked when a peer switches to
// or falls back from a local network endpoint. onLocal is called with the endpoint that
// became active; onFallback is called when we give up on the active local endpoint and
// resume the normal public/relay monitoring logic.
func (pm *PeerMonitor) SetLocalConnectionCallbacks(onLocal func(siteId int, endpoint string), onFallback func(siteId int)) {
pm.mutex.Lock()
defer pm.mutex.Unlock()
pm.localSwitchCallback = onLocal
pm.localFallbackCallback = onFallback
}
// RapidTestPeer performs a rapid connectivity test for a newly added peer.
// This is designed to quickly determine if holepunch is viable within ~1-2 seconds.
// Returns true if the connection is viable (holepunch works), false if it should relay.
@@ -295,6 +370,126 @@ func (pm *PeerMonitor) RapidTestPeer(siteID int, endpoint string) bool {
return false
}
// RapidTestLocalEndpoints performs a rapid connectivity test of local candidate endpoints
// for a newly added peer, so local viability is known within the same ~1-2 second window as
// RapidTestPeer's public-endpoint test (rather than waiting for the next backoff-loop tick,
// which could be tens of seconds away). Candidates are tried in order (best-to-worst) and
// the first reachable one wins. Returns the winning endpoint, or "" if none are reachable.
func (pm *PeerMonitor) RapidTestLocalEndpoints(siteID int, endpoints []string) string {
if pm.holepunchTester == nil || len(endpoints) == 0 {
return ""
}
pm.mutex.Lock()
timeout := pm.rapidTestTimeout
pm.mutex.Unlock()
for _, endpoint := range endpoints {
result := pm.holepunchTester.TestEndpoint(endpoint, timeout)
if !result.Success {
continue
}
logger.Info("Rapid test: local endpoint %s for site %d SUCCEEDED (RTT: %v)", endpoint, siteID, result.RTT)
pm.mutex.Lock()
// Peer may have been removed while we were testing.
stillTracked := false
if _, tracked := pm.localEndpoints[siteID]; tracked {
stillTracked = true
pm.localActiveEndpoint[siteID] = endpoint
pm.localFailures[siteID] = 0
}
pm.mutex.Unlock()
if stillTracked {
pm.sendLocal(siteID, endpoint)
}
return endpoint
}
logger.Info("Rapid test: no local endpoint reachable for site %d", siteID)
return ""
}
// remainingLocalCandidates returns all of endpoints except exclude, preserving order.
func remainingLocalCandidates(endpoints []string, exclude string) []string {
remaining := make([]string, 0, len(endpoints))
for _, ep := range endpoints {
if ep != exclude {
remaining = append(remaining, ep)
}
}
return remaining
}
// rapidTestOnLocalFallback runs a fast (~1-2 second) test of the public endpoint, racing it
// against any remaining untried local candidates, immediately after we fall back from a dead
// active local endpoint. Without this, the peer would sit on the public endpoint - which may
// itself be unreachable - relying on the normal checkHolepunchEndpoints loop to notice, which
// can take tens of seconds if the holepunch backoff interval had climbed while the local
// endpoint was stable. If neither the public endpoint nor a local candidate is reachable, relay
// is requested immediately. Mirrors PeerManager.performRapidInitialTest's race, but is triggered
// by local-endpoint failure rather than initial peer setup.
func (pm *PeerMonitor) rapidTestOnLocalFallback(siteID int, publicEndpoint string, remainingLocal []string) {
if pm.holepunchTester == nil {
return
}
var wg sync.WaitGroup
var localWinner string
var holepunchViable bool
if len(remainingLocal) > 0 {
wg.Add(1)
go func() {
defer wg.Done()
localWinner = pm.RapidTestLocalEndpoints(siteID, remainingLocal)
}()
}
if publicEndpoint != "" {
wg.Add(1)
go func() {
defer wg.Done()
holepunchViable = pm.RapidTestPeer(siteID, publicEndpoint)
}()
}
wg.Wait()
pm.mutex.Lock()
_, stillTracked := pm.localEndpoints[siteID]
noLocalActiveYet := pm.localActiveEndpoint[siteID] == ""
switchCb := pm.localSwitchCallback
pm.mutex.Unlock()
if !stillTracked {
return // peer was removed while we were testing
}
if localWinner != "" {
// RapidTestLocalEndpoints already recorded the new active endpoint and notified the
// server, but doesn't move the WireGuard peer itself - do that here, unless a
// concurrent checkLocalEndpoints tick already beat us to activating something.
if noLocalActiveYet && switchCb != nil {
switchCb(siteID, localWinner)
}
logger.Info("Rapid fallback test: local connection %s viable for site %d", localWinner, siteID)
return
}
if !holepunchViable {
logger.Warn("Rapid fallback test: site %d unreachable on public endpoint after local fallback, requesting relay", siteID)
if pm.wsClient != nil {
pm.sendRelay(siteID)
}
} else {
logger.Info("Rapid fallback test: site %d reachable on public endpoint after local fallback", siteID)
}
}
// UpdatePeerEndpoint updates the monitor endpoint for a peer
func (pm *PeerMonitor) UpdatePeerEndpoint(siteID int, monitorPeer string) {
pm.mutex.Lock()
@@ -328,15 +523,18 @@ func (pm *PeerMonitor) removePeerUnlocked(siteID int) {
// RemovePeer stops monitoring a peer and removes it from the monitor
func (pm *PeerMonitor) RemovePeer(siteID int) {
pm.mutex.Lock()
defer pm.mutex.Unlock()
// remove the holepunch endpoint info
delete(pm.holepunchEndpoints, siteID)
delete(pm.holepunchStatus, siteID)
delete(pm.relayedPeers, siteID)
delete(pm.holepunchFailures, siteID)
delete(pm.localEndpoints, siteID)
delete(pm.localActiveEndpoint, siteID)
delete(pm.localFailures, siteID)
pm.removePeerUnlocked(siteID)
pm.mutex.Unlock()
}
func (pm *PeerMonitor) RemoveHolepunchEndpoint(siteID int) {
@@ -377,10 +575,22 @@ func (pm *PeerMonitor) handleConnectionStatusChange(siteID int, status Connectio
pm.mutex.Lock()
previousStatus, exists := pm.wgConnectionStatus[siteID]
pm.wgConnectionStatus[siteID] = status.Connected
if status.Connected && status.RTT > 0 {
pm.wgConnectionRTT[siteID] = status.RTT
}
isRelayed := pm.relayedPeers[siteID]
localEndpoint := pm.localActiveEndpoint[siteID]
endpoint := pm.holepunchEndpoints[siteID]
pm.mutex.Unlock()
isLocal := localEndpoint != ""
if isLocal {
// Report the active local endpoint rather than the public one; local and relay
// are mutually exclusive.
endpoint = localEndpoint
isRelayed = false
}
// Log status changes
if !exists || previousStatus != status.Connected {
if status.Connected {
@@ -392,24 +602,32 @@ func (pm *PeerMonitor) handleConnectionStatusChange(siteID int, status Connectio
// Update API with connection status
if pm.apiServer != nil {
pm.apiServer.UpdatePeerStatus(siteID, status.Connected, status.RTT, endpoint, isRelayed)
pm.apiServer.UpdatePeerStatus(siteID, status.Connected, status.RTT, endpoint, isRelayed, isLocal)
}
// Notify route optimizer of status change
if pm.statusChangeCallback != nil {
pm.statusChangeCallback(siteID)
}
}
// sendRelay sends a relay message to the server
// sendRelay sends a relay message to the server with retry, keyed by chainId
func (pm *PeerMonitor) sendRelay(siteID int) error {
if pm.wsClient == nil {
return fmt.Errorf("websocket client is nil")
}
err := pm.wsClient.SendMessage("olm/wg/relay", map[string]interface{}{
"siteId": siteID,
})
if err != nil {
logger.Error("Failed to send registration message: %v", err)
return err
}
logger.Info("Sent relay message")
chainId := generateChainId()
stopFunc, _ := pm.wsClient.SendMessageInterval("olm/wg/relay", map[string]interface{}{
"siteId": siteID,
"chainId": chainId,
}, 2*time.Second, 10)
pm.relaySendMu.Lock()
pm.relaySends[chainId] = stopFunc
pm.relaySendMu.Unlock()
logger.Info("Sent relay message for site %d (chain %s)", siteID, chainId)
return nil
}
@@ -419,23 +637,121 @@ func (pm *PeerMonitor) RequestRelay(siteID int) error {
return pm.sendRelay(siteID)
}
// sendUnRelay sends an unrelay message to the server
// sendUnRelay sends an unrelay message to the server with retry, keyed by chainId
func (pm *PeerMonitor) sendUnRelay(siteID int) error {
if pm.wsClient == nil {
return fmt.Errorf("websocket client is nil")
}
err := pm.wsClient.SendMessage("olm/wg/unrelay", map[string]interface{}{
"siteId": siteID,
})
if err != nil {
logger.Error("Failed to send registration message: %v", err)
return err
}
logger.Info("Sent unrelay message")
chainId := generateChainId()
stopFunc, _ := pm.wsClient.SendMessageInterval("olm/wg/unrelay", map[string]interface{}{
"siteId": siteID,
"chainId": chainId,
}, 2*time.Second, 10)
pm.relaySendMu.Lock()
pm.relaySends[chainId] = stopFunc
pm.relaySendMu.Unlock()
logger.Info("Sent unrelay message for site %d (chain %s)", siteID, chainId)
return nil
}
// sendLocal notifies the server that this peer switched to a local network endpoint, with
// retry keyed by chainId. This is informational (e.g. so the server can relay the information
// to newt) - olm does not wait for an acknowledgement before using the local connection, but
// it does stop retrying once the server acks via CancelLocalSend, same as relay/unrelay.
func (pm *PeerMonitor) sendLocal(siteID int, endpoint string) {
if pm.wsClient == nil {
return
}
chainId := generateChainId()
stopFunc, _ := pm.wsClient.SendMessageInterval("olm/wg/local", map[string]interface{}{
"siteId": siteID,
"endpoint": endpoint,
"chainId": chainId,
}, 2*time.Second, 10)
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)
}
// sendUnLocal notifies the server that this peer fell back from its local network endpoint,
// with retry keyed by chainId.
func (pm *PeerMonitor) sendUnLocal(siteID int) {
if pm.wsClient == nil {
return
}
chainId := generateChainId()
stopFunc, _ := pm.wsClient.SendMessageInterval("olm/wg/unlocal", map[string]interface{}{
"siteId": siteID,
"chainId": chainId,
}, 2*time.Second, 10)
pm.localSendMu.Lock()
pm.localSends[chainId] = stopFunc
pm.localSendMu.Unlock()
logger.Info("Sent 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.
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)
}
logger.Info("Cancelled all local-connection senders")
return
}
if stop, ok := pm.localSends[chainId]; ok {
stop()
delete(pm.localSends, 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.
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)
}
logger.Info("Cancelled all relay senders")
return
}
if stop, ok := pm.relaySends[chainId]; ok {
stop()
delete(pm.relaySends, chainId)
logger.Info("Cancelled relay sender for chain %s", chainId)
} else {
logger.Warn("CancelRelaySend: no active sender for chain %s", chainId)
}
}
// Stop stops monitoring all peers
func (pm *PeerMonitor) Stop() {
// Stop holepunch monitor first (outside of mutex to avoid deadlock)
@@ -474,6 +790,25 @@ func (pm *PeerMonitor) IsPeerRelayed(siteID int) bool {
return pm.relayedPeers[siteID]
}
// SetStatusChangeCallback registers a callback that is invoked whenever a peer's
// WireGuard connection status changes (connected/disconnected). The callback must
// be non-blocking (e.g., send to a buffered channel).
func (pm *PeerMonitor) SetStatusChangeCallback(cb func(siteId int)) {
pm.mutex.Lock()
defer pm.mutex.Unlock()
pm.statusChangeCallback = cb
}
// GetConnectionQuality returns the current connection quality metrics for a peer.
func (pm *PeerMonitor) GetConnectionQuality(siteId int) (connected bool, relayed bool, rtt time.Duration) {
pm.mutex.Lock()
defer pm.mutex.Unlock()
connected = pm.wgConnectionStatus[siteId]
relayed = pm.relayedPeers[siteId]
rtt = pm.wgConnectionRTT[siteId]
return
}
// startHolepunchMonitor starts the holepunch connection monitoring
// Note: This function assumes the mutex is already held by the caller (called from Start())
func (pm *PeerMonitor) startHolepunchMonitor() error {
@@ -534,11 +869,12 @@ func (pm *PeerMonitor) runHolepunchMonitor() {
pm.holepunchCurrentInterval = pm.holepunchMinInterval
currentInterval := pm.holepunchCurrentInterval
pm.mutex.Unlock()
timer.Reset(currentInterval)
logger.Debug("Holepunch monitor interval updated, reset to %v", currentInterval)
case <-timer.C:
anyStatusChanged := pm.checkHolepunchEndpoints()
localChanged := pm.checkLocalEndpoints()
anyStatusChanged := pm.checkHolepunchEndpoints() || localChanged
pm.mutex.Lock()
if anyStatusChanged {
@@ -560,6 +896,140 @@ func (pm *PeerMonitor) runHolepunchMonitor() {
}
}
// checkLocalEndpoints tests local network endpoints for sites that have them configured.
// For a site not currently using a local endpoint, it probes each candidate in order
// (candidates are ordered best-to-worst by the server) and switches to the first one that
// succeeds. For a site already using a local endpoint, it re-tests that endpoint and falls
// back to the normal public/relay logic after a few consecutive failures.
// Returns true if any site's local-connection status changed.
func (pm *PeerMonitor) checkLocalEndpoints() bool {
pm.mutex.Lock()
if !pm.running {
pm.mutex.Unlock()
return false
}
if pm.holepunchTester == nil {
pm.mutex.Unlock()
return false
}
candidates := make(map[int][]string, len(pm.localEndpoints))
for siteID, eps := range pm.localEndpoints {
if len(eps) > 0 {
candidates[siteID] = eps
}
}
active := make(map[int]string, len(pm.localActiveEndpoint))
for siteID, ep := range pm.localActiveEndpoint {
active[siteID] = ep
}
timeout := pm.localTestTimeout
maxAttempts := pm.holepunchMaxAttempts
pm.mutex.Unlock()
anyChanged := false
for siteID, endpoints := range candidates {
if activeEndpoint := active[siteID]; activeEndpoint != "" {
// Already using a local endpoint - verify it's still working.
result := pm.holepunchTester.TestEndpoint(activeEndpoint, timeout)
pm.mutex.Lock()
if _, stillTracked := pm.localEndpoints[siteID]; !stillTracked {
pm.mutex.Unlock()
continue // peer was removed while we were testing
}
if result.Success {
pm.localFailures[siteID] = 0
pm.mutex.Unlock()
continue
}
pm.localFailures[siteID]++
failureCount := pm.localFailures[siteID]
pm.mutex.Unlock()
if failureCount >= maxAttempts {
logger.Warn("Local endpoint %s for site %d failed %d times, falling back to public/relay logic", activeEndpoint, siteID, failureCount)
pm.mutex.Lock()
pm.localActiveEndpoint[siteID] = ""
pm.localFailures[siteID] = 0
pm.holepunchFailures[siteID] = 0 // don't immediately re-trigger relay on stale failures
// The holepunch backoff timer keeps climbing while a local endpoint is
// active (checkHolepunchEndpoints skips those sites but backoff still
// applies), so reset it here to avoid the resumed public/relay logic
// being stuck polling at a stale, heavily-backed-off interval.
pm.holepunchCurrentInterval = pm.holepunchMinInterval
publicEndpoint := pm.holepunchEndpoints[siteID]
remainingLocal := remainingLocalCandidates(pm.localEndpoints[siteID], activeEndpoint)
pm.mutex.Unlock()
anyChanged = true
pm.deactivateLocalEndpoint(siteID)
// Don't wait out the next backed-off checkHolepunchEndpoints tick to find out
// whether the public endpoint is reachable - rapidly test it (and any untried
// local candidates) now so a total connectivity loss triggers relay within
// ~1-2 seconds instead of potentially tens of seconds.
go pm.rapidTestOnLocalFallback(siteID, publicEndpoint, remainingLocal)
}
continue
}
// Not currently using a local endpoint - probe candidates in order.
for _, endpoint := range endpoints {
result := pm.holepunchTester.TestEndpoint(endpoint, timeout)
pm.mutex.Lock()
if _, stillTracked := pm.localEndpoints[siteID]; !stillTracked {
pm.mutex.Unlock()
break // peer was removed while we were testing
}
if !result.Success {
pm.mutex.Unlock()
continue
}
pm.localActiveEndpoint[siteID] = endpoint
pm.localFailures[siteID] = 0
pm.mutex.Unlock()
logger.Info("Local endpoint %s for site %d is reachable (RTT: %v), switching to local connection", endpoint, siteID, result.RTT)
anyChanged = true
pm.activateLocalEndpoint(siteID, endpoint)
break
}
}
return anyChanged
}
// activateLocalEndpoint invokes the switch callback and notifies the server that a local
// endpoint became active for the given site.
func (pm *PeerMonitor) activateLocalEndpoint(siteID int, endpoint string) {
pm.mutex.Lock()
cb := pm.localSwitchCallback
pm.mutex.Unlock()
if cb != nil {
cb(siteID, endpoint)
}
pm.sendLocal(siteID, endpoint)
}
// deactivateLocalEndpoint invokes the fallback callback and notifies the server that the
// given site fell back from its local endpoint.
func (pm *PeerMonitor) deactivateLocalEndpoint(siteID int) {
pm.mutex.Lock()
cb := pm.localFallbackCallback
pm.mutex.Unlock()
if cb != nil {
cb(siteID)
}
pm.sendUnLocal(siteID)
}
// checkHolepunchEndpoints tests all holepunch endpoints
// Returns true if any endpoint's status changed
func (pm *PeerMonitor) checkHolepunchEndpoints() bool {
@@ -571,6 +1041,9 @@ func (pm *PeerMonitor) checkHolepunchEndpoints() bool {
}
endpoints := make(map[int]string, len(pm.holepunchEndpoints))
for siteID, endpoint := range pm.holepunchEndpoints {
if pm.localActiveEndpoint[siteID] != "" {
continue // using a local connection, skip public/relay monitoring
}
endpoints[siteID] = endpoint
}
timeout := pm.holepunchTimeout
@@ -628,8 +1101,10 @@ func (pm *PeerMonitor) checkHolepunchEndpoints() bool {
wgConnected := pm.wgConnectionStatus[siteID]
pm.mutex.Unlock()
// Update API - use holepunch endpoint and relay status
pm.apiServer.UpdatePeerStatus(siteID, wgConnected, result.RTT, endpoint, isRelayed)
// Update API - use holepunch endpoint and relay status. Sites with an active
// local endpoint are filtered out of this loop above, so isLocal is always
// false here.
pm.apiServer.UpdatePeerStatus(siteID, wgConnected, result.RTT, endpoint, isRelayed, false)
}
// Handle relay logic based on holepunch status
@@ -677,6 +1152,26 @@ func (pm *PeerMonitor) Close() {
// Stop holepunch monitor first (outside of mutex to avoid deadlock)
pm.stopHolepunchMonitor()
// 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 local-connection senders
pm.localSendMu.Lock()
for chainId, stop := range pm.localSends {
if stop != nil {
stop()
}
delete(pm.localSends, chainId)
}
pm.localSendMu.Unlock()
pm.mutex.Lock()
defer pm.mutex.Unlock()
+20 -12
View File
@@ -10,17 +10,26 @@ import (
"golang.zx2c4.com/wireguard/wgctrl/wgtypes"
)
// ConfigurePeer sets up or updates a peer within the WireGuard device
func ConfigurePeer(dev *device.Device, siteConfig SiteConfig, privateKey wgtypes.Key, relay bool, persistentKeepalive int) error {
var endpoint string
if relay && siteConfig.RelayEndpoint != "" {
endpoint = formatEndpoint(siteConfig.RelayEndpoint)
// ConfigurePeer sets up or updates a peer within the WireGuard device.
// If siteConfig.ActiveLocalEndpoint is set, it takes priority over both the relay and the
// public endpoint since it's a directly-reachable address on the site host's local network.
func ConfigurePeer(dev *device.Device, siteConfig SiteConfig, privateKey wgtypes.Key, relay bool, persistentKeepalive int, publicDNS []string) error {
var siteHost string
if siteConfig.ActiveLocalEndpoint != "" {
// Local endpoints are already literal ip:port pairs on the local network, no DNS resolution needed.
siteHost = siteConfig.ActiveLocalEndpoint
} else {
endpoint = formatEndpoint(siteConfig.Endpoint)
}
siteHost, err := util.ResolveDomain(endpoint)
if err != nil {
return fmt.Errorf("failed to resolve endpoint for site %d: %v", siteConfig.SiteId, err)
var endpoint string
if relay && siteConfig.RelayEndpoint != "" {
endpoint = formatEndpoint(siteConfig.RelayEndpoint)
} else {
endpoint = formatEndpoint(siteConfig.Endpoint)
}
var err error
siteHost, err = util.ResolveDomainUpstream(endpoint, publicDNS)
if err != nil {
return fmt.Errorf("failed to resolve endpoint for site %d: %v", siteConfig.SiteId, err)
}
}
// Split off the CIDR of the server IP which is just a string and add /32 for the allowed IP
@@ -66,8 +75,7 @@ func ConfigurePeer(dev *device.Device, siteConfig SiteConfig, privateKey wgtypes
config := configBuilder.String()
logger.Debug("Configuring peer with config: %s", config)
err = dev.IpcSet(config)
if err != nil {
if err := dev.IpcSet(config); err != nil {
return fmt.Errorf("failed to configure WireGuard peer: %v", err)
}
+22 -10
View File
@@ -8,16 +8,21 @@ type PeerAction struct {
// UpdatePeerData represents the data needed to update a peer
type SiteConfig struct {
SiteId int `json:"siteId"`
Name string `json:"name,omitempty"`
Endpoint string `json:"endpoint,omitempty"`
RelayEndpoint string `json:"relayEndpoint,omitempty"`
PublicKey string `json:"publicKey,omitempty"`
ServerIP string `json:"serverIP,omitempty"`
ServerPort uint16 `json:"serverPort,omitempty"`
RemoteSubnets []string `json:"remoteSubnets,omitempty"` // optional, array of subnets that this site can access
AllowedIps []string `json:"allowedIps,omitempty"` // optional, array of allowed IPs for the peer
Aliases []Alias `json:"aliases,omitempty"` // optional, array of alias configurations
SiteId int `json:"siteId"`
Name string `json:"name,omitempty"`
Endpoint string `json:"endpoint,omitempty"`
LocalEndpoints []string `json:"localEndpoints,omitempty"` // optional, ip:port endpoints on the site host's local network interfaces, ordered best-to-worst
RelayEndpoint string `json:"relayEndpoint,omitempty"`
PublicKey string `json:"publicKey,omitempty"`
ServerIP string `json:"serverIP,omitempty"`
ServerPort uint16 `json:"serverPort,omitempty"`
RemoteSubnets []string `json:"remoteSubnets,omitempty"` // optional, array of subnets that this site can access
AllowedIps []string `json:"allowedIps,omitempty"` // optional, array of allowed IPs for the peer
Aliases []Alias `json:"aliases,omitempty"` // optional, array of alias configurations
// ActiveLocalEndpoint tracks the local network endpoint currently in use for this
// peer, if any. Not part of the wire protocol; set internally by the PeerManager.
ActiveLocalEndpoint string `json:"-"`
}
type Alias struct {
@@ -41,6 +46,13 @@ type UnRelayPeerData struct {
Endpoint string `json:"endpoint"`
}
// LocalPeerAckData represents the server's acknowledgement of an "olm/wg/local" or
// "olm/wg/unlocal" message. olm has already applied the local connection switch by the time
// it sends the notification, so the ack is only used to stop the retry sender.
type LocalPeerAckData struct {
SiteId int `json:"siteId"`
}
// PeerAdd represents the data needed to add remote subnets to a peer
type PeerAdd struct {
SiteId int `json:"siteId"`
+81 -6
View File
@@ -2,6 +2,7 @@ package websocket
import (
"bytes"
"compress/gzip"
"crypto/tls"
"crypto/x509"
"encoding/json"
@@ -21,6 +22,14 @@ import (
"github.com/gorilla/websocket"
)
// writeDeadline bounds how long a websocket write may block before it is
// treated as a failure. Without this, a write to a TCP connection whose
// underlying network interface has disappeared (e.g. laptop sleep/resume,
// Wi-Fi roam) can sit buffered in the kernel for minutes without erroring,
// which prevents the ping monitor from ever detecting the dead connection
// and reconnecting.
const writeDeadline = 10 * time.Second
// AuthError represents an authentication/authorization error (401/403)
type AuthError struct {
StatusCode int
@@ -82,7 +91,7 @@ type Client struct {
isDisconnected bool // Flag to track if client is intentionally disconnected
reconnectMux sync.RWMutex
pingInterval time.Duration
pingTimeout time.Duration
pongWait time.Duration // read deadline window; if no pong/message arrives within it, the connection is considered dead
onConnect func() error
onTokenUpdate func(token string, exitNodes []ExitNode)
onAuthError func(statusCode int, message string) // Callback for auth errors
@@ -158,7 +167,7 @@ func (c *Client) OnAuthError(callback func(statusCode int, message string)) {
}
// NewClient creates a new websocket client
func NewClient(ID, secret, userToken, orgId, endpoint string, pingInterval time.Duration, pingTimeout time.Duration, opts ...ClientOption) (*Client, error) {
func NewClient(ID, secret, userToken, orgId, endpoint string, pingInterval time.Duration, opts ...ClientOption) (*Client, error) {
config := &Config{
ID: ID,
Secret: secret,
@@ -167,6 +176,16 @@ func NewClient(ID, secret, userToken, orgId, endpoint string, pingInterval time.
OrgID: orgId,
}
// Read deadline window: must exceed pingInterval so a healthy connection
// (which gets a pong/message at least every pingInterval) is never torn
// down, but a dead/half-open one — including one where writes keep
// "succeeding" because small pings fit in the kernel send buffer even
// under total packet loss — is detected within ~2 ping cycles.
pongWait := pingInterval * 2
if pongWait < 20*time.Second {
pongWait = 20 * time.Second
}
client := &Client{
config: config,
baseURL: endpoint, // default value
@@ -175,7 +194,7 @@ func NewClient(ID, secret, userToken, orgId, endpoint string, pingInterval time.
reconnectInterval: 3 * time.Second,
isConnected: false,
pingInterval: pingInterval,
pingTimeout: pingTimeout,
pongWait: pongWait,
clientType: "olm",
pingDone: make(chan struct{}),
}
@@ -269,6 +288,9 @@ func (c *Client) SendMessage(messageType string, data interface{}) error {
c.writeMux.Lock()
defer c.writeMux.Unlock()
if err := c.conn.SetWriteDeadline(time.Now().Add(writeDeadline)); err != nil {
return err
}
return c.conn.WriteJSON(msg)
}
@@ -388,6 +410,7 @@ func (c *Client) getToken() (string, []ExitNode, error) {
tokenData := map[string]interface{}{
"olmId": c.config.ID,
"secret": c.config.Secret,
"userToken": c.config.UserToken,
"orgId": c.config.OrgID,
}
jsonData, err := json.Marshal(tokenData)
@@ -582,6 +605,18 @@ func (c *Client) establishConnection() error {
c.conn = conn
c.setConnected(true)
// Arm a read deadline and refresh it whenever a pong arrives. Combined with
// the protocol-level ping sent alongside the app-level one in sendPing,
// this detects a dead or half-open connection (e.g. the route disappearing
// on sleep/resume, or total packet loss) that a write-side check alone
// misses: small periodic pings fit in the kernel send buffer and keep
// "succeeding" even when nothing is actually reaching the peer.
_ = c.conn.SetReadDeadline(time.Now().Add(c.pongWait))
c.conn.SetPongHandler(func(appData string) error {
_ = c.conn.SetReadDeadline(time.Now().Add(c.pongWait))
return nil
})
// Note: ping monitor is NOT started here - it will be started when
// StartPingMonitor() is called after registration completes
@@ -697,7 +732,17 @@ func (c *Client) sendPing() {
logger.Debug("websocket: Sending ping: %+v", pingMsg)
c.writeMux.Lock()
err := c.conn.WriteJSON(pingMsg)
err := c.conn.SetWriteDeadline(time.Now().Add(writeDeadline))
if err == nil {
err = c.conn.WriteJSON(pingMsg)
}
if err == nil {
// Protocol-level ping: a standards-compliant server replies with a
// PONG, which refreshes the read deadline via SetPongHandler. This is
// what actually detects a half-open connection where writes still
// "succeed" (buffered by the kernel) but nothing is reaching the peer.
_ = c.conn.WriteControl(websocket.PingMessage, nil, time.Now().Add(writeDeadline))
}
c.writeMux.Unlock()
if err != nil {
// Check if we're shutting down before logging error and reconnecting
@@ -802,8 +847,14 @@ func (c *Client) readPumpWithDisconnectDetection() {
case <-c.done:
return
default:
var msg WSMessage
err := c.conn.ReadJSON(&msg)
messageType, p, err := c.conn.ReadMessage()
if err == nil {
// Any inbound traffic means the peer is alive — extend the
// read deadline (also covers servers that answer the
// app-level "olm/ping" with a message rather than a
// protocol pong).
_ = c.conn.SetReadDeadline(time.Now().Add(c.pongWait))
}
if err != nil {
// Check if we're shutting down or explicitly disconnected before logging error
select {
@@ -828,6 +879,30 @@ func (c *Client) readPumpWithDisconnectDetection() {
}
}
// Decompress binary frames (gzip-compressed JSON)
var data []byte
if messageType == websocket.BinaryMessage {
gr, gzErr := gzip.NewReader(bytes.NewReader(p))
if gzErr != nil {
logger.Error("websocket: failed to create gzip reader: %v", gzErr)
continue
}
data, gzErr = io.ReadAll(gr)
gr.Close()
if gzErr != nil {
logger.Error("websocket: failed to decompress message: %v", gzErr)
continue
}
} else {
data = p
}
var msg WSMessage
if err = json.Unmarshal(data, &msg); err != nil {
logger.Error("websocket: failed to parse message: %v", err)
continue
}
// Update config version from incoming message
c.setConfigVersion(msg.ConfigVersion)